Files
Mutual_Fund/tests/test_retrieve.py
T

103 lines
3.9 KiB
Python

import unittest
from unittest.mock import AsyncMock
from rag.retrieve import rag_retrieve, retrieve_candidates, retrieve_with_status
class RetrievalTests(unittest.IsolatedAsyncioTestCase):
async def test_reads_topk_threshold_and_searches_three_business_collections(self):
milvus = AsyncMock()
milvus.search.return_value = [
[{"id": "chunk-1", "distance": 0.9, "entity": {"doc_id": "d1"}}]
]
embedder = AsyncMock(return_value=[[0.0] * 768])
values = {
"agent.customer.rag.topk.faq": "3",
"agent.customer.rag.threshold.faq": "0.75",
"agent.customer.rag.topk.funddoc": "5",
"agent.customer.rag.threshold.funddoc": "0.7",
"agent.customer.rag.topk.policy": "5",
"agent.customer.rag.threshold.policy": "0.7",
}
await retrieve_candidates(
"基金是什么", None, milvus_client=milvus, embedder=embedder,
config_getter=values.get,
)
self.assertEqual(milvus.search.await_count, 3)
calls = {call.kwargs["collection_name"]: call.kwargs for call in milvus.search.await_args_list}
self.assertEqual(calls["fin_faq"]["limit"], 3)
self.assertEqual(calls["fin_faq"]["filter"], "")
self.assertEqual(calls["fin_faq"]["data"], [[0.0] * 768])
async def test_customer_memory_is_searched_only_for_a_customer(self):
milvus = AsyncMock()
milvus.search.return_value = [[]]
embedder = AsyncMock(return_value=[[0.0] * 768])
values = {"agent.customer.rag.topk.memory": "5", "agent.customer.rag.threshold.memory": "0.6"}
await retrieve_candidates(
"风险", "customer-7", milvus_client=milvus, embedder=embedder,
config_getter=values.get,
)
memory_call = milvus.search.await_args_list[-1].kwargs
self.assertEqual(memory_call["collection_name"], "customer_memory")
self.assertEqual(memory_call["filter"], 'customer_id == "customer-7"')
async def test_public_retrieve_returns_milvus_sources_and_filters_low_scores(self):
milvus = AsyncMock()
milvus.search.side_effect = [
[[
{"id": "c1", "distance": 0.90, "entity": {
"doc_id": "doc-1", "title": "FAQ", "section_title": "",
"text": "基金正文",
}},
{"id": "c2", "distance": 0.50, "entity": {
"doc_id": "doc-low", "title": "低分", "text": "不应返回",
}},
]],
[[]],
[[]],
]
embedder = AsyncMock(return_value=[[0.0] * 768])
sources = await rag_retrieve(
"基金", None, milvus_client=milvus, embedder=embedder,
config_getter={}.get,
)
self.assertEqual(sources, [{
"doc_id": "doc-1", "title": "FAQ", "section_title": None,
"chunk_text": "基金正文", "score": 0.90,
}])
async def test_milvus_failure_returns_empty_sources_without_fake_hit(self):
milvus = AsyncMock()
milvus.search.side_effect = TimeoutError("Milvus timeout")
embedder = AsyncMock(return_value=[[0.0] * 768])
result = await retrieve_with_status(
"基金", None, milvus_client=milvus, embedder=embedder,
config_getter={}.get,
)
self.assertEqual(result, {"status": "milvus_unavailable", "sources": []})
async def test_embedding_failure_returns_failed_status_without_mysql_fallback(self):
milvus = AsyncMock()
embedder = AsyncMock(side_effect=RuntimeError("embedding down"))
result = await retrieve_with_status(
"基金", None, milvus_client=milvus, embedder=embedder,
config_getter={}.get,
)
self.assertEqual(result, {"status": "embedding_failed", "sources": []})
milvus.search.assert_not_awaited()
if __name__ == "__main__":
unittest.main()