feat(workflow): add per-agent vision configuration

This commit is contained in:
Xin Wang
2026-07-18 00:00:28 +08:00
parent bdf3d3dd9c
commit 28c6380d8a
23 changed files with 553 additions and 34 deletions

View File

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