Compare commits

..

8 Commits

45 changed files with 5048 additions and 277 deletions

View File

@ -16,7 +16,7 @@ API_VERSION=1.0.0
MAX_UPLOAD_SIZE_MB=5 MAX_UPLOAD_SIZE_MB=5
# LibreOffice soffice service port (used by docker/soffice service) # LibreOffice soffice service port (used by docker/soffice service)
SOFFICE_HOST=localhost SOFFICE_HOST=rag-soffice #localhost
SOFFICE_PORT=8003 SOFFICE_PORT=8003
# ============================================ # ============================================
@ -25,32 +25,15 @@ SOFFICE_PORT=8003
# CHROMA_SERVER_HOST: ChromaDB 服务器地址 # CHROMA_SERVER_HOST: ChromaDB 服务器地址
# - 使用 host 网络模式: localhost # - 使用 host 网络模式: localhost
# - 远程服务器: 192.168.1.100 或 chromadb.example.com # - 远程服务器: 192.168.1.100 或 chromadb.example.com
CHROMA_SERVER_HOST=localhost CHROMA_SERVER_HOST=rag-chromadb #localhost
CHROMA_SERVER_PORT=8002 CHROMA_SERVER_PORT=8000 #8002
CHROMA_COLLECTION_NAME=rag_collection CHROMA_COLLECTION_NAME=rag_collection
# ============================================ # ============================================
# LLM 配置 (用于文本生成) # Ollama 配置
# ============================================ # ============================================
# LLM provider: ollama, openai 等 # OLLAMA_BASE_URL: Ollama 服务地址
LLM_PROVIDER=ollama OLLAMA_BASE_URL=http://host.docker.internal:11434
LLM_BASE_URL=http://localhost:11434/v1
LLM_MODEL=qwen3:8b
LLM_API_KEY= # 如使用openai等需要API Key的服务
# ============================================
# Embedding 配置 (用于向量检索)
# ============================================
# Embedding provider: ollama, openai 等
EMBEDDING_PROVIDER=ollama
EMBEDDING_BASE_URL=http://localhost:11434
EMBEDDING_MODEL=qwen3-embedding:0.6b
EMBEDDING_API_KEY= # 如使用openai等需要API Key的服务
# ============================================
# 兼容旧版本配置 (已废弃,仍可用但推荐使用上面的配置)
# ============================================
OLLAMA_BASE_URL=http://localhost:11434
OLLAMA_MODEL=qwen3:8b OLLAMA_MODEL=qwen3:8b
OLLAMA_EMBEDDING_MODEL=qwen3-embedding:0.6b OLLAMA_EMBEDDING_MODEL=qwen3-embedding:0.6b

60
.env.zkxdocker Normal file
View File

@ -0,0 +1,60 @@
# ============================================
# RAG API 环境变量配置文件
# ============================================
# 复制此文件为 .env 并根据实际情况修改
# 所有配置都可以通过此文件统一管理,方便不同机器之间移植
# docker-compose.yml 会自动读取此文件中的配置
# ============================================
# API 配置
# ============================================
API_HOST=0.0.0.0
API_PORT=8001
API_TITLE=RAG API
API_VERSION=1.0.0
# 文件上传大小限制单位MB默认5MB
MAX_UPLOAD_SIZE_MB=5
# LibreOffice soffice service port (used by docker/soffice service)
SOFFICE_HOST=rag-soffice #localhost
SOFFICE_PORT=8003
# ============================================
# ChromaDB 配置
# ============================================
# CHROMA_SERVER_HOST: ChromaDB 服务器地址
# - 使用 host 网络模式: localhost
# - 远程服务器: 192.168.1.100 或 chromadb.example.com
CHROMA_SERVER_HOST=rag-chromadb #localhost
CHROMA_SERVER_PORT=8000 #8002
CHROMA_COLLECTION_NAME=rag_collection
# ============================================
# Ollama 配置
# ============================================
# OLLAMA_BASE_URL: Ollama 服务地址
OLLAMA_BASE_URL=http://host.docker.internal:11434
OLLAMA_MODEL=qwen3:8b
OLLAMA_EMBEDDING_MODEL=qwen3-embedding:0.6b
# ============================================
# RAG 配置
# ============================================
EMBEDDING_DIMENSION=768
CHUNK_SIZE=1024
CHUNK_OVERLAP=200
TOP_K=5
# ============================================
# 同步配置
# ============================================
SYNC_INTERVAL=300
AUTO_SYNC=true
# ============================================
# NLTK 配置
# 推荐:如果你已在仓库内保存了 NLTK 数据(离线使用),可以设置为相对路径。例如:
# ./nltk_data/
# 或者使用本地绝对路径: /path/to/nltk_data
# ============================================
NLTK_DATA=./nltk_data/

60
.env.zkxlocal Normal file
View File

@ -0,0 +1,60 @@
# ============================================
# RAG API 环境变量配置文件
# ============================================
# 复制此文件为 .env 并根据实际情况修改
# 所有配置都可以通过此文件统一管理,方便不同机器之间移植
# docker-compose.yml 会自动读取此文件中的配置
# ============================================
# API 配置
# ============================================
API_HOST=0.0.0.0
API_PORT=8001
API_TITLE=RAG API
API_VERSION=1.0.0
# 文件上传大小限制单位MB默认5MB
MAX_UPLOAD_SIZE_MB=5
# LibreOffice soffice service port (used by docker/soffice service)
# SOFFICE_HOST=rag-soffice #localhost
SOFFICE_PORT=8003
# ============================================
# ChromaDB 配置
# ============================================
# CHROMA_SERVER_HOST: ChromaDB 服务器地址
# - 使用 host 网络模式: localhost
# - 远程服务器: 192.168.1.100 或 chromadb.example.com
# CHROMA_SERVER_HOST=rag-chromadb #localhost
CHROMA_SERVER_PORT=8000
CHROMA_COLLECTION_NAME=rag_collection
# ============================================
# Ollama 配置
# ============================================
# OLLAMA_BASE_URL: Ollama 服务地址
# OLLAMA_BASE_URL=http://host.docker.internal:11434 #http://localhost:11434
OLLAMA_MODEL=qwen3:8b
OLLAMA_EMBEDDING_MODEL=qwen3-embedding:0.6b
# ============================================
# RAG 配置
# ============================================
EMBEDDING_DIMENSION=768
CHUNK_SIZE=1024
CHUNK_OVERLAP=200
TOP_K=5
# ============================================
# 同步配置
# ============================================
SYNC_INTERVAL=10
AUTO_SYNC=true
# ============================================
# NLTK 配置
# 推荐:如果你已在仓库内保存了 NLTK 数据(离线使用),可以设置为相对路径。例如:
# ./nltk_data/
# 或者使用本地绝对路径: /path/to/nltk_data
# ============================================
NLTK_DATA=./nltk_data/

2
.gitignore vendored
View File

@ -40,7 +40,7 @@ llamaindex/
!.vscode/extensions.json !.vscode/extensions.json
!.vscode/tasks.json !.vscode/tasks.json
# Ignore user-specific VSCode files # Ignore user-specific VSCode files
.vscode/launch.json !.vscode/launch.json
.vscode/*.code-workspace .vscode/*.code-workspace
# JetBrains IDEs # JetBrains IDEs

17
.vscode/launch.json vendored Normal file
View File

@ -0,0 +1,17 @@
{
// 使 IntelliSense
//
// 访: https://go.microsoft.com/fwlink/?linkid=830387
"version": "0.2.0",
"configurations": [
{
"name": "Python: RAG Main",
"type": "debugpy",
"request": "launch",
"program": "${workspaceFolder}/main.py", // api/main.py
"cwd": "${workspaceFolder}", // S:\research\RAG
"console": "integratedTerminal",
"justMyCode": false
}
]
}

View File

@ -29,6 +29,7 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
libssl-dev \ libssl-dev \
libcrypto++-dev \ libcrypto++-dev \
libgmp-dev \ libgmp-dev \
git \
&& rm -rf /var/lib/apt/lists/* && rm -rf /var/lib/apt/lists/*
# 配置 pip 镜像源(加速 Python 包安装) # 配置 pip 镜像源(加速 Python 包安装)
@ -40,7 +41,7 @@ COPY requirements.txt /app/
# 安装 Python 依赖 # 安装 Python 依赖
RUN pip install --upgrade pip && \ RUN pip install --upgrade pip && \
pip install -r requirements.txt pip install -r requirements.txt -i https://pypi.tuna.tsinghua.edu.cn/simple
# 复制项目文件 # 复制项目文件
COPY . /app/ COPY . /app/

View File

@ -5,13 +5,14 @@
## 功能特性 ## 功能特性
- 🔍 **智能检索**: 使用 LlamaIndex 和 ChromaDB 实现高效的向量检索 - 🔍 **智能检索**: 使用 LlamaIndex 和 ChromaDB 实现高效的向量检索
- 💾 **数据同步**: 自动同步 MySQL 数据库数据到 ChromaDB 向量库 - 💾 **数据同步**: 自动同步 MySQL 数据库、本地/远程文件夹和 Git 代码库数据到 ChromaDB 向量库
- 🌊 **流式输出**: 基于 FastAPI 的流式响应,支持实时对话 - 🌊 **流式输出**: 基于 FastAPI 的流式响应,支持实时对话
- 🤖 **本地 LLM**: 集成 Ollama 本地部署的大模型 - 🤖 **本地 LLM**: 集成 Ollama 本地部署的大模型
- ⚡ **高并发**: 支持多用户同时访问 - ⚡ **高并发**: 支持多用户同时访问
- 🔄 **自动同步**: 支持定时自动同步和手动触发同步 - 🔄 **自动同步**: 支持定时自动同步和手动触发同步
- 🐳 **Docker 部署**: 使用 Docker Compose 一键部署 - 🐳 **Docker 部署**: 使用 Docker Compose 一键部署
- ⚙️ **统一配置**: 所有配置统一在 `.env` 文件中管理,方便不同机器之间移植 - ⚙️ **统一配置**: 所有配置统一在 `.env` 文件中管理,方便不同机器之间移植
- 🧑‍💻 **Git 集成**: 支持 Git 代码库的自动同步和检索,包括连接测试和分支管理
## 快速开始 ## 快速开始
@ -29,7 +30,7 @@
# 检查 Ollama 是否运行 # 检查 Ollama 是否运行
curl http://localhost:11434/api/tags curl http://localhost:11434/api/tags
# 下载所需的模型(如果未下载) # 下载所需的模型(如果未下载)(可以使用更小的模型)
ollama pull qwen3:235b # LLM模型用于文本生成 ollama pull qwen3:235b # LLM模型用于文本生成
ollama pull qwen3-embedding:8b # Embedding模型用于向量化 ollama pull qwen3-embedding:8b # Embedding模型用于向量化
``` ```
@ -168,7 +169,20 @@ ollama pull qwen3-embedding:8b # Embedding模型用于向量化
- 点击"测试SSH连接",检查 SSH 连接是否成功。 - 点击"测试SSH连接",检查 SSH 连接是否成功。
- 点击右上角"保存"按钮,保存文件夹配置 - 点击右上角"保存"按钮,保存文件夹配置
3. 更新数据源配置 3. **Git代码库类型 (git)**
- 点击"新增数据源"-"选择类型"-"Git代码库"
- Git仓库配置
- Git仓库URL必填Git代码库的URL地址
- 分支必填要同步的Git分支默认main
- 协议选择https或ssh协议
- HTTPS Token如果使用https协议填写访问令牌
- SSH密钥如果使用ssh协议填写SSH私钥
- 点击"测试Git连接",检查 Git 连接是否成功。
- 点击右上角"保存"按钮保存Git仓库配置
4. 更新数据源配置
- 点击左侧数据源列表中的数据源 - 点击左侧数据源列表中的数据源
- 修改配置后点击保存,后台会自动删除原来同步的数据并重新同步 - 修改配置后点击保存,后台会自动删除原来同步的数据并重新同步
@ -313,6 +327,27 @@ docker-compose restart rag-api
- 查看同步服务日志: `docker-compose logs rag-api | grep sync` - 查看同步服务日志: `docker-compose logs rag-api | grep sync`
- 手动触发同步: 在配置管理界面中点击"同步"按钮 - 手动触发同步: 在配置管理界面中点击"同步"按钮
### 8. Git连接失败
**错误**: `Git连接失败``Failed to connect to Git repository`
**解决**:
- 确保 Git 仓库 URL 正确
- 检查网络连接是否正常
- 验证 Git 凭证HTTPS Token 或 SSH 密钥)是否有效
- 确保目标 Git 仓库存在且可访问
- 查看详细错误信息: `docker-compose logs rag-api | grep git`
### 9. Git同步失败
**错误**: `Git同步失败``Failed to sync Git repository`
**解决**:
- 检查 Git 仓库是否有访问权限
- 验证本地磁盘空间是否充足
- 查看同步服务日志获取详细错误信息: `docker-compose logs rag-api | grep sync`
- 尝试手动触发同步: 在配置管理界面中点击"同步"按钮
## 性能优化建议 ## 性能优化建议
### 1. 调整配置参数 ### 1. 调整配置参数
@ -356,6 +391,9 @@ docker-compose restart rag-api
**启动步骤** **启动步骤**
```bash ```bash
# 0. 确保 .env 文件中的host配置准确
cp .env.zkxlocal .env
# 1. 创建虚拟环境 # 1. 创建虚拟环境
uv venv --python 3.13.9 uv venv --python 3.13.9
source .venv/bin/activate # Windows: venv\Scripts\activate source .venv/bin/activate # Windows: venv\Scripts\activate
@ -377,4 +415,6 @@ curl http://localhost:8003/health
# 7. 启动 RAG API 服务 # 7. 启动 RAG API 服务
python main.py python main.py
# 8. 如要调试,使用.vscode/launch.json 启动调试会话
``` ```

View File

@ -258,6 +258,10 @@ class QueryRequest(BaseModel):
query: str = Field(..., description="User query string", min_length=1) query: str = Field(..., description="User query string", min_length=1)
top_k: Optional[int] = Field(None, description="Number of documents to retrieve", ge=1, le=20) top_k: Optional[int] = Field(None, description="Number of documents to retrieve", ge=1, le=20)
stream: bool = Field(True, description="Whether to stream the response") stream: bool = Field(True, description="Whether to stream the response")
repo: Optional[str] = Field(None, description="Git repository name (optional)")
branch: Optional[str] = Field(None, description="Git branch name (optional)")
is_code_related: Optional[bool] = Field(None, description="Whether the query is code related (optional)")
history: Optional[List[Dict[str, str]]] = Field(None, description="Conversation history (optional)")
class RetrieveRequest(BaseModel): class RetrieveRequest(BaseModel):
@ -606,11 +610,23 @@ async def query(request: QueryRequest):
raise HTTPException(status_code=503, detail="RAG engine not initialized") raise HTTPException(status_code=503, detail="RAG engine not initialized")
try: try:
# Build conversation history
history_str = ""
if request.history:
history_parts = []
for msg in request.history:
role = msg.get('role', 'user')
content = msg.get('content', '')
if role == 'user':
history_parts.append(f"用户: {content}")
else:
history_parts.append(f"助手: {content}")
history_str = "\n".join(history_parts)
if request.stream: if request.stream:
# Stream response # Stream response
return StreamingResponse( return StreamingResponse(
rag_engine.query_stream(request.query, None, request.top_k), rag_engine.query_stream(request.query, history_str, request.top_k),
media_type="text/event-stream", media_type="text/event-stream",
headers={ headers={
"X-Accel-Buffering": "no", "X-Accel-Buffering": "no",
@ -620,7 +636,7 @@ async def query(request: QueryRequest):
) )
else: else:
# Return complete response (run in thread pool for better concurrency) # Return complete response (run in thread pool for better concurrency)
response = await rag_engine.query(request.query, None, request.top_k) response = await rag_engine.query(request.query, history_str, request.top_k)
return response return response
except Exception as e: except Exception as e:
logger.error(f"Error processing query: {e}") logger.error(f"Error processing query: {e}")
@ -1365,6 +1381,18 @@ async def create_config(config: Dict[str, Any]):
config["host"].lower(), config["host"].lower(),
folder_path folder_path
]) ])
elif config_type == "git":
# Git配置需要仓库URL
if not config.get("git_url"):
raise HTTPException(status_code=400, detail="Git配置必须包含仓库URL")
# 添加Git仓库URL到唯一标识符
# 替换URL中的特殊字符为下划线
git_url = config["git_url"].lower().replace("/", "_").replace(":", "_").replace(".", "_")
# 截取URL的一部分作为唯一标识
git_url_part = git_url[:100] # 限制长度
unique_id_parts.extend([
git_url_part
])
else: else:
raise HTTPException(status_code=400, detail=f"不支持的配置类型: {config_type}") raise HTTPException(status_code=400, detail=f"不支持的配置类型: {config_type}")
@ -1403,6 +1431,13 @@ async def create_config(config: Dict[str, Any]):
status_code=409, status_code=409,
detail=f"已存在相同服务器和路径的文件夹配置。如需调整,请点击配置列表中的配置并修改配置内容。" detail=f"已存在相同服务器和路径的文件夹配置。如需调整,请点击配置列表中的配置并修改配置内容。"
) )
elif config_type == 'git':
# For git configs, same source means same git url
if existing_config_data.get('git_url') == config.get('git_url'):
raise HTTPException(
status_code=409,
detail=f"已存在相同Git仓库的配置。如需调整请点击配置列表中的配置并修改配置内容。"
)
except sqlite3.OperationalError as e: except sqlite3.OperationalError as e:
# 表不存在的情况,会在后面创建表 # 表不存在的情况,会在后面创建表
logger.warning(f"SQLite table error: {e}. This is expected if the table doesn't exist yet.") logger.warning(f"SQLite table error: {e}. This is expected if the table doesn't exist yet.")
@ -1433,7 +1468,7 @@ async def create_config(config: Dict[str, Any]):
global sync_manager global sync_manager
if sync_manager is not None: if sync_manager is not None:
# Create appropriate data source config object # Create appropriate data source config object
from config import BaseDataSourceConfig, DatabaseDataSourceConfig, FolderDataSourceConfig from config import BaseDataSourceConfig, DatabaseDataSourceConfig, FolderDataSourceConfig, GitDataSourceConfig
if config_type == "database": if config_type == "database":
source_config = DatabaseDataSourceConfig( source_config = DatabaseDataSourceConfig(
@ -1470,6 +1505,20 @@ async def create_config(config: Dict[str, Any]):
recursive=config.get("recursive", True), recursive=config.get("recursive", True),
ignore_patterns=config.get("ignore_patterns") ignore_patterns=config.get("ignore_patterns")
) )
elif config_type == "git":
source_config = GitDataSourceConfig(
name=config_id,
git_url=config.get("git_url"),
branch=config.get("branch", "main"),
protocol=config.get("protocol", "https"),
ssh_key=config.get("ssh_key"),
https_token=config.get("https_token"),
local_repo_path=config.get("local_repo_path"),
poll_interval=config.get("poll_interval", 300),
support_lang=config.get("support_lang"),
latest_commit_id=config.get("latest_commit_id"),
last_sync_time=config.get("last_sync_time")
)
else: else:
logger.warning(f"Unknown config type: {config_type}") logger.warning(f"Unknown config type: {config_type}")
# Create a base config as fallback # Create a base config as fallback
@ -1540,6 +1589,55 @@ async def test_folder_connection(connection_data: Dict[str, Any]):
raise HTTPException(status_code=500, detail=f"SSH连接失败: {str(e)}") raise HTTPException(status_code=500, detail=f"SSH连接失败: {str(e)}")
@app.post("/git/test-connection")
async def test_git_connection(connection_data: Dict[str, Any]):
"""
Test Git connection for Git repository configuration
Args:
connection_data: Connection data including git_url, protocol, branch, https_token, ssh_key
Returns:
Success message if connection is successful
"""
try:
git_url = connection_data.get("git_url")
protocol = connection_data.get("protocol", "https")
branch = connection_data.get("branch", "main")
https_token = connection_data.get("https_token")
ssh_key = connection_data.get("ssh_key")
if not git_url:
raise HTTPException(status_code=400, detail="Git仓库URL是必填项")
# Import GitTool here to avoid circular imports
from utils.git_tool import GitTool
# Create a temporary GitTool instance to test connection
git_tool = GitTool(
git_url=git_url,
branch=branch,
protocol=protocol,
https_token=https_token,
ssh_key=ssh_key,
local_repo_path=None # 测试连接不需要本地路径
)
# Try to test connection
success = git_tool.test_connection()
if success:
return {"message": "Git连接成功"}
else:
raise HTTPException(status_code=500, detail="Git连接失败")
except HTTPException:
raise
except Exception as e:
logger.error(f"Error testing Git connection: {e}")
raise HTTPException(status_code=500, detail=f"Git连接失败: {str(e)}")
@app.post("/folder-configs/remote") @app.post("/folder-configs/remote")
async def create_remote_folder_config(config: Dict[str, Any]): async def create_remote_folder_config(config: Dict[str, Any]):
""" """

140
config.py
View File

@ -111,6 +111,35 @@ class FolderDataSourceConfig(BaseDataSourceConfig):
self.ignore_patterns = ignore_patterns # 忽略的文件模式列表 self.ignore_patterns = ignore_patterns # 忽略的文件模式列表
class GitDataSourceConfig(BaseDataSourceConfig):
"""Git data source configuration"""
def __init__(
self,
name: str,
git_url: str,
branch: str = "main",
protocol: str = "https", # https 或 ssh
ssh_key: Optional[str] = None, # SSH私钥
https_token: Optional[str] = None, # HTTPS令牌
local_repo_path: Optional[str] = None, # 本地存储路径
poll_interval: int = 300, # 轮询间隔(秒)
support_lang: Optional[List[str]] = None, # 支持的编程语言
latest_commit_id: Optional[str] = None, # 最新commit ID
last_sync_time: Optional[str] = None # 最后同步时间
):
super().__init__(name, "git")
self.git_url = git_url # Git仓库地址
self.branch = branch # 分支名称
self.protocol = protocol # 协议类型
self.ssh_key = ssh_key # SSH私钥加密存储
self.https_token = https_token # HTTPS令牌加密存储
self.local_repo_path = local_repo_path # 本地存储路径
self.poll_interval = poll_interval # 轮询间隔
self.support_lang = support_lang # 支持的编程语言
self.latest_commit_id = latest_commit_id # 最新commit ID
self.last_sync_time = last_sync_time # 最后同步时间
class Settings(BaseSettings): class Settings(BaseSettings):
""" """
Application settings Application settings
@ -138,28 +167,17 @@ class Settings(BaseSettings):
CHROMA_DB_PATH: str = "./chroma_db" # Only used for PersistentClient mode CHROMA_DB_PATH: str = "./chroma_db" # Only used for PersistentClient mode
CHROMA_COLLECTION_NAME: str = "rag_collection" CHROMA_COLLECTION_NAME: str = "rag_collection"
# LLM Settings # Ollama Settings
# LLM provider: "ollama", "vllm", "openai", "deepseek" etc. # Configure OLLAMA_BASE_URL in .env file based on your deployment
LLM_PROVIDER: str = "ollama" # - Local: http://localhost:11434
LLM_BASE_URL: str = "http://localhost:11434" # - Remote: http://192.168.1.100:11434
LLM_MODEL: str = "qwen3:8b"
LLM_API_KEY: Optional[str] = None
# Embedding Settings
# Embedding provider: "ollama", "vllm", "openai" etc.
EMBEDDING_PROVIDER: str = "ollama"
EMBEDDING_BASE_URL: str = "http://localhost:11434"
EMBEDDING_MODEL: str = "qwen3-embedding:0.6b"
EMBEDDING_API_KEY: Optional[str] = None
# Legacy Ollama Settings (for backward compatibility)
OLLAMA_BASE_URL: str = "http://localhost:11434" OLLAMA_BASE_URL: str = "http://localhost:11434"
OLLAMA_MODEL: str = "qwen3:8b" OLLAMA_MODEL: str = "qwen3:1.7b" # LLM model for text generation
OLLAMA_EMBEDDING_MODEL: str = "qwen3-embedding:0.6b" OLLAMA_EMBEDDING_MODEL: str = "qwen3-embedding:0.6b" # Embedding model for vectorization
# RAG Settings # RAG Settings
EMBEDDING_DIMENSION: int = 768 EMBEDDING_DIMENSION: int = 768
CHUNK_SIZE: int = 1024 CHUNK_SIZE: int = 4000
CHUNK_OVERLAP: int = 200 CHUNK_OVERLAP: int = 200
TOP_K: int = 5 # Number of documents to retrieve TOP_K: int = 5 # Number of documents to retrieve
@ -187,6 +205,12 @@ class Settings(BaseSettings):
SOFFICE_HOST: str = "127.0.0.1" SOFFICE_HOST: str = "127.0.0.1"
SOFFICE_PORT: int = 8003 SOFFICE_PORT: int = 8003
# Git 相关配置
GIT_LOCAL_STORAGE_ROOT: str = "./git_repos" # Git仓库本地存储根目录
GIT_DEFAULT_BRANCH: str = "main" # 默认分支
GIT_POLL_INTERVAL: int = 300 # 默认轮询间隔(秒)
GIT_MAX_REPO_SIZE_MB: int = 500 # 最大仓库大小MB
# Pydantic v2 configuration # Pydantic v2 configuration
model_config = SettingsConfigDict( model_config = SettingsConfigDict(
env_file=".env", env_file=".env",
@ -195,71 +219,6 @@ class Settings(BaseSettings):
extra="ignore" # Ignore extra fields in .env file that are not defined in Settings extra="ignore" # Ignore extra fields in .env file that are not defined in Settings
) )
def get_llm_config(self) -> Dict[str, Any]:
"""
Get LLM configuration with backward compatibility for legacy OLLAMA_* settings.
Returns:
Dictionary with provider, base_url, model, and api_key
"""
if self.LLM_PROVIDER == "ollama":
return {
"provider": "ollama",
"base_url": self.LLM_BASE_URL or self.OLLAMA_BASE_URL,
"model": self.LLM_MODEL or self.OLLAMA_MODEL,
"api_key": self.LLM_API_KEY
}
elif self.LLM_PROVIDER == "openai":
return {
"provider": "openai",
"base_url": self.LLM_BASE_URL or "https://api.openai.com/v1",
"model": self.LLM_MODEL,
"api_key": self.LLM_API_KEY
}
else:
return {
"provider": self.LLM_PROVIDER,
"base_url": self.LLM_BASE_URL,
"model": self.LLM_MODEL,
"api_key": self.LLM_API_KEY
}
def get_embedding_config(self) -> Dict[str, Any]:
"""
Get Embedding configuration with backward compatibility for legacy OLLAMA_* settings.
Returns:
Dictionary with provider, base_url, model, and api_key
"""
if self.EMBEDDING_PROVIDER == "ollama":
return {
"provider": "ollama",
"base_url": self.EMBEDDING_BASE_URL or self.OLLAMA_BASE_URL,
"model": self.EMBEDDING_MODEL or self.OLLAMA_EMBEDDING_MODEL,
"api_key": self.EMBEDDING_API_KEY
}
elif self.EMBEDDING_PROVIDER == "vllm":
return {
"provider": "vllm",
"base_url": self.EMBEDDING_BASE_URL,
"model": self.EMBEDDING_MODEL,
"api_key": self.EMBEDDING_API_KEY
}
elif self.EMBEDDING_PROVIDER == "openai":
return {
"provider": "openai",
"base_url": self.EMBEDDING_BASE_URL or "https://api.openai.com/v1",
"model": self.EMBEDDING_MODEL,
"api_key": self.EMBEDDING_API_KEY
}
else:
return {
"provider": self.EMBEDDING_PROVIDER,
"base_url": self.EMBEDDING_BASE_URL,
"model": self.EMBEDDING_MODEL,
"api_key": self.EMBEDDING_API_KEY
}
def get_data_sources(self) -> List[BaseDataSourceConfig]: def get_data_sources(self) -> List[BaseDataSourceConfig]:
""" """
Get list of data source configurations Get list of data source configurations
@ -336,6 +295,21 @@ class Settings(BaseSettings):
recursive=ds_config.get('recursive', True), recursive=ds_config.get('recursive', True),
ignore_patterns=ds_config.get('ignore_patterns', None) ignore_patterns=ds_config.get('ignore_patterns', None)
)) ))
elif source_type == 'git':
# Create git data source
configs.append(GitDataSourceConfig(
name=name, # 使用数据库表中的name列
git_url=ds_config.get('git_url'),
branch=ds_config.get('branch', 'main'),
protocol=ds_config.get('protocol', 'https'),
ssh_key=ds_config.get('ssh_key'),
https_token=ds_config.get('https_token'),
local_repo_path=ds_config.get('local_repo_path'),
poll_interval=ds_config.get('poll_interval', 300),
support_lang=ds_config.get('support_lang'),
latest_commit_id=ds_config.get('latest_commit_id'),
last_sync_time=ds_config.get('last_sync_time')
))
else: else:
from loguru import logger from loguru import logger
logger.warning(f"Unknown data source type: {source_type}, skipping") logger.warning(f"Unknown data source type: {source_type}, skipping")

View File

@ -2,6 +2,7 @@
Database utilities for RAG system Database utilities for RAG system
""" """
import sqlite3 import sqlite3
import json
from pathlib import Path from pathlib import Path
from datetime import datetime from datetime import datetime
from typing import Tuple, Optional from typing import Tuple, Optional
@ -90,6 +91,75 @@ def update_data_source_update_at(source_name: str, update_at: datetime) -> bool:
return False return False
def add_git_datasource(user_id: str, repo_config: dict):
"""
新增Git仓库配置
Args:
user_id: 用户ID
repo_config: 仓库配置
"""
conn, cursor = get_db_connection()
try:
cursor.execute("""
INSERT INTO datasource (user_id, name, type, git_config, create_time)
VALUES (?, ?, 'git', ?, datetime('now'))
""", (user_id, repo_config["name"], json.dumps(repo_config)))
conn.commit()
logger.info(f"新增Git数据源: {repo_config['name']}")
finally:
conn.close()
def update_git_sync_status(user_id: str, repo_id: str, sync_status: dict):
"""
更新Git仓库同步状态
Args:
user_id: 用户ID
repo_id: 仓库ID
sync_status: 同步状态
"""
conn, cursor = get_db_connection()
try:
cursor.execute("""
UPDATE datasource SET git_config = json_set(git_config, '$.latest_commit_id', ?, '$.last_sync_time', ?)
WHERE user_id = ? AND id = ?
""", (sync_status["latest_commit_id"], sync_status["last_sync_time"], user_id, repo_id))
conn.commit()
logger.info(f"更新Git同步状态: {repo_id}")
finally:
conn.close()
def update_git_repo_config(user_id: str, repo_id: str, config: dict):
"""
更新Git仓库配置
Args:
user_id: 用户ID
repo_id: 仓库ID
config: 配置信息
"""
conn, cursor = get_db_connection()
try:
# 获取当前配置
cursor.execute("SELECT git_config FROM datasource WHERE user_id = ? AND id = ?", (user_id, repo_id))
result = cursor.fetchone()
if result:
current_config = json.loads(result[0])
# 更新配置
current_config.update(config)
cursor.execute("""
UPDATE datasource SET git_config = ?
WHERE user_id = ? AND id = ?
""", (json.dumps(current_config), user_id, repo_id))
conn.commit()
logger.info(f"更新Git仓库配置: {repo_id}")
finally:
conn.close()
def init_session_db(): def init_session_db():
""" """
Initialize session database with users and sessions tables Initialize session database with users and sessions tables

View File

@ -52,7 +52,10 @@ services:
container_name: rag-api container_name: rag-api
# 使用宿主机网络模式可以直接访问宿主机上的服务Ollama、MySQL 等) # 使用宿主机网络模式可以直接访问宿主机上的服务Ollama、MySQL 等)
# 注意:使用 host 网络模式时,不能使用 ports 映射,容器直接使用宿主机的网络 # 注意:使用 host 网络模式时,不能使用 ports 映射,容器直接使用宿主机的网络
network_mode: host # network_mode: host # 注释/删除host网络模式Windows下无效)
# 添加端口映射Windows下开发
ports:
- "${API_PORT:-8001}:8001"
# 自动读取 .env 文件(如果存在) # 自动读取 .env 文件(如果存在)
env_file: env_file:
- .env - .env
@ -96,7 +99,7 @@ services:
# RAG 配置 # RAG 配置
- EMBEDDING_DIMENSION=${EMBEDDING_DIMENSION:-768} - EMBEDDING_DIMENSION=${EMBEDDING_DIMENSION:-768}
- CHUNK_SIZE=${CHUNK_SIZE:-1024} - CHUNK_SIZE=${CHUNK_SIZE:-4000}
- CHUNK_OVERLAP=${CHUNK_OVERLAP:-200} - CHUNK_OVERLAP=${CHUNK_OVERLAP:-200}
- TOP_K=${TOP_K:-5} - TOP_K=${TOP_K:-5}

View File

@ -25,7 +25,7 @@ RUN apt-get update \
WORKDIR /app WORKDIR /app
COPY requirements.txt /app/requirements.txt COPY requirements.txt /app/requirements.txt
RUN pip install --no-cache-dir -r requirements.txt RUN pip install --no-cache-dir -r requirements.txt -i https://pypi.tuna.tsinghua.edu.cn/simple
COPY docker/soffice/app.py /app/app.py COPY docker/soffice/app.py /app/app.py
COPY docker/soffice/README.md /app/README.md COPY docker/soffice/README.md /app/README.md

View File

@ -0,0 +1,828 @@
# RAG智能问答助手——Git 代码库二次开发
# 新增Git代码库作为数据源的二次开发实现细节
本次二次开发核心是在原仓库**MySQL/达梦数据库/文件夹**三种数据源基础上,新增**Git代码库**数据源类型,需适配**代码拉取-解析切片-向量化存储-增量同步-代码专属检索/生成**全流程。以下结合原仓库的代码架构、文件结构,按**工程结构调整、核心模块开发、前后端适配、配置/部署修改、核心方法实现**五个维度给出可落地的实现细节完全复用原仓库的LlamaIndex/ChromaDB/Ollama基础能力仅做针对性扩展。
## 一、原仓库核心架构复用与工程结构调整
原仓库已实现**通用同步基类、ChromaDB向量操作、FastAPI接口、前端配置管理**等基础能力,本次开发仅需**新增Git相关模块、扩展原有基类/接口、适配代码场景的解析/检索逻辑**,不改动原核心代码,保证兼容性。
### 1. 原仓库核心复用模块
|原仓库模块/文件|复用功能|扩展点|
|---|---|---|
|`sync/base_sync.py`|同步基类、通用向量化、ChromaDB基础写入|新增**代码库专属的抽象方法**(如`git_clone`/`detect_git_update`让GitSync子类实现|
|`config.py`|全局环境变量读取、配置管理|新增Git代码库相关的全局配置本地存储根目录、默认分支等|
|`main.py`|FastAPI接口、流式响应、路由注册|新增Git数据源的配置路由、Git仓库手动同步路由|
|`static/config/`|前端数据源配置界面|新增Git类型的配置表单仓库地址、协议、SSH密钥等|
|`db_utils.py`|数据库工具、配置持久化|新增Git仓库同步配置的存储逻辑最后同步commit ID、分支、轮询频率等|
|原ChromaDB操作逻辑|向量增删改查、embedding配对|重构**存储结构**,适配函数级代码的元数据/业务数据如函数唯一ID、仓库名、分支等|
### 2. 新增/修改的文件结构
在原仓库基础上新增**Git同步、代码解析**专属模块,修改少量核心文件,新增文件如下(按目录分类):
```Plain Text
# 核心同步模块新增
sync/
├── git_sync.py # Git代码库同步子类继承BaseSync实现代码拉取/增量同步/函数解析
└── ast_parser.py # 代码AST解析工具类实现跨语言函数级切片核心
# 前端配置界面新增Git配置表单
static/config/
├── js/git_config.js # Git配置的前端逻辑凭证验证、仓库地址解析
└── components/
└── git-form.html # Git数据源配置的HTML组件嵌入原config/index.html
# 工具类新增
utils/
├── git_tool.py # Git命令封装工具类clone/fetch/merge/日志解析封装subprocess执行Git命令
└── func_id_generator.py # 函数全局唯一ID生成工具类按用户/仓库/分支/文件/函数生成)
# 原文件修改(仅扩展,不改动原有逻辑)
sync/base_sync.py # 扩展基类,新增代码向量化专属方法
config.py # 新增Git相关全局配置
main.py # 新增Git数据源路由
db_utils.py # 新增Git同步配置持久化
Dockerfile # 安装git命令容器内需要执行Git操作
.env.example # 新增Git相关环境变量配置项
```
## 二、核心模块开发(按技术流程拆解)
按**代码拉取与存储→代码解析与向量化→增量同步→代码专属检索/生成**的技术流程,结合原仓库代码实现核心功能,每个环节均给出**原仓库对接点+代码实现思路**。
### 阶段1代码拉取与存储Git仓库专属
核心实现**用户Git配置验证、多协议克隆、本地结构化存储**,封装为`git_tool.py`工具类,在`git_sync.py`中调用,复用原仓库的**数据源配置管理**能力。
#### 1. 原仓库对接点
- 前端配置界面提交的Git配置信息仓库地址、协议、SSH私钥/HTTPS令牌、分支通过原仓库的`/api/config/datasource`路由接收,新增`type: git`标识,与`mysql/folder`区分;
- 配置信息通过`db_utils.py`持久化到原仓库的配置库SQLite/MySQL新增`git_config`字段存储Git专属配置commit ID、存储路径、轮询间隔等
#### 2. 核心实现细节
##### 1Git工具类封装`utils/git_tool.py`
封装所有Git原生命令避免硬编码处理**SSH/HTTPS/git**多协议,实现**克隆、远程更新检测、增量拉取、文件变更解析**等核心功能,示例核心方法:
```Python
import subprocess
import os
from config import settings # 原仓库的全局配置
class GitTool:
def __init__(self, user_id: str, repo_id: str, git_config: dict):
self.user_id = user_id
self.repo_id = repo_id
self.git_url = git_config["git_url"] # 用户配置的Git仓库地址
self.branch = git_config.get("branch", settings.GIT_DEFAULT_BRANCH)
self.ssh_key = git_config.get("ssh_key") # SSH私钥base64加密存储
# 本地结构化存储路径(按文档规范:根目录/用户ID/仓库ID
self.local_repo_path = os.path.join(settings.GIT_LOCAL_STORAGE_ROOT, user_id, repo_id)
self._init_git_env() # 初始化Git环境SSH密钥配置
# 初始化Git SSH环境核心解决容器内SSH密钥验证
def _init_git_env(self):
if self.ssh_key:
# 解密SSH私钥写入临时文件配置Git SSH
os.environ["GIT_SSH_COMMAND"] = f"ssh -i /tmp/ssh_key_{self.user_id} -o StrictHostKeyChecking=no"
with open(f"/tmp/ssh_key_{self.user_id}", "w") as f:
f.write(self.ssh_key)
os.chmod(f"/tmp/ssh_key_{self.user_id}", 0o600)
# 克隆Git仓库适配多协议复用文档的git clone命令
def clone_repo(self) -> bool:
if not os.path.exists(self.local_repo_path):
os.makedirs(os.path.dirname(self.local_repo_path), exist_ok=True)
# 执行git clone --depth=none --single-branch --branch <分支> <地址> <本地路径>
cmd = [
"git", "clone", "--depth=none", "--single-branch",
"--branch", self.branch, self.git_url, self.local_repo_path
]
res = subprocess.run(cmd, capture_output=True, text=True)
if res.returncode != 0:
raise Exception(f"Git克隆失败: {res.stderr}")
# 克隆后校验git fsck + 语言检测)
self._check_repo_integrity()
return True
return False
# 仓库完整性校验git fsck+ 支持的编程语言检测
def _check_repo_integrity(self):
# 执行git fsck
subprocess.run(["git", "fsck"], cwd=self.local_repo_path, check=True)
# 扫描文件类型,记录支持的编程语言(如.py/.java/.go存入配置库
from utils.lang_detect import detect_support_lang # 简单的文件后缀检测工具
support_lang = detect_support_lang(self.local_repo_path)
from db_utils import update_git_repo_config
update_git_repo_config(self.user_id, self.repo_id, {"support_lang": support_lang})
# 远程更新检测(对比本地/远程commit ID复用文档逻辑
def detect_remote_update(self) -> tuple[bool, str, str]:
# 拉取远程commit记录仅拉取不拉取文件
subprocess.run(["git", "fetch", "origin", f"{self.branch}:{self.branch}"], cwd=self.local_repo_path, check=True)
# 获取本地/远程commit ID
local_commit = subprocess.run(["git", "rev-parse", "HEAD"], cwd=self.local_repo_path, capture_output=True, text=True).stdout.strip()
remote_commit = subprocess.run(["git", "rev-parse", f"origin/{self.branch}"], cwd=self.local_repo_path, capture_output=True, text=True).stdout.strip()
return local_commit != remote_commit, local_commit, remote_commit
# 增量拉取代码+解析文件变更(新增/修改/删除)
def incremental_pull(self) -> dict:
# 快进合并到远程最新版本
subprocess.run(["git", "merge", "--ff-only", f"origin/{self.branch}"], cwd=self.local_repo_path, check=True)
# 提取增量commit的文件变更
delta_commits = subprocess.run(["git", "log", "--pretty=format:%H", f"{self.local_commit}..{self.remote_commit}"], cwd=self.local_repo_path, capture_output=True, text=True).stdout.split()
# 解析文件变更为ADD/MODIFY/DELETE
delta_files = self._parse_delta_files(delta_commits)
return delta_files
# 解析文件变更集复用文档的parse_delta_files.py逻辑
def _parse_delta_files(self, delta_commits: list) -> dict:
add_files, modify_files, delete_files = [], [], []
for commit in delta_commits:
# git show --name-status 获取文件变更
res = subprocess.run(["git", "show", "--name-status", commit], cwd=self.local_repo_path, capture_output=True, text=True).stdout
for line in res.splitlines():
if not line: continue
status, file_path = line.split("\t", 1)
file_path = os.path.join(self.local_repo_path, file_path)
if status == "A": add_files.append(file_path)
elif status == "M": modify_files.append(file_path)
elif status == "D": delete_files.append(file_path)
# 去重并返回
return {
"ADD": list(set(add_files)),
"MODIFY": list(set(modify_files)),
"DELETE": list(set(delete_files))
}
```
##### 2Git配置持久化修改`db_utils.py`
原仓库已实现数据源配置的持久化,新增**Git仓库专属配置字段**,存储:
- 仓库基础配置:`git_url`/`protocol`/`branch`/`ssh_key`AES256加密/`https_token`(加密);
- 同步状态配置:`local_repo_path`/`latest_commit_id`/`support_lang`/`poll_interval`(轮询间隔)/`last_sync_time`
- 示例新增方法:
```Python
# 新增Git仓库配置存储
def add_git_datasource(user_id: str, repo_config: dict):
# 原仓库的配置表新增git_config字段存储json格式的配置
conn = get_sqlite_conn() # 原仓库的SQLite连接方法
cursor = conn.cursor()
cursor.execute("""
INSERT INTO datasource (user_id, name, type, git_config, create_time)
VALUES (?, ?, 'git', ?, datetime('now'))
""", (user_id, repo_config["name"], json.dumps(repo_config)))
conn.commit()
conn.close()
# 更新Git仓库同步状态最后同步commit ID、时间
def update_git_sync_status(user_id: str, repo_id: str, sync_status: dict):
conn = get_sqlite_conn()
cursor = conn.cursor()
cursor.execute("""
UPDATE datasource SET git_config = json_set(git_config, '$.latest_commit_id', ?, '$.last_sync_time', ?)
WHERE user_id = ? AND id = ?
""", (sync_status["latest_commit_id"], sync_status["last_sync_time"], user_id, repo_id))
conn.commit()
conn.close()
```
### 阶段2代码解析与向量化代码场景核心改造
核心实现**函数级AST切片、LLM标准化生成函数描述、ChromaDB函数级存储**,是本次开发的**核心改造点**,需扩展原仓库的`base_sync.py`向量化逻辑,新增`ast_parser.py`和函数ID生成工具。
#### 1. 原仓库对接点
- 复用原仓库的**Ollama向量化能力**`qwen3-embedding:8b`),仅修改向量化的**源数据**从普通文本→LLM生成的函数描述
- 复用原仓库的ChromaDB基础操作`add/delete/query`**重构ChromaDB的存储结构**,适配函数级代码的元数据/业务数据;
- 继承`sync/base_sync.py`的`BaseSync`基类,实现`extract_data`(函数切片)、`vectorize_data`(函数描述向量化)、`save_to_chroma`(函数级存入)方法。
#### 2. 核心实现细节
##### 1AST函数级切片`sync/ast_parser.py`
摒弃正则,采用**编程语言专属AST解析库**,实现跨语言函数提取,输出**标准化函数字典**(复用文档的格式),示例核心方法:
```Python
import ast
import libcst # Python AST解析支持代码修改
from typing import List, Dict
class ASTParser:
def __init__(self, file_path: str, lang: str):
self.file_path = file_path
self.lang = lang # 编程语言python/java/go等
self.func_list: List[Dict] = [] # 提取的函数列表
# 统一入口:根据语言调用对应解析方法
def parse_functions(self) -> List[Dict]:
if not os.path.exists(self.file_path):
raise Exception(f"文件不存在: {self.file_path}")
with open(self.file_path, "r", encoding="utf-8") as f:
self.code = f.read()
# 按语言解析
if self.lang == "python":
self._parse_python()
# 后续扩展java/go此处先实现Python
return self.func_list
# Python函数解析基于ast+libcst
def _parse_python(self):
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:
raise Exception(f"Python代码语法错误: {e}")
# 提取Python函数的标准化信息
def _extract_python_func_info(self, node) -> Dict:
# 提取函数名、参数、返回值、函数体等
func_name = node.name
params = [arg.arg for arg in node.args.args] # 简化参数提取,可扩展类型注解
return_type = ast.unparse(node.returns) if node.returns else "None"
# 提取函数体代码
func_body = libcst.parse_module(self.code).code_for_node(node)
return {
"file_path": self.file_path,
"func_name": func_name,
"params": params,
"return_type": return_type,
"func_body": func_body,
"class_name": None # 类方法需额外解析,此处简化
}
```
##### 2函数全局唯一ID生成`utils/func_id_generator.py`
为每个函数生成**全局唯一ID**核心用于增量同步时精准定位ChromaDB条目复用文档的ID格式`用户ID_仓库ID_分支_文件相对路径_类名_函数名`
```Python
import os
from config import settings
def generate_func_unique_id(user_id: str, repo_id: str, branch: str, file_path: str, class_name: str, func_name: str) -> str:
# 将本地绝对路径转为仓库根目录的相对路径
local_repo_root = os.path.join(settings.GIT_LOCAL_STORAGE_ROOT, user_id, repo_id)
rel_file_path = os.path.relpath(file_path, local_repo_root).replace(os.sep, "_")
# 类名为None则拼接空字符串
class_name = class_name if class_name else "None"
# 生成唯一ID
unique_id = f"{user_id}_{repo_id}_{branch}_{rel_file_path}_{class_name}_{func_name}"
# 替换特殊字符避免ChromaDB主键冲突
unique_id = unique_id.replace("/", "_").replace("\\", "_").replace(":", "_")
return unique_id
```
##### 3LLM生成标准化函数描述修改`sync/base_sync.py`
复用原仓库的Ollama LLM调用能力`qwen3:235b`**新增标准化Prompt模板**(复用文档),为每个函数生成描述,示例方法:
```Python
# 在sync/base_sync.py的BaseSync类中新增方法
def generate_func_desc(self, func_info: Dict) -> str:
"""调用LLM生成标准化函数描述"""
# 文档中的标准化Prompt模板
prompt = f"""
### 任务要求
你是资深程序员,需要为给定的代码函数生成**简洁、准确、结构化的自然语言描述**,用于代码语义检索,严格遵循以下规则:
1. 描述仅包含「函数功能+入参作用+返回值意义」,无额外冗余内容;
2. 语言为中文,字数控制在50-100字;
3. 若为类中的方法,需体现方法与类的关联;
4. 不添加代码、注释、表情,仅纯自然语言描述。
### 待描述函数信息
文件路径:{func_info['file_path']}
所属类:{func_info['class_name']}
函数名:{func_info['func_name']}
参数:{func_info['params']}
返回值类型:{func_info['return_type']}
函数代码:
{func_info['func_body']}
### 输出示例
示例1(全局函数):该函数为工具函数,接收两个整数类型的参数a和b,实现两数相加的功能,返回相加后的整数结果。
### 请输出你的描述
""".strip()
# 调用原仓库的Ollama LLM调用方法
from utils.ollama_client import call_ollama # 原仓库的Ollama客户端
desc = call_ollama(prompt, model=self.ollama_model)
return desc.strip()
```
##### 4ChromaDB函数级存储重构`base_sync.py`的`save_to_chroma`,贴合检索需求)
核心贴合你的需求:**按func_desc检索、返回对应func_body**改造核心是明确「func_desc向量化生成embedding检索核心+ 业务数据关联存储返回func_body依据复用原仓库ChromaDB客户端仅重构入参结构确保检索时通过func_desc匹配精准返回对应func_body具体改造如下
**核心改造原仓库的ChromaDB存储结构**,按文档要求设计**向量字段+元数据字段+业务数据字段**复用原仓库的ChromaDB客户端仅修改入参结构示例方法
```Python
# 重构sync/base_sync.py的save_to_chroma方法完全贴合「按func_desc检索、返回func_body」需求
def save_to_chroma(self, func_data_list: List[Dict], embeddings: List[List[float]]):
"""
核心设计:
1. 检索核心embeddings仅基于func_desc生成与检索逻辑完全对齐
2. 关联存储将func_body及关键信息存入metadatas结构化存储便于检索后直接提取
3. 检索匹配documents仅存func_desc确保检索时仅匹配函数描述提升精准度
"""
# 初始化ChromaDB客户端复用原仓库配置不做修改
import chromadb
client = chromadb.HttpClient(host=settings.CHROMA_SERVER_HOST, port=settings.CHROMA_SERVER_PORT)
# 按用户隔离集合(复用原仓库多用户隔离逻辑,避免数据冲突)
collection = client.get_or_create_collection(name=f"code_rag_{self.user_id}")
# 构造ChromaDB入参核心重构贴合需求
# 1. 唯一ID沿用函数全局唯一ID用于精准定位和增量更新复用原生成逻辑
ids = [func["func_unique_id"] for func in func_data_list]
# 2. 元数据核心关联存储存入func_body及关键信息作为检索后返回func_body的直接依据
metadatas = [
{
"user_id": self.user_id,
"repo_id": self.repo_id,
"branch": self.branch,
"file_path": func["file_path"],
"func_name": func["func_name"],
"func_body": func["func_body"], # 关键存储func_body检索后直接提取返回
"latest_commit_id": self.latest_commit_id
} for func in func_data_list
]
# 3. 检索匹配字段仅存func_desc确保检索时仅基于函数描述进行向量匹配提升精准度
documents = [func["func_desc"] for func in func_data_list]
# 4. 向量核心embeddings仅基于func_desc生成与documents完全对应检索核心
# embeddings由外部传入对应vectorize_data方法中func_desc的向量化结果
# 批量存入ChromaDB复用原仓库批量操作逻辑不做修改
collection.add(
ids=ids,
embeddings=embeddings,
metadatas=metadatas,
documents=documents
)
# 补充检索逻辑对应调整后续检索时通过func_desc生成embedding查询从metadatas提取func_body
# 此处提前预留检索逻辑适配说明,确保存储与检索闭环
```
##### 5GitSync子类实现`sync/git_sync.py`
继承原仓库的`BaseSync`基类,整合**Git拉取、AST切片、LLM描述、向量化、ChromaDB存储**全流程,实现基类的抽象方法:
```Python
from sync.base_sync import BaseSync
from utils.git_tool import GitTool
from sync.ast_parser import ASTParser
from utils.func_id_generator import generate_func_unique_id
from config import settings
class GitSync(BaseSync):
def __init__(self, user_id: str, repo_id: str, git_config: dict):
super().__init__(user_id)
self.repo_id = repo_id
self.git_config = git_config
self.branch = git_config.get("branch", settings.GIT_DEFAULT_BRANCH)
self.git_tool = GitTool(user_id, repo_id, git_config)
self.latest_commit_id = git_config.get("latest_commit_id")
# 实现基类的extract_data拉取代码+AST函数切片
def extract_data(self) -> List[Dict]:
# 1. 克隆/拉取代码
self.git_tool.clone_repo()
# 2. 获取仓库支持的编程语言
support_lang = self.git_config.get("support_lang", ["python"])
# 3. 遍历仓库文件AST切片提取函数
func_data_list = []
for root, _, files in os.walk(self.git_tool.local_repo_path):
for file in files:
file_path = os.path.join(root, file)
# 匹配支持的编程语言
lang = self._get_file_lang(file_path)
if lang not in support_lang:
continue
# 4. AST解析函数
ast_parser = ASTParser(file_path, lang)
func_list = ast_parser.parse_functions()
# 5. 为每个函数生成唯一ID+LLM描述
for func in func_list:
func["func_unique_id"] = generate_func_unique_id(
self.user_id, self.repo_id, self.branch,
file_path, func["class_name"], func["func_name"]
)
func["func_desc"] = self.generate_func_desc(func) # 调用基类的LLM描述方法
func["latest_commit_id"] = self.latest_commit_id
func_data_list.append(func)
return func_data_list
# 实现基类的vectorize_data函数描述向量化复用原仓库Ollama
def vectorize_data(self, func_data_list: List[Dict]) -> List[List[float]]:
func_descs = [func["func_desc"] for func in func_data_list]
# 调用原仓库的向量化方法复用qwen3-embedding:8b
return self._ollama_embedding(func_descs)
# 实现基类的run_sync整合全流程
def run_sync(self):
# 1. 提取函数数据
func_data_list = self.extract_data()
if not func_data_list:
return "无函数数据可同步"
# 2. 函数描述向量化
embeddings = self.vectorize_data(func_data_list)
# 3. 存入ChromaDB
self.save_to_chroma(func_data_list, embeddings)
# 4. 更新同步状态最后commit ID
from db_utils import update_git_sync_status
update_git_sync_status(self.user_id, self.repo_id, {
"latest_commit_id": self.latest_commit_id,
"last_sync_time": datetime.now().strftime("%Y-%m-%d %H:%M:%S")
})
return f"同步成功,共处理{len(func_data_list)}个函数"
# 辅助方法:根据文件后缀判断编程语言
def _get_file_lang(self, file_path: str) -> str:
suffix = os.path.splitext(file_path)[1].lower()
lang_map = {".py": "python", ".java": "java", ".go": "go", ".js": "javascript"}
return lang_map.get(suffix, "unknown")
```
### 阶段3代码增量同步Git仓库核心亮点
核心实现**定期轮询、Git增量拉取、函数级增删改识别、ChromaDB精准增量更新**,复用原仓库的**定时同步框架**`sync_service.py`),在`git_sync.py`中新增增量同步方法,**全程避免全量解析/向量化**。
#### 1. 原仓库对接点
- 复用原仓库的**定时同步能力**`SYNC_INTERVAL`/`AUTO_SYNC`为Git数据源新增**自定义轮询间隔**(用户可配置);
- 复用原仓库的**手动同步路由**,新增`/api/sync/git`路由支持手动触发Git仓库增量同步
- 基于ChromaDB的`delete`+`add`实现增量更新原仓库已支持ChromaDB的增删操作
#### 2. 核心实现细节(在`git_sync.py`中新增增量同步方法)
ps. 「 2. 增量拉取代码解析文件变更集ADD/MODIFY/DELETE」更细粒度的处理 新增一个path_change_files识别仅「路径变、内容不变」的文件在增量更新时仅更新元数据复用原有向量和业务数据。
```Python
# 在GitSync类中新增增量同步方法
def run_incremental_sync(self) -> str:
"""Git仓库增量同步检测更新→增量拉取→函数级变更→ChromaDB增量更新"""
# 1. 检测远程更新
has_update, local_commit, remote_commit = self.git_tool.detect_remote_update()
if not has_update:
return "无远程更新,无需同步"
self.latest_commit_id = remote_commit # 更新为最新commit ID
# 2. 增量拉取代码解析文件变更集ADD/MODIFY/DELETE
delta_files = self.git_tool.incremental_pull(local_commit, remote_commit)
add_files, modify_files, delete_files = delta_files["ADD"], delta_files["MODIFY"], delta_files["DELETE"]
# 3. 函数级增删改识别(核心)
func_change = self._detect_func_change(add_files, modify_files, delete_files)
add_funcs, modify_funcs, delete_func_ids = func_change["ADD"], func_change["MODIFY"], func_change["DELETE"]
# 4. ChromaDB增量更新复用文档的先删后加逻辑
self._chroma_incremental_update(add_funcs, modify_funcs, delete_func_ids)
# 5. 更新同步状态
from db_utils import update_git_sync_status
update_git_sync_status(self.user_id, self.repo_id, {
"latest_commit_id": remote_commit,
"last_sync_time": datetime.now().strftime("%Y-%m-%d %H:%M:%S")
})
return f"增量同步成功:新增{len(add_funcs)}个函数,修改{len(modify_funcs)}个函数,删除{len(delete_func_ids)}个函数"
# 函数级增删改识别:文件变更→函数变更
def _detect_func_change(self, add_files: list, modify_files: list, delete_files: list) -> dict:
add_funcs, modify_funcs = [], []
# 初始化ChromaDB客户端获取当前仓库的函数数据
client = chromadb.HttpClient(host=settings.CHROMA_SERVER_HOST, port=settings.CHROMA_SERVER_PORT)
collection = client.get_collection(name=f"code_rag_{self.user_id}")
# 过滤当前仓库的所有函数ID
repo_funcs = collection.query(
where={"$and": [{"user_id": self.user_id}, {"repo_id": self.repo_id}]},
ids_only=True
)
repo_func_ids = set(repo_funcs["ids"])
# 处理新增/修改文件重解析→对比函数ID
all_change_files = add_files + modify_files
support_lang = self.git_config.get("support_lang", ["python"])
for file in all_change_files:
lang = self._get_file_lang(file)
if lang not in support_lang:
continue
# AST解析最新函数
ast_parser = ASTParser(file, lang)
latest_funcs = ast_parser.parse_functions()
# 生成函数唯一ID
for func in latest_funcs:
func["func_unique_id"] = generate_func_unique_id(
self.user_id, self.repo_id, self.branch,
file, func["class_name"], func["func_name"]
)
func["func_desc"] = self.generate_func_desc(func)
func["latest_commit_id"] = self.latest_commit_id
# 新增函数ID不在仓库函数ID中
if func["func_unique_id"] not in repo_func_ids:
add_funcs.append(func)
# 修改函数ID存在代码不一致
else:
modify_funcs.append(func)
# 处理删除文件过滤ChromaDB中该文件的所有函数ID
delete_func_ids = []
for file in delete_files:
rel_file_path = os.path.relpath(file, self.git_tool.local_repo_path).replace(os.sep, "_")
# 按文件路径过滤函数ID
del_funcs = collection.query(
where={"$and": [{"user_id": self.user_id}, {"repo_id": self.repo_id}, {"file_path": file}]},
ids_only=True
)
delete_func_ids.extend(del_funcs["ids"])
return {"ADD": add_funcs, "MODIFY": modify_funcs, "DELETE": list(set(delete_func_ids))}
# ChromaDB增量更新删→增修改函数先删后加
def _chroma_incremental_update(self, add_funcs: list, modify_funcs: list, delete_func_ids: list):
client = chromadb.HttpClient(host=settings.CHROMA_SERVER_HOST, port=settings.CHROMA_SERVER_PORT)
collection = client.get_collection(name=f"code_rag_{self.user_id}")
# 1. 删除函数:批量删除
if delete_func_ids:
collection.delete(ids=delete_func_ids)
# 2. 新增函数:向量化+批量添加
if add_funcs:
embeddings = self.vectorize_data(add_funcs)
self.save_to_chroma(add_funcs, embeddings)
# 3. 修改函数先删后加ChromaDB无更新API
if modify_funcs:
# 删除旧版本
old_func_ids = [func["func_unique_id"] for func in modify_funcs]
collection.delete(ids=old_func_ids)
# 添加新版本
embeddings = self.vectorize_data(modify_funcs)
self.save_to_chroma(modify_funcs, embeddings)
```
#### 3. 定时同步整合(修改`sync_service.py`
原仓库的`sync_service.py`实现了定时同步的核心逻辑,新增**Git数据源的同步调度**,在`sync_service.py`的main函数里启动所有的同步`GitSync`/`MySQLSync`/`FolderSync`
### 阶段4代码专属检索与生成查询阶段改造
核心实现**代码意图识别、查询优化、ChromaDB元数据过滤、代码专属Prompt生成**,修改原仓库的`main.py`中`/api/chat/stream`接口逻辑,复用原仓库的**流式响应**能力同时确保检索函数code_retrieve与`git_sync.py`存储逻辑完全适配,形成「存储-检索」闭环。
#### 1. 原仓库对接点
- 复用原仓库的**Ollama生成能力**和**流式响应逻辑**,仅修改**检索逻辑**和**Prompt模板**
- 复用原仓库的ChromaDB`query`方法,新增**元数据过滤条件**user_id/repo_id/branch
- 前端聊天界面新增**代码仓库/分支选择器**,传递仓库/分支参数到后端。
#### 2. 核心实现细节(修改`main.py`的聊天接口)
```Python
# 新增代码专属检索ChromaDB元数据过滤+向量匹配)- 已适配git_sync.py存储逻辑
def code_retrieve(user_id: str, repo_id: str, branch: str, query: str, top_k: int) -> list:
"""ChromaDB检索代码函数
核心逻辑通过用户查询生成embedding匹配存储的func_desc向量从metadatas中提取func_body贴合存储逻辑
适配性说明与git_sync.py存储逻辑对应
1. 向量匹配与git_sync.py中vectorize_data方法一致均调用BaseSync._ollama_embedding生成embedding确保检索与存储的向量逻辑统一
2. 元数据过滤where条件user_id/repo_id/branch与git_sync.py.save_to_chroma存入的metadatas字段完全对应确保数据隔离精准
3. 数据提取从metadatas提取func_body/file_path/func_name均为git_sync.py中明确存入的字段无字段缺失
4. 集合命名collection命名code_rag_{user_id}与git_sync.py中存储时的集合命名规则完全一致避免集合错乱。
返回包含func_body及溯源信息的列表供后续生成回答使用
"""
import chromadb
client = chromadb.HttpClient(host=settings.CHROMA_SERVER_HOST, port=settings.CHROMA_SERVER_PORT)
collection = client.get_collection(name=f"code_rag_{self.user_id}")
# 调用原仓库的向量化方法生成查询embedding与git_sync.py中func_desc向量化逻辑完全一致
from sync.base_sync import BaseSync
embedding = BaseSync(user_id)._ollama_embedding([query])[0]
# 元数据过滤:仅检索当前用户/仓库/分支与git_sync.py存入的metadatas字段精准对应
results = collection.query(
query_embeddings=[embedding], # 基于用户查询生成的embedding匹配git_sync.py存储的func_desc向量
n_results=top_k,
where={
"$and": [
{"user_id": user_id},
{"repo_id": repo_id},
{"branch": branch}
]
},
include=["metadatas"] # 明确指定获取metadatas对应git_sync.py中存储func_body的核心位置无需额外获取documents
)
# 从metadatas中提取func_body字段与git_sync.py存入的metadatas完全匹配确保能精准提取
retrieved_funcs = []
for metadata in results["metadatas"][0]: # results["metadatas"]是二维列表外层对应查询次数内层对应top-k结果
func_info = f"文件路径:{metadata['file_path']}
函数名:{metadata['func_name']}
函数代码:{metadata['func_body']}"
retrieved_funcs.append(func_info)
return retrieved_funcs # 返回提取的func_body列表替代原有的documents列表与git_sync.py存储逻辑闭环
# 适配性补充说明与git_sync.py存储逻辑强关联
def retrieve_storage_compatibility_check() -> bool:
"""校验code_retrieve与git_sync.py存储逻辑的适配性可用于启动时自检"""
# 1. 校验集合命名规则一致(确保检索与存储使用同一集合)
from sync.git_sync import GitSync
mock_sync = GitSync(user_id="test", repo_id="test_repo", git_config={})
sync_collection_name = f"code_rag_{mock_sync.user_id}"
retrieve_collection_name = f"code_rag_test"
if sync_collection_name != retrieve_collection_name:
raise Exception("适配异常code_retrieve与git_sync.py的ChromaDB集合命名规则不一致")
# 2. 校验元数据字段一致确保检索时提取的字段均在git_sync.py中已存入
sync_metadata_fields = ["user_id", "repo_id", "branch", "file_path", "func_name", "func_body"]
retrieve_extract_fields = ["file_path", "func_name", "func_body"]
for field in retrieve_extract_fields:
if field not in sync_metadata_fields:
raise Exception(f"适配异常code_retrieve提取的{field}字段未在git_sync.py存储逻辑中定义")
# 3. 校验向量化逻辑一致确保检索与存储的embedding生成方法统一
sync_embedding_logic = "基于func_desc调用BaseSync._ollama_embedding"
retrieve_embedding_logic = "基于用户query调用BaseSync._ollama_embedding"
if not sync_embedding_logic.split("调用")[1] == retrieve_embedding_logic.split("调用")[1]:
raise Exception("适配异常code_retrieve与git_sync.py的向量化方法不一致")
return True
# 新增构建代码专属Prompt
def build_code_prompt(optimized_query: str, retrieved_funcs: list) -> str:
"""复用文档的代码回答Prompt模板适配新的retrieved_funcsfunc_body列表"""
retrieved_context = "\n".join([f"{i+1}. {func}" for i, func in enumerate(retrieved_funcs)])
prompt = f"""
### 角色
你是资深程序员,负责解答用户关于指定代码仓库的技术问题,回答必须严格基于提供的代码上下文,不得编造代码/信息。
### 核心规则
1. 回答需「先给出核心结论,再补充详细解释」,逻辑清晰;
2. 若询问函数功能,需结合函数代码说明功能、参数作用、返回值意义;
3. 若询问实现逻辑,需逐行/分模块解析代码的执行流程;
4. 若询问使用方式,需给出具体的调用示例(基于函数参数);
5. 若提供的代码中无相关答案,明确告知「未检索到相关代码,无法解答」,不做猜测;
6. 代码相关的回答需附带「所属文件路径+函数名」,方便用户溯源。
### 检索到的相关代码(共{len(retrieved_funcs)}个,含完整函数体)
{retrieved_context}
### 用户问题
{optimized_query}
### 请输出你的回答
""".strip()
return prompt
```
## 三、前后端适配新增Git配置+代码聊天界面)
### 1. 后端接口扩展(修改`main.py`
在原仓库的数据源配置路由中,**新增Git类型的配置支持**,无需新增独立路由,仅在入参中判断`type: git`,示例:
```Python
# 修改/api/config/datasource的POST接口
@router.post("/config/datasource")
async def add_datasource(ds_config: DataSourceConfig):
if ds_config.type == "git":
from db_utils import add_git_datasource
add_git_datasource(ds_config.user_id, ds_config.git_config)
return {"code": 200, "msg": "Git数据源配置成功"}
elif ds_config.type == "mysql":
# 原仓库的MySQL配置逻辑
elif ds_config.type == "folder":
# 原仓库的文件夹配置逻辑
```
新增**Git仓库手动同步路由**
```Python
# 新增Git手动同步路由
@router.post("/sync/git")
async def sync_git(user_id: str, repo_id: str):
from db_utils import get_git_datasource
git_config = get_git_datasource(user_id, repo_id)
git_sync = GitSync(user_id, repo_id, git_config)
res = git_sync.run_incremental_sync()
return {"code": 200, "msg": res}
```
### 2. 前端适配(修改`static/config/`和`static/chat/`
#### 1配置界面新增Git配置表单
在`static/config/index.html`中嵌入`git-form.html`组件,实现**Git仓库地址、协议选择、SSH密钥/HTTPS令牌、分支、轮询间隔**的配置,新增:
- 协议选择SSH/HTTPS/git动态显示对应的凭证输入框SSH私钥/HTTPS令牌
- **Git配置验证按钮**:调用`/api/config/verify/git`接口,验证仓库地址和凭证的有效性;
- 分支输入框,默认填充`main/master`。
#### 2聊天界面
无需特别改动。
暂时对代码问答和基于其他知识库的问答不做区分。先简单粗暴把相关代码搜出来即可,让模型判断是否在回答中用搜出来的代码进行增强。
- 知识库检索:按照现在处理,即只是后台无差别检索。
- 围绕知识库回答按照现在处理即“根据参考消息x”。
## 四、配置与部署修改
### 1. 环境变量配置(修改`.env.example`
新增Git代码库相关的全局配置项所有配置通过`config.py`读取,示例:
```TOML
# Git代码库配置
GIT_LOCAL_STORAGE_ROOT=/opt/rag-code-repo # 本地Git仓库存储根目录
GIT_DEFAULT_BRANCH=main # 默认克隆分支
GIT_DEFAULT_POLL_INTERVAL=300 # 默认轮询间隔(秒)
CODE_TOP_K=5 # 代码检索top-k值
LIGHT_LLM_MODEL=qwen2:0.5b # 代码意图识别的轻量LLM模型
# SSH密钥加密配置
AES_KEY=xxxxxxxxxxxxxxxx # AES256加密密钥用于加密SSH/HTTPS凭证
```
### 2. Docker部署修改修改`Dockerfile`和`docker-compose.yml`
原仓库的Docker容器内需要执行Git命令因此**修改Dockerfile安装git**
```Dockerfile
# 原仓库的Dockerfile新增
RUN apt-get update && apt-get install -y git && apt-get clean
# 新建SSH临时目录
RUN mkdir -p /tmp/ssh && chmod 777 /tmp/ssh
```
`docker-compose.yml`中**挂载Git本地存储目录**,实现数据持久化:
```YAML
# 新增卷挂载
volumes:
- ./chroma_db_data:/chroma_db_data
- ./data:/data
- ./rag-code-repo:/opt/rag-code-repo # Git仓库存储目录
```
## 五、异常处理与性能优化(补充)
### 1. 核心异常处理(新增到`git_tool.py`和`git_sync.py`
按文档要求处理Git操作的常见异常示例
- **SSH密钥验证失败**:捕获`subprocess.CalledProcessError`,返回凭证无效提示;
- **Git强制推送**:检测到版本分叉时,执行`git fetch --force`,触发全量同步;
- **网络故障**:采用**指数退避重试**机制重试3次失败后暂停轮询
- **ChromaDB写入失败**捕获ChromaDB的API异常回滚操作记录告警日志。
### 2. 性能优化(复用文档建议)
- **异步分批次解析**大仓库首次解析时按文件分片借助Celery实现异步解析
- **Redis缓存**:将高频检索的函数向量/代码缓存到Redis减少ChromaDB查询压力
- **解析白名单**:支持用户配置需要解析的目录(如`src/`),忽略`node_modules/`/`dist/`等无效目录;
- **LLM描述缓存**函数代码未变更时复用原有LLM描述避免重复调用。
## 六、二次开发后整体流程验证
1. **前端配置**用户新增Git数据源填写仓库地址、SSH密钥、分支点击验证并保存
2. **首次同步**系统自动克隆Git仓库AST切片提取函数LLM生成描述向量化后存入ChromaDB
3. **增量同步**定时轮询远程Git仓库检测到新commit后增量拉取代码识别函数级增删改精准更新ChromaDB
4. **代码查询**用户在代码聊天界面提问系统识别代码意图优化查询后检索ChromaDB基于代码生成精准回答并流式返回。
本次开发完全**复用原仓库的核心架构和能力**仅做Git代码库的专属扩展保证了代码的兼容性和可维护性同时实现了文档中要求的**企业级代码RAG核心能力**。
> (注:文档部分内容可能由 AI 生成)

35
docs/手工debug记录.md Normal file
View File

@ -0,0 +1,35 @@
1. 没找到id。
错因:字典用的字段错了。
解法debug找到纠正。
2. metadata类型错误。
self.collection.add(
ids=valid_batch_ids,
embeddings=valid_batch_embeddings,
documents=valid_batch_texts,
metadatas=valid_batch_metadatas
)
Traceback (most recent call last):
File "<string>", line 1, in <module>
File "s:\research\RAG\.venv\Lib\site-packages\chromadb\api\models\Collection.py", line 106, in add
self._client._add(
~~~~~~~~~~~~~~~~~^
collection_id=self.id,
^^^^^^^^^^^^^^^^^^^^^^
...<6 lines>...
database=self.database,
^^^^^^^^^^^^^^^^^^^^^^^
)
^
File "s:\research\RAG\.venv\Lib\site-packages\chromadb\api\rust.py", line 452, in _add
return self.bindings.add(
~~~~~~~~~~~~~~~~~^
ids,
^^^^
...<6 lines>...
database,
^^^^^^^^^
)
^
TypeError: argument 'metadatas': Cannot convert Python object to MetadataValue
错因是因为valid_batch_metadatas 中包含了 None 值,而 ChromaDB 不允许 None 值作为元数据。
解法:将 None 值设置为 "None" 字符串。

@ -0,0 +1 @@
Subproject commit c22e97f255c92b61ad6014ef74ec662d9f5f6325

@ -0,0 +1 @@
Subproject commit 678dedbbf94be54b3c9c258368e28bb8e7736d62

@ -0,0 +1 @@
Subproject commit 1e967fc87c2761167e0a0a0e84dd2d213c0e1186

@ -0,0 +1 @@
Subproject commit c22e97f255c92b61ad6014ef74ec662d9f5f6325

@ -0,0 +1 @@
Subproject commit 31c3f4bc08e848cb34760d3cdf9065d385ecca14

BIN
install/V1.0.0.0.zip Normal file

Binary file not shown.

BIN
install/soffice.tar.gz Normal file

Binary file not shown.

View File

@ -5,7 +5,6 @@ from llama_index.core.query_engine import RetrieverQueryEngine
from llama_index.core.response_synthesizers import ResponseMode from llama_index.core.response_synthesizers import ResponseMode
from llama_index.core.base.response.schema import StreamingResponse from llama_index.core.base.response.schema import StreamingResponse
from llama_index.llms.ollama import Ollama from llama_index.llms.ollama import Ollama
from llama_index.llms.openai import OpenAI
from llama_index.core import PromptTemplate from llama_index.core import PromptTemplate
from typing import AsyncIterator, Optional, Tuple from typing import AsyncIterator, Optional, Tuple
import asyncio import asyncio
@ -14,6 +13,8 @@ from loguru import logger
from config import settings from config import settings
from .vector_store import VectorStoreManager from .vector_store import VectorStoreManager
from .chunk_handler import OptimizedDeltaThinkFilter from .chunk_handler import OptimizedDeltaThinkFilter
from utils.query_processor import QueryProcessor
from utils.code_prompt_manager import generate_dynamic_code_prompt
# Single module-level prompt string for easy editing in one place # Single module-level prompt string for easy editing in one place
@ -83,71 +84,42 @@ class RAGEngine:
"""Main RAG engine for query processing""" """Main RAG engine for query processing"""
@staticmethod @staticmethod
def check_llm_connection(provider: str, base_url: str, model: str, api_key: Optional[str] = None) -> Tuple[bool, str]: def check_ollama_connection() -> Tuple[bool, str]:
""" """
Check if LLM server is accessible and connection can be established Check if Ollama server is accessible and connection can be established
Args:
provider: LLM provider (ollama, vllm, openai)
base_url: Base URL for the LLM service
model: Model name
api_key: Optional API key for authentication
Returns: Returns:
Tuple of (is_connected: bool, error_message: str) Tuple of (is_connected: bool, error_message: str)
If connected, error_message will be empty string If connected, error_message will be empty string
""" """
try: try:
headers = {} logger.info(f"Checking Ollama connection to {settings.OLLAMA_BASE_URL}...")
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
if provider == "ollama": # Test connection by calling Ollama API
logger.info(f"Checking Ollama connection to {base_url}...") with httpx.Client(timeout=10.0) as client:
with httpx.Client(timeout=10.0) as client: response = client.get(f"{settings.OLLAMA_BASE_URL}/api/tags")
response = client.get(f"{base_url}/api/tags") if response.status_code == 200:
if response.status_code == 200: models = response.json().get("models", [])
models = response.json().get("models", []) model_names = [m.get("name", "unknown") for m in models]
model_names = [m.get("name", "unknown") for m in models] logger.info(f"✓ Ollama connection successful (found {len(models)} model(s): {', '.join(model_names[:3])}{'...' if len(model_names) > 3 else ''})")
logger.info(f"✓ Ollama connection successful (found {len(models)} model(s): {', '.join(model_names[:3])}{'...' if len(model_names) > 3 else ''})")
return True, ""
else:
error_message = f"Ollama API returned status {response.status_code}: {response.text}"
logger.error(f"✗ Ollama connection failed: {error_message}")
return False, error_message
elif provider == "openai":
logger.info(f"Checking OpenAI-compatible API connection to {base_url}...")
if not api_key:
logger.warning("OpenAI provider requires API key, skipping connection check")
return True, "" return True, ""
with httpx.Client(timeout=10.0, headers=headers) as client: else:
response = client.get(f"{base_url}/models") error_message = f"Ollama API returned status {response.status_code}: {response.text}"
if response.status_code == 200: logger.error(f"✗ Ollama connection failed: {error_message}")
models = response.json().get("data", []) return False, error_message
model_names = [m.get("id", "unknown") for m in models]
logger.info(f"✓ OpenAI-compatible API connection successful (found {len(models)} model(s): {', '.join(model_names[:3])}{'...' if len(model_names) > 3 else ''})")
return True, ""
else:
error_message = f"OpenAI API returned status {response.status_code}: {response.text}"
logger.error(f"✗ OpenAI connection failed: {error_message}")
return False, error_message
else:
logger.warning(f"Unknown LLM provider: {provider}, skipping connection check")
return True, ""
except httpx.ConnectError as e: except httpx.ConnectError as e:
error_message = f"Cannot connect to {provider} server at {base_url}. " \ error_message = f"Cannot connect to Ollama server at {settings.OLLAMA_BASE_URL}. " \
f"Please check if server is running and accessible." f"Please check if Ollama server is running and accessible."
logger.error(f"{provider} connection failed: {error_message}") logger.error(f"✗ Ollama connection failed: {error_message}")
return False, error_message return False, error_message
except httpx.TimeoutException: except httpx.TimeoutException:
error_message = f"Connection to {provider} server at {base_url} timed out. " \ error_message = f"Connection to Ollama server at {settings.OLLAMA_BASE_URL} timed out. " \
f"Please check if server is running and accessible." f"Please check if Ollama server is running and accessible."
logger.error(f"{provider} connection failed: {error_message}") logger.error(f"✗ Ollama connection failed: {error_message}")
return False, error_message return False, error_message
except Exception as e: except Exception as e:
error_message = f"Unexpected error while checking {provider} connection: {str(e)}" error_message = f"Unexpected error while checking Ollama connection: {str(e)}"
logger.error(f"{provider} connection check failed: {error_message}") logger.error(f"Ollama connection check failed: {error_message}")
return False, error_message return False, error_message
def __init__( def __init__(
@ -156,59 +128,40 @@ class RAGEngine:
prompt_template: Optional[PromptTemplate] = None, prompt_template: Optional[PromptTemplate] = None,
system_prompt: Optional[str] = None, system_prompt: Optional[str] = None,
temperature: float = 0.7, temperature: float = 0.7,
request_timeout: float = 120.0, request_timeout: float = 1200.0,
): ):
self.vector_store_manager = vector_store_manager self.vector_store_manager = vector_store_manager
# Configurable LLM / prompt parameters
self._user_prompt_template = prompt_template self._user_prompt_template = prompt_template
self._system_prompt = system_prompt self._system_prompt = system_prompt
self._temperature = temperature self._temperature = temperature
self._request_timeout = request_timeout self._request_timeout = request_timeout
llm_config = settings.get_llm_config() # Check Ollama connection before initializing
provider = llm_config["provider"] is_connected, error_message = self.check_ollama_connection()
base_url = llm_config["base_url"]
model = llm_config["model"]
api_key = llm_config["api_key"]
is_connected, error_message = self.check_llm_connection(provider, base_url, model, api_key)
if not is_connected: if not is_connected:
error_msg = ( error_msg = (
f"Error: Cannot connect to {provider} server.\n" f"Error: Cannot connect to Ollama server.\n"
f" {error_message}\n" f" {error_message}\n"
f"Connection info: {base_url}\n" f"Connection info: {settings.OLLAMA_BASE_URL}\n"
f"Please check:\n" f"Please check:\n"
f" 1. {provider} server is running\n" f" 1. Ollama server is running\n"
f" 2. {provider} server is accessible from this host\n" f" 2. Ollama server is accessible from this host\n"
f" 3. LLM_BASE_URL is correctly configured\n" f" 3. OLLAMA_BASE_URL is correctly configured\n"
f" 4. Firewall rules allow connection" f" 4. Firewall rules allow connection to Ollama port"
) )
logger.error(error_msg) logger.error(error_msg)
raise RuntimeError(error_msg) raise RuntimeError(error_msg)
if provider == "ollama": self.llm = Ollama(
self.llm = Ollama( model=settings.OLLAMA_MODEL,
model=model, base_url=settings.OLLAMA_BASE_URL,
base_url=base_url, temperature=self._temperature,
temperature=self._temperature, request_timeout=self._request_timeout,
request_timeout=self._request_timeout, )
)
elif provider == "openai":
self.llm = OpenAI(
model=model,
base_url=base_url,
api_key=api_key,
temperature=self._temperature,
timeout=self._request_timeout,
)
else:
self.llm = Ollama(
model=model,
base_url=base_url,
temperature=self._temperature,
request_timeout=self._request_timeout,
)
self._llm_provider = provider # 初始化查询处理器
self.query_processor = QueryProcessor(llm=self.llm)
def extract_text_from_chunk(self, chunk) -> Optional[str]: def extract_text_from_chunk(self, chunk) -> Optional[str]:
if hasattr(chunk, 'delta'): if hasattr(chunk, 'delta'):
@ -239,29 +192,127 @@ class RAGEngine:
Response text chunks Response text chunks
""" """
try: try:
# Create query engine with streaming mode # 1. 处理查询整合意图识别、filter生成和查询转换
retriever = self.vector_store_manager.get_retriever(top_k=top_k or settings.TOP_K) logger.info(f"开始查询处理: {query}")
# query index process_result = self.query_processor.process_query(query, history)
retrieved_nodes = await retriever.aretrieve(query)
# 2. 构建上下文 # 提取结果
intent_result = process_result['intent']
filters = process_result['filters']
transformed_queries = [process_result['transformed']['rewritten']] + process_result['transformed']['sub_queries']
logger.info(f"代码意图识别结果: {intent_result.category}")
logger.info(f"过滤条件: {filters}")
logger.info(f"查询转换结果: {transformed_queries}")
# 4. 使用融合检索,对每个转换后的查询进行检索
logger.info("使用融合检索策略")
all_hybrid_results = []
for transformed_query in transformed_queries:
hybrid_results = await self.vector_store_manager.ahybrid_search(
query=transformed_query,
top_k=settings.TOP_K,
filters=filters
)
all_hybrid_results.extend(hybrid_results)
# 去重并按得分排序
seen_doc_ids = set()
unique_hybrid_results = []
for doc_id, score, metadata in all_hybrid_results:
if doc_id not in seen_doc_ids:
seen_doc_ids.add(doc_id)
unique_hybrid_results.append((doc_id, score, metadata))
# 按得分排序
unique_hybrid_results.sort(key=lambda x: x[1], reverse=True)
# 限制结果数量
hybrid_results = unique_hybrid_results[:top_k or settings.TOP_K]
# 美化输出融合检索结果
logger.info("融合检索结果:")
for i, (doc_id, score, metadata) in enumerate(hybrid_results):
func_name = metadata.get('func_name', 'N/A')
file_path = metadata.get('file_path', 'N/A')
lang = metadata.get('lang', 'N/A')
logger.info(f" [{i+1}] 相似度: {score:.4f}")
logger.info(f" 函数: {func_name}")
logger.info(f" 文件: {file_path}")
logger.info(f" 语言: {lang}")
logger.info(" " + "-" * 50)
# 根据文档ID获取完整的节点信息
retrieved_nodes = []
for doc_id, score, metadata in hybrid_results:
# 从向量存储中获取文档内容
doc_chunks = self.vector_store_manager.get_document_by_id(doc_id)
for chunk in doc_chunks:
# 创建节点对象
from llama_index.core.schema import TextNode
node = TextNode(
text=chunk['text'],
node_id=chunk['id'],
metadata=chunk['metadata']
)
retrieved_nodes.append(node)
# 3. 构建上下文
context_parts = [] context_parts = []
for i, node in enumerate(retrieved_nodes[:top_k or settings.TOP_K], 1): # 限制数量 max_nodes = top_k or settings.TOP_K
# 构建原始的检索结果列表,包含 text 和 metadata
retrieved_results = []
for i, node in enumerate(retrieved_nodes[:max_nodes], 1):
text = node.text if hasattr(node, 'text') else str(node) text = node.text if hasattr(node, 'text') else str(node)
# 清理和截断
text = text.strip() text = text.strip()
if len(text) > 400:
text = text[:400] + "..." # 获取原始 metadata
context_parts.append(f"【参考信息{i}{text}") metadata = getattr(node, 'metadata', {})
# 保存原始的 text 和 metadata
retrieved_results.append({
'id': i,
'text': text,
'metadata': metadata
})
# 构建简单的上下文字符串,只包含基本信息
context_parts = []
for result in retrieved_results:
metadata = result['metadata']
file_path = metadata.get('file_path', '')
func_name = metadata.get('func_name', '')
context_part = f"【参考信息{result['id']}"
if file_path:
context_part += f"(来源:{file_path}"
if func_name:
context_part += f"\n函数:{func_name}"
context_part += f"\n{result['text']}"
context_parts.append(context_part)
context_str = "\n\n".join(context_parts) if context_parts else "未找到相关参考信息" context_str = "\n\n".join(context_parts) if context_parts else "未找到相关参考信息"
logger.info(f"上下文: {context_str}")
if history is not None: # 4. 生成优化的Prompt
qa_prompt = QA_PROMPT_HISTORY logger.info("生成优化的Prompt")
filled_prompt = qa_prompt.format(history=history, context_str=context_str, query_str=query) if intent_result.is_code_related:
# 对于代码相关问题使用代码专用Prompt
filled_prompt = generate_dynamic_code_prompt(
user_query=query,
intent_result=intent_result.to_dict(),
code_context=context_str,
retrieved_results=retrieved_results,
conversation_history=history
)
logger.info(f"使用代码专用Prompt类型: {intent_result.category}")
else: else:
qa_prompt = QA_PROMPT_NO_HISTORY # 对于非代码问题使用通用Prompt
filled_prompt = qa_prompt.format(context_str=context_str, query_str=query) if history:
qa_prompt = QA_PROMPT_HISTORY
filled_prompt = qa_prompt.format(history=history, context_str=context_str, query_str=query)
else:
qa_prompt = QA_PROMPT_NO_HISTORY
filled_prompt = qa_prompt.format(context_str=context_str, query_str=query)
logger.info("使用通用Prompt")
stream_response = await self.llm.astream_complete( stream_response = await self.llm.astream_complete(
prompt=filled_prompt prompt=filled_prompt
@ -272,7 +323,6 @@ class RAGEngine:
async for chunk in stream_response: async for chunk in stream_response:
# 提取文本内容 # 提取文本内容
# text_chunk = self.extract_text_from_chunk(chunk)
delta, full_text, has_output = think_filter.process_delta_robust(chunk) delta, full_text, has_output = think_filter.process_delta_robust(chunk)
if delta is not None: if delta is not None:
@ -281,10 +331,12 @@ class RAGEngine:
await asyncio.sleep(0.001) # slight delay to yield control await asyncio.sleep(0.001) # slight delay to yield control
logger.info(f"响应完成,长度: {len(full_response)}字符") logger.info(f"响应完成,长度: {len(full_response)}字符")
print(full_response)
except Exception as e: except Exception as e:
logger.error(f"Error in RAG query: {e}") logger.error(f"Error in RAG query: {e}")
async def query(self, query: str, history: str, top_k: Optional[int] = None) -> str: async def query(self, query: str, history: str, top_k: Optional[int] = None) -> str:
""" """
Query the RAG system and return complete response Query the RAG system and return complete response
@ -298,29 +350,113 @@ class RAGEngine:
Complete response string Complete response string
""" """
try: try:
# Create query engine with streaming mode # 1. 处理查询整合意图识别、filter生成和查询转换
retriever = self.vector_store_manager.get_retriever(top_k=top_k or settings.TOP_K) logger.info(f"开始查询处理: {query[:50]}...")
# query index process_result = self.query_processor.process_query(query, history)
retrieved_nodes = await retriever.aretrieve(query)
# 2. 构建上下文 # 提取结果
context_parts = [] intent_result = process_result['intent']
for i, node in enumerate(retrieved_nodes[:top_k or settings.TOP_K], 1): # 限制数量 filters = process_result['filters']
transformed_queries = [process_result['transformed']['rewritten']] + process_result['transformed']['sub_queries']
logger.info(f"代码意图识别结果: {intent_result.category}")
logger.info(f"过滤条件: {filters}")
logger.info(f"查询转换完成,生成了 {len(transformed_queries)} 个转换后的查询")
# 2. 使用融合检索,对每个转换后的查询进行检索
logger.info("使用融合检索策略")
all_hybrid_results = []
for transformed_query in transformed_queries:
hybrid_results = await self.vector_store_manager.ahybrid_search(
query=transformed_query,
top_k=top_k or settings.TOP_K,
filters=filters
)
all_hybrid_results.extend(hybrid_results)
# 去重并按得分排序
seen_doc_ids = set()
unique_hybrid_results = []
for doc_id, score, metadata in all_hybrid_results:
if doc_id not in seen_doc_ids:
seen_doc_ids.add(doc_id)
unique_hybrid_results.append((doc_id, score, metadata))
# 按得分排序
unique_hybrid_results.sort(key=lambda x: x[1], reverse=True)
# 限制结果数量
hybrid_results = unique_hybrid_results[:top_k or settings.TOP_K]
# 根据文档ID获取完整的节点信息
retrieved_nodes = []
for doc_id, score, metadata in hybrid_results:
# 从向量存储中获取文档内容
doc_chunks = self.vector_store_manager.get_document_by_id(doc_id)
for chunk in doc_chunks:
# 创建节点对象
from llama_index.core.schema import TextNode
node = TextNode(
text=chunk['text'],
node_id=chunk['id'],
metadata=chunk['metadata']
)
retrieved_nodes.append(node)
# 3. 构建原始的检索结果列表,包含 text 和 metadata
retrieved_results = []
for i, node in enumerate(retrieved_nodes[:top_k or settings.TOP_K], 1):
text = node.text if hasattr(node, 'text') else str(node) text = node.text if hasattr(node, 'text') else str(node)
# 清理和截断
text = text.strip() text = text.strip()
if len(text) > 400:
text = text[:400] + "..." # 获取原始 metadata
context_parts.append(f"【参考信息{i}{text}") metadata = getattr(node, 'metadata', {})
# 保存原始的 text 和 metadata
retrieved_results.append({
'id': i,
'text': text,
'metadata': metadata
})
# 构建简单的上下文字符串,只包含基本信息
context_parts = []
for result in retrieved_results:
metadata = result['metadata']
file_path = metadata.get('file_path', '')
func_name = metadata.get('func_name', '')
context_part = f"【参考信息{result['id']}"
if file_path:
context_part += f"(来源:{file_path}"
if func_name:
context_part += f"\n函数:{func_name}"
context_part += f"\n{result['text']}"
context_parts.append(context_part)
context_str = "\n\n".join(context_parts) if context_parts else "未找到相关参考信息" context_str = "\n\n".join(context_parts) if context_parts else "未找到相关参考信息"
if history is not None: # 4. 生成优化的Prompt
qa_prompt = QA_PROMPT_HISTORY logger.info("生成优化的Prompt")
filled_prompt = qa_prompt.format(history=history, context_str=context_str, query_str=query) if intent_result.is_code_related:
# 对于代码相关问题使用代码专用Prompt
filled_prompt = generate_dynamic_code_prompt(
user_query=query,
intent_result=intent_result.to_dict(),
code_context=context_str,
retrieved_results=retrieved_results,
conversation_history=history
)
logger.info(f"使用代码专用Prompt类型: {intent_result.category}")
else: else:
qa_prompt = QA_PROMPT_NO_HISTORY # 对于非代码问题使用通用Prompt
filled_prompt = qa_prompt.format(context_str=context_str, query_str=query) if history:
qa_prompt = QA_PROMPT_HISTORY
filled_prompt = qa_prompt.format(history=history, context_str=context_str, query_str=query)
else:
qa_prompt = QA_PROMPT_NO_HISTORY
filled_prompt = qa_prompt.format(context_str=context_str, query_str=query)
logger.info("使用通用Prompt")
response = await self.llm.acomplete( response = await self.llm.acomplete(
prompt=filled_prompt prompt=filled_prompt

View File

@ -5,15 +5,19 @@ import os
import time import time
from datetime import datetime, date from datetime import datetime, date
from concurrent.futures import ThreadPoolExecutor, as_completed from concurrent.futures import ThreadPoolExecutor, as_completed
from typing import List, Dict, Any from typing import List, Dict, Any, Tuple
import chromadb import chromadb
from chromadb.config import Settings as ChromaSettings from chromadb.config import Settings as ChromaSettings
from llama_index.vector_stores.chroma import ChromaVectorStore from llama_index.vector_stores.chroma import ChromaVectorStore
from llama_index.core import VectorStoreIndex, StorageContext from llama_index.core import VectorStoreIndex, StorageContext
from llama_index.embeddings.ollama import OllamaEmbedding from llama_index.embeddings.ollama import OllamaEmbedding
from llama_index.embeddings.openai import OpenAIEmbedding
from loguru import logger from loguru import logger
from config import settings from config import settings
import numpy as np
import re
import string
from rank_bm25 import BM25Okapi
from utils.query_processor import MetadataFilter
class VectorStoreManager: class VectorStoreManager:
@ -107,31 +111,11 @@ class VectorStoreManager:
metadata={"hnsw:space": "cosine"} metadata={"hnsw:space": "cosine"}
) )
# Initialize embedding model based on provider # Initialize embedding model
embed_config = settings.get_embedding_config() self.embed_model = OllamaEmbedding(
provider = embed_config["provider"] model_name=settings.OLLAMA_EMBEDDING_MODEL,
base_url = embed_config["base_url"] base_url=settings.OLLAMA_BASE_URL
model = embed_config["model"] )
api_key = embed_config["api_key"]
if provider == "ollama":
self.embed_model = OllamaEmbedding(
model_name=model,
base_url=base_url
)
elif provider == "openai":
self.embed_model = OpenAIEmbedding(
model=model,
base_url=base_url,
api_key=api_key
)
else:
self.embed_model = OllamaEmbedding(
model_name=model,
base_url=base_url
)
logger.info(f"Initialized {provider} embedding model: {model}")
# Create ChromaVectorStore # Create ChromaVectorStore
self.vector_store = ChromaVectorStore(chroma_collection=self.collection) self.vector_store = ChromaVectorStore(chroma_collection=self.collection)
@ -716,10 +700,17 @@ class VectorStoreManager:
logger.error(f"写入剩余文档到 ChromaDB 失败: {write_error}") logger.error(f"写入剩余文档到 ChromaDB 失败: {write_error}")
total_elapsed = time.time() - start_time total_elapsed = time.time() - start_time
logger.info( if total_added > 0:
f"[Embedding进度] ✓ 单个生成完成: 成功写入 {total_added}/{total_count} 个文档 " avg_time = total_elapsed/total_added*1000
f"(耗时: {total_elapsed:.1f}秒, 平均: {total_elapsed/total_added*1000:.1f}ms/个, 跳过: {skipped_count}个)" logger.info(
) f"[Embedding进度] ✓ 单个生成完成: 成功写入 {total_added}/{total_count} 个文档 "
f"(耗时: {total_elapsed:.1f}秒, 平均: {avg_time:.1f}ms/个, 跳过: {skipped_count}个)"
)
else:
logger.info(
f"[Embedding进度] ✓ 单个生成完成: 成功写入 {total_added}/{total_count} 个文档 "
f"(耗时: {total_elapsed:.1f}秒, 跳过: {skipped_count}个)"
)
if total_added == 0: if total_added == 0:
raise ValueError(f"Failed to add any valid documents to ChromaDB: {e}") raise ValueError(f"Failed to add any valid documents to ChromaDB: {e}")
@ -836,12 +827,13 @@ class VectorStoreManager:
logger.warning(f"Error checking document count: {e}") logger.warning(f"Error checking document count: {e}")
return False return False
def get_retriever(self, top_k: int = None): def get_retriever(self, top_k: int = None, filters: dict = None):
""" """
Get a retriever for querying the vector store Get a retriever for querying the vector store
Args: Args:
top_k: Number of documents to retrieve (defaults to settings.TOP_K) top_k: Number of documents to retrieve (defaults to settings.TOP_K)
filters: Metadata filters to apply before vector search
Returns: Returns:
VectorStoreRetriever instance VectorStoreRetriever instance
@ -898,3 +890,314 @@ class VectorStoreManager:
except Exception as e: except Exception as e:
logger.error(f"获取db_source({target_db_source}) metadata中content_column失败: {e}") logger.error(f"获取db_source({target_db_source}) metadata中content_column失败: {e}")
return '' return ''
def _preprocess_text(self, text: str) -> List[str]:
"""
预处理文本用于BM25搜索
Args:
text: 原始文本
Returns:
分词后的文本列表
"""
text = text.lower() # 转换为小写
text = text.translate(str.maketrans('', '', string.punctuation)) # 移除标点符号
tokens = re.findall(r'\b\w+\b', text) # 分词
return tokens
def keyword_search(self, query: str, top_k: int = 5, filters: dict = None) -> List[Tuple[str, float, Dict[str, Any]]]:
"""
使用BM25算法进行关键词搜索
Args:
query: 搜索查询
top_k: 返回结果数量
filters: metadata过滤条件
Returns:
排序后的结果列表每个元素包含(doc_id, score, metadata)
"""
try:
# 获取文档
results = self.collection.get(
include=['documents', 'metadatas']
)
documents = results.get('documents', [])
metadatas = results.get('metadatas', [])
ids = results.get('ids', [])
if not documents:
return []
# 应用过滤条件
filtered_documents = []
filtered_metadatas = []
filtered_ids = []
for doc, meta, doc_id in zip(documents, metadatas, ids):
if not filters:
# 没有过滤条件,直接添加
filtered_documents.append(doc)
filtered_metadatas.append(meta)
filtered_ids.append(doc_id)
else:
# 应用过滤条件
match = True
for key, value in filters.items():
if key not in meta:
match = False
break
meta_value = meta[key]
if isinstance(meta_value, str) and isinstance(value, str):
# 对于字符串类型,使用大小写不敏感的模糊匹配
if value.lower() not in meta_value.lower():
match = False
break
else:
# 对于其他类型,使用精确匹配
if meta_value != value:
match = False
break
if match:
filtered_documents.append(doc)
filtered_metadatas.append(meta)
filtered_ids.append(doc_id)
if not filtered_documents:
return []
tokenized_docs = [self._preprocess_text(doc) for doc in filtered_documents] # 预处理文档
bm25 = BM25Okapi(tokenized_docs) # 初始化BM25
tokenized_query = self._preprocess_text(query) # 预处理查询
scores = bm25.get_scores(tokenized_query) # 计算BM25得分
sorted_indices = np.argsort(scores)[::-1][:top_k] # 排序并获取top_k结果
# 构建结果列表
search_results = []
for idx in sorted_indices:
if scores[idx] > 0: # 只返回得分大于0的结果
doc_id = filtered_ids[idx]
score = float(scores[idx])
metadata = filtered_metadatas[idx]
# 由于已经在获取文档时应用了过滤条件,这里不需要再次应用
search_results.append((doc_id, score, metadata))
return search_results
except Exception as e:
logger.error(f"关键词搜索失败: {e}")
return []
def hybrid_search(self, query: str, top_k: int = 5, vector_weight: float = 0.6, keyword_weight: float = 0.4, filters: dict = None) -> List[Tuple[str, float, Dict[str, Any]]]:
"""
融合检索结合向量搜索和关键词搜索
Args:
query: 搜索查询
top_k: 返回结果数量
vector_weight: 向量搜索权重
keyword_weight: 关键词搜索权重
filters: metadata过滤条件
Returns:
排序后的结果列表每个元素包含(doc_id, score, metadata)
"""
try:
# 1. 执行向量搜索
# 获取更多结果以确保有足够的候选
vector_retriever = self.get_retriever(top_k=top_k * 4) # 获取更多结果
vector_nodes = vector_retriever.retrieve(query)
vector_results = {}
for node in vector_nodes:
if hasattr(node, 'id_'):
doc_id = node.id_
elif hasattr(node, 'node_id'):
doc_id = node.node_id
else:
continue
vector_results[doc_id] = {
'score': node.score if hasattr(node, 'score') else 0.5,
'metadata': node.metadata if hasattr(node, 'metadata') else {},
'text': node.text if hasattr(node, 'text') else ''
}
# 2. 执行关键词搜索
keyword_results = self.keyword_search(query, top_k=top_k * 4, filters=filters)
keyword_scores = {}
keyword_metadata = {}
for doc_id, score, metadata in keyword_results:
keyword_scores[doc_id] = score
keyword_metadata[doc_id] = metadata
# 3. 归一化得分
# 归一化向量得分
if vector_results:
vector_scores = list(vector_results.values())
vector_min = min(item['score'] for item in vector_scores)
vector_max = max(item['score'] for item in vector_scores)
vector_range = vector_max - vector_min if vector_max > vector_min else 1
for doc_id in vector_results:
vector_results[doc_id]['normalized_score'] = (vector_results[doc_id]['score'] - vector_min) / vector_range
# 归一化关键词得分
if keyword_scores:
keyword_min = min(keyword_scores.values())
keyword_max = max(keyword_scores.values())
keyword_range = keyword_max - keyword_min if keyword_max > keyword_min else 1
for doc_id in keyword_scores:
keyword_scores[doc_id] = (keyword_scores[doc_id] - keyword_min) / keyword_range
# 4. 融合得分
hybrid_results = {}
# 合并向量搜索结果
for doc_id, info in vector_results.items():
# 应用过滤条件
if filters:
metadata = info['metadata']
match = True
for key, value in filters.items():
if key not in metadata:
match = False
break
meta_value = metadata[key]
if isinstance(meta_value, str) and isinstance(value, str):
# 对于字符串类型,使用大小写不敏感的模糊匹配
if value.lower() not in meta_value.lower():
match = False
break
else:
# 对于其他类型,使用精确匹配
if meta_value != value:
match = False
break
if not match:
continue
vector_score = info.get('normalized_score', 0)
keyword_score = keyword_scores.get(doc_id, 0)
# 计算融合得分
hybrid_score = vector_weight * vector_score + keyword_weight * keyword_score
hybrid_results[doc_id] = {
'score': hybrid_score,
'metadata': info['metadata'],
'text': info['text']
}
# 合并关键词搜索结果(不在向量搜索结果中的)
for doc_id, score in keyword_scores.items():
if doc_id not in hybrid_results:
# 应用过滤条件
if filters:
metadata = keyword_metadata.get(doc_id, {})
match = True
for key, value in filters.items():
if key not in metadata:
match = False
break
meta_value = metadata[key]
if isinstance(meta_value, str) and isinstance(value, str):
# 对于字符串类型,使用大小写不敏感的模糊匹配
if value.lower() not in meta_value.lower():
match = False
break
else:
# 对于其他类型,使用精确匹配
if meta_value != value:
match = False
break
if not match:
continue
hybrid_score = keyword_weight * score
hybrid_results[doc_id] = {
'score': hybrid_score,
'metadata': keyword_metadata.get(doc_id, {}),
'text': ''
}
# 5. 排序并获取top_k结果
sorted_results = sorted(
hybrid_results.items(),
key=lambda x: x[1]['score'],
reverse=True
)[:top_k]
# 6. 构建最终结果
final_results = []
for doc_id, info in sorted_results:
metadata = info['metadata']
final_results.append((doc_id, info['score'], metadata))
return final_results
except Exception as e:
logger.error(f"融合检索失败: {e}")
# 失败时回退到向量搜索
vector_retriever = self.get_retriever(top_k=top_k)
vector_nodes = vector_retriever.retrieve(query)
fallback_results = []
for node in vector_nodes:
if hasattr(node, 'id_'):
doc_id = node.id_
elif hasattr(node, 'node_id'):
doc_id = node.node_id
else:
continue
# 应用过滤条件
if filters:
metadata = node.metadata if hasattr(node, 'metadata') else {}
match = True
for key, value in filters.items():
if key not in metadata:
match = False
break
meta_value = metadata[key]
if isinstance(meta_value, str) and isinstance(value, str):
# 对于字符串类型,使用大小写不敏感的模糊匹配
if value.lower() not in meta_value.lower():
match = False
break
else:
# 对于其他类型,使用精确匹配
if meta_value != value:
match = False
break
if not match:
continue
score = node.score if hasattr(node, 'score') else 0.5
metadata = node.metadata if hasattr(node, 'metadata') else {}
fallback_results.append((doc_id, score, metadata))
return fallback_results
async def ahybrid_search(self, query: str, top_k: int = 5, vector_weight: float = 0.6, keyword_weight: float = 0.4, filters: dict = None) -> List[Tuple[str, float, Dict[str, Any]]]:
"""
异步融合检索结合向量搜索和关键词搜索
Args:
query: 搜索查询
top_k: 返回结果数量
vector_weight: 向量搜索权重
keyword_weight: 关键词搜索权重
filters: metadata过滤条件
Returns:
排序后的结果列表每个元素包含(doc_id, score, metadata)
"""
# 由于BM25搜索是CPU密集型的这里使用同步方法
# 在实际生产环境中,可以使用线程池来异步执行
return self.hybrid_search(query, top_k, vector_weight, keyword_weight, filters)

View File

@ -19,6 +19,7 @@
<option value="">选择类型</option> <option value="">选择类型</option>
<option value="database">数据库</option> <option value="database">数据库</option>
<option value="folder">文件夹</option> <option value="folder">文件夹</option>
<option value="git">Git代码库</option>
</select> </select>
</div> </div>
</div> </div>
@ -75,6 +76,7 @@
<select id="configType" name="type" required> <select id="configType" name="type" required>
<option value="database">数据库 (database)</option> <option value="database">数据库 (database)</option>
<option value="folder">文件夹 (folder)</option> <option value="folder">文件夹 (folder)</option>
<option value="git">Git代码库 (git)</option>
</select> </select>
</div> </div>
<div class="form-group" id="dbTypeGroup" style="display: none;"> <div class="form-group" id="dbTypeGroup" style="display: none;">

View File

@ -224,6 +224,7 @@ function generateConfigForm(config) {
<select id="formType" disabled> <select id="formType" disabled>
<option value="database" ${config.type === 'database' ? 'selected' : ''}>数据库 (database)</option> <option value="database" ${config.type === 'database' ? 'selected' : ''}>数据库 (database)</option>
<option value="folder" ${config.type === 'folder' ? 'selected' : ''}>文件夹 (folder)</option> <option value="folder" ${config.type === 'folder' ? 'selected' : ''}>文件夹 (folder)</option>
<option value="git" ${config.type === 'git' ? 'selected' : ''}>Git代码库 (git)</option>
</select> </select>
</div> </div>
`; `;
@ -285,6 +286,48 @@ function generateConfigForm(config) {
bindFolderEvents(); bindFolderEvents();
} }
// Git代码库配置
if (config.type === 'git') {
const gitSection = document.createElement('div');
gitSection.innerHTML = `
<h3 class="section-title">Git代码库配置</h3>
<div class="form-group">
<label for="formGitUrl">Git仓库地址 <span class="required">*</span></label>
<input type="text" id="formGitUrl" value="${config.git_url || ''}" required>
</div>
<div class="form-group">
<label for="formBranch">分支名称</label>
<input type="text" id="formBranch" value="${config.branch || 'main'}">
</div>
<div class="form-group">
<label for="formProtocol">协议类型</label>
<select id="formProtocol">
<option value="https" ${config.protocol === 'https' ? 'selected' : ''}>HTTPS</option>
<option value="ssh" ${config.protocol === 'ssh' ? 'selected' : ''}>SSH</option>
</select>
</div>
<div class="form-group" id="httpsTokenGroup">
<label for="formHttpsToken">HTTPS令牌</label>
<input type="password" id="formHttpsToken" value="${config.https_token || ''}">
</div>
<div class="form-group" id="sshKeyGroup" style="display: none;">
<label for="formSshKey">SSH私钥</label>
<textarea id="formSshKey" rows="10" value="${config.ssh_key || ''}">${config.ssh_key || ''}</textarea>
</div>
<div class="form-group">
<label for="formPollInterval">轮询间隔</label>
<input type="number" id="formPollInterval" value="${config.poll_interval || 300}" min="60" max="3600">
</div>
<div class="form-actions">
<button type="button" id="testGitConnectionBtn" class="btn secondary" style="margin-top: 10px;">🔌 测试Git连接</button>
</div>
`;
formElement.appendChild(gitSection);
// 绑定事件
bindGitEvents();
}
if (config.type === 'database') { if (config.type === 'database') {
// 数据库连接配置(放在前面,方便先测试连接) // 数据库连接配置(放在前面,方便先测试连接)
const connectionSection = document.createElement('div'); const connectionSection = document.createElement('div');
@ -1085,6 +1128,16 @@ async function handleAddConfigDirectly() {
username: '', username: '',
password: '' password: ''
}; };
} else if (configType === 'git') {
tempConfig = {
type: configType,
git_url: '',
branch: 'main',
protocol: 'https',
https_token: '',
ssh_key: '',
poll_interval: 300
};
} else { } else {
alert('不支持的配置类型'); alert('不支持的配置类型');
return; return;
@ -1217,12 +1270,17 @@ async function saveConfig() {
// 文件夹folder_主机_文件夹路径替换特殊字符 // 文件夹folder_主机_文件夹路径替换特殊字符
const folderName = formData.folder_path ? formData.folder_path.replace(/[\\/:*?"<>|]/g, '_') : 'unknown'; const folderName = formData.folder_path ? formData.folder_path.replace(/[\\/:*?"<>|]/g, '_') : 'unknown';
generatedName = `folder_${formData.host || 'unknown'}_${folderName}`; generatedName = `folder_${formData.host || 'unknown'}_${folderName}`;
} else if (formData.type === 'git') {
// Git配置git_仓库地址替换特殊字符
const repoName = formData.git_url ? formData.git_url.replace(/[\\/:*?"<>|]/g, '_') : 'unknown';
generatedName = `git_${repoName}_${formData.branch || 'main'}`;
} else { } else {
// 不支持的配置类型 // 不支持的配置类型
alert('不支持的配置类型'); alert('不支持的配置类型');
return; return;
} }
formData.name = generatedName;
} }
// 根据不同类型检查特定字段 // 根据不同类型检查特定字段
@ -1243,6 +1301,10 @@ async function saveConfig() {
if (!formData.port) missingFields.push('端口'); if (!formData.port) missingFields.push('端口');
if (!formData.username) missingFields.push('用户名'); if (!formData.username) missingFields.push('用户名');
if (!formData.password) missingFields.push('密码'); if (!formData.password) missingFields.push('密码');
} else if (formData.type === 'git') {
if (!formData.git_url) missingFields.push('Git仓库地址');
if (!formData.branch) missingFields.push('分支名称');
if (!formData.protocol) missingFields.push('协议类型');
} }
// 如果有缺失的字段,提示用户 // 如果有缺失的字段,提示用户
@ -1417,6 +1479,16 @@ function collectFormData() {
} }
// Git配置
if (formData.type === 'git') {
formData.git_url = document.getElementById('formGitUrl').value;
formData.branch = document.getElementById('formBranch').value;
formData.protocol = document.getElementById('formProtocol').value;
formData.https_token = document.getElementById('formHttpsToken').value;
formData.ssh_key = document.getElementById('formSshKey').value;
formData.poll_interval = parseInt(document.getElementById('formPollInterval').value);
}
return formData; return formData;
} }
@ -1468,6 +1540,81 @@ async function confirmDeleteConfig() {
} }
} }
// Git事件绑定函数
function bindGitEvents() {
// 协议切换逻辑
const protocolSelect = document.getElementById('formProtocol');
const httpsTokenGroup = document.getElementById('httpsTokenGroup');
const sshKeyGroup = document.getElementById('sshKeyGroup');
protocolSelect.addEventListener('change', function() {
const protocol = this.value;
if (protocol === 'https') {
httpsTokenGroup.style.display = 'block';
sshKeyGroup.style.display = 'none';
} else if (protocol === 'ssh') {
httpsTokenGroup.style.display = 'none';
sshKeyGroup.style.display = 'block';
}
});
// 触发一次change事件确保初始状态正确
protocolSelect.dispatchEvent(new Event('change'));
// 测试Git连接
document.getElementById('testGitConnectionBtn')?.addEventListener('click', async () => {
try {
const gitUrl = document.getElementById('formGitUrl').value;
const branch = document.getElementById('formBranch').value;
const protocol = document.getElementById('formProtocol').value;
const httpsToken = document.getElementById('formHttpsToken').value;
const sshKey = document.getElementById('formSshKey').value;
if (!gitUrl) {
alert('请填写Git仓库地址');
return;
}
// 禁用按钮
const btn = document.getElementById('testGitConnectionBtn');
const originalText = btn.textContent;
btn.textContent = '🔌 测试中...';
btn.disabled = true;
// 发送请求
const response = await fetch('/git/test-connection', {
method: 'POST',
headers: {
'Content-Type': 'application/json'
},
body: JSON.stringify({
git_url: gitUrl,
branch: branch,
protocol: protocol,
https_token: httpsToken,
ssh_key: sshKey
})
});
if (response.ok) {
const result = await response.json();
alert('Git连接测试成功');
} else {
const errorData = await response.json();
throw new Error(errorData.detail || '连接失败');
}
} catch (error) {
alert('Git连接测试失败: ' + error.message);
} finally {
// 恢复按钮
const btn = document.getElementById('testGitConnectionBtn');
btn.textContent = '🔌 测试Git连接';
btn.disabled = false;
}
});
}
// 为文件夹配置添加事件绑定 // 为文件夹配置添加事件绑定
function bindFolderEvents() { function bindFolderEvents() {
// 测试SSH连接 // 测试SSH连接

235
sync/ast_parser.py Normal file
View File

@ -0,0 +1,235 @@
"""
代码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()

View File

@ -134,7 +134,7 @@ class BaseSync(ABC):
for doc in docs: for doc in docs:
try: try:
llamaindex_doc = self.doc_to_llamaindex_doc(doc) llamaindex_doc = self.doc_to_llamaindex_doc(doc)
if len(llamaindex_doc.text.strip()) >= 100: if len(llamaindex_doc.text.strip()) >= 10:
documents.append(llamaindex_doc) documents.append(llamaindex_doc)
else: else:
logger.warning(f"跳过过短文档 (id: {llamaindex_doc.id_}),内容长度: {len(llamaindex_doc.text)} 字符") logger.warning(f"跳过过短文档 (id: {llamaindex_doc.id_}),内容长度: {len(llamaindex_doc.text)} 字符")
@ -164,7 +164,7 @@ class BaseSync(ABC):
from config import settings from config import settings
node_parser = SentenceSplitter( node_parser = SentenceSplitter(
chunk_size=settings.CHUNK_SIZE, chunk_size=settings.CHUNK_SIZE, #NOTE: 从settings中获取默认1024
chunk_overlap=settings.CHUNK_OVERLAP chunk_overlap=settings.CHUNK_OVERLAP
) )
@ -285,7 +285,7 @@ def get_sync_class(source_type: str, db_type: str = 'mysql') -> type[BaseSync]:
Get the appropriate sync class based on data source type Get the appropriate sync class based on data source type
Args: Args:
source_type: Type of data source (database, folder) source_type: Type of data source (database, folder, git)
db_type: Type of database (mysql, dameng, etc.) - only used when source_type is 'database' db_type: Type of database (mysql, dameng, etc.) - only used when source_type is 'database'
Returns: Returns:
@ -297,6 +297,7 @@ def get_sync_class(source_type: str, db_type: str = 'mysql') -> type[BaseSync]:
from sync.mysql_sync import MySQLSync from sync.mysql_sync import MySQLSync
from sync.folder_sync import FolderSync from sync.folder_sync import FolderSync
from sync.dameng_sync import DaMengSync from sync.dameng_sync import DaMengSync
from sync.git_sync import GitSync
if source_type == 'database': if source_type == 'database':
# 根据数据库类型选择相应的同步类 # 根据数据库类型选择相应的同步类
@ -307,5 +308,7 @@ def get_sync_class(source_type: str, db_type: str = 'mysql') -> type[BaseSync]:
return MySQLSync return MySQLSync
elif source_type == 'folder': elif source_type == 'folder':
return FolderSync return FolderSync
elif source_type == 'git':
return GitSync
else: else:
raise ValueError(f"Unsupported data source type: {source_type}") raise ValueError(f"Unsupported data source type: {source_type}")

301
sync/git_sync.py Normal file
View File

@ -0,0 +1,301 @@
"""
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)
# 创建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", "")[:200], # 进一步限制文档字符串长度
"start_line": doc.get("start_line"),
"end_line": doc.get("end_line"),
"repo_id": self.config.name,
"branch": self.config.branch,
"func_body": doc["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()}")
return ". ".join(parts)
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

View File

@ -197,6 +197,33 @@ class SyncService:
chunked_docs = self.syncer.chunk_documents(processed_docs) chunked_docs = self.syncer.chunk_documents(processed_docs)
all_chunked_docs.extend(chunked_docs) all_chunked_docs.extend(chunked_docs)
# Update last sync time
self.last_sync_time = datetime.now()
elif self.source_config.type == "git":
if not force:
new_documents = []
for doc in documents:
# 检查服务运行状态:仅在非手动同步时检查
if not is_manual and not self._running:
logger.info(f"Sync interrupted during document filtering: {self.source_name}")
return all_chunked_docs, total_docs, skipped_docs_count # 返回结果,不退出程序
doc_id = doc.get('id')
if not self.vector_store_manager.document_exists(doc_id):
new_documents.append(doc)
else:
skipped_docs_count += 1
if not new_documents:
self.last_sync_time = datetime.now()
return all_chunked_docs, total_docs, skipped_docs_count
documents = new_documents
# Process and chunk documents
processed_docs = self.syncer.process_documents(documents)
chunked_docs = self.syncer.chunk_documents(processed_docs)
all_chunked_docs.extend(chunked_docs)
# Update last sync time # Update last sync time
self.last_sync_time = datetime.now() self.last_sync_time = datetime.now()
else: else:
@ -358,7 +385,7 @@ class SyncService:
return return
max_restart_attempts = 10 # Maximum number of restart attempts max_restart_attempts = 10 # Maximum number of restart attempts
restart_delay = 60 # Wait 60 seconds before restarting after an error restart_delay = 10 # Wait 60 seconds before restarting after an error
restart_count = 0 restart_count = 0
self._running = True self._running = True

View File

@ -0,0 +1,725 @@
"""
代码Prompt管理模块
用于生成和管理代码相关问答的专属Prompt模板
支持基于代码意图分类的动态Prompt选择和生成
"""
import json
import re
import sys
import os
from typing import Dict, Any, Optional, List, Tuple
from enum import Enum
from loguru import logger
# 添加项目根目录到 Python 模块搜索路径
current_dir = os.path.dirname(os.path.abspath(__file__))
project_root = os.path.dirname(current_dir)
if project_root not in sys.path:
sys.path.insert(0, project_root)
from utils.query_processor import CodeIntentCategory, PromptTemplateType
from utils.prompt import (
INTENT_DETECTION_TEMPLATE,
CODE_EXPLANATION_TEMPLATE,
CODE_DEBUGGING_TEMPLATE,
CODE_GENERATION_TEMPLATE,
ALGORITHM_EXPLANATION_TEMPLATE,
CODE_OPTIMIZATION_TEMPLATE,
GENERAL_QA_TEMPLATE,
)
class CodePromptManager:
"""代码Prompt管理类"""
def __init__(self):
"""
初始化代码Prompt管理器
"""
self._templates = self._load_templates()
self._context_cache: Dict[str, List[Dict[str, str]]] = {}
logger.info("代码Prompt管理器初始化完成")
def _load_templates(self) -> Dict[PromptTemplateType, str]:
"""
加载Prompt模板
Returns:
Dict[PromptTemplateType, str]: Prompt模板字典
"""
return {
PromptTemplateType.CODE_EXPLANATION: CODE_EXPLANATION_TEMPLATE,
PromptTemplateType.CODE_DEBUGGING: CODE_DEBUGGING_TEMPLATE,
PromptTemplateType.CODE_GENERATION: CODE_GENERATION_TEMPLATE,
PromptTemplateType.ALGORITHM_EXPLANATION: ALGORITHM_EXPLANATION_TEMPLATE,
PromptTemplateType.CODE_OPTIMIZATION: CODE_OPTIMIZATION_TEMPLATE,
PromptTemplateType.GENERAL_QA: GENERAL_QA_TEMPLATE
}
def _map_intent_to_prompt_type(self, intent_category: str) -> PromptTemplateType:
"""
将代码意图分类映射到Prompt类型
Args:
intent_category: 代码意图分类字符串
Returns:
PromptTemplateType: 对应的Prompt类型
"""
mapping = {
# 代码解释与逻辑类 -> CODE_EXPLANATION
"logic_explanation": PromptTemplateType.CODE_EXPLANATION,
"entity_introduction": PromptTemplateType.CODE_EXPLANATION,
"code_structure": PromptTemplateType.CODE_EXPLANATION,
# 代码生成与实现类 -> CODE_GENERATION
"code_generation": PromptTemplateType.CODE_GENERATION,
"boilerplate_implementation": PromptTemplateType.CODE_GENERATION,
# 调试、优化与理论类
"error_debugging": PromptTemplateType.CODE_DEBUGGING,
"code_optimization": PromptTemplateType.CODE_OPTIMIZATION,
"algorithm_theory": PromptTemplateType.ALGORITHM_EXPLANATION,
# 非代码问题 -> GENERAL_QA
"general_technical": PromptTemplateType.GENERAL_QA,
"non_technical": PromptTemplateType.GENERAL_QA,
"unknown": PromptTemplateType.GENERAL_QA
}
return mapping.get(intent_category, PromptTemplateType.GENERAL_QA)
def _build_conversation_history(self, history: Optional[Any]) -> str:
"""
构建对话历史字符串
Args:
history: 对话历史可以是字符串或字典列表
Returns:
str: 格式化的对话历史
"""
if not history:
return ""
# 如果是字符串,直接返回
if isinstance(history, str):
return history
# 如果是字典列表,格式化为字符串
if isinstance(history, list):
history_str = []
for item in history:
if isinstance(item, dict):
role = item.get('role', 'user')
content = item.get('content', '')
if role == 'user':
history_str.append(f"用户: {content}")
else:
history_str.append(f"助手: {content}")
return "\n".join(history_str)
# 其他类型,转换为字符串
return str(history)
def _extract_code_from_context(self, code_context: str) -> str:
"""
从上下文中提取代码
Args:
code_context: 代码上下文
Returns:
str: 提取的代码
"""
if not code_context:
return ""
# 尝试提取代码块
code_blocks = re.findall(r'```[\w]*\n[\s\S]*?```', code_context)
if code_blocks:
# 提取所有代码块并合并
extracted_code = []
for block in code_blocks:
# 提取语言标记
lang_match = re.match(r'```([\w]*)\n', block)
language = lang_match.group(1) if lang_match else ""
# 去除代码块标记
code = re.sub(r'```[\w]*\n|```', '', block)
code = code.strip()
if code:
if language:
extracted_code.append(f"语言: {language}\n{code}")
else:
extracted_code.append(code)
return "\n\n".join(extracted_code)
# 如果没有代码块标记,尝试提取看起来像代码的部分
# 查找连续的多行代码(以缩进或常见代码关键字开头)
lines = code_context.split('\n')
code_lines = []
in_code = False
for line in lines:
# 检查是否是代码行
line_stripped = line.strip()
if (line_stripped and
(line.startswith(' ') or line.startswith('\t') or # 缩进
line_stripped.startswith('def ') or line_stripped.startswith('class ') or # Python关键字
line_stripped.startswith('import ') or line_stripped.startswith('from ') or # 导入
line_stripped.startswith('if ') or line_stripped.startswith('for ') or # 控制流
line_stripped.startswith('while ') or line_stripped.startswith('try ') or
line_stripped.startswith('except ') or line_stripped.startswith('finally ') or
line_stripped.startswith('return ') or line_stripped.startswith('print(') or
line_stripped.startswith('// ') or line_stripped.startswith('# ') or # 注释
line_stripped.endswith(';') or # 分号结尾如Java、C++等)
line_stripped.startswith('{') or line_stripped.startswith('}') or # 大括号
re.match(r'^[\w_]+\s*=\s*', line_stripped) or # 变量赋值
re.match(r'^[\w_]+\s*\(.*\)\s*\{{?', line_stripped))): # 函数定义
code_lines.append(line)
in_code = True
elif in_code and line.strip() == '':
# 保留代码中的空行
code_lines.append(line)
elif in_code and len(code_lines) > 3:
# 如果已经收集了多行代码,并且遇到非代码行,停止收集
break
else:
# 非代码行,重置
code_lines = []
in_code = False
if len(code_lines) > 3:
return "\n".join(code_lines)
return code_context
def _format_code_block(self, code: str, language: str = "") -> str:
"""
格式化代码块提高显示质量
Args:
code: 代码内容
language: 代码语言
Returns:
str: 格式化的代码块
"""
if not code:
return ""
# 添加语言标记
lang_tag = language if language else ""
# 确保代码块格式正确
formatted_code = f"```{lang_tag}\n{code}\n```"
return formatted_code
def _build_enhanced_context(self, retrieved_results: List[Dict[str, Any]], intent_result: Optional[Dict[str, Any]]) -> str:
"""
根据意图和检索结果构建增强的上下文
Args:
retrieved_results: 检索结果列表每个元素包含 idtext metadata
intent_result: 意图识别结果
Returns:
str: 增强的上下文
"""
if not retrieved_results:
return "未找到相关参考信息"
context_parts = []
intent = intent_result.get('intent', '') if intent_result else ''
for result in retrieved_results:
metadata = result.get('metadata', {})
text = result.get('text', '')
result_id = result.get('id', 1)
# 提取所有 metadata 字段
func_id = metadata.get('func_id', '')
func_name = metadata.get('func_name', '')
class_name = metadata.get('class_name', 'None')
file_path = metadata.get('file_path', '')
lang = metadata.get('lang', '')
params = metadata.get('params', 0)
return_type = metadata.get('return_type', 'None')
docstring = metadata.get('docstring', '')
start_line = metadata.get('start_line', '')
end_line = metadata.get('end_line', '')
repo_id = metadata.get('repo_id', '')
branch = metadata.get('branch', '')
func_body = metadata.get('func_body', '')
# 根据意图构建不同的上下文
if intent == "code_understanding":
# 代码理解意图,强调语言、函数名、类名、参数、返回类型和函数体
context_part = f"【参考信息{result_id}】这是由{lang}实现的函数{func_name}"
if class_name and class_name != "None":
context_part += f",属于{class_name}"
context_part += f",它接收{params}个参数,返回类型为{return_type}"
if docstring:
context_part += f"。函数说明:{docstring}"
context_part += f"\n文件路径:{file_path},位置:{start_line}-{end_line}\n"
context_part += f"具体实现:\n{func_body}\n"
context_part += f"仓库:{repo_id},分支:{branch}\n"
context_part += f"原始文本:\n{text}"
elif intent == "code_modification":
# 代码修改意图,强调文件路径、位置和函数体
context_part = f"【参考信息{result_id}】需要修改的代码位于文件:{file_path},位置:{start_line}-{end_line}"
context_part += f"\n函数名:{func_name}"
if class_name and class_name != "None":
context_part += f"{class_name}类)"
context_part += f",由{lang}实现\n"
context_part += f"函数签名:接收{params}个参数,返回类型为{return_type}\n"
if docstring:
context_part += f"函数说明:{docstring}\n"
context_part += f"具体实现:\n{func_body}\n"
context_part += f"仓库:{repo_id},分支:{branch}\n"
context_part += f"原始文本:\n{text}"
elif intent == "functionality_question":
# 功能询问意图,强调函数名、文档、参数和返回类型
context_part = f"【参考信息{result_id}】函数{func_name}"
if class_name and class_name != "None":
context_part += f"{class_name}类)"
context_part += f"的功能说明:\n{docstring}\n"
context_part += f"{lang}实现,接收{params}个参数,返回类型为{return_type}\n"
context_part += f"文件路径:{file_path},位置:{start_line}-{end_line}\n"
context_part += f"具体实现:\n{func_body}\n"
context_part += f"仓库:{repo_id},分支:{branch}\n"
context_part += f"原始文本:\n{text}"
else:
# 其他意图,综合所有信息
context_part = f"【参考信息{result_id}】(来源:{file_path}"
context_part += f"\n函数:{func_name}"
if class_name and class_name != "None":
context_part += f"{class_name}类)"
context_part += f",语言:{lang}\n"
context_part += f"参数:{params}个,返回类型:{return_type}\n"
if docstring:
context_part += f"说明:{docstring}\n"
context_part += f"位置:{start_line}-{end_line}\n"
context_part += f"仓库:{repo_id},分支:{branch}\n"
context_part += f"实现:\n{func_body}\n"
context_part += f"原始文本:\n{text}"
context_parts.append(context_part)
return "\n\n".join(context_parts)
def generate_prompt(
self,
user_query: str,
intent_category: CodeIntentCategory,
code_context: Optional[str] = None,
conversation_history: Optional[List[Dict[str, str]]] = None,
error_message: Optional[str] = None,
target_language: Optional[str] = None,
user_requirement: Optional[str] = None
) -> str:
"""
生成代码专用Prompt
Args:
user_query: 用户问题
intent_category: 代码意图分类
code_context: 代码上下文
conversation_history: 对话历史
error_message: 错误信息仅Bug修复场景
target_language: 目标编程语言仅代码生成场景
user_requirement: 用户需求仅代码生成场景
Returns:
str: 生成的Prompt
"""
try:
logger.info(f"生成代码Prompt意图分类: {intent_category}")
# 映射意图到Prompt类型
prompt_type = self._map_intent_to_prompt_type(intent_category)
logger.info(f"选择Prompt类型: {prompt_type}")
# 获取对应模板
template = self._templates.get(prompt_type)
if not template:
logger.warning(f"未找到对应Prompt模板: {prompt_type}")
template = self._templates[PromptTemplateType.GENERAL_QA]
# 准备参数
params = {
"user_query": user_query,
"code_context": code_context or "",
"conversation_history": self._build_conversation_history(conversation_history),
"error_message": error_message or "",
"target_language": target_language or "根据上下文判断",
"user_requirement": user_requirement or user_query,
"algorithm_code": self._extract_code_from_context(code_context) if code_context else ""
}
# 填充模板
prompt = template
for key, value in params.items():
placeholder = f"{{{key}}}"
prompt = prompt.replace(placeholder, value)
logger.info(f"Prompt生成完成长度: {len(prompt)}字符")
return prompt
except Exception as e:
logger.error(f"生成Prompt失败: {e}")
# 返回通用模板
return self._templates[PromptTemplateType.GENERAL_QA].format(
user_query=user_query,
code_context=code_context or "",
conversation_history=self._build_conversation_history(conversation_history),
error_message="",
target_language="根据上下文判断",
user_requirement=user_query,
algorithm_code=""
)
def generate_dynamic_prompt(
self,
user_query: str,
intent_result: Optional[Dict[str, Any]] = None,
code_context: Optional[str] = None,
conversation_history: Optional[List[Dict[str, str]]] = None,
**kwargs
) -> str:
"""
生成动态Prompt基于意图识别结果
Args:
user_query: 用户问题
intent_result: 意图识别结果
code_context: 代码上下文
conversation_history: 对话历史
**kwargs: 其他参数
Returns:
str: 生成的动态Prompt
"""
try:
if intent_result:
# 从意图结果中提取分类
category_str = intent_result.get('category', 'unknown')
# 直接使用category_str因为CodeIntentCategory是一个普通的类不是枚举类型
intent_category = category_str
else:
# 默认使用通用分类
intent_category = CodeIntentCategory.UNKNOWN
# 提取其他参数
error_message = kwargs.get('error_message')
target_language = kwargs.get('target_language')
user_requirement = kwargs.get('user_requirement')
retrieved_results = kwargs.get('retrieved_results', [])
# 根据意图和 retrieved_results 构建增强的上下文
enhanced_context = code_context
if retrieved_results:
enhanced_context = self._build_enhanced_context(retrieved_results, intent_result)
# 生成Prompt
return self.generate_prompt(
user_query=user_query,
intent_category=intent_category,
code_context=enhanced_context,
conversation_history=conversation_history,
error_message=error_message,
target_language=target_language,
user_requirement=user_requirement
)
except Exception as e:
logger.error(f"生成动态Prompt失败: {e}")
# 返回通用Prompt
return self._templates[PromptTemplateType.GENERAL_QA].format(
user_query=user_query,
code_context=code_context or "",
conversation_history=self._build_conversation_history(conversation_history),
error_message="",
target_language="根据上下文判断",
user_requirement=user_query,
algorithm_code=""
)
def optimize_prompt(
self,
prompt: str,
max_length: int = 4000,
preserve_structure: bool = True
) -> str:
"""
优化Prompt长度
Args:
prompt: 原始Prompt
max_length: 最大长度
preserve_structure: 是否保留结构
Returns:
str: 优化后的Prompt
"""
if len(prompt) <= max_length:
return prompt
logger.warning(f"Prompt过长 ({len(prompt)} > {max_length}),需要优化")
if preserve_structure:
# 保留结构,只优化内容部分
# 1. 保留角色设定和核心指令
# 2. 精简分析要求
# 3. 缩短代码上下文
# 提取角色设定和核心指令
role_match = re.search(r'# 角色设定[\s\S]*?# 核心指令[\s\S]*?\n', prompt)
if role_match:
role_section = role_match.group(0)
else:
role_section = ""
# 提取分析要求
req_match = re.search(r'# 分析要求[\s\S]*?(?=# |$)', prompt)
if req_match:
req_section = req_match.group(0)
# 精简分析要求
req_lines = req_section.split('\n')
# 只保留前3条要求
req_section = '\n'.join(req_lines[:4]) # 保留标题和前3条
else:
req_section = ""
# 提取其他部分
rest_match = re.search(r'# (代码上下文|错误信息|对话历史|用户问题|输出格式)[\s\S]*$', prompt)
if rest_match:
rest_section = rest_match.group(0)
# 缩短代码上下文
if '# 代码上下文' in rest_section:
code_match = re.search(r'# 代码上下文[\s\S]*?(?=# |$)', rest_section)
if code_match:
code_section = code_match.group(0)
# 只保留前500个字符
if len(code_section) > 600:
code_lines = code_section.split('\n')
if len(code_lines) > 3:
# 保留标题和前几行
code_section = '\n'.join(code_lines[:2]) + '\n...\n(代码已截断)'
rest_section = rest_section.replace(code_match.group(0), code_section)
else:
rest_section = ""
optimized = role_section + '\n' + req_section + '\n' + rest_section
if len(optimized) > max_length:
# 进一步缩短
optimized = optimized[:max_length - 3] + '...'
else:
# 直接截断
optimized = prompt[:max_length - 3] + '...'
logger.info(f"Prompt优化完成长度: {len(optimized)}字符")
return optimized
def save_prompt_template(
self,
template_type: PromptTemplateType,
template_content: str,
description: Optional[str] = None
) -> bool:
"""
保存自定义Prompt模板
Args:
template_type: Prompt类型
template_content: 模板内容
description: 模板描述
Returns:
bool: 保存是否成功
"""
try:
# 这里可以扩展为持久化存储
# 目前只是在内存中更新
self._templates[template_type] = template_content
logger.info(f"保存Prompt模板成功: {template_type.value}")
return True
except Exception as e:
logger.error(f"保存Prompt模板失败: {e}")
return False
def get_prompt_template(self, template_type: PromptTemplateType) -> Optional[str]:
"""
获取Prompt模板
Args:
template_type: Prompt类型
Returns:
Optional[str]: 模板内容
"""
return self._templates.get(template_type)
def list_available_templates(self) -> List[Dict[str, Any]]:
"""
列出可用的Prompt模板
Returns:
List[Dict[str, Any]]: 模板列表
"""
templates = []
for template_type, content in self._templates.items():
templates.append({
"type": template_type.value,
"name": template_type.name,
"length": len(content),
"sample": content[:100] + "..." if len(content) > 100 else content
})
return templates
# 全局Prompt管理器实例
_prompt_manager = None
def get_prompt_manager() -> CodePromptManager:
"""
获取全局Prompt管理器实例
Returns:
CodePromptManager: Prompt管理器实例
"""
global _prompt_manager
if _prompt_manager is None:
_prompt_manager = CodePromptManager()
return _prompt_manager
def generate_code_prompt(
user_query: str,
intent_category: CodeIntentCategory,
code_context: Optional[str] = None,
conversation_history: Optional[List[Dict[str, str]]] = None,
**kwargs
) -> str:
"""
生成代码专用Prompt
Args:
user_query: 用户问题
intent_category: 代码意图分类
code_context: 代码上下文
conversation_history: 对话历史
**kwargs: 其他参数
Returns:
str: 生成的Prompt
"""
manager = get_prompt_manager()
return manager.generate_prompt(
user_query=user_query,
intent_category=intent_category,
code_context=code_context,
conversation_history=conversation_history,
**kwargs
)
def generate_dynamic_code_prompt(
user_query: str,
intent_result: Optional[Dict[str, Any]] = None,
code_context: Optional[str] = None,
conversation_history: Optional[List[Dict[str, str]]] = None,
**kwargs
) -> str:
"""
生成动态代码Prompt
Args:
user_query: 用户问题
intent_result: 意图识别结果
code_context: 代码上下文
conversation_history: 对话历史
**kwargs: 其他参数
Returns:
str: 生成的动态Prompt
"""
manager = get_prompt_manager()
return manager.generate_dynamic_prompt(
user_query=user_query,
intent_result=intent_result,
code_context=code_context,
conversation_history=conversation_history,
**kwargs
)
if __name__ == "__main__":
"""测试代码"""
import asyncio
from utils.code_intent import CodeIntentDetector
async def test_prompt_generation():
"""测试Prompt生成"""
print("=" * 80)
print("测试代码Prompt生成")
print("=" * 80)
# 初始化管理器
manager = CodePromptManager()
detector = CodeIntentDetector()
# 测试用例
test_cases = [
{
"query": "这个函数是做什么的?如何使用它?",
"code": "def calculate_factorial(n):\n if n <= 1:\n return 1\n return n * calculate_factorial(n-1)",
"category": CodeIntentCategory.ENTITY_INTRODUCTION
},
{
"query": "为什么会报语法错误?",
"code": "for i in range(10)\n print(i)",
"error": "SyntaxError: invalid syntax",
"category": CodeIntentCategory.ERROR_DEBUGGING
},
{
"query": "如何实现快速排序算法?",
"category": CodeIntentCategory.CODE_GENERATION
},
{
"query": "如何优化这段代码的性能?",
"code": "def slow_function():\n result = []\n for i in range(100000):\n result.append(i * 2)\n return result",
"category": CodeIntentCategory.CODE_OPTIMIZATION
}
]
for i, test_case in enumerate(test_cases):
print(f"\n测试用例 {i+1}: {test_case['query']}")
print("-" * 60)
# 生成Prompt
prompt = manager.generate_prompt(
user_query=test_case['query'],
intent_category=test_case['category'],
code_context=test_case.get('code'),
error_message=test_case.get('error')
)
# 打印结果
print(f"Prompt类型: {manager._map_intent_to_prompt_type(test_case['category']).value}")
print(f"Prompt长度: {len(prompt)}字符")
print("\nPrompt内容:")
print(prompt[:300] + "..." if len(prompt) > 300 else prompt)
print("-" * 60)
print("\n" + "=" * 80)
print("测试完成")
print("=" * 80)
asyncio.run(test_prompt_generation())

116
utils/func_id_generator.py Normal file
View File

@ -0,0 +1,116 @@
"""
函数全局唯一ID生成工具类
按用户/仓库/分支/文件/函数生成唯一ID
"""
import os
from typing import Optional, Dict
from config import settings
def generate_func_unique_id(
user_id: str,
repo_id: str,
branch: str,
file_path: str,
class_name: Optional[str],
func_name: str
) -> str:
"""
生成函数全局唯一ID
Args:
user_id: 用户ID
repo_id: 仓库ID
branch: 分支名
file_path: 文件路径
class_name: 类名
func_name: 函数名
Returns:
str: 函数唯一ID
"""
# 类名为None则使用空字符串
class_name = class_name if class_name else "None"
# 直接使用文件路径的绝对路径部分,确保唯一性
# 替换路径分隔符为下划线
file_path = file_path.split(os.sep)[3:]
file_path = "_".join(file_path)
normalized_file_path = file_path.replace(os.sep, "_")
# 生成唯一ID
unique_id = f"{user_id}_{repo_id}_{branch}_{normalized_file_path}_{class_name}_{func_name}"
# 替换特殊字符避免ChromaDB主键冲突
unique_id = unique_id.replace("/", "_").replace("\\", "_").replace(":", "_").replace(" ", "_")
return unique_id
def parse_func_unique_id(func_id: str) -> Dict[str, str]:
"""
解析函数唯一ID
Args:
func_id: 函数唯一ID
Returns:
Dict[str, str]: 解析后的信息
"""
parts = func_id.split("_")
if len(parts) < 6:
raise Exception(f"无效的函数ID格式: {func_id}")
# 解析各部分
user_id = parts[0]
repo_id = parts[1]
branch = parts[2]
# 解析文件路径(可能包含下划线)
# 从第3个部分开始到倒数第2个部分结束
file_path_parts = parts[3:-2]
file_path = "_".join(file_path_parts).replace("_", os.sep)
class_name = parts[-2]
if class_name == "None":
class_name = None
func_name = parts[-1]
return {
"user_id": user_id,
"repo_id": repo_id,
"branch": branch,
"file_path": file_path,
"class_name": class_name,
"func_name": func_name
}
def get_repo_path_from_func_id(func_id: str) -> str:
"""
从函数ID获取仓库路径
Args:
func_id: 函数唯一ID
Returns:
str: 仓库路径
"""
info = parse_func_unique_id(func_id)
return os.path.join(settings.GIT_LOCAL_STORAGE_ROOT, info["user_id"], info["repo_id"])
def get_file_path_from_func_id(func_id: str) -> str:
"""
从函数ID获取文件路径
Args:
func_id: 函数唯一ID
Returns:
str: 文件路径
"""
info = parse_func_unique_id(func_id)
repo_path = get_repo_path_from_func_id(func_id)
return os.path.join(repo_path, info["file_path"])

324
utils/git_tool.py Normal file
View File

@ -0,0 +1,324 @@
"""
Git命令封装工具类
实现Git仓库的克隆更新检测增量拉取等功能
"""
import subprocess
import os
from typing import Tuple, Dict, List
from loguru import logger
from config import settings
class GitTool:
def __init__(self, user_id: str = "test", repo_id: str = "test", git_config: dict = None,
git_url: str = None, branch: str = None, protocol: str = None,
https_token: str = None, ssh_key: str = None, local_repo_path: str = None):
"""
初始化Git工具类
Args:
user_id: 用户ID
repo_id: 仓库ID
git_config: Git配置信息
git_url: Git仓库URL直接参数优先级高于git_config
branch: Git分支直接参数优先级高于git_config
protocol: Git协议直接参数优先级高于git_config
https_token: HTTPS令牌直接参数优先级高于git_config
ssh_key: SSH密钥直接参数优先级高于git_config
local_repo_path: 本地仓库路径直接参数优先级高于git_config
"""
self.user_id = user_id
self.repo_id = repo_id
# 优先使用直接参数如果没有则使用git_config
if git_config:
self.git_url = git_url or git_config.get("git_url")
self.branch = branch or git_config.get("branch", settings.GIT_DEFAULT_BRANCH)
self.protocol = protocol or git_config.get("protocol", "https")
self.ssh_key = ssh_key or git_config.get("ssh_key")
self.https_token = https_token or git_config.get("https_token")
# 本地结构化存储路径
self.local_repo_path = local_repo_path or git_config.get("local_repo_path")
else:
self.git_url = git_url
self.branch = branch or settings.GIT_DEFAULT_BRANCH
self.protocol = protocol or "https"
self.ssh_key = ssh_key
self.https_token = https_token
self.local_repo_path = local_repo_path
if not self.local_repo_path:
self.local_repo_path = os.path.join(settings.GIT_LOCAL_STORAGE_ROOT, user_id, repo_id)
# 初始化Git环境
self._init_git_env()
def _init_git_env(self):
"""
初始化Git环境SSH密钥配置
"""
if self.ssh_key:
# 解密SSH私钥写入临时文件配置Git SSH
ssh_key_path = f"/tmp/ssh_key_{self.user_id}_{self.repo_id}"
with open(ssh_key_path, "w") as f:
f.write(self.ssh_key)
if not self.ssh_key.endswith('\n'):
f.write('\n')
os.chmod(ssh_key_path, 0o600)
os.environ["GIT_SSH_COMMAND"] = f"ssh -i {ssh_key_path} -o StrictHostKeyChecking=no"
def clone_repo(self) -> bool:
"""
克隆Git仓库
Returns:
bool: 是否成功克隆
"""
if not os.path.exists(self.local_repo_path):
os.makedirs(os.path.dirname(self.local_repo_path), exist_ok=True)
# 执行git clone命令
cmd = [
"git", "clone", "--single-branch",
"--branch", self.branch, self.git_url, self.local_repo_path
]
logger.info(f"执行Git克隆命令: {' '.join(cmd)}")
res = subprocess.run(cmd, capture_output=True, text=True)
if res.returncode != 0:
logger.error(f"Git克隆失败: {res.stderr}")
raise Exception(f"Git克隆失败: {res.stderr}")
# 克隆后校验
self._check_repo_integrity()
logger.info(f"Git仓库克隆成功: {self.local_repo_path}")
return True
logger.info(f"Git仓库已存在: {self.local_repo_path}")
return False
def _check_repo_integrity(self):
"""
仓库完整性校验git fsck+ 支持的编程语言检测
"""
# 执行git fsck
try:
subprocess.run(["git", "fsck"], cwd=self.local_repo_path, check=True, capture_output=True, text=True)
logger.info(f"Git仓库完整性校验成功: {self.local_repo_path}")
except subprocess.CalledProcessError as e:
logger.warning(f"Git仓库完整性校验失败: {e.stderr}")
# 扫描文件类型,记录支持的编程语言
support_lang = self._detect_support_lang()
logger.info(f"检测到支持的编程语言: {support_lang}")
return support_lang
def _detect_support_lang(self) -> List[str]:
"""
检测仓库支持的编程语言
Returns:
List[str]: 支持的编程语言列表
"""
lang_extensions = {
"python": [".py"],
"java": [".java"],
"go": [".go"],
"javascript": [".js", ".jsx"],
"typescript": [".ts", ".tsx"],
"c": [".c", ".h"],
"cpp": [".cpp", ".hpp", ".cc"],
"csharp": [".cs"],
"rust": [".rs"],
"php": [".php"],
"ruby": [".rb"],
"swift": [".swift"],
"kotlin": [".kt"],
"scala": [".scala"]
}
support_lang = []
for root, dirs, files in os.walk(self.local_repo_path):
# 跳过.git目录
if ".git" in dirs:
dirs.remove(".git")
# 跳过其他常见的非代码目录
dirs_to_skip = ["node_modules", "venv", "dist", "build", "__pycache__"]
dirs[:] = [d for d in dirs if d not in dirs_to_skip]
for file in files:
for lang, extensions in lang_extensions.items():
if any(file.endswith(ext) for ext in extensions):
if lang not in support_lang:
support_lang.append(lang)
break
return support_lang
def detect_remote_update(self) -> Tuple[bool, str, str]:
"""
远程更新检测
Returns:
Tuple[bool, str, str]: (是否有更新, 本地commit ID, 远程commit ID)
"""
# 确保仓库存在
if not os.path.exists(self.local_repo_path):
raise Exception(f"Git仓库不存在: {self.local_repo_path}")
# 拉取远程commit记录
try:
subprocess.run(["git", "fetch", "origin", f"{self.branch}:{self.branch}"],
cwd=self.local_repo_path, check=True, capture_output=True, text=True)
except subprocess.CalledProcessError as e:
logger.error(f"Git fetch失败: {e.stderr}")
raise
# 获取本地/远程commit ID
local_commit = subprocess.run(["git", "rev-parse", "HEAD"],
cwd=self.local_repo_path, capture_output=True, text=True).stdout.strip()
remote_commit = subprocess.run(["git", "rev-parse", f"origin/{self.branch}"],
cwd=self.local_repo_path, capture_output=True, text=True).stdout.strip()
has_update = local_commit != remote_commit
logger.info(f"Git更新检测: 本地={local_commit[:7]}, 远程={remote_commit[:7]}, 有更新={has_update}")
return has_update, local_commit, remote_commit
def incremental_pull(self, local_commit: str, remote_commit: str) -> Dict[str, List[str]]:
"""
增量拉取代码+解析文件变更
Args:
local_commit: 本地commit ID
remote_commit: 远程commit ID
Returns:
Dict[str, List[str]]: 文件变更集
"""
# 快进合并到远程最新版本
try:
subprocess.run(["git", "merge", "--ff-only", f"origin/{self.branch}"],
cwd=self.local_repo_path, check=True, capture_output=True, text=True)
logger.info(f"Git快进合并成功: {self.branch}")
except subprocess.CalledProcessError as e:
logger.error(f"Git合并失败: {e.stderr}")
raise
# 提取增量commit的文件变更
delta_commits = subprocess.run(["git", "log", "--pretty=format:%H", f"{local_commit}..{remote_commit}"],
cwd=self.local_repo_path, capture_output=True, text=True).stdout.split()
# 解析文件变更为ADD/MODIFY/DELETE
delta_files = self._parse_delta_files(delta_commits)
logger.info(f"Git增量变更: ADD={len(delta_files['ADD'])}, MODIFY={len(delta_files['MODIFY'])}, DELETE={len(delta_files['DELETE'])}")
return delta_files
def _parse_delta_files(self, delta_commits: List[str]) -> Dict[str, List[str]]:
"""
解析文件变更集
Args:
delta_commits: 增量commit列表
Returns:
Dict[str, List[str]]: 文件变更集
"""
add_files, modify_files, delete_files = [], [], []
for commit in delta_commits:
# git show --name-status 获取文件变更
res = subprocess.run(["git", "show", "--name-status", commit],
cwd=self.local_repo_path, capture_output=True, text=True).stdout
for line in res.splitlines():
if not line:
continue
# 解析状态和文件路径
if "\t" in line:
status, file_path = line.split("\t", 1)
full_path = os.path.join(self.local_repo_path, file_path)
if status == "A":
add_files.append(full_path)
elif status == "M":
modify_files.append(full_path)
elif status == "D":
delete_files.append(full_path)
# 去重并返回
return {
"ADD": list(set(add_files)),
"MODIFY": list(set(modify_files)),
"DELETE": list(set(delete_files))
}
def get_current_commit(self) -> str:
"""
获取当前commit ID
Returns:
str: 当前commit ID
"""
if not os.path.exists(self.local_repo_path):
raise Exception(f"Git仓库不存在: {self.local_repo_path}")
commit_id = subprocess.run(["git", "rev-parse", "HEAD"],
cwd=self.local_repo_path, capture_output=True, text=True).stdout.strip()
return commit_id
def get_repo_info(self) -> Dict[str, str]:
"""
获取仓库信息
Returns:
Dict[str, str]: 仓库信息
"""
if not os.path.exists(self.local_repo_path):
raise Exception(f"Git仓库不存在: {self.local_repo_path}")
# 获取仓库URL
remote_url = subprocess.run(["git", "config", "--get", "remote.origin.url"],
cwd=self.local_repo_path, capture_output=True, text=True).stdout.strip()
# 获取当前分支
current_branch = subprocess.run(["git", "rev-parse", "--abbrev-ref", "HEAD"],
cwd=self.local_repo_path, capture_output=True, text=True).stdout.strip()
# 获取当前commit
current_commit = self.get_current_commit()
return {
"remote_url": remote_url,
"current_branch": current_branch,
"current_commit": current_commit,
"local_path": self.local_repo_path
}
def test_connection(self) -> bool:
"""
测试Git连接
Returns:
bool: 是否连接成功
"""
if not self.git_url:
raise Exception("Git仓库URL未设置")
logger.info(f"测试Git连接: {self.git_url}")
# 尝试执行git ls-remote命令来测试连接
try:
cmd = ["git", "ls-remote", "--heads", self.git_url, f"refs/heads/{self.branch}"]
logger.info(f"执行Git连接测试命令: {' '.join(cmd)}")
res = subprocess.run(cmd, capture_output=True, text=True)
if res.returncode == 0:
# 检查输出是否包含预期的分支信息
if self.branch in res.stdout:
logger.info("Git连接测试成功")
return True
else:
logger.warning(f"Git连接测试失败分支 {self.branch} 不存在")
return False
else:
logger.error(f"Git连接测试失败: {res.stderr}")
return False
except Exception as e:
logger.error(f"Git连接测试异常: {e}")
return False

23
utils/prompt/__init__.py Normal file
View File

@ -0,0 +1,23 @@
"""Prompt模板包"""
from .intent_detector import INTENT_DETECTION_TEMPLATE
from .answer_generator import (
CODE_EXPLANATION_TEMPLATE,
CODE_DEBUGGING_TEMPLATE,
CODE_GENERATION_TEMPLATE,
ALGORITHM_EXPLANATION_TEMPLATE,
CODE_OPTIMIZATION_TEMPLATE,
GENERAL_QA_TEMPLATE
)
from .integrated_query_processing import INTEGRATED_QUERY_PROCESSING_TEMPLATE
__all__ = [
"INTENT_DETECTION_TEMPLATE",
"CODE_EXPLANATION_TEMPLATE",
"CODE_DEBUGGING_TEMPLATE",
"CODE_GENERATION_TEMPLATE",
"ALGORITHM_EXPLANATION_TEMPLATE",
"CODE_OPTIMIZATION_TEMPLATE",
"GENERAL_QA_TEMPLATE",
"INTEGRATED_QUERY_PROCESSING_TEMPLATE"
]

View File

@ -0,0 +1,17 @@
"""答案生成相关的Prompt模板"""
from .code_explanation import CODE_EXPLANATION_TEMPLATE
from .code_generation import CODE_GENERATION_TEMPLATE
from .algorithm_explanation import ALGORITHM_EXPLANATION_TEMPLATE
from .code_debugging import CODE_DEBUGGING_TEMPLATE
from .code_optimization import CODE_OPTIMIZATION_TEMPLATE
from .general_qa import GENERAL_QA_TEMPLATE
__all__ = [
"CODE_EXPLANATION_TEMPLATE",
"CODE_DEBUGGING_TEMPLATE",
"CODE_GENERATION_TEMPLATE",
"ALGORITHM_EXPLANATION_TEMPLATE",
"CODE_OPTIMIZATION_TEMPLATE",
"GENERAL_QA_TEMPLATE"
]

View File

@ -0,0 +1,45 @@
"""算法解释相关的Prompt模板"""
ALGORITHM_EXPLANATION_TEMPLATE = """# 角色设定
你是一位算法专家擅长深入解析算法原理和实现能够根据不同类型的问题提供精准的专业解答
## 核心指令
请基于提供的算法代码上下文和问题类型为用户提供详细专业的算法分析
## 分析要求(根据问题类型调整重点)
### 算法解释algorithm_explanation
- **算法原理**详细解释算法的基本原理和设计思想
- **算法步骤**逐步说明算法的执行流程
- **时间复杂度**分析时间复杂度最好平均最坏情况
- **空间复杂度**分析空间复杂度说明内存使用情况
- **优缺点**分析算法的优势和局限性
- **适用场景**说明算法的适用场景和典型应用
- **比较分析**与其他同类算法进行对比
### 算法实现algorithm_implementation
- **算法选择**选择最适合该问题的算法
- **实现细节**提供完整的算法实现代码
- **复杂度分析**分析实现的时间和空间复杂度
- **边界处理**考虑边界情况和特殊输入
- **优化建议**提供可能的优化方向
- **测试用例**提供测试算法的示例用例
### 数据结构data_structure
- **结构定义**详细说明数据结构的定义和特点
- **操作方法**说明数据结构支持的操作及其复杂度
- **实现方式**提供数据结构的实现代码
- **适用场景**说明数据结构的适用场景和典型应用
- **性能对比**与其他数据结构进行性能对比
- **使用示例**提供数据结构的使用示例
## 算法代码
{algorithm_code}
## 对话历史
{conversation_history}
## 用户问题
{user_query}
请开始你的分析"""

View File

@ -0,0 +1,46 @@
"""代码调试相关的Prompt模板"""
CODE_DEBUGGING_TEMPLATE = """# 角色设定
你是一位专业的代码调试专家擅长快速定位和解决代码中的错误能够根据不同类型的错误提供精准的诊断和解决方案
## 核心指令
请基于提供的代码错误信息上下文和错误类型为用户提供详细的错误分析和解决方案
## 分析要求(根据错误类型调整重点)
### 语法错误syntax_error
- **错误定位**准确指出语法错误的具体位置行号列号
- **错误原因**解释违反了哪条语法规则
- **修正方案**提供修正后的完整代码
- **预防建议**说明如何避免类似的语法错误
- **常见模式**列出该语法错误的常见触发场景
### 运行时错误runtime_error
- **错误分析**详细分析异常类型和错误信息
- **堆栈跟踪**解释错误堆栈中的关键信息
- **根本原因**深入分析导致错误的根本原因
- **修复方案**提供具体的修复代码和实施步骤
- **异常处理**建议如何添加异常处理来预防此类错误
- **测试建议**说明如何测试修复是否有效
### 调试问题debugging
- **调试方法**提供适合该问题的调试策略
- **断点设置**建议在哪些位置设置断点
- **日志分析**说明如何通过日志分析问题
- **变量检查**建议检查哪些关键变量的值
- **逐步排查**提供逐步排查问题的流程
- **工具推荐**推荐适合的调试工具和技巧
## 代码上下文
{code_context}
## 错误信息
{error_message}
## 对话历史
{conversation_history}
## 用户问题
{user_query}
请开始你的分析"""

View File

@ -0,0 +1,41 @@
"""代码解释相关的Prompt模板"""
CODE_EXPLANATION_TEMPLATE = """# 角色设定
你是一位资深的代码分析专家擅长深入解析代码结构和功能能够根据不同类型的问题提供精准的专业解答
## 核心指令
请基于提供的代码上下文和用户问题类型为用户提供详细专业的代码分析
## 分析要求(根据问题类型调整重点)
### 逻辑解释问题logic_explanation
- **功能说明**详细解释代码的功能和用途
- **逻辑分析**说明代码的执行流程和核心逻辑
- **实现原理**解释代码的实现原理和技术细节
- **使用示例**提供实际可运行的代码示例
- **注意事项**指出使用时需要注意的要点和常见错误
### 实体介绍问题entity_introduction
- **实体结构**说明函数API等实体的结构和组成
- **参数分析**说明每个参数的类型含义和默认值
- **返回值说明**解释返回值的类型含义和可能的取值
- **使用示例**提供实际可运行的代码示例
- **注意事项**指出使用时需要注意的要点和常见错误
### 代码结构问题code_structure
- **项目结构**详细说明项目的目录和文件组织
- **模块划分**说明各个模块的功能和职责
- **依赖关系**解释模块之间的依赖和调用关系
- **架构设计**说明项目的整体架构和设计思路
- **文件说明**解释关键文件的作用和内容
## 代码上下文
{code_context}
## 对话历史
{conversation_history}
## 用户问题
{user_query}
请开始你的分析"""

View File

@ -0,0 +1,58 @@
"""代码生成相关的Prompt模板"""
CODE_GENERATION_TEMPLATE = '''# 角色设定
你是一位经验丰富的代码生成专家擅长根据不同类型的需求编写高质量可维护的代码
## 核心指令
请基于用户的需求上下文和生成类型为用户提供完整专业可直接使用的代码
## 代码生成要求(根据生成类型调整重点)
### 完整代码生成code_generation
- **需求分析**深入理解用户的功能需求
- **架构设计**设计合理的代码结构和模块划分
- **完整实现**提供完整可运行的代码包括所有必要的导入
- **最佳实践**遵循目标语言的编码规范和最佳实践
- **错误处理**添加适当的异常处理和边界检查
- **代码注释**添加清晰的注释解释关键逻辑
- **使用示例**提供如何使用该代码的示例
### 函数实现function_implementation
- **函数签名**设计清晰的函数名参数和返回值
- **参数验证**添加参数类型检查和验证逻辑
- **边界处理**考虑边界情况和特殊输入
- **错误处理**使用适当的异常处理机制
- **文档字符串**添加详细的docstring说明函数用途
- **类型提示**使用类型注解提高代码可读性
- **单元测试**提供简单的测试用例
### 类实现class_implementation
- **类设计**设计合理的类结构和方法划分
- **构造函数**实现__init__方法正确初始化属性
- **封装性**合理使用私有属性和公共方法
- **方法实现**实现所有必要的方法确保功能完整
- **特殊方法**根据需要实现__str____repr__等特殊方法
- **文档字符串**为类和主要方法添加docstring
- **使用示例**提供类的使用示例
## 通用代码质量要求
1. **语法正确性**确保代码语法完全正确可直接运行
2. **代码风格**遵循PEP 8Python或其他语言的编码规范
3. **可读性**使用有意义的变量名和函数名添加必要的注释
4. **可维护性**代码结构清晰易于理解和修改
5. **性能考虑**在保证正确性的前提下考虑性能优化
6. **安全性**注意常见的安全问题如SQL注入XSS等
## 目标语言
{target_language}
## 代码上下文
{code_context}
## 对话历史
{conversation_history}
## 用户需求
{user_requirement}
请开始生成代码'''

View File

@ -0,0 +1,76 @@
"""代码优化相关的Prompt模板"""
CODE_OPTIMIZATION_TEMPLATE = """你是一个专业的代码优化专家。请根据用户的问题和检索到的代码上下文,提供专业的优化建议。
## 意图上下文
当前意图{{intent_type}}
检索策略{{retrieval_strategy}}
## 对话历史
{{conversation_history}}
## 检索上下文
{{code_context}}
## 响应指南
1. **必须使用标准 Markdown 格式**
2. **确保流式输出体验**首句直接入题段落间使用 \n\n 分隔
3. **中英文之间自动添加空格**
4. **代码块必须闭合且标注语言**
5. **针对 {{intent_type}} 采用对应的回答模板**
## 回答模板
### 代码优化问题
当用户询问如何优化代码时请按以下结构回答
#### 1. 代码分析
- 指出当前代码存在的问题
- 分析性能瓶颈
- 说明可优化的地方
#### 2. 优化建议
- 提供具体的优化方案
- 说明优化原理
- 给出优化后的代码示例
#### 3. 性能对比
- 对比优化前后的性能
- 说明优化的效果
- 给出具体的性能指标
#### 4. 最佳实践
- 提供相关的编程最佳实践
- 说明代码规范
- 给出可维护性建议
### 性能调优问题
当用户询问性能调优时请按以下结构回答
#### 1. 性能分析
- 分析当前性能问题
- 定位性能瓶颈
- 说明影响性能的因素
#### 2. 调优策略
- 提供具体的调优方案
- 说明调优的原理
- 给出调优的步骤
#### 3. 优化效果
- 说明调优后的效果
- 给出性能提升的数据
- 对比调优前后的差异
#### 4. 注意事项
- 说明调优时的注意事项
- 提供避免问题的建议
- 给出监控和评估方法
## 代码示例要求
- 代码块必须使用 ```python ```javascript 等标注语言
- 代码必须完整可运行
- 添加必要的注释说明
- 保持代码风格一致
现在请根据用户问题和检索上下文提供专业的代码优化建议"""

View File

@ -0,0 +1,99 @@
"""并发编程相关的Prompt模板"""
CONCURRENCY_EXPLANATION_TEMPLATE = '''你是一个专业的并发编程专家。请根据用户的问题和检索到的代码上下文,提供专业的并发编程指导。
## 意图上下文
当前意图{{intent_type}}
检索策略{{retrieval_strategy}}
## 对话历史
{{conversation_history}}
## 检索上下文
{{code_context}}
## 响应指南
1. **必须使用标准 Markdown 格式**
2. **确保流式输出体验**首句直接入题段落间使用 \n\n 分隔
3. **中英文之间自动添加空格**
4. **代码块必须闭合且标注语言**
5. **针对 {{intent_type}} 采用对应的回答模板**
## 回答模板
### 并发问题
当用户询问并发问题时请按以下结构回答
#### 1. 并发模型
- 说明并发的基本概念
- 解释并发模型
- 对比不同并发模型
#### 2. 实现方式
- 提供具体的实现代码
- 说明实现原理
- 给出使用示例
#### 3. 同步机制
- 说明同步的必要性
- 提供同步方法
- 解释同步原理
#### 4. 竞争条件处理
- 说明竞争条件的概念
- 提供避免竞争条件的方法
- 给出并发设计模式
### 线程问题
当用户询问线程问题时请按以下结构回答
#### 1. 线程基础
- 说明线程的概念
- 解释线程的生命周期
- 说明线程的创建和管理
#### 2. 线程同步
- 说明线程同步的必要性
- 提供同步方法信号量等
- 给出同步示例代码
#### 3. 死锁预防
- 说明死锁的概念
- 提供死锁预防方法
- 给出避免死锁的最佳实践
#### 4. 线程安全
- 说明线程安全的概念
- 提供线程安全的实现方法
- 给出线程安全的编程建议
### 异步编程问题
当用户询问异步编程时请按以下结构回答
#### 1. 异步编程模型
- 说明异步编程的概念
- 解释异步与同步的区别
- 说明异步的优势
#### 2. 异步实现
- 提供异步编程的代码示例
- 说明 async/await 的使用
- 解释事件循环的原理
#### 3. 异步 I/O
- 说明异步 I/O 的概念
- 提供异步 I/O 的实现方法
- 给出异步 I/O 的使用示例
#### 4. 错误处理
- 说明异步编程中的错误处理
- 提供异常处理的方法
- 给出调试和测试建议
## 代码示例要求
- 代码块必须使用 ```python ```javascript 等标注语言
- 代码必须完整可运行
- 添加必要的注释说明
- 展示并发/异步的完整流程
现在请根据用户问题和检索上下文提供专业的并发编程指导'''

View File

@ -0,0 +1,99 @@
"""测试部署相关的Prompt模板"""
DEPLOYMENT_EXPLANATION_TEMPLATE = """你是一个专业的测试和部署专家。请根据用户的问题和检索到的代码上下文,提供专业的测试和部署指导。
## 意图上下文
当前意图{{intent_type}}
检索策略{{retrieval_strategy}}
## 对话历史
{{conversation_history}}
## 检索上下文
{{code_context}}
## 响应指南
1. **必须使用标准 Markdown 格式**
2. **确保流式输出体验**首句直接入题段落间使用 \n\n 分隔
3. **中英文之间自动添加空格**
4. **代码块必须闭合且标注语言**
5. **针对 {{intent_type}} 采用对应的回答模板**
## 回答模板
### 测试问题
当用户询问测试问题时请按以下结构回答
#### 1. 测试类型
- 说明不同类型的测试单元测试集成测试端到端测试
- 解释各种测试的适用场景
- 提供测试策略建议
#### 2. 测试框架
- 推荐适合的测试框架
- 说明框架的特点和优势
- 提供框架的使用示例
#### 3. 测试实现
- 提供具体的测试代码
- 说明测试的编写方法
- 给出测试的最佳实践
#### 4. Mock 和测试数据
- 说明 Mock 的使用场景
- 提供 Mock 的实现方法
- 给出测试数据的准备策略
### 部署问题
当用户询问部署问题时请按以下结构回答
#### 1. 部署策略
- 说明不同的部署方式手动部署自动化部署
- 解释部署的流程
- 提供部署策略建议
#### 2. 环境配置
- 说明开发测试生产环境的配置
- 提供环境变量的管理方法
- 给出配置文件的组织方式
#### 3. CI/CD 流程
- 说明 CI/CD 的概念
- 提供主流 CI/CD 工具的使用方法
- 给出 CI/CD 流程的配置示例
#### 4. 监控和告警
- 说明监控的重要性
- 提供监控工具的推荐
- 给出告警策略的配置方法
### 配置问题
当用户询问配置问题时请按以下结构回答
#### 1. 配置文件格式
- 说明不同配置文件的格式JSONYAMLINI
- 解释各种格式的优缺点
- 提供格式选择的建议
#### 2. 环境变量管理
- 说明环境变量的使用场景
- 提供环境变量的管理方法
- 给出环境变量的最佳实践
#### 3. 参数配置
- 说明参数配置的原则
- 提供参数验证的方法
- 给出参数管理的建议
#### 4. 配置验证
- 说明配置验证的重要性
- 提供配置验证的方法
- 给出配置错误的处理建议
## 代码示例要求
- 代码块必须使用 ```bash```yaml```python 等标注语言
- 配置文件必须完整且格式正确
- 添加必要的注释说明
- 提供可执行的命令或脚本
现在请根据用户问题和检索上下文提供专业的测试和部署指导"""

View File

@ -0,0 +1,58 @@
"""通用代码相关的Prompt模板"""
GENERAL_QA_TEMPLATE = """# 角色设定
你是一位专业的代码顾问能够回答各种代码相关问题擅长根据不同类型的问题提供精准的专业解答
## 核心指令
请基于提供的代码上下文和问题类型为用户提供全面准确专业的回答
## 回答要求(根据问题类型调整重点)
### 测试问题testing
- **测试方法**说明适合的测试方法单元测试集成测试等
- **测试框架**推荐适合的测试框架如pytestunittest等
- **测试用例**提供具体的测试用例示例
- **Mock技术**说明如何mock外部依赖
- **覆盖率**解释测试覆盖率的概念和如何提高覆盖率
- **最佳实践**提供测试的最佳实践和常见陷阱
### 部署问题deployment
- **部署策略**说明适合的部署方式容器化云部署等
- **环境配置**详细说明环境变量的配置方法
- **CI/CD流程**解释持续集成和持续部署的流程
- **依赖管理**说明如何管理生产环境的依赖
- **监控告警**建议部署后的监控和告警方案
- **回滚策略**说明如何处理部署失败的情况
### 配置问题configuration
- **配置文件**说明配置文件的格式和位置
- **环境变量**解释如何设置和使用环境变量
- **参数配置**详细说明各个配置参数的含义和取值
- **配置验证**提供验证配置是否正确的方法
- **常见问题**列出配置相关的常见错误和解决方案
- **最佳实践**提供配置管理的最佳实践
### 通用知识问题general_knowledge
- **概念解释**详细解释相关概念和术语
- **原理说明**深入说明技术原理和机制
- **应用场景**说明技术的适用场景和典型应用
- **发展趋势**介绍技术的发展趋势和未来方向
- **学习资源**推荐相关的学习资源和文档
- **实践建议**提供实际应用的建议和注意事项
### 非技术问题non_technical
- **友好回应**保持友好自然的对话风格
- **相关信息**提供与问题相关的有用信息
- **引导澄清**如果问题模糊引导用户明确需求
- **保持自然**避免过度技术化保持对话的自然流畅
## 代码上下文
{code_context}
## 对话历史
{conversation_history}
## 用户问题
{user_query}
请开始你的回答"""

View File

@ -0,0 +1,148 @@
"""集成查询处理Prompt模板
整合意图识别Metadata过滤条件提取和查询转换功能
"""
INTEGRATED_QUERY_PROCESSING_TEMPLATE = """### 角色定义
你是一个全面的查询处理助手需要完成以下三个任务
1. 代码意图识别分析用户问题的意图类型
2. 元数据过滤条件提取从问题中提取显示限制的过滤条件
3. 查询转换重写查询生成更广泛的查询分解复杂查询
### 对话历史
{history_str}
### 当前用户问题
{query}
---
### 任务1代码意图识别
请分析用户问题的意图判断其属于以下分类之一
- logic_explanation解释既有代码的底层逻辑
- entity_introduction介绍具体的代码实体
- code_structure询问项目组织
- code_generation请求从零编写完整代码或功能块
- boilerplate_implementation请求提供标准算法/模板
- error_debugging排查 Bug 或异常
- code_optimization改进既有代码的性能或质量
- algorithm_theory算法原理或复杂度分析
- general_technical通用技术咨询
- non_technical非技术问题
- unknown未知类型
#### 分类决策树 (判定逻辑)
在判定分类前请严格执行以下优先级逻辑
1. **上下文回溯**如果 query 中提到的实体函数变量类名在对话历史或上下文代码中出现过优先判定为代码解释/架构类
2. **句式辨析**
- **[实体/功能] 是怎么实现的/怎么做的** -> 倾向于代码解释语态为"对既有状态的追溯"
- **怎么实现 [功能]/ 帮我写一个...** -> 倾向于代码生成语态为"对未知实现的请求"
3. **理论深度**若问题涉及性能瓶颈数学原理或复杂度优先归类为算法与优化类
#### 语义微调示例 (Few-Shot)
- **输入**: "find_median 是如何实现的?"
**判定**: logic_explanation | **原因**: 指向特定函数名且询问其现状
- **输入**: "如何实现查找中位数的算法?"
**判定**: code_generation | **原因**: 泛指功能实现表现为编程请求
- **输入**: "这段代码能跑快一点吗?"
**判定**: code_optimization | **原因**: 基于现有代码的性能改进请求
- **输入**: "什么是深度优先搜索?"
**判定**: algorithm_theory | **原因**: 概念性理论询问
---
### 任务2严格元数据过滤条件提取
请从用户问题中提取显示限制的metadata过滤条件只提取与以下key一致的条件
- func_id: 函数ID
- func_name: 函数名
- class_name: 类名
- file_path: 文件路径
- lang: 编程语言
- params: 参数数量
- return_type: 返回类型
- docstring: 文档字符串
- start_line: 开始行号
- end_line: 结束行号
- repo_id: 仓库ID
- branch: 分支名
- func_body: 函数体
**重要规则**
- 只提取查询中**明确提到**的条件不要进行任何推测
- 只有当查询中明确使用了与某个key相关的词汇时才提取该key的value
- **value必须为小写**
- **一个key只对应一个value**
- **value的字符串长度尽可能短**
- 例如对于查询"在 algorithms 目录下用java实现的排序算法"
只提取 {{"lang": "java"}}不要提取其他任何key
**强制性约束**
1. **零推测原则**仅提取用户明确指定的属性限定若用户说计算斐波那契的函数由于未指定函数名文件名或语言提取结果应为空 `{{}}`
2. **关键词触发**
- 提取 `file_path`原文必须包含路径特征 .py, /path, 文件夹等
- 提取 `func_name` / `class_name`原文必须包含名为叫作或明显的标识符引用
- 提取 `return_type` / `params`原文必须明确提到返回类型为...参数个数为...
3. **格式规范**value 一律小写保持极简严禁包含任何描述性文字
---
### 任务3查询转换
请完成以下三个转换
#### 3.1 重写查询
将查询重写为更具体详细且对RAG系统中的信息检索更有效的形式
- 更具体和详细
- 如果适用包含来自对话历史的相关上下文
- 保持原始意图
- 适合向量搜索
#### 3.2 生成更广泛的查询
生成给定用户查询的更广泛版本以帮助在RAG系统中检索更全面的上下文信息
- 涵盖与原始查询相关的更一般方面
- 能够帮助检索相关的背景信息
- 保持原始查询的核心主题
- 适合向量搜索
#### 3.3 分解查询
将复杂用户查询分解为更简单更集中的子查询这些子查询可用于RAG系统中的全面信息检索
- 2-5个更简单的子查询
- 每个子查询应关注原始查询的特定方面
- 所有子查询一起应涵盖整个原始查询
- 每个子查询应适合向量搜索
---
### 输出格式要求
请以JSON格式返回所有结果包含以下字段
{{
"intent": {{
"is_code_related": true/false,
"category": "分类名称",
"confidence": 0.0-1.0,
"keywords": ["关键词列表"],
"reasoning": "分类理由",
"requires_code_context": true/false,
"suggested_search_terms": ["搜索词列表"]
}},
"filters": {{
"file_path": "value",
"lang": "value",
...
}},
"transformed": {{
"rewritten": "重写后的查询",
"backward": "更广泛的查询",
"sub_queries": ["子查询1", "子查询2", ...]
}}
}}
### 输出规则
1. 必须输出有效的JSON格式不要包含其他内容
2. confidence表示分类的置信度范围0.0-1.0
3. keywords从问题和对话历史中提取的关键词最多8个
4. reasoning简要说明为什么这样分类要考虑对话历史的内容
5. requires_code_context表示是否需要代码上下文来回答
6. suggested_search_terms建议的检索词最多5个要考虑对话历史中提到的技术或库
7. 只提取明确提到的信息不要进行推测
8. 确保所有字段都有合理的值
现在请分析用户问题和对话历史并输出JSON结果"""

418
utils/query_processor.py Normal file
View File

@ -0,0 +1,418 @@
#!/usr/bin/env python3
"""
查询处理器
整合意图识别filter生成和查询转换功能使用一次LLM调用生成所有结果
"""
import json
from typing import Dict, Any, Optional, List
from loguru import logger
from config import settings
from llama_index.llms.ollama import Ollama
class CodeIntentCategory:
"""代码意图分类枚举 - 细粒度分类体系"""
# 代码解释与逻辑类 (Existing Code Focus)
LOGIC_EXPLANATION = "logic_explanation" # 解释既有代码的底层逻辑
ENTITY_INTRODUCTION = "entity_introduction" # 介绍具体的代码实体函数定义、类属性、API参数
CODE_STRUCTURE = "code_structure" # 询问项目组织
# 代码生成与实现类 (New Code Focus)
CODE_GENERATION = "code_generation" # 请求从零编写完整代码或功能块
BOILERPLATE_IMPLEMENTATION = "boilerplate_implementation" # 请求提供标准算法/模板
# 调试、优化与理论类
ERROR_DEBUGGING = "error_debugging" # 排查 Bug 或异常
CODE_OPTIMIZATION = "code_optimization" # 改进既有代码的性能或质量
ALGORITHM_THEORY = "algorithm_theory" # 算法原理或复杂度分析
# 非代码类
GENERAL_TECHNICAL = "general_technical" # 通用技术咨询
NON_TECHNICAL = "non_technical" # 非技术问题
UNKNOWN = "unknown" # 未知类型
class PromptTemplateType:
"""Prompt模板类型枚举"""
CODE_EXPLANATION = "code_explanation" # 代码解释模板(逻辑解释、实体介绍、代码结构)
CODE_GENERATION = "code_generation" # 代码生成模板(代码生成、模板实现)
CODE_DEBUGGING = "code_debugging" # 代码调试模板(错误调试)
CODE_OPTIMIZATION = "code_optimization" # 代码优化模板(代码优化)
ALGORITHM_EXPLANATION = "algorithm_explanation" # 算法解释模板(算法理论)
GENERAL_QA = "general_qa" # 通用问答模板(通用技术咨询、非技术问题)
class CodeIntentResult:
"""代码意图识别结果"""
def __init__(
self,
is_code_related: bool,
category: str,
confidence: float,
prompt_template_type: str,
keywords: List[str],
reasoning: str,
requires_code_context: bool,
suggested_search_terms: List[str],
):
self.is_code_related = is_code_related # 是否与代码相关True/False
self.category = category # 代码意图分类
self.confidence = confidence # 置信度分数范围0-1之间
self.prompt_template_type = prompt_template_type # Prompt模板类型
self.keywords = keywords # 相关关键词列表
self.reasoning = reasoning # 解释或理由
self.requires_code_context = requires_code_context # 是否需要代码上下文True/False
self.suggested_search_terms = suggested_search_terms # 建议搜索条款列表
def to_dict(self) -> Dict[str, Any]:
"""转换为字典格式"""
return {
"is_code_related": self.is_code_related,
"category": self.category,
"confidence": self.confidence,
"prompt_template_type": self.prompt_template_type,
"keywords": self.keywords,
"reasoning": self.reasoning,
"requires_code_context": self.requires_code_context,
"suggested_search_terms": self.suggested_search_terms,
}
def to_json(self) -> str:
"""转换为JSON格式"""
return json.dumps(self.to_dict(), ensure_ascii=False, indent=2)
class QueryProcessor:
"""
查询处理器
整合意图识别filter生成和查询转换功能
"""
def __init__(self, llm: Optional[Ollama] = None):
"""
初始化查询处理器
Args:
llm: LLM实例如果为None则使用默认配置
"""
if llm is None:
self.llm = Ollama(
model=settings.OLLAMA_MODEL,
base_url=settings.OLLAMA_BASE_URL,
temperature=0.1, # 低温度确保输出稳定
request_timeout=1200.0
)
else:
self.llm = llm
logger.info("查询处理器初始化完成")
def _build_integrated_prompt(self, query: str, history: Optional[str] = None) -> str:
"""
构建集成Prompt一次调用生成所有结果
Args:
query: 用户问题
history: 对话历史
Returns:
str: 集成Prompt
"""
from utils.prompt.integrated_query_processing import INTEGRATED_QUERY_PROCESSING_TEMPLATE
history_str = "" if not history else history
prompt = INTEGRATED_QUERY_PROCESSING_TEMPLATE.format(
history_str=history_str,
query=query
)
return prompt
def _map_to_prompt_template_type(self, category: str) -> str:
"""
根据分类确定Prompt模板类型
Args:
category: 意图分类
Returns:
str: Prompt模板类型
"""
# 代码解释与逻辑类 -> CODE_EXPLANATION
if category in [CodeIntentCategory.LOGIC_EXPLANATION, CodeIntentCategory.ENTITY_INTRODUCTION, CodeIntentCategory.CODE_STRUCTURE]:
return PromptTemplateType.CODE_EXPLANATION
# 代码生成与实现类 -> CODE_GENERATION
elif category in [CodeIntentCategory.CODE_GENERATION, CodeIntentCategory.BOILERPLATE_IMPLEMENTATION]:
return PromptTemplateType.CODE_GENERATION
# 调试、优化与理论类
elif category == CodeIntentCategory.ERROR_DEBUGGING:
return PromptTemplateType.CODE_DEBUGGING
elif category == CodeIntentCategory.CODE_OPTIMIZATION:
return PromptTemplateType.CODE_OPTIMIZATION
elif category == CodeIntentCategory.ALGORITHM_THEORY:
return PromptTemplateType.ALGORITHM_EXPLANATION
# 非代码问题 -> GENERAL_QA
elif category in [CodeIntentCategory.GENERAL_TECHNICAL, CodeIntentCategory.NON_TECHNICAL, CodeIntentCategory.UNKNOWN]:
return PromptTemplateType.GENERAL_QA
else:
return PromptTemplateType.GENERAL_QA
def _parse_llm_response(self, response_text: str) -> Optional[Dict[str, Any]]:
"""
解析LLM响应
Args:
response_text: LLM响应文本
Returns:
解析后的字典如果解析失败返回None
"""
try:
response_text = response_text.strip()
# 尝试提取JSON部分
json_start = response_text.find('{')
json_end = response_text.rfind('}')
if json_start == -1 or json_end == -1:
logger.warning(f"未找到JSON格式响应: {response_text}")
return None
json_str = response_text[json_start:json_end + 1]
logger.debug(f"提取的JSON字符串: {json_str}")
result = json.loads(json_str)
return result
except json.JSONDecodeError as e:
logger.error(f"JSON解析失败: {e}, 响应: {response_text}")
return None
except Exception as e:
logger.error(f"解析响应失败: {e}")
return None
def process_query(self, query: str, history: Optional[str] = None) -> Dict[str, Any]:
"""
处理查询一次调用生成所有结果
Args:
query: 用户问题
history: 对话历史
Returns:
包含意图识别filter生成和查询转换结果的字典
"""
try:
# 构建集成Prompt
prompt = self._build_integrated_prompt(query, history)
# 调用LLM
response = self.llm.complete(prompt)
response_text = response.text
logger.debug(f"LLM响应: {response_text}")
# 解析响应
parsed_result = self._parse_llm_response(response_text)
if parsed_result is None:
logger.error("解析LLM响应失败使用默认结果")
return self._get_default_result(query)
# 验证并处理结果
result = {
"intent": None,
"filters": {},
"transformed": {
"rewritten": query,
"backward": query,
"sub_queries": [query]
}
}
# 处理意图识别结果
try:
if "intent" in parsed_result:
intent_data = parsed_result["intent"]
# 确保所有必需字段都存在
intent_data.setdefault("is_code_related", False)
intent_data.setdefault("category", CodeIntentCategory.UNKNOWN)
intent_data.setdefault("confidence", 0.5)
intent_data.setdefault("keywords", [])
intent_data.setdefault("reasoning", "")
intent_data.setdefault("requires_code_context", False)
intent_data.setdefault("suggested_search_terms", [])
# 确定Prompt模板类型
prompt_template_type = self._map_to_prompt_template_type(intent_data["category"])
# 构建CodeIntentResult对象
intent_result = CodeIntentResult(
is_code_related=intent_data["is_code_related"],
category=intent_data["category"],
confidence=intent_data["confidence"],
prompt_template_type=prompt_template_type,
keywords=intent_data["keywords"],
reasoning=intent_data["reasoning"],
requires_code_context=intent_data["requires_code_context"],
suggested_search_terms=intent_data["suggested_search_terms"]
)
result["intent"] = intent_result
except Exception as e:
logger.error(f"处理意图识别结果失败: {e}")
# 处理过滤条件
try:
if "filters" in parsed_result:
filters = parsed_result["filters"]
if isinstance(filters, dict):
# 验证并过滤结果确保只包含有效的metadata key
valid_keys = ['func_id', 'func_name', 'class_name', 'file_path', 'lang', 'params', 'return_type', 'docstring', 'start_line', 'end_line', 'repo_id', 'branch', 'func_body']
filtered_filters = {}
for key, value in filters.items():
if key in valid_keys and value:
# 确保value是字符串类型
if isinstance(value, str):
filtered_filters[key] = value
result["filters"] = filtered_filters
except Exception as e:
logger.error(f"处理过滤条件失败: {e}")
# 处理查询转换结果
try:
if "transformed" in parsed_result:
transformed = parsed_result["transformed"]
if isinstance(transformed, dict):
result["transformed"].update({
"rewritten": transformed.get("rewritten", query),
"backward": transformed.get("backward", query),
"sub_queries": transformed.get("sub_queries", [query])
})
except Exception as e:
logger.error(f"处理查询转换结果失败: {e}")
logger.info("查询处理完成")
return result
except Exception as e:
logger.error(f"查询处理失败: {e}")
return self._get_default_result(query)
def _get_default_result(self, query: str) -> Dict[str, Any]:
"""
获取默认结果当处理失败时使用
Args:
query: 用户问题
Returns:
默认结果
"""
logger.warning(f"使用默认查询处理结果: {query}")
# 构建默认的意图识别结果
default_intent = CodeIntentResult(
is_code_related=False,
category=CodeIntentCategory.UNKNOWN,
confidence=0.0,
prompt_template_type=PromptTemplateType.GENERAL_QA,
keywords=[],
reasoning="查询处理失败,使用默认结果",
requires_code_context=False,
suggested_search_terms=[query]
)
return {
"intent": default_intent,
"filters": {},
"transformed": {
"rewritten": query,
"backward": query,
"sub_queries": [query]
}
}
# 辅助函数
def create_query_processor(llm: Optional[Ollama] = None) -> QueryProcessor:
"""
创建查询处理器实例
Args:
llm: LLM实例如果为None则使用默认配置
Returns:
QueryProcessor: 处理器实例
"""
return QueryProcessor(llm)
class MetadataFilter:
"""
Metadata过滤工具类
"""
@staticmethod
def apply_filter(metadata: Dict[str, Any], filters: Dict[str, Any]) -> bool:
"""
应用过滤条件到metadata
Args:
metadata: 文档的metadata
filters: 过滤条件
Returns:
bool: 如果metadata符合过滤条件返回True否则返回False
"""
if not filters:
return True
for key, value in filters.items():
if key not in metadata:
return False
metadata_value = metadata[key]
if isinstance(metadata_value, str) and isinstance(value, str):
# 对于字符串类型,使用大小写不敏感的模糊匹配
if value.lower() not in metadata_value.lower():
return False
else:
if metadata_value != value:
return False
return True
if __name__ == "__main__":
"""测试查询处理器"""
processor = QueryProcessor()
test_queries = [
"在 data_structures 目录下二叉搜索树Binary Search Tree的删除操作依赖于哪些辅助方法来寻找后继节点",
"如何实现快速排序算法?",
"kth_number的时间复杂度是多少",
"今天天气怎么样?"
]
for query in test_queries:
print(f"\n{'='*60}")
print(f"问题: {query}")
print('='*60)
result = processor.process_query(query)
print("意图识别结果:")
if result["intent"]:
print(result["intent"].to_json())
print("\n过滤条件:")
import json
print(json.dumps(result["filters"], ensure_ascii=False, indent=2))
print("\n查询转换结果:")
print(json.dumps(result["transformed"], ensure_ascii=False, indent=2))

150
utils/query_transformer.py Normal file
View File

@ -0,0 +1,150 @@
"""
Query transformer for RAG system
Implements three query transformation techniques:
1. Query rewriting: Make queries more specific and detailed
2. Backward prompt generation: Generate broader queries for context retrieval
3. Sub-query decomposition: Break down complex queries into simpler components
"""
from typing import List, Dict, Optional, Tuple
from loguru import logger
from config import settings
from llama_index.llms.ollama import Ollama
from utils.prompt.query_trans import REWRITE_PROMPT, BACKWARD_PROMPT, DECOMPOSE_PROMPT
class QueryTransformer:
"""
Query transformer for RAG system
"""
def __init__(self, llm: Ollama):
"""
Initialize query transformer
Args:
llm: Ollama LLM instance for generating transformations
"""
self.llm = llm
async def rewrite_query(self, query: str, history: str = "") -> str:
"""
Rewrite query to be more specific and detailed
Args:
query: Original user query
history: Chat history (optional)
Returns:
Rewritten query
"""
try:
prompt = REWRITE_PROMPT.format(query=query, history=history)
response = await self.llm.acomplete(prompt=prompt)
rewritten_query = response.text.strip()
logger.info(f"Rewritten query: {rewritten_query}")
return rewritten_query
except Exception as e:
logger.error(f"Error in query rewriting: {e}")
return query
async def backward_query(self, query: str, history: str = "") -> str:
"""
Generate broader query for context retrieval
Args:
query: Original user query
history: Chat history (optional)
Returns:
Broader query for context retrieval
"""
try:
prompt = BACKWARD_PROMPT.format(query=query, history=history)
response = await self.llm.acomplete(prompt=prompt)
backward_query = response.text.strip()
logger.info(f"Backward query: {backward_query}")
return backward_query
except Exception as e:
logger.error(f"Error in backward query generation: {e}")
return query
async def decompose_query(self, query: str, history: str = "") -> List[str]:
"""
Decompose complex query into simpler sub-queries
Args:
query: Original user query
history: Chat history (optional)
Returns:
List of sub-queries
"""
try:
prompt = DECOMPOSE_PROMPT.format(query=query, history=history)
response = await self.llm.acomplete(prompt=prompt)
sub_queries_text = response.text.strip()
# Parse sub-queries from response
sub_queries = []
for line in sub_queries_text.split('\n'):
line = line.strip()
if line and (line.startswith('1.') or line.startswith('2.') or line.startswith('3.') or line.startswith('4.') or line.startswith('5.')):
sub_query = line.split('.', 1)[1].strip()
if sub_query:
sub_queries.append(sub_query)
# If no valid sub-queries found, return original query as single sub-query
if not sub_queries:
sub_queries = [query]
logger.info(f"Original query: {query}")
logger.info(f"Decomposed sub-queries: {sub_queries}")
return sub_queries
except Exception as e:
logger.error(f"Error in query decomposition: {e}")
return [query]
async def transform_query(self, query: str, history: str = "") -> Dict[str, any]:
"""
Transform query using all three techniques
Args:
query: Original user query
history: Chat history (optional)
Returns:
Dict with transformed queries
"""
try:
logger.info(f"Transforming query: {query}")
# Run all transformations in parallel
from asyncio import gather
rewritten, backward, sub_queries = await gather(
self.rewrite_query(query, history),
self.backward_query(query, history),
self.decompose_query(query, history)
)
# Combine all transformed queries for comprehensive retrieval
all_transformed_queries = [rewritten, backward] + sub_queries
return {
"original": query,
"rewritten": rewritten,
"backward": backward,
"sub_queries": sub_queries,
"all_transformed": all_transformed_queries
}
except Exception as e:
logger.error(f"Error in query transformation: {e}")
return {
"original": query,
"rewritten": query,
"backward": query,
"sub_queries": [query],
"all_transformed": [query]
}