feat: route workflow image inputs natively
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user