2ab8b56a01
此提交实现了完整的知识库管理系统: 1. 新增Vue3 + Antd Vue前端管理后台,包含登录、文档管理、检索、类目设置等完整页面 2. 重构后端LLM调用抽象层,支持Ollama与OpenAI兼容服务动态切换 3. 调整默认嵌入模型配置为本地bge-m3模式 4. 优化入库任务去重逻辑与缓存清理机制 5. 完善Docker镜像构建与docker-compose部署配置 6. 修复多项测试用例与兼容性问题 7. 新增运行时配置API,支持动态调整系统参数
246 lines
11 KiB
Python
246 lines
11 KiB
Python
"""query 解析与分类路由模块
|
||
|
||
分层 RAG 在线侧第一步:用 LLM(Ollama 或 OpenAI 兼容服务,由
|
||
runtime_settings.models.query 决定)将用户 query 解析为结构化 JSON
|
||
(命中类目+置信度、rewrite 后 query、关键词),再由纯函数做路由决策:
|
||
- 高置信且命中类目数 <= 上限 → 按类目过滤检索
|
||
- 低置信 / 解析失败 / 命中类目过多 → 全库兜底(不丢召回)
|
||
|
||
解析(LLM 调用)与决策(纯函数)分离,便于单元测试。
|
||
"""
|
||
|
||
import json
|
||
import re
|
||
from hashlib import sha256
|
||
|
||
import structlog
|
||
from pydantic import BaseModel, Field, ValidationError
|
||
|
||
from app.config import settings
|
||
from app.models.knowledge import UNCATEGORIZED, TaxonomyCategory, load_taxonomy
|
||
from app.services.llm import LLMClient, create_llm_client
|
||
from app.services.redis import RedisCache, get_cache
|
||
|
||
logger = structlog.get_logger()
|
||
|
||
# 从 LLM 输出中提取第一个 {...} JSON 块(贪婪匹配到最后的 },兼容嵌套对象)
|
||
_JSON_BLOCK_RE = re.compile(r"\{.*\}", re.DOTALL)
|
||
|
||
|
||
class CategoryHit(BaseModel):
|
||
"""query 命中的类目及置信度"""
|
||
|
||
name: str = Field(description="类目名称,必须来自 taxonomy")
|
||
confidence: float = Field(ge=0, le=1, description="命中置信度,范围 [0, 1]")
|
||
|
||
|
||
class ParsedQuery(BaseModel):
|
||
"""query 解析结果"""
|
||
|
||
raw_query: str = Field(description="原始 query")
|
||
rewrite: str = Field(description="rewrite 后的 query")
|
||
keywords: list[str] = Field(default_factory=list, description="提取的关键词")
|
||
categories: list[CategoryHit] = Field(default_factory=list, description="命中类目列表")
|
||
entities: list[str] = Field(default_factory=list, description="提取的实体(人名/产品/技术等)")
|
||
intent: str = Field(default="", description="查询意图分类(如 查询/对比/操作/定义/排查)")
|
||
time_range: str = Field(default="", description="时间范围(如 '2023年'、'最近一个月'),空表示无")
|
||
parse_failed: bool = Field(default=False, description="LLM 输出解析是否失败")
|
||
|
||
|
||
class RouteDecision(BaseModel):
|
||
"""路由决策结果"""
|
||
|
||
fallback: bool = Field(description="是否走全库兜底")
|
||
filter_categories: list[str] | None = Field(default=None, description="过滤类目名列表,None 表示全库不过滤")
|
||
reason: str = Field(description="决策原因:parse_failed | low_confidence | too_many_categories | routed")
|
||
parsed: ParsedQuery = Field(description="对应的 query 解析结果")
|
||
|
||
|
||
def _extract_json(raw: str) -> dict | None:
|
||
"""从 LLM 输出中提取 JSON 对象
|
||
|
||
先尝试直接解析;失败则用正则提取第一个 {...} 块再解析。
|
||
返回 None 表示无法提取出合法的 JSON 对象。
|
||
"""
|
||
text = raw.strip()
|
||
try:
|
||
data = json.loads(text)
|
||
return data if isinstance(data, dict) else None
|
||
except json.JSONDecodeError:
|
||
pass
|
||
|
||
match = _JSON_BLOCK_RE.search(text)
|
||
if not match:
|
||
return None
|
||
try:
|
||
data = json.loads(match.group(0))
|
||
return data if isinstance(data, dict) else None
|
||
except json.JSONDecodeError:
|
||
return None
|
||
|
||
|
||
def decide_route(parsed: ParsedQuery, threshold: float, max_categories: int) -> RouteDecision:
|
||
"""路由决策(纯函数)
|
||
|
||
- 解析失败 → 全库兜底(parse_failed)
|
||
- 无 confidence >= threshold 的类目 → 全库兜底(low_confidence)
|
||
- 命中类目数 > max_categories → 全库兜底(too_many_categories)
|
||
- 否则按类目过滤检索,类目按 confidence 降序排列(routed)
|
||
"""
|
||
if parsed.parse_failed:
|
||
return RouteDecision(fallback=True, filter_categories=None, reason="parse_failed", parsed=parsed)
|
||
|
||
hits = [c for c in parsed.categories if c.confidence >= threshold]
|
||
if not hits:
|
||
return RouteDecision(fallback=True, filter_categories=None, reason="low_confidence", parsed=parsed)
|
||
|
||
if len(hits) > max_categories:
|
||
return RouteDecision(fallback=True, filter_categories=None, reason="too_many_categories", parsed=parsed)
|
||
|
||
hits.sort(key=lambda c: c.confidence, reverse=True)
|
||
return RouteDecision(
|
||
fallback=False,
|
||
filter_categories=[c.name for c in hits],
|
||
reason="routed",
|
||
parsed=parsed,
|
||
)
|
||
|
||
|
||
class QueryParser:
|
||
"""query 解析器:调用 LLM 将 query 解析为结构化 JSON"""
|
||
|
||
def __init__(
|
||
self,
|
||
ollama: LLMClient | None = None,
|
||
taxonomy: list[TaxonomyCategory] | None = None,
|
||
cache: RedisCache | None = None,
|
||
) -> None:
|
||
# 默认按 runtime_settings.models.query 选择 LLM 实现;
|
||
# 测试可通过 ollama 参数注入替身。
|
||
self.ollama = ollama or create_llm_client("query")
|
||
self.taxonomy = taxonomy if taxonomy is not None else load_taxonomy()
|
||
# 解析结果缓存,缺省用全局单例;RedisCache 全操作容错,缓存不可用时退化为无缓存行为
|
||
self.cache = cache if cache is not None else get_cache()
|
||
# 可作为路由命中类目的名字集合(uncategorized 不可作为路由命中类目)
|
||
self._routable_names = {c.name for c in self.taxonomy if c.name != UNCATEGORIZED}
|
||
|
||
async def parse(self, query: str) -> ParsedQuery:
|
||
"""调用 LLM 将 query 解析为结构化结果
|
||
|
||
解析失败(输出非 JSON / 必填字段缺失或类型错误)时,
|
||
返回 parse_failed=True 的兜底结果(rewrite 为原 query)。
|
||
"""
|
||
prompt = self._build_prompt(query)
|
||
raw = await self.ollama.generate(prompt, json_mode=True)
|
||
|
||
data = _extract_json(raw)
|
||
if data is None:
|
||
logger.warning("query 解析失败:LLM 输出非合法 JSON", query=query, output=raw[:200])
|
||
return ParsedQuery(raw_query=query, rewrite=query, parse_failed=True)
|
||
|
||
rewrite = data.get("rewrite")
|
||
keywords = data.get("keywords")
|
||
categories = data.get("categories")
|
||
if not isinstance(rewrite, str) or not isinstance(keywords, list) or not isinstance(categories, list):
|
||
logger.warning("query 解析失败:JSON 必填字段缺失或类型错误", query=query, output=raw[:200])
|
||
return ParsedQuery(raw_query=query, rewrite=query, parse_failed=True)
|
||
|
||
# 扩展字段:实体/意图/时间范围(缺失或类型错误时降级为默认值,不触发 parse_failed)
|
||
entities = data.get("entities", [])
|
||
if not isinstance(entities, list):
|
||
logger.warning("query 解析 entities 非列表,置空", query=query)
|
||
entities = []
|
||
intent = data.get("intent", "")
|
||
if not isinstance(intent, str):
|
||
logger.warning("query 解析 intent 非字符串,置空", query=query)
|
||
intent = ""
|
||
time_range = data.get("time_range", "")
|
||
if not isinstance(time_range, str):
|
||
logger.warning("query 解析 time_range 非字符串,置空", query=query)
|
||
time_range = ""
|
||
|
||
return ParsedQuery(
|
||
raw_query=query,
|
||
rewrite=rewrite,
|
||
keywords=[str(k) for k in keywords],
|
||
categories=self._validate_categories(categories, query),
|
||
entities=[str(e) for e in entities],
|
||
intent=intent,
|
||
time_range=time_range,
|
||
)
|
||
|
||
async def parse_and_route(self, query: str) -> RouteDecision:
|
||
"""解析 query 并做路由决策(threshold / max_categories 取自 settings)
|
||
|
||
parse() 的 LLM 解析结果按 query 哈希缓存;路由决策每次现算(纯函数,
|
||
阈值取最新 settings),缓存数据损坏时回退为重新走 LLM 解析。
|
||
"""
|
||
cache_key = f"qparse:{sha256(query.encode()).hexdigest()[:16]}"
|
||
parsed = await self._get_cached_parsed(cache_key, query)
|
||
if parsed is None:
|
||
parsed = await self.parse(query)
|
||
await self.cache.set_json(cache_key, parsed.model_dump())
|
||
return decide_route(
|
||
parsed,
|
||
threshold=settings.classify_confidence_threshold,
|
||
max_categories=settings.classify_max_categories,
|
||
)
|
||
|
||
async def _get_cached_parsed(self, cache_key: str, query: str) -> ParsedQuery | None:
|
||
"""读取缓存的解析结果;未命中或缓存数据无法重建时返回 None"""
|
||
cached = await self.cache.get_json(cache_key)
|
||
if cached is None:
|
||
return None
|
||
try:
|
||
parsed = ParsedQuery.model_validate(cached)
|
||
except ValidationError:
|
||
logger.warning("query 解析缓存数据损坏,按未命中处理", query=query, cache_key=cache_key)
|
||
return None
|
||
logger.info("query 解析缓存命中", query=query, cache_key=cache_key)
|
||
return parsed
|
||
|
||
def _build_prompt(self, query: str) -> str:
|
||
"""构造解析 prompt:列出 taxonomy 类目名+描述,要求模型只输出 JSON"""
|
||
category_lines = [f"- {c.name}: {c.description}" for c in self.taxonomy if c.name != UNCATEGORIZED]
|
||
category_block = "\n".join(category_lines)
|
||
return (
|
||
"你是搜索查询分析助手。请分析用户 query,完成以下事项:\n"
|
||
"1. 判断 query 意图命中以下哪些知识类目,并给出每个类目的置信度(0~1 之间的小数);\n"
|
||
"2. 将 query 改写为更适合检索的形式;\n"
|
||
"3. 提取 query 的关键词;\n"
|
||
"4. 提取 query 中的实体(人名、产品名、技术名词、组织机构等);\n"
|
||
"5. 判断 query 的查询意图,从 [查询, 对比, 操作, 定义, 排查] 中选最接近的一个;\n"
|
||
"6. 提取 query 中的时间范围(如 '2023年'、'上个月'、'Q1'),无则留空字符串。\n\n"
|
||
f"可选类目:\n{category_block}\n\n"
|
||
"要求:\n"
|
||
"- categories 中的 name 只能从上面的类目名中选择,不要输出其他名称;\n"
|
||
"- 若没有明显命中的类目,categories 返回空列表;\n"
|
||
"- entities 没有则返回空数组;\n"
|
||
"- 只输出 JSON,不要输出任何其他内容。\n\n"
|
||
'输出格式:{"categories": [{"name": "类目名", "confidence": 0.0}], '
|
||
'"rewrite": "改写后的 query", "keywords": ["关键词"], '
|
||
'"entities": ["实体"], "intent": "查询", "time_range": ""}\n\n'
|
||
f"用户 query:{query}"
|
||
)
|
||
|
||
def _validate_categories(self, categories: list, query: str) -> list[CategoryHit]:
|
||
"""校验并过滤 LLM 输出的类目条目
|
||
|
||
丢弃:类目名不在 taxonomy 可路由类目中的条目(含 uncategorized)、
|
||
confidence 缺失或越界(不在 [0, 1])的条目。
|
||
"""
|
||
hits: list[CategoryHit] = []
|
||
for item in categories:
|
||
if not isinstance(item, dict):
|
||
continue
|
||
name = item.get("name")
|
||
confidence = item.get("confidence")
|
||
if not isinstance(name, str) or name not in self._routable_names:
|
||
logger.warning("丢弃不在 taxonomy 中的类目", query=query, name=name)
|
||
continue
|
||
if not isinstance(confidence, (int, float)) or not 0 <= confidence <= 1:
|
||
logger.warning("丢弃 confidence 越界的类目", query=query, name=name, confidence=confidence)
|
||
continue
|
||
hits.append(CategoryHit(name=name, confidence=float(confidence)))
|
||
return hits
|