RAG/sync/ast_parser.py

236 lines
7.0 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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