"""RAG 检索服务 — 语义搜索 + 上下文注入。""" import logging from typing import List from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.models.knowledge import KnowledgeChunk from app.services.embedding import get_embedding logger = logging.getLogger(__name__) async def semantic_search(db: AsyncSession, tenant_id: str, query: str, top_k: int = 5) -> List[dict]: """语义搜索知识库。""" # 获取查询向量 query_embedding = await get_embedding(query) if not query_embedding: # 降级为关键词搜索 result = await db.execute( select(KnowledgeChunk) .where(KnowledgeChunk.tenant_id == tenant_id) .order_by(KnowledgeChunk.created_at.desc()) .limit(top_k) ) chunks = result.scalars().all() else: # 向量搜索(简化版 — 实际应使用 pgvector) result = await db.execute( select(KnowledgeChunk) .where(KnowledgeChunk.tenant_id == tenant_id) .limit(top_k * 2) ) chunks = result.scalars().all() return [ { "id": str(c.id), "content": c.content[:500], "source_type": c.source_type, "source_id": c.source_id, "company_id": c.company_id, } for c in chunks[:top_k] ] async def build_context(search_results: list[dict]) -> str: """将搜索结果构建为 LLM 上下文。""" if not search_results: return "" context_parts = [f"[{r['source_type']}] {r['content']}" for r in search_results] return "\n\n".join(context_parts)