236 lines
7.0 KiB
Python
236 lines
7.0 KiB
Python
"""
|
||
代码AST解析工具类
|
||
实现跨语言函数级切片
|
||
"""
|
||
import ast
|
||
import os
|
||
from typing import List, Dict, Optional
|
||
from loguru import logger
|
||
|
||
|
||
class ASTParser:
|
||
def __init__(self, file_path: str, lang: str):
|
||
"""
|
||
初始化AST解析器
|
||
|
||
Args:
|
||
file_path: 文件路径
|
||
lang: 编程语言
|
||
"""
|
||
self.file_path = file_path
|
||
self.lang = lang
|
||
self.func_list: List[Dict] = [] # 提取的函数列表
|
||
|
||
def parse_functions(self) -> List[Dict]:
|
||
"""
|
||
统一入口:根据语言调用对应解析方法
|
||
|
||
Returns:
|
||
List[Dict]: 函数信息列表
|
||
"""
|
||
if not os.path.exists(self.file_path):
|
||
raise Exception(f"文件不存在: {self.file_path}")
|
||
|
||
# 读取文件内容
|
||
try:
|
||
with open(self.file_path, "r", encoding="utf-8") as f:
|
||
self.code = f.read()
|
||
except Exception as e:
|
||
logger.error(f"读取文件失败: {e}")
|
||
raise
|
||
|
||
# 按语言解析
|
||
if self.lang == "python":
|
||
self._parse_python()
|
||
elif self.lang == "java":
|
||
self._parse_java()
|
||
elif self.lang == "go":
|
||
self._parse_go()
|
||
elif self.lang == "javascript" or self.lang == "typescript":
|
||
self._parse_javascript()
|
||
else:
|
||
logger.warning(f"暂不支持的编程语言: {self.lang}")
|
||
|
||
logger.info(f"解析文件 {self.file_path},提取到 {len(self.func_list)} 个函数")
|
||
return self.func_list
|
||
|
||
def _parse_python(self):
|
||
"""
|
||
解析Python代码
|
||
"""
|
||
try:
|
||
tree = ast.parse(self.code)
|
||
for node in ast.walk(tree):
|
||
# 提取函数定义(普通函数/类方法/异步函数)
|
||
if isinstance(node, ast.FunctionDef) or isinstance(node, ast.AsyncFunctionDef):
|
||
func_info = self._extract_python_func_info(node)
|
||
self.func_list.append(func_info)
|
||
except SyntaxError as e:
|
||
logger.error(f"Python代码语法错误: {e}")
|
||
raise Exception(f"Python代码语法错误: {e}")
|
||
|
||
def _extract_python_func_info(self, node) -> Dict:
|
||
"""
|
||
提取Python函数的标准化信息
|
||
|
||
Args:
|
||
node: AST节点
|
||
|
||
Returns:
|
||
Dict: 函数信息
|
||
"""
|
||
# 提取函数名
|
||
func_name = node.name
|
||
|
||
# 提取参数
|
||
params = []
|
||
for arg in node.args.args:
|
||
param_info = {
|
||
"name": arg.arg,
|
||
"type": None
|
||
}
|
||
# 提取类型注解
|
||
if arg.annotation:
|
||
try:
|
||
param_info["type"] = ast.unparse(arg.annotation)
|
||
except Exception:
|
||
pass
|
||
params.append(param_info)
|
||
|
||
# 提取返回值类型
|
||
return_type = None
|
||
if node.returns:
|
||
try:
|
||
return_type = ast.unparse(node.returns)
|
||
except Exception:
|
||
pass
|
||
|
||
# 提取函数体代码
|
||
func_body = self._get_func_body(node)
|
||
|
||
# 提取所属类名
|
||
class_name = None
|
||
parent = node
|
||
while hasattr(parent, "parent"):
|
||
parent = parent.parent
|
||
if isinstance(parent, ast.ClassDef):
|
||
class_name = parent.name
|
||
break
|
||
|
||
# 提取函数文档字符串
|
||
docstring = ast.get_docstring(node)
|
||
|
||
return {
|
||
"file_path": self.file_path,
|
||
"lang": "python",
|
||
"func_name": func_name,
|
||
"class_name": class_name,
|
||
"params": params,
|
||
"return_type": return_type,
|
||
"func_body": func_body,
|
||
"docstring": docstring,
|
||
"start_line": node.lineno,
|
||
"end_line": node.end_lineno
|
||
}
|
||
|
||
def _parse_java(self):
|
||
"""
|
||
解析Java代码
|
||
注意:这里使用简单的正则解析,实际项目中建议使用专业的Java解析库
|
||
"""
|
||
logger.warning("Java解析功能暂未完全实现,使用简单的正则解析")
|
||
# TODO: 实现Java代码的AST解析
|
||
|
||
def _parse_go(self):
|
||
"""
|
||
解析Go代码
|
||
注意:这里使用简单的正则解析,实际项目中建议使用专业的Go解析库
|
||
"""
|
||
logger.warning("Go解析功能暂未完全实现,使用简单的正则解析")
|
||
# TODO: 实现Go代码的AST解析
|
||
|
||
def _parse_javascript(self):
|
||
"""
|
||
解析JavaScript/TypeScript代码
|
||
注意:这里使用简单的正则解析,实际项目中建议使用专业的JS解析库
|
||
"""
|
||
logger.warning("JavaScript解析功能暂未完全实现,使用简单的正则解析")
|
||
# TODO: 实现JavaScript代码的AST解析
|
||
|
||
def _get_func_body(self, node) -> str:
|
||
"""
|
||
获取函数体代码
|
||
|
||
Args:
|
||
node: AST节点
|
||
|
||
Returns:
|
||
str: 函数体代码
|
||
"""
|
||
try:
|
||
# 使用ast.unparse获取函数体代码
|
||
return ast.unparse(node)
|
||
except Exception:
|
||
# 降级方案:根据行号提取代码
|
||
lines = self.code.splitlines()
|
||
start_line = node.lineno - 1 # 转换为0-based索引
|
||
end_line = node.end_lineno # 转换为0-based索引
|
||
if start_line >= 0 and end_line <= len(lines):
|
||
return "\n".join(lines[start_line:end_line])
|
||
return ""
|
||
|
||
@staticmethod
|
||
def detect_language(file_path: str) -> Optional[str]:
|
||
"""
|
||
根据文件扩展名检测编程语言
|
||
|
||
Args:
|
||
file_path: 文件路径
|
||
|
||
Returns:
|
||
Optional[str]: 编程语言
|
||
"""
|
||
ext = os.path.splitext(file_path)[1].lower()
|
||
|
||
lang_map = {
|
||
".py": "python",
|
||
".java": "java",
|
||
".go": "go",
|
||
".js": "javascript",
|
||
".ts": "typescript",
|
||
".jsx": "javascript",
|
||
".tsx": "typescript",
|
||
".c": "c",
|
||
".cpp": "cpp",
|
||
".h": "c",
|
||
".hpp": "cpp",
|
||
".cs": "csharp",
|
||
".rs": "rust",
|
||
".php": "php",
|
||
".rb": "ruby",
|
||
".swift": "swift",
|
||
".kt": "kotlin",
|
||
".scala": "scala"
|
||
}
|
||
|
||
return lang_map.get(ext)
|
||
|
||
@staticmethod
|
||
def parse_file(file_path: str) -> List[Dict]:
|
||
"""
|
||
静态方法:解析文件
|
||
|
||
Args:
|
||
file_path: 文件路径
|
||
|
||
Returns:
|
||
List[Dict]: 函数信息列表
|
||
"""
|
||
lang = ASTParser.detect_language(file_path)
|
||
if not lang:
|
||
logger.warning(f"无法检测文件类型: {file_path}")
|
||
return []
|
||
|
||
parser = ASTParser(file_path, lang)
|
||
return parser.parse_functions()
|