"""`KnowledgeRetrievalService` 单测(Task 6):路由、白名单、维度失败关闭、降级与可见性过滤。 全部使用 fake client / fake embedder / fake session —— **不连真 Milvus、不连 MySQL**。 """ from __future__ import annotations from datetime import UTC, datetime, timedelta from typing import Any import pytest from app.core.errors import ForbiddenAgentError, RecoverableAgentError, ValidationAgentError from app.core.knowledge_contracts import VECTOR_DIM, KnowledgeQuery from app.service.knowledge_retrieval_service import DEFAULT_ROUTES, KnowledgeRetrievalService # --- fakes ------------------------------------------------------------------- class FakeEmbedding: def __init__(self, vector: list[float]) -> None: self.vector = vector class FakeEmbedder: def __init__(self, vector: list[float] | None = None) -> None: self.vector = vector if vector is not None else [0.01] * VECTOR_DIM self.calls: list[tuple[list[Any], str]] = [] async def embed( self, endpoints: list[Any], text_value: str, *, max_attempts: int = 2 ) -> FakeEmbedding: self.calls.append((endpoints, text_value)) return FakeEmbedding(list(self.vector)) class FakeClient: """记录调用参数并按集合返回预置命中行;`fail=True` 时模拟 Milvus 不可用。""" def __init__(self, hits: list[dict[str, Any]] | None = None, *, fail: bool = False) -> None: self.hits = hits or [] self.fail = fail self.calls: list[dict[str, Any]] = [] async def search( self, *, collection: str, vector: list[float], top_k: int, **kwargs: Any ) -> list[dict[str, Any]]: self.calls.append({"collection": collection, "top_k": top_k, "vector": vector}) if self.fail: raise RecoverableAgentError("知识向量检索失败:milvus down") return [dict(row) for row in self.hits] class FakeResult: def __init__(self, rows: list[dict[str, Any]]) -> None: self._rows = rows def mappings(self) -> FakeResult: return self def all(self) -> list[dict[str, Any]]: return self._rows class FakeSession: """按 SQL 文本分发:命中回表查询返回 `hits_rows`,降级 LIKE 返回 `degraded_rows`。""" def __init__( self, hits_rows: list[dict[str, Any]], degraded_rows: list[dict[str, Any]] ) -> None: self.hits_rows = hits_rows self.degraded_rows = degraded_rows self.statements: list[str] = [] self.params: list[dict[str, Any]] = [] async def __aenter__(self) -> FakeSession: return self async def __aexit__(self, *args: object) -> None: return None async def execute(self, statement: Any, params: dict[str, Any] | None = None) -> FakeResult: sql = str(statement) self.statements.append(sql) self.params.append(params or {}) if "LIKE :pattern" in sql: return FakeResult(self.degraded_rows) return FakeResult(self.hits_rows) class SessionFactory: def __init__(self, session: FakeSession) -> None: self.session = session def __call__(self) -> FakeSession: return self.session def metadata_row(**overrides: Any) -> dict[str, Any]: row: dict[str, Any] = { "id": 101, "milvus_collection": "fin_faq_collection", "title": "场内基金申购费率", "version": "v1", "tags": '["费率","申购"]', "content_text": "场内基金申购费率说明正文。", "effective_date": None, "expire_date": None, "review_status": "published", "status": "active", } row.update(overrides) return row def service( *, hits: list[dict[str, Any]] | None = None, hits_rows: list[dict[str, Any]] | None = None, degraded_rows: list[dict[str, Any]] | None = None, fail: bool = False, vector: list[float] | None = None, ) -> tuple[KnowledgeRetrievalService, FakeClient, FakeEmbedder, FakeSession]: client = FakeClient(hits, fail=fail) embedder = FakeEmbedder(vector) session = FakeSession(hits_rows or [], degraded_rows or []) svc = KnowledgeRetrievalService( client, embedder=embedder, session_factory=SessionFactory(session) ) return svc, client, embedder, session def query(intents: tuple[str, ...] = ("faq",), top_k: int | None = None) -> KnowledgeQuery: """`top_k=None` 时用契约默认 5,便于断言路由默认 top_k 仍然生效(取二者较大值)。""" payload: dict[str, Any] = {"query": "申购费率是多少", "intents": intents} if top_k is not None: payload["top_k"] = top_k return KnowledgeQuery(**payload) # --- ① 路由 ------------------------------------------------------------------ async def test_faq_intent_routes_to_faq_collection_with_top_k_three() -> None: hits = [ { "knowledge_id": "101", "title": "场内基金申购费率", "snippet": "正文", "tags": "费率,申购", "version": "v1", "score": 0.83, } ] svc, client, _, _ = service(hits=hits, hits_rows=[metadata_row()]) result = await svc.search(query(), embedding_endpoints=[object()]) assert client.calls[0]["collection"] == "fin_faq_collection" assert client.calls[0]["top_k"] == DEFAULT_ROUTES["faq"][1] == 3 assert result.searched_collections == ("fin_faq_collection",) assert result.degraded is False assert result.hits[0].knowledge_id == "101" assert result.hits[0].collection == "fin_faq_collection" assert result.hits[0].tags == ("费率", "申购") assert result.hits[0].score == pytest.approx(0.83) async def test_sparse_intent_label_falls_back_to_collection_name() -> None: """`intent` 是稀疏标签:缺字段与空串都必须按集合名推断,而不是只判 `is None`。""" hits = [ {"knowledge_id": "101", "snippet": "无标签", "score": 0.9}, {"knowledge_id": "102", "snippet": "空串标签", "intent": "", "score": 0.8}, {"knowledge_id": "103", "snippet": "显式标签", "intent": "chitchat", "score": 0.7}, ] rows = [metadata_row(id=101), metadata_row(id=102), metadata_row(id=103)] svc, _, _, _ = service(hits=hits, hits_rows=rows) result = await svc.search(query(), embedding_endpoints=[object()]) assert [hit.intent for hit in result.hits] == ["faq", "faq", "chitchat"] # --- ② 白名单 ---------------------------------------------------------------- async def test_collection_outside_allowlist_rejected_before_any_network_call() -> None: svc, client, embedder, _ = service() svc.config.routes["evil"] = ("sanguo_faq", 3) with pytest.raises(ForbiddenAgentError): await svc.search(query(intents=("evil",)), embedding_endpoints=[object()]) assert client.calls == [] # 未发生任何 Milvus 调用 assert embedder.calls == [] # 也未发生向量化调用 async def test_unknown_intent_rejected_without_network_call() -> None: svc, client, embedder, _ = service() with pytest.raises(ForbiddenAgentError): await svc.search(query(intents=("unknown_intent",)), embedding_endpoints=[object()]) assert client.calls == [] assert embedder.calls == [] # --- ③ 维度失败关闭 ---------------------------------------------------------- async def test_dimension_mismatch_fails_closed_instead_of_empty_result() -> None: svc, client, _, _ = service(vector=[0.1] * 768) with pytest.raises(ValidationAgentError): await svc.search(query(), embedding_endpoints=[object()]) assert client.calls == [] # 维度校验先于任何检索调用 def test_assert_vector_dim_accepts_exact_dimension() -> None: svc, _, _, _ = service() svc.assert_vector_dim([0.0] * VECTOR_DIM) # 不抛异常 # --- ④ 降级 ------------------------------------------------------------------ async def test_milvus_failure_degrades_to_mysql_like_with_flag() -> None: degraded_rows = [metadata_row()] svc, client, _, session = service(degraded_rows=degraded_rows, fail=True) result = await svc.search(query(intents=("faq",)), embedding_endpoints=[object()]) assert client.calls[0]["collection"] == "fin_faq_collection" # 确实尝试过 Milvus assert result.degraded is True assert result.degradation_reason assert result.hits and result.hits[0].knowledge_id == "101" assert any("LIKE :pattern" in sql for sql in session.statements) # 降级 SQL 仍限定白名单集合 like_params = [ params for sql, params in zip(session.statements, session.params, strict=True) if "LIKE :pattern" in sql ] assert set(like_params[0]["collections"]) <= { "fin_faq_collection", "fin_product_collection", "fin_policy_collection", } # --- ⑤ 降级路径的可见性过滤 -------------------------------------------------- async def test_degraded_path_still_applies_published_and_effective_window() -> None: today = datetime.now(UTC).date() rows = [ metadata_row(id=201), # 可见 metadata_row(id=202, review_status="pending"), # 未发布 metadata_row(id=203, status="inactive"), # 已下线 metadata_row(id=204, effective_date=today + timedelta(days=30)), # 尚未生效 metadata_row(id=205, expire_date=today - timedelta(days=1)), # 已过期 metadata_row(id=206, expire_date=today), # 当天到期视为失效 metadata_row(id=207, effective_date=today - timedelta(days=1)), # 已生效 ] svc, _, _, _ = service(degraded_rows=rows, fail=True) result = await svc.search(query(), embedding_endpoints=[object()]) assert result.degraded is True assert [hit.knowledge_id for hit in result.hits] == ["201", "207"] async def test_vector_hits_also_pass_published_and_effective_window() -> None: today = datetime.now(UTC).date() hits = [ {"knowledge_id": "101", "snippet": "向量正文", "score": 0.9}, {"knowledge_id": "102", "snippet": "过期正文", "score": 0.8}, {"knowledge_id": "103", "snippet": "外部集合", "score": 0.7}, ] rows = [ metadata_row(id=101), metadata_row(id=102, expire_date=today - timedelta(days=1)), metadata_row(id=103, milvus_collection="sanguo_faq"), ] svc, _, _, _ = service(hits=hits, hits_rows=rows) result = await svc.search(query(), embedding_endpoints=[object()]) assert result.degraded is False assert [hit.knowledge_id for hit in result.hits] == ["101"] assert result.hits[0].snippet == "场内基金申购费率说明正文。"