RAG/sync/ast_parser.py

236 lines
7.0 KiB
Python
Raw Permalink Normal View History

"""
代码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()