95 lines
3.9 KiB
Python
95 lines
3.9 KiB
Python
import unittest
|
|
from unittest.mock import AsyncMock
|
|
|
|
from rag.intent import Intent
|
|
from service.customer_agent.chat import AnonymousCustomerAgent, QueryTooLongError
|
|
|
|
|
|
class AnonymousAgentTests(unittest.IsolatedAsyncioTestCase):
|
|
def _build(self, intent, sources=None, answer="LLM answer"):
|
|
context = AsyncMock()
|
|
rag = AsyncMock(return_value=sources or [])
|
|
recognize = AsyncMock(return_value=intent)
|
|
generate = AsyncMock(return_value=answer)
|
|
audit = AsyncMock()
|
|
config = {
|
|
"agent.customer.template.guide_purchase": "请前往开户页面办理",
|
|
"agent.customer.template.guide_advisor": "如需推荐,请联系投顾",
|
|
"agent.customer.template.fallback_human": "请转人工客服",
|
|
}
|
|
service = AnonymousCustomerAgent(
|
|
context=context,
|
|
rag_retrieve=rag,
|
|
intent_recognize=recognize,
|
|
generate_answer=generate,
|
|
audit_writer=audit,
|
|
config_getter=config.get,
|
|
)
|
|
return service, context, rag, recognize, generate, audit
|
|
|
|
async def test_knowledge_question_uses_rag_without_customer_memory(self):
|
|
service, context, rag, _, generate, _ = self._build(
|
|
Intent.KNOWLEDGE_QA,
|
|
sources=[{"doc_id": "d1", "chunk_text": "基金正文", "score": 0.9}],
|
|
)
|
|
|
|
result = await service.handle("s1", "基金是什么", trace_id="trace-1")
|
|
|
|
rag.assert_awaited_once_with("基金是什么", None)
|
|
generate.assert_awaited_once()
|
|
self.assertEqual(result["sources"][0]["doc_id"], "d1")
|
|
self.assertEqual(result["trace_id"], "trace-1")
|
|
self.assertEqual(context.append.await_count, 2)
|
|
|
|
async def test_recommendation_returns_fixed_advisor_guidance_without_llm(self):
|
|
service, _, rag, _, generate, _ = self._build(Intent.WANT_ADVISOR)
|
|
|
|
result = await service.handle("s1", "给我推荐一只基金", trace_id="trace-2")
|
|
|
|
self.assertEqual(result["answer"], "如需推荐,请联系投顾")
|
|
rag.assert_not_awaited()
|
|
generate.assert_not_awaited()
|
|
self.assertEqual(result["sources"], [])
|
|
|
|
async def test_sensitive_input_is_audited_and_no_private_data_is_loaded(self):
|
|
service, _, rag, _, generate, audit = self._build(Intent.KNOWLEDGE_QA, sources=[])
|
|
|
|
await service.handle("s1", "我的手机号是13812345678", trace_id="trace-3")
|
|
|
|
audit.assert_awaited_once()
|
|
self.assertEqual(audit.await_args.kwargs["action"], "anon_sensitive_input")
|
|
self.assertEqual(audit.await_args.kwargs["trace_id"], "trace-3")
|
|
rag.assert_awaited_once_with("我的手机号是13812345678", None)
|
|
generate.assert_not_awaited()
|
|
|
|
async def test_query_over_2000_characters_is_rejected(self):
|
|
service, *_ = self._build(Intent.KNOWLEDGE_QA)
|
|
|
|
with self.assertRaises(QueryTooLongError):
|
|
await service.handle("s1", "x" * 2001, trace_id="trace-4")
|
|
|
|
async def test_milvus_failure_returns_human_fallback(self):
|
|
service, _, rag, _, generate, _ = self._build(Intent.KNOWLEDGE_QA)
|
|
rag.side_effect = TimeoutError("Milvus down")
|
|
|
|
result = await service.handle("s1", "基金是什么", trace_id="trace-5")
|
|
|
|
self.assertEqual(result["answer"], "请转人工客服")
|
|
generate.assert_not_awaited()
|
|
|
|
async def test_llm_failure_returns_human_fallback_without_fake_sources(self):
|
|
service, _, _, _, generate, _ = self._build(
|
|
Intent.KNOWLEDGE_QA,
|
|
sources=[{"doc_id": "d1", "chunk_text": "正文", "score": 0.8}],
|
|
)
|
|
generate.side_effect = RuntimeError("LLM down")
|
|
|
|
result = await service.handle("s1", "基金是什么", trace_id="trace-6")
|
|
|
|
self.assertEqual(result["answer"], "请转人工客服")
|
|
self.assertEqual(result["sources"][0]["doc_id"], "d1")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|