Files
ai-video-fullstack/backend/tests/test_workflow_router.py
2026-08-03 10:55:57 +08:00

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