92 lines
3.6 KiB
Python
92 lines
3.6 KiB
Python
# -*- coding: utf-8 -*-
|
|
from __future__ import annotations
|
|
|
|
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
from unittest import mock
|
|
|
|
import dsh_agent
|
|
|
|
|
|
class DeepSeekDeskTest(unittest.TestCase):
|
|
def setUp(self):
|
|
self.temp = tempfile.TemporaryDirectory()
|
|
self.desk = dsh_agent.DeepSeekDesk(Path(self.temp.name))
|
|
|
|
def tearDown(self):
|
|
self.desk.close()
|
|
self.temp.cleanup()
|
|
|
|
def test_create_session_and_reuse_session_id(self):
|
|
first = self.desk.create_session()
|
|
self.assertTrue(first.id.startswith("desk-"))
|
|
self.assertEqual(first.title, "新对话")
|
|
with mock.patch.object(
|
|
self.desk,
|
|
"_run_model",
|
|
return_value=dsh_agent.RunOutcome(
|
|
session_id=first.id,
|
|
text="已经帮您拟好回复。",
|
|
backend=dsh_agent.BACKEND_COMPAT,
|
|
finish_reason="completed",
|
|
),
|
|
):
|
|
result = self.desk.run("帮我回客户说下午再联系", session_id=first.id)
|
|
self.assertEqual(result.session_id, first.id)
|
|
stored = self.desk.get_session(first.id)
|
|
self.assertEqual(stored.title, "帮我回客户说下午再联系")
|
|
self.assertEqual(len(stored.messages), 2)
|
|
self.assertEqual(stored.messages[0]["role"], "user")
|
|
self.assertEqual(stored.messages[1]["content"], "已经帮您拟好回复。")
|
|
|
|
def test_independent_tasks_use_new_session_ids(self):
|
|
one = self.desk.create_session()
|
|
two = self.desk.create_session()
|
|
self.assertNotEqual(one.id, two.id)
|
|
|
|
def test_compat_backend_sends_system_prompt_and_history(self):
|
|
session = self.desk.create_session()
|
|
captured = {}
|
|
|
|
def fake_completion(messages, tools=None):
|
|
captured["messages"] = messages
|
|
return {"content": "草稿已写好"}
|
|
|
|
with (
|
|
mock.patch.object(self.desk, "_start_harness", side_effect=RuntimeError("no runtime")),
|
|
mock.patch("ai_chat._chat_completion", side_effect=fake_completion),
|
|
):
|
|
first = self.desk.run("客户问血糖高怎么办", session_id=session.id)
|
|
second = self.desk.run("那饮食呢", session_id=session.id)
|
|
self.assertEqual(first.backend, dsh_agent.BACKEND_COMPAT)
|
|
self.assertEqual(second.text, "草稿已写好")
|
|
roles = [item["role"] for item in captured["messages"]]
|
|
self.assertEqual(roles[0], "system")
|
|
self.assertIn("DeepSeek Agent", captured["messages"][0]["content"])
|
|
self.assertEqual(roles[-1], "user")
|
|
self.assertEqual(captured["messages"][-1]["content"], "那饮食呢")
|
|
self.assertTrue(any(item.get("content") == "客户问血糖高怎么办" for item in captured["messages"]))
|
|
|
|
def test_sdk_backend_reuses_harness_run(self):
|
|
session = self.desk.create_session()
|
|
harness = mock.Mock()
|
|
harness.run.return_value = mock.Mock(
|
|
final_response="SDK 回复",
|
|
finish_reason="completed",
|
|
)
|
|
with mock.patch.object(self.desk, "_start_harness", return_value=harness):
|
|
result = self.desk.run("检查这段话术", session_id=session.id)
|
|
self.assertEqual(result.backend, dsh_agent.BACKEND_SDK)
|
|
self.assertEqual(result.text, "SDK 回复")
|
|
harness.run.assert_called_once()
|
|
self.assertEqual(harness.run.call_args.kwargs["session_id"], session.id)
|
|
|
|
def test_empty_prompt_is_rejected(self):
|
|
with self.assertRaises(ValueError):
|
|
self.desk.run(" ")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|