Files
huihuiSquare/digital-avatar-app/backend/tests/test_chat_orchestration.py
2026-07-23 17:21:49 +08:00

96 lines
3.2 KiB
Python

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