nex_docus/backend/app/services/rag_service.py

316 lines
12 KiB
Python

"""
知识库 RAG 服务 - ZVec 向量检索 + 大模型生成
"""
import logging
from typing import AsyncIterator, List, Dict, Any
from pathlib import Path
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select
from app.models.llm_model_config import LLMModelConfig
from app.core.config import settings
from app.services.llm_provider_service import LLMProviderService
from app.services.zvec_service import zvec_service
from app.services.storage import storage_service
logger = logging.getLogger(__name__)
# 单个命中分块注入上下文时,围绕命中位置向前后各扩展的字符数。
# 分块后每次只注入命中段落及其上下文,而非整篇文档,从根本上避免长文档被截断。
CHUNK_CONTEXT_WINDOW = 1200
MAX_CHUNKS_PER_DOCUMENT = 3
class RAGService:
"""RAG 知识库检索和生成服务"""
@staticmethod
async def retrieve_documents(
db: AsyncSession,
project_id: int,
query: str,
top_k: int = 5,
) -> List[Dict[str, Any]]:
"""通过 ZVec 向量检索相关文档,并读取文档内容"""
# 多取一些分块,再按文件去重,避免高频命中文件挤掉其他相关文档。
matched = await zvec_service.search_similar(db, project_id, query, top_k * 4)
if not matched:
return []
from app.models.project import Project
stmt = select(Project).where(Project.id == project_id)
result = await db.execute(stmt)
project = result.scalar_one_or_none()
if not project:
return []
# 读取文件内容做缓存,避免同一文件多个分块命中时重复读盘
file_content_cache: Dict[str, str] = {}
async def _read_content(path: str) -> str:
if path in file_content_cache:
return file_content_cache[path]
full_path = storage_service.get_secure_path(project.storage_key, path)
text = await storage_service.read_file(full_path)
file_content_cache[path] = text
return text
docs_by_path: Dict[str, Dict[str, Any]] = {}
for item in matched:
file_path = item["file_path"]
chunk_text = item.get("chunk_text") or ""
try:
content = await _read_content(file_path)
excerpt = RAGService._resolve_excerpt(
content=content,
chunk_text=chunk_text,
chunk_index=item.get("chunk_index", 0),
)
snippet = RAGService._extract_chunk_context(content, chunk_text, excerpt)
doc = {
"file_path": file_path,
"file_name": Path(file_path).name,
"chunk_index": item.get("chunk_index", 0),
"anchor_text": RAGService._build_anchor(excerpt or chunk_text),
"score": item.get("score", 0.0),
"excerpt": excerpt,
"content": snippet,
"_merged_chunks": 1,
}
except Exception:
doc = {
"file_path": file_path,
"file_name": Path(file_path).name,
"chunk_index": item.get("chunk_index", 0),
"anchor_text": RAGService._build_anchor(chunk_text),
"score": item.get("score", 0.0),
"excerpt": "",
"content": "",
"_merged_chunks": 1,
}
existing = docs_by_path.get(file_path)
if existing is None:
docs_by_path[file_path] = doc
continue
if existing["_merged_chunks"] >= MAX_CHUNKS_PER_DOCUMENT:
continue
existing["score"] = max(existing.get("score", 0.0), doc.get("score", 0.0))
existing["excerpt"] = RAGService._merge_text_blocks(
existing.get("excerpt", ""), doc.get("excerpt", "")
)
existing["content"] = RAGService._merge_text_blocks(
existing.get("content", ""), doc.get("content", "")
)
existing["_merged_chunks"] += 1
docs_with_content = list(docs_by_path.values())
docs_with_content = docs_with_content[:top_k]
for citation_id, doc in enumerate(docs_with_content, 1):
doc.pop("_merged_chunks", None)
doc["citation_id"] = citation_id
return docs_with_content
@staticmethod
def _merge_text_blocks(existing: str, incoming: str) -> str:
"""合并同一文件的命中内容,避免重复分块生成多个引用。"""
current = (existing or "").strip()
candidate = (incoming or "").strip()
if not candidate or candidate in current:
return current
if not current or current in candidate:
return candidate
return f"{current}\n\n{candidate}"
@staticmethod
def _resolve_excerpt(content: str, chunk_text: str, chunk_index: int) -> str:
"""返回引用预览片段。
document_vector.chunk_text 是向量命中的分块文本,应直接作为引用预览。
仅在 chunk_text 为空时,才按 chunk_index 从原文反推作为兜底。
"""
text = (chunk_text or "").strip()
if text:
return text
return RAGService._extract_chunk_by_index(
content,
chunk_index,
settings.CHUNK_SIZE,
settings.CHUNK_OVERLAP,
) or text
@staticmethod
def _build_anchor(text: str, max_chars: int = 120) -> str:
for raw_line in (text or "").splitlines():
line = raw_line.strip()
if line:
return line[:max_chars]
return (text or "").strip()[:max_chars]
@staticmethod
def _extract_chunk_by_index(
content: str,
chunk_index: int,
chunk_size: int,
overlap: int,
) -> str:
"""按向量化时的分块参数,从原文切出命中的原始分块。"""
text = content or ""
if not text:
return ""
chunk_size = max(1, int(chunk_size or 1))
overlap = max(0, min(int(overlap or 0), chunk_size - 1))
step = chunk_size - overlap
start = max(0, int(chunk_index or 0)) * step
if start >= len(text):
return ""
return text[start:start + chunk_size]
@staticmethod
def _extract_chunk_context(content: str, anchor: str, fallback: str = "") -> str:
"""围绕命中锚点截取上下文窗口。
用锚点在原文中定位命中位置,向前后各扩展 CHUNK_CONTEXT_WINDOW 字符,
避免注入整篇长文档。定位失败时回退为命中的原始分块。
"""
text = content or ""
if not text:
return ""
pos = text.find(anchor) if anchor else -1
if pos < 0:
return fallback or text[:CHUNK_CONTEXT_WINDOW * 2]
start = max(0, pos - CHUNK_CONTEXT_WINDOW)
end = min(len(text), pos + len(anchor) + CHUNK_CONTEXT_WINDOW)
return text[start:end]
@staticmethod
async def generate_response(
db: AsyncSession,
query: str,
project_id: int,
llm_config_id: int,
retrieved_docs: List[Dict[str, Any]],
conversation_history: List[Dict[str, str]],
) -> str:
"""基于检索文档生成对话回复"""
llm_config, system_prompt, messages = await RAGService._prepare_generation(
db, query, llm_config_id, retrieved_docs, conversation_history
)
response_text = await LLMProviderService.generate_text(
provider=llm_config.provider,
endpoint_url=llm_config.endpoint_url,
api_key=llm_config.api_key,
llm_model_name=llm_config.llm_model_name,
timeout=llm_config.llm_timeout,
temperature=float(llm_config.llm_temperature),
top_p=float(llm_config.llm_top_p),
max_tokens=llm_config.llm_max_tokens,
system_prompt=system_prompt,
messages=messages,
)
return response_text
@staticmethod
async def generate_response_stream(
db: AsyncSession,
query: str,
project_id: int,
llm_config_id: int,
retrieved_docs: List[Dict[str, Any]],
conversation_history: List[Dict[str, str]],
) -> AsyncIterator[str]:
"""基于检索文档流式生成对话回复"""
llm_config, system_prompt, messages = await RAGService._prepare_generation(
db, query, llm_config_id, retrieved_docs, conversation_history
)
async for chunk in LLMProviderService.generate_text_stream(
provider=llm_config.provider,
endpoint_url=llm_config.endpoint_url,
api_key=llm_config.api_key,
llm_model_name=llm_config.llm_model_name,
timeout=llm_config.llm_timeout,
temperature=float(llm_config.llm_temperature),
top_p=float(llm_config.llm_top_p),
max_tokens=llm_config.llm_max_tokens,
system_prompt=system_prompt,
messages=messages,
):
yield chunk
@staticmethod
async def _prepare_generation(
db: AsyncSession,
query: str,
llm_config_id: int,
retrieved_docs: List[Dict[str, Any]],
conversation_history: List[Dict[str, str]],
):
stmt = select(LLMModelConfig).where(LLMModelConfig.config_id == llm_config_id)
result = await db.execute(stmt)
llm_config = result.scalar_one_or_none()
if not llm_config:
raise ValueError(f"LLM配置不存在: {llm_config_id}")
context_text = RAGService._build_context(retrieved_docs)
citation_hint = RAGService._build_citation_hint(retrieved_docs)
system_prompt = f"""你是一个知识库助手。基于用户提供的知识库文档,回答用户的问题。
如果知识库中没有相关信息,请明确说明。
知识库文档内容:
{context_text}
引用规则:
{citation_hint}
回答要求:
1. 仅基于知识库内容回答,不要编造。
2. 如果使用了某条知识,请在对应句子后标注引用编号,例如 [1] 或 [1][3]。
3. 如果没有找到依据,请直接说明未检索到相关内容。
请基于以上知识库内容,用中文回答用户的问题。"""
messages = list(conversation_history)
messages.append({"role": "user", "content": query})
return llm_config, system_prompt, messages
@staticmethod
def _build_context(retrieved_docs: List[Dict[str, Any]]) -> str:
"""构建上下文文本"""
if not retrieved_docs:
return "(未找到相关知识库内容)"
context_parts = []
for i, doc_info in enumerate(retrieved_docs, 1):
content = doc_info.get("content", "").strip()
if not content:
continue
citation_id = doc_info.get("citation_id", i)
context_parts.append(
f"--- [{citation_id}] 文档: {doc_info['file_path']} ---\n{content}"
)
return "\n\n".join(context_parts) if context_parts else "(未找到相关知识库内容)"
@staticmethod
def _build_citation_hint(retrieved_docs: List[Dict[str, Any]]) -> str:
"""构建引用提示"""
if not retrieved_docs:
return "未检索到文档时不要标注引用。"
return "\n".join(
f"[{doc.get('citation_id', index)}] {doc.get('file_name') or doc.get('file_path')}"
for index, doc in enumerate(retrieved_docs, 1)
)
rag_service = RAGService()