88 lines
3.0 KiB
Python
88 lines
3.0 KiB
Python
import pytest
|
|
|
|
from app.core.contracts import RequestContext
|
|
from app.core.errors import RecoverableAgentError
|
|
from app.core.knowledge_contracts import KnowledgeHit, KnowledgeQuery
|
|
from app.service.knowledge_config import KnowledgeRuntimeConfig
|
|
from app.service.knowledge_retrieval_service import KnowledgeRetrievalService
|
|
|
|
# ruff: noqa: E501
|
|
|
|
|
|
class FakeEmbedder:
|
|
async def embed(self, text: str) -> list[float]:
|
|
assert text == "开户"
|
|
return [0.1] * 1024
|
|
|
|
|
|
class FakeVectorStore:
|
|
def __init__(self) -> None:
|
|
self.calls: list[tuple[str, int]] = []
|
|
|
|
async def search(self, collection: str, vector: list[float], top_k: int) -> list[dict[str, object]]:
|
|
assert len(vector) == 1024
|
|
self.calls.append((collection, top_k))
|
|
return []
|
|
|
|
|
|
class FakeAuthority:
|
|
async def filter_published(self, hits: tuple[object, ...]) -> list[object]:
|
|
return []
|
|
|
|
async def search_keyword(self, query: object, collections: tuple[str, ...], top_k: int) -> list[object]:
|
|
return []
|
|
|
|
|
|
class BrokenVectorStore:
|
|
async def search(self, collection: str, vector: list[float], top_k: int) -> list[dict[str, object]]:
|
|
raise RecoverableAgentError("知识检索不可用")
|
|
|
|
|
|
class FallbackAuthority:
|
|
def __init__(self) -> None:
|
|
self.calls: list[tuple[tuple[str, ...], int]] = []
|
|
|
|
async def filter_published(self, hits: tuple[object, ...]) -> list[object]:
|
|
return []
|
|
|
|
async def search_keyword(self, query: KnowledgeQuery, collections: tuple[str, ...], top_k: int) -> list[KnowledgeHit]:
|
|
self.calls.append((collections, top_k))
|
|
return [KnowledgeHit(
|
|
knowledge_id="101", collection="fin_policy_collection", snippet="确认规则",
|
|
answer="工作日确认", score=1.0,
|
|
)]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_search_uses_faq_collection_for_faq_only() -> None:
|
|
vector_store = FakeVectorStore()
|
|
service = KnowledgeRetrievalService(
|
|
FakeEmbedder(), vector_store, KnowledgeRuntimeConfig(), FakeAuthority()
|
|
)
|
|
|
|
result = await service.search(
|
|
KnowledgeQuery(query="开户", intents=("faq",)),
|
|
RequestContext(user_id="visitor-1", trace_id="trace", roles=("visitor",), data_scope="public"),
|
|
)
|
|
|
|
assert vector_store.calls == [("fin_faq_collection", 3)]
|
|
assert result.searched_collections == ("fin_faq_collection",)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_milvus_failure_falls_back_to_published_active_unexpired_knowledge() -> None:
|
|
authority = FallbackAuthority()
|
|
service = KnowledgeRetrievalService(
|
|
FakeEmbedder(), BrokenVectorStore(), KnowledgeRuntimeConfig(), authority
|
|
)
|
|
|
|
result = await service.search(
|
|
KnowledgeQuery(query="开户", intents=("policy_explain",)),
|
|
RequestContext(user_id="visitor-1", trace_id="trace", roles=("visitor",), data_scope="public"),
|
|
)
|
|
|
|
assert authority.calls == [(("fin_policy_collection",), 5)]
|
|
assert result.degraded is True
|
|
assert result.degradation_reason == "milvus_unavailable"
|
|
assert result.hits[0].answer == "工作日确认"
|