Files
QMDSearch/app/core/query_parser.py
T
kplam 2ab8b56a01 feat: 完成全量功能开发,包括前端管理后台与后端服务优化
此提交实现了完整的知识库管理系统:
1. 新增Vue3 + Antd Vue前端管理后台,包含登录、文档管理、检索、类目设置等完整页面
2. 重构后端LLM调用抽象层,支持Ollama与OpenAI兼容服务动态切换
3. 调整默认嵌入模型配置为本地bge-m3模式
4. 优化入库任务去重逻辑与缓存清理机制
5. 完善Docker镜像构建与docker-compose部署配置
6. 修复多项测试用例与兼容性问题
7. 新增运行时配置API,支持动态调整系统参数
2026-07-31 12:05:25 +08:00

246 lines
11 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""query 解析与分类路由模块
分层 RAG 在线侧第一步:用 LLMOllama 或 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