feat: add avatar chat and knowledge workflow
This commit is contained in:
1
digital-avatar-app/backend/tests/__init__.py
Normal file
1
digital-avatar-app/backend/tests/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
|
||||
95
digital-avatar-app/backend/tests/test_chat_orchestration.py
Normal file
95
digital-avatar-app/backend/tests/test_chat_orchestration.py
Normal file
@@ -0,0 +1,95 @@
|
||||
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()
|
||||
32
digital-avatar-app/backend/tests/test_embeddings.py
Normal file
32
digital-avatar-app/backend/tests/test_embeddings.py
Normal file
@@ -0,0 +1,32 @@
|
||||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
|
||||
import embeddings
|
||||
|
||||
|
||||
class TextExtractionTests(unittest.TestCase):
|
||||
def write_text(self, suffix, content):
|
||||
handle = tempfile.NamedTemporaryFile(suffix=suffix, delete=False)
|
||||
handle.close()
|
||||
self.addCleanup(lambda: os.path.exists(handle.name) and os.unlink(handle.name))
|
||||
with open(handle.name, "w", encoding="utf-8") as stream:
|
||||
stream.write(content)
|
||||
return handle.name
|
||||
|
||||
def test_extracts_utf8_markdown(self):
|
||||
path = self.write_text(".md", "# 退款规则\n\n七日内可申请退款。")
|
||||
self.assertEqual(embeddings.extract_text(path, ".md"), "# 退款规则\n\n七日内可申请退款。")
|
||||
|
||||
def test_extracts_utf8_text(self):
|
||||
path = self.write_text(".txt", "客服热线:400-123-4567")
|
||||
self.assertEqual(embeddings.extract_text(path, ".txt"), "客服热线:400-123-4567")
|
||||
|
||||
def test_rejects_unsupported_extension(self):
|
||||
path = self.write_text(".csv", "not supported")
|
||||
with self.assertRaises(ValueError):
|
||||
embeddings.extract_text(path, ".csv")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user