RAG/sync/git_sync.py

364 lines
13 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.

"""
Git代码库同步子类
继承BaseSync实现代码拉取/增量同步/函数解析
"""
import os
from typing import List, Dict, Any, Set, Optional
from datetime import datetime
from loguru import logger
from config import BaseDataSourceConfig, GitDataSourceConfig, settings
from sync.base_sync import BaseSync
from sync.ast_parser import ASTParser
from utils.git_tool import GitTool
from utils.func_id_generator import generate_func_unique_id
class GitSync(BaseSync):
def __init__(self, config: GitDataSourceConfig, vector_store_manager=None):
"""
初始化Git同步器
Args:
config: Git数据源配置
vector_store_manager: 向量存储管理器
"""
super().__init__(config, vector_store_manager)
self.config = config
# 初始化Git工具
self.git_tool = GitTool(
user_id="default", # 暂时使用默认用户ID
repo_id=config.name,
git_config={
"git_url": config.git_url,
"branch": config.branch,
"ssh_key": config.ssh_key,
"https_token": config.https_token,
"local_repo_path": config.local_repo_path
}
)
def fetch_all_documents(self) -> List[Dict[str, Any]]:
"""
获取所有文档(函数)
Returns:
List[Dict[str, Any]]: 函数信息列表
"""
# 克隆仓库
self.git_tool.clone_repo()
# 扫描仓库文件
func_list = []
support_lang = self.git_tool._detect_support_lang()
for root, dirs, files in os.walk(self.git_tool.local_repo_path):
# 跳过.git目录
if ".git" in dirs:
dirs.remove(".git")
for file in files:
file_path = os.path.join(root, file)
# 检测文件语言
lang = ASTParser.detect_language(file_path)
if lang and lang in support_lang:
# 解析文件中的函数
parser = ASTParser(file_path, lang)
try:
functions = parser.parse_functions()
# 为每个函数生成doc_id并设置到字典中
for func in functions:
# 生成唯一的文档ID
func_id = generate_func_unique_id(
user_id="default",
repo_id=self.config.name,
branch=self.config.branch,
file_path=func["file_path"],
class_name=func.get("class_name"),
func_name=func["func_name"]
)
func['id'] = func_id
func_list.extend(functions)
except Exception as e:
logger.error(f"解析文件失败 {file_path}: {e}")
logger.info(f"获取到 {len(func_list)} 个函数")
return func_list
def doc_to_llamaindex_doc(self, doc: Dict) -> 'Document':
"""
转换函数信息为LlamaIndex Document
Args:
doc: 函数信息
Returns:
Document: LlamaIndex Document对象
"""
from llama_index.core import Document
# 生成函数唯一ID
func_id = doc.get('id')
# 生成函数描述
func_desc = self.generate_func_desc(doc)
# 获取函数体
func_body = doc.get("func_body", "")
# 创建Document对象
document = Document(
text=func_desc, # 使用描述作为文本
id_=func_id,
metadata={
"func_id": func_id,
"func_name": doc["func_name"],
"class_name": doc.get("class_name") if doc.get("class_name")!=None else "None",
"file_path": doc["file_path"],
"lang": doc["lang"],
"params": len(doc.get("params", [])), # 只存储参数数量,不存储完整参数列表
"return_type": doc.get("return_type") if doc.get("return_type")!=None else "None",
"docstring": (doc.get("docstring") or "")[:200], # 进一步限制文档字符串长度
"start_line": doc.get("start_line"),
"end_line": doc.get("end_line"),
"repo_id": self.config.name,
"branch": self.config.branch,
"func_body": func_body[:1000] # 限制函数体长度避免metadata过长
}
)
return document
def generate_func_desc(self, func_info: Dict) -> str:
"""
生成函数描述
Args:
func_info: 函数信息
Returns:
str: 函数描述
"""
# 构建函数描述
parts = []
# 函数类型
if func_info.get("class_name"):
parts.append(f"{func_info['class_name']}类的{func_info['func_name']}方法")
else:
parts.append(f"{func_info['func_name']}函数")
# 参数信息
params = func_info.get("params", [])
if params:
param_str = []
for param in params:
if param.get("type"):
param_str.append(f"{param['name']}: {param['type']}")
else:
param_str.append(param['name'])
parts.append(f"接收参数: {', '.join(param_str)}")
# 返回值信息
return_type = func_info.get("return_type")
if return_type:
parts.append(f"返回类型: {return_type}")
# 文档字符串
docstring = func_info.get("docstring")
if docstring:
parts.append(f"功能描述: {docstring.strip()}")
else:
# 如果没有文档字符串,使用本地大模型根据函数体生成描述
parts.append(f"功能描述: {self._generate_docstring_from_body(func_info.get('func_body', ''))}")
return ". ".join(parts)
def _generate_docstring_from_body(self, func_body: str) -> str:
"""
使用本地大模型根据函数体生成文档字符串
Args:
func_body: 函数体代码
Returns:
str: 生成的文档字符串
"""
if not func_body:
return "无文档字符串"
try:
from config import settings
from llama_index.llms.ollama import Ollama
# 初始化Ollama LLM
llm = Ollama(
model=settings.OLLAMA_MODEL,
base_url=settings.OLLAMA_BASE_URL,
temperature=0.3, # 降低温度,生成更确定的结果
request_timeout=300.0
)
# 构建提示词
prompt = f"""
请为以下函数生成简洁的文档字符串,描述其功能、参数和返回值:
{func_body}
要求:
1. 语言简洁明了不超过100字
2. 只返回文档字符串内容,不要包含其他内容
3. 重点描述函数的核心功能
"""
# 生成文档字符串
response = llm.complete(prompt)
generated_docstring = response.text.strip()
# 限制长度
if len(generated_docstring) > 200:
generated_docstring = generated_docstring[:200] + "..."
logger.debug(f"生成的文档字符串: {generated_docstring}")
return generated_docstring
except Exception as e:
logger.error(f"生成文档字符串失败: {e}")
# 降级方案:返回基于函数名的简单描述
return "执行相关操作的函数"
def fetch_new_documents(self, last_sync_time=None) -> List[Dict[str, Any]]:
"""
获取新文档(增量同步)
Args:
last_sync_time: 上次同步时间
Returns:
List[Dict[str, Any]]: 新函数信息列表
"""
# 检测远程更新
has_update, local_commit, remote_commit = self.git_tool.detect_remote_update()
if not has_update:
logger.info("Git仓库无更新")
return []
# 增量拉取
delta_files = self.git_tool.incremental_pull(local_commit, remote_commit)
# 解析新增/修改的文件
func_list = []
for file_path in delta_files.get("ADD", []) + delta_files.get("MODIFY", []):
lang = ASTParser.detect_language(file_path)
if lang:
parser = ASTParser(file_path, lang)
try:
functions = parser.parse_functions()
# 为每个函数生成doc_id并设置到字典中
for func in functions:
# 生成唯一的文档ID
func_id = generate_func_unique_id(
user_id="default",
repo_id=self.config.name,
branch=self.config.branch,
file_path=func["file_path"],
class_name=func.get("class_name"),
func_name=func["func_name"]
)
func['id'] = func_id
func_list.extend(functions)
except Exception as e:
logger.error(f"解析文件失败 {file_path}: {e}")
logger.info(f"增量同步获取到 {len(func_list)} 个函数")
return func_list
def get_synced_document_ids(self) -> Set[str]:
"""
获取已同步的文档ID
Returns:
Set[str]: 文档ID集合
"""
# 从向量存储中获取已同步的函数ID
if not self.vector_store_manager:
return set()
try:
# 获取所有已存在的文档ID
all_doc_ids = self.vector_store_manager.get_existing_doc_ids()
# 过滤出与当前Git仓库相关的文档ID
synced_ids = set()
# 获取所有文档的元数据,用于过滤
results = self.vector_store_manager.collection.get(include=['metadatas'])
metadatas = results.get('metadatas', [])
ids = results.get('ids', [])
for doc_id, metadata in zip(ids, metadatas):
if metadata and metadata.get('repo_id') == self.config.name:
synced_ids.add(doc_id)
logger.info(f"获取到 {len(synced_ids)} 个已同步的Git函数ID")
return synced_ids
except Exception as e:
logger.error(f"获取已同步文档ID失败: {e}")
return set()
def generate_doc_id(self, identifier: str) -> str:
"""
生成唯一的文档ID
Args:
identifier: 文档的唯一标识符(文件路径等)
Returns:
str: 唯一的文档ID
"""
from utils.func_id_generator import generate_func_unique_id
# 对于Git数据源使用函数唯一ID生成器
# 假设identifier是文件路径
return generate_func_unique_id(
user_id="default",
repo_id=self.config.name,
branch=self.config.branch,
file_path=identifier,
class_name="",
func_name=identifier.split('/')[-1].split('.')[0]
)
@staticmethod
def check_data_source_exists(config: BaseDataSourceConfig) -> bool:
"""
检查数据源是否存在
Args:
config: 数据源配置
Returns:
bool: 是否存在
"""
try:
# 尝试克隆仓库
git_tool = GitTool(
user_id="default",
repo_id=config.name,
git_config={
"git_url": config.git_url,
"branch": config.branch,
"ssh_key": config.ssh_key,
"https_token": config.https_token
}
)
git_tool.clone_repo()
logger.info(f"Git数据源检查成功: {config.name}")
return True
except Exception as e:
logger.error(f"Git数据源检查失败: {e}")
return False