149 lines
5.3 KiB
Python
149 lines
5.3 KiB
Python
from __future__ import annotations
|
|
|
|
import unittest
|
|
from types import SimpleNamespace
|
|
from unittest.mock import patch
|
|
|
|
from models import AssistantConfig
|
|
from services.workflow.models import RouteStatus
|
|
from services.workflow_router import WorkflowLLMRouter
|
|
|
|
|
|
class WorkflowLLMRouterTest(unittest.IsolatedAsyncioTestCase):
|
|
async def test_uses_required_tool_choice_without_developer_messages(self):
|
|
requests = []
|
|
|
|
class FakeCompletions:
|
|
async def create(self, **kwargs):
|
|
requests.append(kwargs)
|
|
return SimpleNamespace(
|
|
choices=[
|
|
SimpleNamespace(
|
|
message=SimpleNamespace(
|
|
tool_calls=[
|
|
SimpleNamespace(
|
|
function=SimpleNamespace(name="goto_age", arguments="{}")
|
|
)
|
|
]
|
|
)
|
|
)
|
|
]
|
|
)
|
|
|
|
class FakeClient:
|
|
def __init__(self, **_kwargs):
|
|
self.chat = SimpleNamespace(completions=FakeCompletions())
|
|
self.closed = False
|
|
|
|
async def close(self):
|
|
self.closed = True
|
|
|
|
cfg = AssistantConfig(
|
|
type="workflow",
|
|
model="deepseek-chat",
|
|
llm_api_key="secret",
|
|
llm_base_url="https://llm.test/v1",
|
|
)
|
|
router = WorkflowLLMRouter(cfg)
|
|
edges = [
|
|
{
|
|
"id": "age",
|
|
"data": {"condition": "用户已经回答姓名", "priority": 10},
|
|
}
|
|
]
|
|
|
|
with patch("services.workflow_router.AsyncOpenAI", FakeClient):
|
|
selected = await router.select_edge(
|
|
node_name="询问姓名",
|
|
node_prompt="询问用户姓名",
|
|
edges=edges,
|
|
history=[{"role": "user", "message": "我叫李白"}],
|
|
variables={"customer_type": "new"},
|
|
edge_name=lambda _edge: "goto_age",
|
|
edge_description=lambda _edge: "用户已经回答姓名",
|
|
)
|
|
|
|
self.assertEqual(selected.status, RouteStatus.MATCHED)
|
|
self.assertEqual(selected.function_name, "goto_age")
|
|
self.assertEqual(requests[0]["tool_choice"], "required")
|
|
self.assertEqual(
|
|
[message["role"] for message in requests[0]["messages"]],
|
|
["system", "user"],
|
|
)
|
|
self.assertNotIn("developer", str(requests[0]["messages"]))
|
|
|
|
async def test_routes_with_the_current_multimodal_user_message(self):
|
|
requests = []
|
|
|
|
class FakeCompletions:
|
|
async def create(self, **kwargs):
|
|
requests.append(kwargs)
|
|
return SimpleNamespace(
|
|
choices=[
|
|
SimpleNamespace(
|
|
message=SimpleNamespace(
|
|
tool_calls=[
|
|
SimpleNamespace(
|
|
function=SimpleNamespace(
|
|
name="goto_confirm",
|
|
arguments="{}",
|
|
)
|
|
)
|
|
]
|
|
)
|
|
)
|
|
]
|
|
)
|
|
|
|
class FakeClient:
|
|
def __init__(self, **_kwargs):
|
|
self.chat = SimpleNamespace(completions=FakeCompletions())
|
|
|
|
async def close(self):
|
|
return None
|
|
|
|
router = WorkflowLLMRouter(
|
|
AssistantConfig(
|
|
type="workflow",
|
|
model="visual-model",
|
|
llm_api_key="secret",
|
|
llm_base_url="https://llm.test/v1",
|
|
)
|
|
)
|
|
image_message = {
|
|
"role": "user",
|
|
"content": [
|
|
{"type": "text", "text": "请检查车牌照片"},
|
|
{
|
|
"type": "image_url",
|
|
"image_url": {"url": "data:image/jpeg;base64,AA=="},
|
|
},
|
|
],
|
|
}
|
|
|
|
with patch("services.workflow_router.AsyncOpenAI", FakeClient):
|
|
selected = await router.select_edge(
|
|
node_name="采集车牌",
|
|
node_prompt="确认车牌照片是否清晰",
|
|
edges=[{"id": "confirm", "data": {"condition": "车牌清晰"}}],
|
|
history=[
|
|
{"role": "user", "message": "之前的消息"},
|
|
{"role": "user", "message": "请检查车牌照片"},
|
|
],
|
|
variables={},
|
|
edge_name=lambda _edge: "goto_confirm",
|
|
edge_description=lambda _edge: "车牌清晰",
|
|
current_user_message=image_message,
|
|
)
|
|
|
|
self.assertEqual(selected.status, RouteStatus.MATCHED)
|
|
content = requests[0]["messages"][1]["content"]
|
|
self.assertIsInstance(content, list)
|
|
self.assertEqual(content[-1], image_message["content"][-1])
|
|
self.assertIn("之前的消息", content[0]["text"])
|
|
self.assertEqual(content[0]["text"].count("请检查车牌照片"), 0)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|