import unittest from types import SimpleNamespace from unittest.mock import Mock from fastapi import HTTPException from models import Avatar, User from routers.chat import _build_prompt, _match_standard_qa, _require_owned_avatar, _resolve_reply class ChatOrchestrationTests(unittest.TestCase): def setUp(self): self.avatar = SimpleNamespace( id="avatar-1", owner_id="huihui-user-1", config={ "replyStyle": "professional", "creativity": 50, "rigor": 80, "humor": 20, "responseLength": "medium", "systemPrompt": "不要编造政策。", }, ) self.qa = SimpleNamespace(question="公司地址?", answer="标准地址", enabled=True) self.disabled_qa = SimpleNamespace(question="公司地址?", answer="错误答案", enabled=False) def test_enabled_qa_wins_without_calling_model(self): fake_model = Mock() result = _resolve_reply( None, self.avatar, " 公司地址? ", [], qa_pairs=[self.disabled_qa, self.qa], search_fn=lambda *_args, **_kwargs: [], model_client=fake_model, ) self.assertEqual(result["source"], "qa") self.assertEqual(result["answer"], "标准地址") fake_model.assert_not_called() def test_knowledge_context_is_sent_to_qwen_after_qa_miss(self): fake_model = Mock(return_value="根据知识库内容回答") knowledge_hit = { "filename": "退款.md", "snippet": "知识库内容:七日内可申请退款。", "score": 0.92, } result = _resolve_reply( None, self.avatar, "退款规则", [], qa_pairs=[], search_fn=lambda *_args, **_kwargs: [knowledge_hit], model_client=fake_model, ) self.assertEqual(result["source"], "knowledge") self.assertIn("知识库内容", fake_model.call_args.kwargs["messages"][0]["content"]) def test_prompt_contains_personality_configuration(self): messages = _build_prompt(self.avatar, [], "你好", []) self.assertIn("严谨度", messages[0]["content"]) self.assertIn("不要编造政策", messages[0]["content"]) def test_chat_rejects_avatar_owned_by_another_user(self): class Query: def __init__(self, value): self.value = value def filter(self, *_args, **_kwargs): return self def first(self): return self.value self_avatar = self.avatar class DB: avatar = self_avatar def query(self, model): return Query( self.avatar if model is Avatar else SimpleNamespace(huihui_user_id="huihui-user-2") ) db = DB() with self.assertRaises(HTTPException) as caught: _require_owned_avatar(db, self.avatar.id, "Bearer other-token") self.assertEqual(caught.exception.status_code, 403) if __name__ == "__main__": unittest.main()