如何从Python脚本(.py/.ipynb)中提取模型输入数据源位置
解决方案
核心思路
用正则表达式匹配目标函数的参数,针对.py和.ipynb两种文件格式分别处理:前者直接读取文本内容匹配,后者先解析JSON提取代码单元格内容再匹配。
1. 处理Python脚本(.py)
直接读取文件文本,用正则匹配pd.read_csv的路径参数和pd.read_sql_query的SQL语句,再从SQL中提取表名。
正则规则
匹配
pd.read_csv的文件路径:regex_csv = r"pd\.read_csv\(['\"](.*?)['\"]"该规则会捕获单/双引号包裹的第一个字符串参数,也就是代码中直接写死的CSV路径。
匹配
pd.read_sql_query的SQL语句并提取表名:
先捕获SQL语句内容:regex_sql = r"pd\.read_sql_query\(['\"](.*?)['\"]"再从SQL中提取表名(支持
库.模式.表、模式.表、表三种常见格式):regex_table = r"FROM\s+(\w+\.\w+\.\w+|\w+\.\w+|\w+)"
示例代码
import re def extract_py_sources(file_path): sources = {"csv_paths": [], "sql_tables": []} with open(file_path, "r", encoding="utf-8", errors="ignore") as f: content = f.read() # 提取pd.read_csv路径 csv_matches = re.findall(regex_csv, content, re.IGNORECASE) sources["csv_paths"].extend(csv_matches) # 提取pd.read_sql_query的表名 sql_matches = re.findall(regex_sql, content, re.IGNORECASE) for sql in sql_matches: table_matches = re.findall(regex_table, sql, re.IGNORECASE) sources["sql_tables"].extend(table_matches) return sources
2. 处理Jupyter Notebook(.ipynb)
.ipynb是JSON格式,需要先解析出所有代码单元格的源代码,再套用和.py文件相同的正则规则匹配。
示例代码
import json def extract_ipynb_sources(file_path): sources = {"csv_paths": [], "sql_tables": []} with open(file_path, "r", encoding="utf-8") as f: nb_data = json.load(f) # 遍历所有代码单元格 for cell in nb_data["cells"]: if cell["cell_type"] == "code": cell_content = "\n".join(cell["source"]) # 提取csv路径 csv_matches = re.findall(regex_csv, cell_content, re.IGNORECASE) sources["csv_paths"].extend(csv_matches) # 提取sql表名 sql_matches = re.findall(regex_sql, cell_content, re.IGNORECASE) for sql in sql_matches: table_matches = re.findall(regex_table, sql, re.IGNORECASE) sources["sql_tables"].extend(table_matches) return sources
3. 整合遍历逻辑
把上述函数和你已有的文件遍历逻辑结合,覆盖所有目标文件:
import glob import os def scan_all_models(root_dir): # 匹配所有.py和.ipynb文件 py_files = glob.glob(os.path.join(root_dir, "**/*.py"), recursive=True) ipynb_files = glob.glob(os.path.join(root_dir, "**/*.ipynb"), recursive=True) all_sources = [] # 处理.py文件 for py_file in py_files: sources = extract_py_sources(py_file) if sources["csv_paths"] or sources["sql_tables"]: all_sources.append({"file": py_file, **sources}) # 处理.ipynb文件(可保留你已有的含csv关键词筛选逻辑) for ipynb_file in ipynb_files: # 可选:跳过无csv相关代码的文件 with open(ipynb_file, "r", encoding="utf-8") as f: nb_data = json.load(f) has_csv = any("csv" in "\n".join(cell["source"]) for cell in nb_data["cells"] if cell["cell_type"] == "code") if not has_csv: continue sources = extract_ipynb_sources(ipynb_file) if sources["csv_paths"] or sources["sql_tables"]: all_sources.append({"file": ipynb_file, **sources}) return all_sources # 调用示例 root = "/path/to/your/models" result = scan_all_models(root) for item in result: print(f"文件: {item['file']}") if item["csv_paths"]: print(f"CSV路径: {', '.join(item['csv_paths'])}") if item["sql_tables"]: print(f"SQL表名: {', '.join(item['sql_tables'])}") print("---")
扩展优化点
- 处理函数别名:如果代码中用了
from pandas import read_csv、import pandas as pd2这类写法,可以扩展正则规则,比如匹配(pd|pandas|read_csv)\.read_csv或者直接匹配read_csv\(; - 解析动态参数:如果数据源是通过变量传递的(比如
pd.read_csv(csv_path)),正则无法直接捕获,需要用Python的ast模块解析语法树,遍历函数调用节点分析参数; - 去重处理:同一个文件中可能重复引用同一数据源,可以用
set去重后转列表。
内容的提问来源于stack exchange,提问作者Joshua_1980
相关产品推荐
相关产品推荐

