292 lines
10 KiB
Python
292 lines
10 KiB
Python
"""`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 == "场内基金申购费率说明正文。"
|