diff --git a/api/core/app/apps/workflow_app_runner.py b/api/core/app/apps/workflow_app_runner.py index 84d740774b1..f8cf6e4fc9c 100644 --- a/api/core/app/apps/workflow_app_runner.py +++ b/api/core/app/apps/workflow_app_runner.py @@ -42,6 +42,7 @@ from core.workflow.node_factory import ( get_default_root_node_id, resolve_workflow_node_class, ) +from core.workflow.nodes.agent.events import NodeRunAgentLogEvent from core.workflow.nodes.human_input.boundary import enrich_graph_pause_reasons from core.workflow.nodes.human_input.pause_reason import HumanInputRequired from core.workflow.system_variables import ( @@ -66,7 +67,6 @@ from graphon.graph_events import ( GraphRunPausedEvent, GraphRunStartedEvent, GraphRunSucceededEvent, - NodeRunAgentLogEvent, NodeRunExceptionEvent, NodeRunFailedEvent, NodeRunHumanInputFormFilledEvent, diff --git a/api/core/workflow/nodes/agent/agent_node.py b/api/core/workflow/nodes/agent/agent_node.py index 2b6745d46a9..90536f7ff82 100644 --- a/api/core/workflow/nodes/agent/agent_node.py +++ b/api/core/workflow/nodes/agent/agent_node.py @@ -1,16 +1,19 @@ from __future__ import annotations from collections.abc import Generator, Mapping, Sequence +from functools import singledispatchmethod from typing import TYPE_CHECKING, Any, override from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY, DifyRunContext from core.workflow.system_variables import SystemVariableKey, get_system_text from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionStatus +from graphon.graph_events import GraphNodeEventBase from graphon.node_events import NodeEventBase, NodeRunResult, StreamCompletedEvent from graphon.nodes.base.node import Node from graphon.nodes.base.variable_template_parser import VariableTemplateParser from .entities import AgentNodeData +from .events import AgentLogEvent, NodeRunAgentLogEvent from .exceptions import ( AgentInvocationError, AgentMessageTransformError, @@ -71,6 +74,29 @@ class AgentNode(Node[AgentNodeData]): ), } + @override + @singledispatchmethod + def _dispatch( # pyrefly: ignore[missing-override-decorator] + self, event: NodeEventBase + ) -> GraphNodeEventBase: + return super()._dispatch(event) + + @_dispatch.register + def _dispatch_agent_log(self, event: AgentLogEvent) -> NodeRunAgentLogEvent: + return NodeRunAgentLogEvent( + id=self.execution_id, + node_id=self._node_id, + node_type=self.node_type, + message_id=event.message_id, + label=event.label, + node_execution_id=event.node_execution_id, + parent_id=event.parent_id, + error=event.error, + status=event.status, + data=event.data, + metadata=event.metadata, + ) + @override def _run(self) -> Generator[NodeEventBase, None, None]: from core.plugin.impl.exc import PluginDaemonClientSideError diff --git a/api/core/workflow/nodes/agent/events.py b/api/core/workflow/nodes/agent/events.py new file mode 100644 index 00000000000..824decb5199 --- /dev/null +++ b/api/core/workflow/nodes/agent/events.py @@ -0,0 +1,34 @@ +from collections.abc import Mapping +from typing import Any + +from pydantic import Field + +from graphon.graph_events import GraphNodeEventBase +from graphon.node_events import NodeEventBase + + +class AgentLogEvent(NodeEventBase): + message_id: str = Field(..., description="id") + label: str = Field(..., description="label") + node_execution_id: str = Field(..., description="node execution id") + parent_id: str | None = Field(..., description="parent id") + error: str | None = Field(..., description="error") + status: str = Field(..., description="status") + data: Mapping[str, Any] = Field(..., description="data") + metadata: Mapping[str, Any] = Field(default_factory=dict, description="metadata") + node_id: str = Field(..., description="node id") + + +class GraphAgentNodeEventBase(GraphNodeEventBase): + pass + + +class NodeRunAgentLogEvent(GraphAgentNodeEventBase): + message_id: str = Field(..., description="message id") + label: str = Field(..., description="label") + node_execution_id: str = Field(..., description="node execution id") + parent_id: str | None = Field(..., description="parent id") + error: str | None = Field(..., description="error") + status: str = Field(..., description="status") + data: Mapping[str, object] = Field(..., description="data") + metadata: Mapping[str, object] = Field(default_factory=dict) diff --git a/api/core/workflow/nodes/agent/message_transformer.py b/api/core/workflow/nodes/agent/message_transformer.py index f44681377dc..d3147a3581c 100644 --- a/api/core/workflow/nodes/agent/message_transformer.py +++ b/api/core/workflow/nodes/agent/message_transformer.py @@ -16,7 +16,6 @@ from graphon.file import File, FileTransferMethod, get_file_type_by_mime_type from graphon.model_runtime.entities.llm_entities import LLMUsage, LLMUsageMetadata from graphon.model_runtime.utils.encoders import jsonable_encoder from graphon.node_events import ( - AgentLogEvent, NodeEventBase, NodeRunResult, StreamChunkEvent, @@ -26,6 +25,7 @@ from graphon.variables.segments import ArrayFileSegment from models import ToolFile from services.tools.builtin_tools_manage_service import BuiltinToolManageService +from .events import AgentLogEvent from .exceptions import AgentNodeError, AgentVariableTypeError, ToolFileNotFoundError _file_access_controller = DatabaseFileAccessController() diff --git a/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_core.py b/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_core.py index f32546d1b23..460d2943624 100644 --- a/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_core.py +++ b/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_core.py @@ -23,6 +23,7 @@ from core.app.entities.queue_entities import ( QueueWorkflowStartedEvent, QueueWorkflowSucceededEvent, ) +from core.workflow.nodes.agent.events import NodeRunAgentLogEvent from core.workflow.nodes.human_input.pause_reason import HumanInputRequired from core.workflow.system_variables import default_system_variables from graphon.entities.pause_reason import HitlRequired @@ -32,7 +33,6 @@ from graphon.graph_events import ( GraphRunPausedEvent, GraphRunStartedEvent, GraphRunSucceededEvent, - NodeRunAgentLogEvent, NodeRunExceptionEvent, NodeRunFailedEvent, NodeRunHumanInputFormFilledEvent, diff --git a/api/tests/unit_tests/core/workflow/nodes/agent/test_agent_node.py b/api/tests/unit_tests/core/workflow/nodes/agent/test_agent_node.py new file mode 100644 index 00000000000..b1596fecd1d --- /dev/null +++ b/api/tests/unit_tests/core/workflow/nodes/agent/test_agent_node.py @@ -0,0 +1,61 @@ +from unittest.mock import MagicMock + +from core.workflow.nodes.agent.agent_node import AgentNode +from core.workflow.nodes.agent.entities import AgentNodeData +from core.workflow.nodes.agent.events import AgentLogEvent, NodeRunAgentLogEvent +from graphon.entities import GraphInitParams +from graphon.enums import BuiltinNodeTypes +from graphon.graph_events import NodeRunStreamChunkEvent +from graphon.node_events import StreamChunkEvent +from graphon.runtime import GraphRuntimeState, VariablePool + + +def test_dispatch_converts_agent_events_and_delegates_other_events() -> None: + node = AgentNode( + node_id="node-id", + data=AgentNodeData(title="Agent"), + graph_init_params=GraphInitParams( + workflow_id="workflow-id", + graph_config={}, + run_context={}, + call_depth=0, + ), + graph_runtime_state=GraphRuntimeState(variable_pool=VariablePool(), start_at=0), + strategy_resolver=MagicMock(), + presentation_provider=MagicMock(), + runtime_support=MagicMock(), + message_transformer=MagicMock(), + ) + node._node_execution_id = "execution-id" + + graph_event = node._dispatch( + AgentLogEvent( + message_id="message-id", + label="label", + node_execution_id="agent-execution-id", + parent_id="parent-id", + error=None, + status="succeeded", + data={"output": "done"}, + metadata={"provider": "test"}, + node_id="source-node-id", + ) + ) + + assert graph_event == NodeRunAgentLogEvent( + id="execution-id", + node_id="node-id", + node_type=BuiltinNodeTypes.AGENT, + message_id="message-id", + label="label", + node_execution_id="agent-execution-id", + parent_id="parent-id", + error=None, + status="succeeded", + data={"output": "done"}, + metadata={"provider": "test"}, + ) + assert isinstance( + node._dispatch(StreamChunkEvent(selector=["node-id", "text"], chunk="hello")), + NodeRunStreamChunkEvent, + )