merge: add workflow vision configuration
This commit is contained in:
@@ -470,6 +470,81 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
|
||||
|
||||
class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_agent_vision_tool_is_scoped_to_effective_stage(self):
|
||||
brain = WorkflowBrain(
|
||||
{
|
||||
"specVersion": 3,
|
||||
"settings": {
|
||||
"globalPrompt": "全局规则",
|
||||
"defaultLlmResourceId": "llm_global",
|
||||
"visionEnabled": True,
|
||||
"visionModelResourceId": "vision_global",
|
||||
},
|
||||
"nodes": [
|
||||
{"id": "start", "type": "start", "data": {}},
|
||||
{
|
||||
"id": "agent",
|
||||
"type": "agent",
|
||||
"data": {
|
||||
"prompt": "观察用户需要展示的物品",
|
||||
"inheritGlobalConfig": True,
|
||||
},
|
||||
},
|
||||
],
|
||||
"edges": [
|
||||
{
|
||||
"id": "begin",
|
||||
"source": "start",
|
||||
"target": "agent",
|
||||
"data": {"mode": "always", "priority": 0},
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
scopes = []
|
||||
vision_function = object()
|
||||
|
||||
async def queue_frame(_frame):
|
||||
pass
|
||||
|
||||
brain._runtime = BrainRuntime(
|
||||
context=LLMContext(messages=[]),
|
||||
llm=FakeLLM(),
|
||||
queue_frame=queue_frame,
|
||||
set_system_prompt=lambda _prompt: None,
|
||||
set_tools=lambda _tools: None,
|
||||
call_end=FakeCallEnd(),
|
||||
set_vision_scope=scopes.append,
|
||||
vision_function=vision_function,
|
||||
)
|
||||
|
||||
await brain._apply_agent_stage("agent")
|
||||
inherited_config = brain._agent_config("agent")
|
||||
self.assertIn(vision_function, inherited_config["functions"])
|
||||
self.assertIn("fetch_user_image", inherited_config["role_message"])
|
||||
self.assertEqual(
|
||||
scopes[-1],
|
||||
{
|
||||
"enabled": True,
|
||||
"vision_model_resource_id": "vision_global",
|
||||
"llm_resource_id": "llm_global",
|
||||
},
|
||||
)
|
||||
|
||||
brain._engine.data("agent").update(
|
||||
{
|
||||
"inheritGlobalConfig": False,
|
||||
"llmResourceId": "llm_agent",
|
||||
"visionEnabled": False,
|
||||
"visionModelResourceId": "",
|
||||
}
|
||||
)
|
||||
await brain._apply_agent_stage("agent")
|
||||
custom_config = brain._agent_config("agent")
|
||||
self.assertNotIn(vision_function, custom_config["functions"])
|
||||
self.assertNotIn("fetch_user_image", custom_config["role_message"])
|
||||
self.assertFalse(scopes[-1]["enabled"])
|
||||
|
||||
async def test_initial_fixed_speech_waits_for_start_greeting_to_finish(self):
|
||||
brain = WorkflowBrain(
|
||||
{
|
||||
|
||||
@@ -1,11 +1,16 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
|
||||
from models import AssistantConfig, RuntimeModelResource
|
||||
from fastapi import HTTPException
|
||||
from routes.assistants import _validate_workflow, _validate_workflow_references
|
||||
from schemas import AssistantUpsert
|
||||
from services.pipecat.service_factory import config_with_resource
|
||||
from services.node_specs import graph_references, normalize_graph, validate_graph
|
||||
from services.runtime_variables import DynamicVariableStore, prepare_dynamic_config
|
||||
from services.vision import config_with_vision_resource
|
||||
from services.workflow_engine import WorkflowEngine
|
||||
|
||||
|
||||
@@ -49,6 +54,28 @@ def valid_graph():
|
||||
|
||||
|
||||
class WorkflowGraphTests(unittest.TestCase):
|
||||
def test_workflow_graph_owns_flat_vision_session_flag(self):
|
||||
graph = valid_graph()
|
||||
graph["settings"].update(
|
||||
{
|
||||
"defaultLlmResourceId": "llm_global",
|
||||
"visionEnabled": True,
|
||||
"visionModelResourceId": "vision_global",
|
||||
}
|
||||
)
|
||||
body = AssistantUpsert(
|
||||
name="视觉工作流",
|
||||
type="workflow",
|
||||
graph=graph,
|
||||
visionEnabled=False,
|
||||
visionModelResourceId="legacy_vision",
|
||||
)
|
||||
|
||||
_validate_workflow(body)
|
||||
|
||||
self.assertTrue(body.vision_enabled)
|
||||
self.assertIsNone(body.vision_model_resource_id)
|
||||
|
||||
def test_agent_entry_mode_defaults_and_validation(self):
|
||||
graph = valid_graph()
|
||||
normalized = normalize_graph(graph)
|
||||
@@ -97,6 +124,8 @@ class WorkflowGraphTests(unittest.TestCase):
|
||||
"defaultLlmResourceId": "llm_global",
|
||||
"defaultAsrResourceId": "asr_global",
|
||||
"defaultTtsResourceId": "tts_global",
|
||||
"visionEnabled": True,
|
||||
"visionModelResourceId": "vision_global",
|
||||
"toolIds": ["tool_global"],
|
||||
"knowledgeBaseId": "kb_global",
|
||||
}
|
||||
@@ -108,6 +137,8 @@ class WorkflowGraphTests(unittest.TestCase):
|
||||
"llmResourceId": "llm_agent",
|
||||
"asrResourceId": "asr_agent",
|
||||
"ttsResourceId": "tts_agent",
|
||||
"visionEnabled": True,
|
||||
"visionModelResourceId": "vision_agent",
|
||||
"toolIds": ["tool_agent"],
|
||||
"knowledgeBaseId": "kb_agent",
|
||||
}
|
||||
@@ -120,9 +151,11 @@ class WorkflowGraphTests(unittest.TestCase):
|
||||
"llm_global",
|
||||
"asr_global",
|
||||
"tts_global",
|
||||
"vision_global",
|
||||
"llm_agent",
|
||||
"asr_agent",
|
||||
"tts_agent",
|
||||
"vision_agent",
|
||||
},
|
||||
)
|
||||
self.assertEqual(refs["tools"], {"tool_global", "tool_agent"})
|
||||
@@ -165,6 +198,8 @@ class WorkflowGraphTests(unittest.TestCase):
|
||||
"defaultLlmResourceId": "llm_global",
|
||||
"defaultAsrResourceId": "asr_global",
|
||||
"defaultTtsResourceId": "tts_global",
|
||||
"visionEnabled": True,
|
||||
"visionModelResourceId": "vision_global",
|
||||
"toolIds": ["tool_global"],
|
||||
"knowledgeBaseId": "kb_global",
|
||||
"knowledgeMode": "on_demand",
|
||||
@@ -180,6 +215,12 @@ class WorkflowGraphTests(unittest.TestCase):
|
||||
engine = WorkflowEngine(graph)
|
||||
inherited = engine.agent_stage_config("agent")
|
||||
self.assertEqual(inherited.llm_resource_id, "llm_global")
|
||||
self.assertTrue(inherited.vision_enabled)
|
||||
self.assertEqual(
|
||||
inherited.vision_model_resource_id,
|
||||
"vision_global",
|
||||
)
|
||||
self.assertTrue(engine.uses_vision())
|
||||
self.assertEqual(inherited.tool_ids, ("tool_global",))
|
||||
self.assertEqual(inherited.knowledge_mode, "on_demand")
|
||||
self.assertFalse(inherited.enable_interrupt)
|
||||
@@ -194,6 +235,8 @@ class WorkflowGraphTests(unittest.TestCase):
|
||||
"llmResourceId": "llm_agent",
|
||||
"toolIds": ["tool_agent"],
|
||||
"knowledgeBaseId": "",
|
||||
"visionEnabled": False,
|
||||
"visionModelResourceId": "",
|
||||
"enableInterrupt": True,
|
||||
"turnConfig": {
|
||||
"bargeIn": {"strategy": "vad"},
|
||||
@@ -203,6 +246,9 @@ class WorkflowGraphTests(unittest.TestCase):
|
||||
)
|
||||
custom = engine.agent_stage_config("agent")
|
||||
self.assertEqual(custom.llm_resource_id, "llm_agent")
|
||||
self.assertFalse(custom.vision_enabled)
|
||||
self.assertIsNone(custom.vision_model_resource_id)
|
||||
self.assertFalse(engine.uses_vision())
|
||||
self.assertEqual(custom.tool_ids, ("tool_agent",))
|
||||
self.assertEqual(custom.knowledge_mode, "disabled")
|
||||
self.assertTrue(custom.enable_interrupt)
|
||||
@@ -211,6 +257,27 @@ class WorkflowGraphTests(unittest.TestCase):
|
||||
"smart_turn",
|
||||
)
|
||||
|
||||
def test_vision_resource_creates_isolated_runtime_config(self):
|
||||
base = AssistantConfig(type="workflow", model="text-only")
|
||||
resource = RuntimeModelResource(
|
||||
id="vision_1",
|
||||
capability="LLM",
|
||||
interface_type="openai-llm",
|
||||
values={
|
||||
"modelId": "vision-model",
|
||||
"apiUrl": "https://vision.test/v1",
|
||||
},
|
||||
secrets={"apiKey": "vision-secret"},
|
||||
support_image_input=True,
|
||||
)
|
||||
|
||||
resolved = config_with_vision_resource(base, resource)
|
||||
|
||||
self.assertEqual(resolved.vision_model, "vision-model")
|
||||
self.assertEqual(resolved.vision_llm_api_key, "vision-secret")
|
||||
self.assertEqual(resolved.vision_llm_base_url, "https://vision.test/v1")
|
||||
self.assertEqual(base.model, "text-only")
|
||||
|
||||
def test_start_agent_action_and_handoff_may_have_no_outgoing_edge(self):
|
||||
terminal_graphs = [
|
||||
{
|
||||
@@ -518,5 +585,46 @@ class WorkflowGraphTests(unittest.TestCase):
|
||||
)
|
||||
|
||||
|
||||
class WorkflowVisionReferenceTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_visual_model_must_support_image_input(self):
|
||||
graph = valid_graph()
|
||||
graph["settings"].update(
|
||||
{
|
||||
"defaultLlmResourceId": "llm_global",
|
||||
"visionEnabled": True,
|
||||
"visionModelResourceId": "vision_global",
|
||||
}
|
||||
)
|
||||
body = AssistantUpsert(
|
||||
name="视觉工作流",
|
||||
type="workflow",
|
||||
graph=graph,
|
||||
)
|
||||
_validate_workflow(body)
|
||||
|
||||
resources = {
|
||||
"llm_global": SimpleNamespace(
|
||||
enabled=True,
|
||||
capability="LLM",
|
||||
support_image_input=False,
|
||||
),
|
||||
"vision_global": SimpleNamespace(
|
||||
enabled=True,
|
||||
capability="LLM",
|
||||
support_image_input=False,
|
||||
),
|
||||
}
|
||||
|
||||
class FakeSession:
|
||||
async def get(self, _model, resource_id):
|
||||
return resources.get(resource_id)
|
||||
|
||||
with self.assertRaisesRegex(HTTPException, "必须支持图片输入"):
|
||||
await _validate_workflow_references(FakeSession(), body)
|
||||
|
||||
resources["vision_global"].support_image_input = True
|
||||
await _validate_workflow_references(FakeSession(), body)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user