feat: add avatar chat and knowledge workflow

This commit is contained in:
stefanfeng
2026-07-23 17:21:49 +08:00
parent 501f548bcc
commit 2d9f26a6f0
15 changed files with 3635 additions and 0 deletions

View File

@@ -0,0 +1 @@

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

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