103 lines
3.9 KiB
Python
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()
|