chore: update gitignore; feat: 新增customer_agent业务模块与api路由
This commit is contained in:
@@ -0,0 +1,102 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user