"""Classifier 文档分类的单元测试(mock OllamaClient,不真实联网)""" import json from app.core.classifier import Classifier 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_classifier(response: str) -> Classifier: return Classifier(ollama=FakeOllama(response), taxonomy=_taxonomy()) # type: ignore[arg-type] class TestClassify: """Classifier.classify 的正常解析与容错""" async def test_classify_valid_json(self): """正常 JSON 响应 → 正确解析主类目/tags/confidence tags 清洗规则:剔除主类目名、最多保留 3 个。 """ response = json.dumps( { "main_category": "技术文档", # "技术文档" 与主类目重复应被剔除;超出 3 个应被截断 "tags": ["架构", "技术文档", "API", "部署", "运维"], "confidence": 0.9, }, ensure_ascii=False, ) classifier = _make_classifier(response) result = await classifier.classify("本文介绍系统架构设计。", title="架构文档") assert result.main_category == "技术文档" assert result.tags == ["架构", "API", "部署"] assert result.confidence == 0.9 async def test_classify_uses_json_mode_and_prompt_contains_taxonomy(self): """以 json_mode 调用 LLM,prompt 包含全部类目名与 uncategorized 用途说明""" classifier = _make_classifier('{"main_category": "财务行政", "tags": [], "confidence": 0.8}') result = await classifier.classify("报销流程说明", title="") assert result.main_category == "财务行政" ollama = classifier.ollama assert ollama.calls[0]["json_mode"] is True prompt = ollama.calls[0]["prompt"] assert "技术文档" in prompt and "产品手册" in prompt and "财务行政" in prompt assert UNCATEGORIZED in prompt assert "跨多个类目" in prompt assert "报销流程说明" in prompt async def test_classify_json_with_surrounding_noise(self): """响应带前后多余文本 → 正则提取 {...} 块成功""" response = ( "好的,分类结果如下:\n" '{"main_category": "产品手册", "tags": ["使用说明"], "confidence": 0.75}\n' "以上是分类结果。" ) classifier = _make_classifier(response) result = await classifier.classify("产品功能使用说明") assert result.main_category == "产品手册" assert result.tags == ["使用说明"] assert result.confidence == 0.75 async def test_classify_non_json_falls_back_to_uncategorized(self): """完全非 JSON 响应 → 归 uncategorized,tags 为空,confidence=0.0""" classifier = _make_classifier("抱歉,我无法完成分类。") result = await classifier.classify("一段总结") assert result.main_category == UNCATEGORIZED assert result.tags == [] assert result.confidence == 0.0 async def test_classify_unknown_category_falls_back(self): """类目名不在 taxonomy → 归 uncategorized,confidence=0.0""" classifier = _make_classifier('{"main_category": "不存在的类目", "tags": ["x"], "confidence": 0.9}') result = await classifier.classify("一段总结") assert result.main_category == UNCATEGORIZED assert result.tags == [] assert result.confidence == 0.0 async def test_classify_low_confidence_soft_recall(self): """置信度低于阈值(默认 0.6)→ 主类目归 uncategorized,候选类目名保留进 tags,confidence 保留原值""" response = json.dumps( {"main_category": "产品手册", "tags": ["手册"], "confidence": 0.4}, ensure_ascii=False, ) classifier = _make_classifier(response) result = await classifier.classify("介于产品和运营之间的内容") assert result.main_category == UNCATEGORIZED # 候选类目名插入 tags 首位,供检索侧软召回 assert result.tags == ["产品手册", "手册"] assert result.confidence == 0.4 async def test_classify_low_confidence_uncategorized_not_duplicated_in_tags(self): """低置信且候选本身为 uncategorized → 不把 uncategorized 塞进 tags""" classifier = _make_classifier('{"main_category": "uncategorized", "tags": [], "confidence": 0.3}') result = await classifier.classify("杂项内容") assert result.main_category == UNCATEGORIZED assert result.tags == [] assert result.confidence == 0.3