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()