"""query 解析与分类路由模块 分层 RAG 在线侧第一步:用 Ollama 小模型将用户 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 from app.services.ollama import OllamaClient 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="命中类目列表") 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 解析器:调用 Ollama 小模型将 query 解析为结构化 JSON""" def __init__( self, ollama: OllamaClient, taxonomy: list[TaxonomyCategory], cache: RedisCache | None = None ) -> None: self.ollama = ollama self.taxonomy = taxonomy # 解析结果缓存,缺省用全局单例;RedisCache 全操作容错,缓存不可用时退化为无缓存行为 self.cache = cache if cache is not None else get_cache() # 可作为路由命中类目的名字集合(uncategorized 不可作为路由命中类目) self._routable_names = {c.name for c in 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) return ParsedQuery( raw_query=query, rewrite=rewrite, keywords=[str(k) for k in keywords], categories=self._validate_categories(categories, query), ) 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\n" f"可选类目:\n{category_block}\n\n" "要求:\n" "- categories 中的 name 只能从上面的类目名中选择,不要输出其他名称;\n" "- 若没有明显命中的类目,categories 返回空列表;\n" "- 只输出 JSON,不要输出任何其他内容。\n\n" '输出格式:{"categories": [{"name": "类目名", "confidence": 0.0}], ' '"rewrite": "改写后的 query", "keywords": ["关键词"]}\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