Files
QMDSearch/tests/test_query_parser.py
kplam dce9e31bde feat: 新增多格式文件上传入库与认证体系
- 新增 JWT 认证模块,支持登录/注册/用户管理
- 新增文件上传接口,支持 .txt/.md/.html/.pdf/.docx 等格式解析入库
- 新增检索结果 AI 总结功能
- 新增文本去重缓存机制
- 新增全局认证夹具简化测试
- 新增配置项与环境变量支持
- 完善文档与测试覆盖
2026-07-30 10:30:15 +08:00

302 lines
12 KiB
Python
Raw Permalink 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 解析与分类路由的单元测试(mock OllamaClient,不真实联网)"""
import json
from app.core.query_parser import CategoryHit, ParsedQuery, QueryParser, decide_route
from app.models.knowledge import UNCATEGORIZED, TaxonomyCategory
def _taxonomy() -> list[TaxonomyCategory]:
"""测试用 taxonomy 类目集"""
return [
TaxonomyCategory(name="技术文档", description="架构设计、API 文档、开发规范等技术资料"),
TaxonomyCategory(name="产品手册", description="产品功能介绍、使用说明"),
TaxonomyCategory(name="财务行政", description="财务制度、报销流程、行政通知"),
TaxonomyCategory(name=UNCATEGORIZED, description="无法归入其他类目的文档"),
]
class FakeOllama:
"""返回固定响应的假 OllamaClient,记录调用参数"""
def __init__(self, response: str) -> None:
self.response = response
self.calls: list[dict] = []
async def generate(self, prompt: str, json_mode: bool = False) -> str:
self.calls.append({"prompt": prompt, "json_mode": json_mode})
return self.response
def _make_parser(response: str) -> QueryParser:
return QueryParser(ollama=FakeOllama(response), taxonomy=_taxonomy()) # type: ignore[arg-type]
class TestQueryParserParse:
"""QueryParser.parse 的 JSON 解析与容错"""
async def test_parse_valid_json(self):
"""正常 JSON 响应 → 正确解析出 categories/rewrite/keywords"""
response = json.dumps(
{
"categories": [{"name": "技术文档", "confidence": 0.9}],
"rewrite": "如何设计系统架构",
"keywords": ["架构", "设计"],
},
ensure_ascii=False,
)
parser = _make_parser(response)
parsed = await parser.parse("怎么做架构设计")
assert parsed.parse_failed is False
assert parsed.raw_query == "怎么做架构设计"
assert parsed.rewrite == "如何设计系统架构"
assert parsed.keywords == ["架构", "设计"]
assert len(parsed.categories) == 1
assert parsed.categories[0].name == "技术文档"
assert parsed.categories[0].confidence == 0.9
async def test_parse_uses_json_mode_and_prompt_contains_taxonomy(self):
"""以 json_mode 调用 LLM,且 prompt 中包含 taxonomy 类目名与描述"""
parser = _make_parser('{"categories": [], "rewrite": "q", "keywords": []}')
await parser.parse("报销流程是什么")
ollama = parser.ollama
assert ollama.calls[0]["json_mode"] is True
prompt = ollama.calls[0]["prompt"]
assert "技术文档" in prompt and "架构设计" in prompt
assert "财务行政" in prompt and "报销流程" in prompt
assert "报销流程是什么" in prompt
async def test_parse_json_with_surrounding_text(self):
"""响应带前后多余文本 → 正则提取第一个 {...} 块成功"""
response = (
"好的,分析结果如下:\n"
'{"categories": [{"name": "财务行政", "confidence": 0.8}], "rewrite": "报销流程", "keywords": ["报销"]}\n'
"以上就是分析结果。"
)
parser = _make_parser(response)
parsed = await parser.parse("怎么报销")
assert parsed.parse_failed is False
assert parsed.rewrite == "报销流程"
assert parsed.keywords == ["报销"]
assert [c.name for c in parsed.categories] == ["财务行政"]
async def test_parse_non_json_response(self):
"""完全非 JSON 响应 → parse_failed=Truerewrite 回退为原 query"""
parser = _make_parser("抱歉,我无法理解这个问题。")
parsed = await parser.parse("blah blah")
assert parsed.parse_failed is True
assert parsed.rewrite == "blah blah"
assert parsed.keywords == []
assert parsed.categories == []
async def test_parse_missing_fields(self):
"""JSON 合法但必填字段缺失 → parse_failed=True"""
parser = _make_parser('{"rewrite": "只有 rewrite"}')
parsed = await parser.parse("q")
assert parsed.parse_failed is True
assert parsed.rewrite == "q"
assert parsed.categories == []
async def test_unknown_category_dropped(self):
"""未知类目名(含 uncategorized)被丢弃,合法类目保留"""
response = json.dumps(
{
"categories": [
{"name": "不存在的类目", "confidence": 0.9},
{"name": UNCATEGORIZED, "confidence": 0.8},
{"name": "产品手册", "confidence": 0.7},
],
"rewrite": "r",
"keywords": [],
},
ensure_ascii=False,
)
parser = _make_parser(response)
parsed = await parser.parse("q")
assert parsed.parse_failed is False
assert [c.name for c in parsed.categories] == ["产品手册"]
async def test_out_of_range_confidence_dropped(self):
"""confidence 越界(>1 或 <0)的条目被丢弃"""
response = json.dumps(
{
"categories": [
{"name": "技术文档", "confidence": 1.5},
{"name": "产品手册", "confidence": -0.2},
{"name": "财务行政", "confidence": 0.6},
],
"rewrite": "r",
"keywords": [],
},
ensure_ascii=False,
)
parser = _make_parser(response)
parsed = await parser.parse("q")
assert [c.name for c in parsed.categories] == ["财务行政"]
async def test_parse_new_fields_populated(self):
"""JSON 含 entities/intent/time_range → 正确填充字段"""
response = json.dumps(
{
"categories": [{"name": "技术文档", "confidence": 0.9}],
"rewrite": "如何设计微服务架构",
"keywords": ["架构"],
"entities": ["微服务", "Kubernetes"],
"intent": "操作",
"time_range": "2023年",
},
ensure_ascii=False,
)
parser = _make_parser(response)
parsed = await parser.parse("怎么做微服务部署")
assert parsed.parse_failed is False
assert parsed.entities == ["微服务", "Kubernetes"]
assert parsed.intent == "操作"
assert parsed.time_range == "2023年"
async def test_parse_new_fields_missing_defaults(self):
"""JSON 缺 entities/intent/time_range → 降级为默认值,parse_failed=False(向后兼容)"""
response = json.dumps(
{
"categories": [{"name": "技术文档", "confidence": 0.9}],
"rewrite": "如何设计架构",
"keywords": ["架构"],
},
ensure_ascii=False,
)
parser = _make_parser(response)
parsed = await parser.parse("怎么做架构设计")
assert parsed.parse_failed is False
assert parsed.entities == []
assert parsed.intent == ""
assert parsed.time_range == ""
async def test_parse_new_fields_wrong_type_defaults(self):
"""entities 非列表 / intent 非字符串 / time_range 非字符串 → 置空,parse_failed=False(容错)"""
response = json.dumps(
{
"categories": [{"name": "技术文档", "confidence": 0.9}],
"rewrite": "如何设计架构",
"keywords": ["架构"],
"entities": "微服务", # 非列表
"intent": 123, # 非字符串
"time_range": ["2023"], # 非字符串
},
ensure_ascii=False,
)
parser = _make_parser(response)
parsed = await parser.parse("怎么做架构设计")
assert parsed.parse_failed is False
assert parsed.entities == []
assert parsed.intent == ""
assert parsed.time_range == ""
class TestDecideRoute:
"""decide_route 纯函数的四个分支"""
def _parsed(self, categories: list[CategoryHit], parse_failed: bool = False) -> ParsedQuery:
return ParsedQuery(raw_query="q", rewrite="q", categories=categories, parse_failed=parse_failed)
def test_parse_failed_fallback(self):
"""解析失败 → 全库兜底,reason=parse_failed"""
decision = decide_route(self._parsed([], parse_failed=True), threshold=0.6, max_categories=3)
assert decision.fallback is True
assert decision.filter_categories is None
assert decision.reason == "parse_failed"
def test_low_confidence_fallback(self):
"""所有类目 confidence 低于阈值 → 全库兜底,reason=low_confidence"""
categories = [CategoryHit(name="技术文档", confidence=0.5), CategoryHit(name="产品手册", confidence=0.3)]
decision = decide_route(self._parsed(categories), threshold=0.6, max_categories=3)
assert decision.fallback is True
assert decision.filter_categories is None
assert decision.reason == "low_confidence"
def test_empty_categories_fallback(self):
"""无命中类目 → 全库兜底,reason=low_confidence"""
decision = decide_route(self._parsed([]), threshold=0.6, max_categories=3)
assert decision.fallback is True
assert decision.filter_categories is None
assert decision.reason == "low_confidence"
def test_too_many_categories_fallback(self):
"""过阈值类目数超过上限 → 全库兜底,reason=too_many_categories"""
categories = [
CategoryHit(name="技术文档", confidence=0.9),
CategoryHit(name="产品手册", confidence=0.8),
CategoryHit(name="财务行政", confidence=0.7),
]
decision = decide_route(self._parsed(categories), threshold=0.6, max_categories=2)
assert decision.fallback is True
assert decision.filter_categories is None
assert decision.reason == "too_many_categories"
def test_routed_sorted_by_confidence_desc(self):
"""正常路由 → 按 confidence 降序输出过滤类目,低于阈值的类目被剔除"""
categories = [
CategoryHit(name="产品手册", confidence=0.7),
CategoryHit(name="财务行政", confidence=0.5), # 低于阈值,应被剔除
CategoryHit(name="技术文档", confidence=0.9),
]
decision = decide_route(self._parsed(categories), threshold=0.6, max_categories=3)
assert decision.fallback is False
assert decision.filter_categories == ["技术文档", "产品手册"]
assert decision.reason == "routed"
class TestParseAndRoute:
"""parse_and_route 便捷方法(threshold/max_categories 取自 settings,默认 0.6/3"""
async def test_parse_and_route_routed(self):
response = json.dumps(
{
"categories": [{"name": "技术文档", "confidence": 0.9}],
"rewrite": "r",
"keywords": ["k"],
},
ensure_ascii=False,
)
parser = _make_parser(response)
decision = await parser.parse_and_route("架构设计文档在哪")
assert decision.fallback is False
assert decision.filter_categories == ["技术文档"]
assert decision.reason == "routed"
assert decision.parsed.raw_query == "架构设计文档在哪"
async def test_parse_and_route_fallback_on_parse_failure(self):
parser = _make_parser("不是 JSON")
decision = await parser.parse_and_route("q")
assert decision.fallback is True
assert decision.filter_categories is None
assert decision.reason == "parse_failed"