feat: route workflow image inputs natively

This commit is contained in:
Xin Wang
2026-08-03 10:55:57 +08:00
parent 2e84de0798
commit f3439b21d1
10 changed files with 412 additions and 33 deletions

View File

@@ -69,8 +69,9 @@ class _MessageContinuation:
token: int
node_id: str
context_messages: list[dict[str, str]]
context_messages: list[dict[str, Any]]
triggering_user_text: str
triggering_user_message: dict[str, Any] | None
task: asyncio.Task[None] | None = None
@@ -88,6 +89,26 @@ class ConfiguredFlowManager(FlowManager):
async def _create_transition_func(self, name, handler):
transition = await super()._create_transition_func(name, handler)
native_vision_handler = getattr(
handler,
"_workflow_native_vision_handler",
None,
)
native_vision_enabled = getattr(
handler,
"_workflow_native_vision_enabled",
None,
)
if callable(native_vision_handler) and callable(native_vision_enabled):
fallback_transition = transition
async def vision_transition(params: FunctionCallParams) -> None:
if native_vision_enabled():
await native_vision_handler(params)
return
await fallback_transition(params)
transition = vision_transition
if not getattr(handler, "_suppress_followup_llm", False):
return transition
@@ -290,14 +311,26 @@ class WorkflowBrain(BaseBrain):
if content and not self._ended:
self._store.record("user", content)
async def on_user_turn_end(self, content: str) -> bool:
async def on_user_turn_end(
self,
content: str,
user_message: dict[str, Any] | None = None,
) -> bool:
"""Route a complete user turn before the active stage may reply."""
if not content or self._ended:
return True
async with self._turn_lock:
return await self._handle_user_turn_end(content)
return await self._handle_user_turn_end(
content,
user_message=user_message,
)
async def _handle_user_turn_end(self, content: str) -> bool:
async def _handle_user_turn_end(
self,
content: str,
*,
user_message: dict[str, Any] | None = None,
) -> bool:
"""Serialized implementation so one user turn cannot transition twice."""
self.record_user_message(content)
self._state.begin_user_turn(content)
@@ -307,7 +340,10 @@ class WorkflowBrain(BaseBrain):
return True
self._state.status = WorkflowStatus.ROUTING
decision = await self._edge_evaluator.evaluate(current)
decision = await self._edge_evaluator.evaluate(
current,
current_user_message=user_message,
)
if decision.status == RouteStatus.ERROR:
await self._require_output().emit_error(
decision.error or "工作流路由失败",
@@ -320,6 +356,7 @@ class WorkflowBrain(BaseBrain):
next_config = await self._follow_edge(
decision.edge,
triggering_user_text=content,
triggering_user_message=user_message,
)
await self._activate_node_config(
next_config,
@@ -387,7 +424,7 @@ class WorkflowBrain(BaseBrain):
def _agent_config(
self,
node_id: str,
leading_messages: list[dict[str, str]] | None = None,
leading_messages: list[dict[str, Any]] | None = None,
) -> NodeConfig:
stage = self._engine.agent_stage_config(node_id)
functions: list[FlowsFunctionSchema] = []
@@ -468,7 +505,7 @@ class WorkflowBrain(BaseBrain):
def _passive_node_config(
self,
node_id: str,
task_messages: list[dict[str, str]] | None = None,
task_messages: list[dict[str, Any]] | None = None,
) -> NodeConfig:
"""Keep a non-conversational terminal node active without ending the call."""
return {
@@ -643,8 +680,9 @@ class WorkflowBrain(BaseBrain):
self,
edge: dict,
*,
leading_messages: list[dict[str, str]] | None = None,
leading_messages: list[dict[str, Any]] | None = None,
triggering_user_text: str = "",
triggering_user_message: dict[str, Any] | None = None,
) -> NodeConfig:
await self._begin_edge_transition(edge)
context_messages = list(leading_messages or [])
@@ -664,14 +702,16 @@ class WorkflowBrain(BaseBrain):
str(edge.get("target") or ""),
leading_messages=context_messages,
triggering_user_text=triggering_user_text,
triggering_user_message=triggering_user_message,
)
async def _resolve_path(
self,
node_id: str,
*,
leading_messages: list[dict[str, str]] | None = None,
leading_messages: list[dict[str, Any]] | None = None,
triggering_user_text: str = "",
triggering_user_message: dict[str, Any] | None = None,
) -> NodeConfig:
context_messages = list(leading_messages or [])
for hop in range(MAX_AUTOMATIC_HOPS):
@@ -684,8 +724,13 @@ class WorkflowBrain(BaseBrain):
triggering_user_text
and self._engine.data(node_id).get("contextPolicy") == "fresh"
):
current_user_message = (
deepcopy(triggering_user_message)
if triggering_user_message
else {"role": "user", "content": triggering_user_text}
)
agent_messages = [
{"role": "user", "content": triggering_user_text},
current_user_message,
*context_messages,
]
return self._agent_config(node_id, agent_messages)
@@ -701,6 +746,7 @@ class WorkflowBrain(BaseBrain):
node_id,
context_messages=context_messages,
triggering_user_text=triggering_user_text,
triggering_user_message=triggering_user_message,
)
return self._passive_node_config(node_id, context_messages)
elif node_type == "handoff":
@@ -736,8 +782,9 @@ class WorkflowBrain(BaseBrain):
self,
node_id: str,
*,
context_messages: list[dict[str, str]],
context_messages: list[dict[str, Any]],
triggering_user_text: str,
triggering_user_message: dict[str, Any] | None,
) -> None:
"""Save the path state without waiting inside the pipeline call stack."""
token = self._next_message_token
@@ -747,6 +794,7 @@ class WorkflowBrain(BaseBrain):
node_id=node_id,
context_messages=[dict(message) for message in context_messages],
triggering_user_text=triggering_user_text,
triggering_user_message=deepcopy(triggering_user_message),
)
self._state.enter(node_id, WorkflowStatus.RUNNING_MESSAGE)
@@ -847,6 +895,7 @@ class WorkflowBrain(BaseBrain):
edge,
leading_messages=context_messages,
triggering_user_text=continuation.triggering_user_text,
triggering_user_message=continuation.triggering_user_message,
)
await self._activate_node_config(
next_config,