Initial commit: QMDSearch 分层信息检索服务
- FastAPI + Qdrant + Redis + Ollama 技术栈 - L1→L2→L3→chunk 四层分层检索(dense + sparse RRF 融合) - 文档三级总结与 2.5 级回退 - query 解析路由与分类 - /admin 管理页面
This commit is contained in:
@@ -0,0 +1,214 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user