feat(workflow): add per-agent vision configuration
This commit is contained in:
@@ -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