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