feat: support custom trace session id for Phoenix tracing (#37056)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
Blackoutta
2026-06-04 08:42:03 +00:00
committed by GitHub
co-authored by autofix-ci[bot]
parent f9320b2c91
commit c8abb11bf0
56 changed files with 1214 additions and 35 deletions
@@ -40,7 +40,7 @@ from core.app.entities.task_entities import (
ChatbotAppStreamResponse,
)
from core.app.layers.pause_state_persist_layer import PauseStateLayerConfig, PauseStatePersistenceLayer
from core.helper.trace_id_helper import extract_external_trace_id_from_args
from core.helper.trace_id_helper import extract_external_trace_id_from_args, extract_trace_session_id_from_args
from core.ops.ops_trace_manager import TraceQueueManager
from core.prompt.utils.get_thread_messages_length import get_thread_messages_length
from core.repositories import DifyCoreRepositoryFactory
@@ -64,6 +64,12 @@ from services.workflow_draft_variable_service import (
logger = logging.getLogger(__name__)
def _extract_trace_session_id_from_debug_args(args: Mapping[str, Any] | Any) -> dict[str, str]:
if isinstance(args, Mapping):
return extract_trace_session_id_from_args(args)
return extract_trace_session_id_from_args({"trace_session_id": getattr(args, "trace_session_id", None)})
class AdvancedChatAppGenerator(MessageBasedAppGenerator):
_dialogue_count: int
@@ -140,6 +146,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
extras = {
"auto_generate_conversation_name": args.get("auto_generate_name", False),
**extract_external_trace_id_from_args(args),
**extract_trace_session_id_from_args(args),
}
# get conversation
@@ -331,7 +338,10 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
user_id=user.id,
stream=streaming,
invoke_from=InvokeFrom.DEBUGGER,
extras={"auto_generate_conversation_name": False},
extras={
"auto_generate_conversation_name": False,
**_extract_trace_session_id_from_debug_args(args),
},
single_iteration_run=AdvancedChatAppGenerateEntity.SingleIterationRunEntity(
node_id=node_id, inputs=args["inputs"]
),
@@ -417,7 +427,10 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
user_id=user.id,
stream=streaming,
invoke_from=InvokeFrom.DEBUGGER,
extras={"auto_generate_conversation_name": False},
extras={
"auto_generate_conversation_name": False,
**_extract_trace_session_id_from_debug_args(args),
},
single_loop_run=AdvancedChatAppGenerateEntity.SingleLoopRunEntity(node_id=node_id, inputs=args.inputs),
)
contexts.plugin_tool_providers.set({})
@@ -131,6 +131,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
user_id=self.application_generate_entity.user_id,
invoke_from=invoke_from,
user_from=user_from,
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
)
elif self.application_generate_entity.single_iteration_run or self.application_generate_entity.single_loop_run:
# Handle single iteration or single loop run
@@ -139,6 +140,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
single_iteration_run=self.application_generate_entity.single_iteration_run,
single_loop_run=self.application_generate_entity.single_loop_run,
user_id=self.application_generate_entity.user_id,
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
)
else:
inputs = self.application_generate_entity.inputs
@@ -199,6 +201,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
user_from=user_from,
invoke_from=invoke_from,
root_node_id=root_node_id,
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
)
db.session.close()
+3 -1
View File
@@ -113,7 +113,9 @@ class AgentAppGenerator(MessageBasedAppGenerator):
user_id=user.id,
stream=streaming,
invoke_from=invoke_from,
extras={"auto_generate_conversation_name": args.get("auto_generate_name", True)},
extras={
"auto_generate_conversation_name": args.get("auto_generate_name", True),
},
call_depth=0,
trace_manager=trace_manager,
agent_id=agent.id,
@@ -20,6 +20,7 @@ from core.app.apps.exc import GenerateTaskStoppedError
from core.app.apps.message_based_app_generator import MessageBasedAppGenerator
from core.app.apps.message_based_app_queue_manager import MessageBasedAppQueueManager
from core.app.entities.app_invoke_entities import AgentChatAppGenerateEntity, InvokeFrom
from core.helper.trace_id_helper import extract_trace_session_id_from_args
from core.ops.ops_trace_manager import TraceQueueManager
from extensions.ext_database import db
from factories import file_factory
@@ -96,7 +97,10 @@ class AgentChatAppGenerator(MessageBasedAppGenerator):
query = query.replace("\x00", "")
inputs = args["inputs"]
extras = {"auto_generate_conversation_name": args.get("auto_generate_name", True)}
extras = {
"auto_generate_conversation_name": args.get("auto_generate_name", True),
**extract_trace_session_id_from_args(args),
}
# get conversation
conversation = None
+5 -1
View File
@@ -20,6 +20,7 @@ from core.app.apps.exc import GenerateTaskStoppedError
from core.app.apps.message_based_app_generator import MessageBasedAppGenerator
from core.app.apps.message_based_app_queue_manager import MessageBasedAppQueueManager
from core.app.entities.app_invoke_entities import ChatAppGenerateEntity, InvokeFrom
from core.helper.trace_id_helper import extract_trace_session_id_from_args
from core.ops.ops_trace_manager import TraceQueueManager
from extensions.ext_database import db
from factories import file_factory
@@ -89,7 +90,10 @@ class ChatAppGenerator(MessageBasedAppGenerator):
query = query.replace("\x00", "")
inputs = args["inputs"]
extras = {"auto_generate_conversation_name": args.get("auto_generate_name", True)}
extras = {
"auto_generate_conversation_name": args.get("auto_generate_name", True),
**extract_trace_session_id_from_args(args),
}
# get conversation
conversation = None
@@ -20,6 +20,7 @@ from core.app.apps.exc import GenerateTaskStoppedError
from core.app.apps.message_based_app_generator import MessageBasedAppGenerator
from core.app.apps.message_based_app_queue_manager import MessageBasedAppQueueManager
from core.app.entities.app_invoke_entities import CompletionAppGenerateEntity, InvokeFrom
from core.helper.trace_id_helper import extract_trace_session_id_from_args
from core.ops.ops_trace_manager import TraceQueueManager
from extensions.ext_database import db
from factories import file_factory
@@ -148,7 +149,9 @@ class CompletionAppGenerator(MessageBasedAppGenerator):
user_id=user.id,
stream=streaming,
invoke_from=invoke_from,
extras={},
extras={
**extract_trace_session_id_from_args(args),
},
trace_manager=trace_manager,
)
+20 -3
View File
@@ -32,7 +32,11 @@ from core.app.entities.task_entities import (
)
from core.app.layers.pause_state_persist_layer import PauseStateLayerConfig, PauseStatePersistenceLayer
from core.db.session_factory import session_factory
from core.helper.trace_id_helper import extract_external_trace_id_from_args, extract_parent_trace_context_from_args
from core.helper.trace_id_helper import (
extract_external_trace_id_from_args,
extract_parent_trace_context_from_args,
extract_trace_session_id_from_args,
)
from core.ops.ops_trace_manager import TraceQueueManager
from core.repositories import DifyCoreRepositoryFactory
from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository
@@ -57,6 +61,12 @@ SKIP_PREPARE_USER_INPUTS_KEY = "_skip_prepare_user_inputs"
logger = logging.getLogger(__name__)
def _extract_trace_session_id_from_debug_args(args: Mapping[str, Any] | Any) -> dict[str, str]:
if isinstance(args, Mapping):
return extract_trace_session_id_from_args(args)
return extract_trace_session_id_from_args({"trace_session_id": getattr(args, "trace_session_id", None)})
class WorkflowAppGenerator(BaseAppGenerator):
@staticmethod
def _should_prepare_user_inputs(args: Mapping[str, Any]) -> bool:
@@ -167,6 +177,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
extras = {
**extract_external_trace_id_from_args(args),
**extract_parent_trace_context_from_args(args),
**extract_trace_session_id_from_args(args),
}
workflow_run_id = str(workflow_run_id or uuid.uuid4())
# FIXME (Yeuoly): we need to remove the SKIP_PREPARE_USER_INPUTS_KEY from the args
@@ -410,7 +421,10 @@ class WorkflowAppGenerator(BaseAppGenerator):
user_id=user.id,
stream=streaming,
invoke_from=InvokeFrom.DEBUGGER,
extras={"auto_generate_conversation_name": False},
extras={
"auto_generate_conversation_name": False,
**_extract_trace_session_id_from_debug_args(args),
},
single_iteration_run=WorkflowAppGenerateEntity.SingleIterationRunEntity(
node_id=node_id, inputs=args["inputs"]
),
@@ -496,7 +510,10 @@ class WorkflowAppGenerator(BaseAppGenerator):
user_id=user.id,
stream=streaming,
invoke_from=InvokeFrom.DEBUGGER,
extras={"auto_generate_conversation_name": False},
extras={
"auto_generate_conversation_name": False,
**_extract_trace_session_id_from_debug_args(args),
},
single_loop_run=WorkflowAppGenerateEntity.SingleLoopRunEntity(node_id=node_id, inputs=args.inputs or {}),
workflow_execution_id=str(uuid.uuid4()),
)
+3
View File
@@ -87,6 +87,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
user_from=user_from,
invoke_from=invoke_from,
root_node_id=self._root_node_id,
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
)
elif self.application_generate_entity.single_iteration_run or self.application_generate_entity.single_loop_run:
graph, variable_pool, graph_runtime_state = self._prepare_single_node_execution(
@@ -94,6 +95,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
single_iteration_run=self.application_generate_entity.single_iteration_run,
single_loop_run=self.application_generate_entity.single_loop_run,
user_id=self.application_generate_entity.user_id,
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
)
else:
inputs = self.application_generate_entity.inputs
@@ -128,6 +130,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
user_from=user_from,
invoke_from=invoke_from,
root_node_id=root_node_id,
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
)
# RUN WORKFLOW
+7
View File
@@ -118,6 +118,7 @@ class WorkflowBasedAppRunner:
tenant_id: str = "",
user_id: str = "",
root_node_id: str | None = None,
trace_session_id: str | None = None,
) -> Graph:
"""
Init graph
@@ -138,6 +139,7 @@ class WorkflowBasedAppRunner:
user_id=user_id,
user_from=user_from,
invoke_from=invoke_from,
trace_session_id=trace_session_id,
)
graph_init_context = DifyGraphInitContext(
workflow_id=workflow_id,
@@ -171,6 +173,7 @@ class WorkflowBasedAppRunner:
single_loop_run: Any | None = None,
*,
user_id: str,
trace_session_id: str | None = None,
) -> tuple[Graph, VariablePool, GraphRuntimeState]:
"""
Prepare graph, variable pool, and runtime state for single node execution
@@ -208,6 +211,7 @@ class WorkflowBasedAppRunner:
node_type_filter_key="iteration_id",
node_type_label="iteration",
user_id=user_id,
trace_session_id=trace_session_id,
)
elif single_loop_run:
graph, variable_pool = self._get_graph_and_variable_pool_for_single_node_run(
@@ -218,6 +222,7 @@ class WorkflowBasedAppRunner:
node_type_filter_key="loop_id",
node_type_label="loop",
user_id=user_id,
trace_session_id=trace_session_id,
)
else:
raise ValueError("Neither single_iteration_run nor single_loop_run is specified")
@@ -236,6 +241,7 @@ class WorkflowBasedAppRunner:
node_type_label: str = "node", # 'iteration' or 'loop' for error messages
*,
user_id: str = "",
trace_session_id: str | None = None,
) -> tuple[Graph, VariablePool]:
"""
Get graph and variable pool for single node execution (iteration or loop).
@@ -301,6 +307,7 @@ class WorkflowBasedAppRunner:
user_id=user_id,
user_from=UserFrom.ACCOUNT,
invoke_from=InvokeFrom.DEBUGGER,
trace_session_id=trace_session_id,
)
graph_init_context = DifyGraphInitContext(
workflow_id=workflow.id,
@@ -54,6 +54,7 @@ class DifyRunContext(BaseModel):
user_id: str
user_from: UserFrom
invoke_from: InvokeFrom
trace_session_id: str | None = None
def build_dify_run_context(
@@ -63,6 +64,7 @@ def build_dify_run_context(
user_id: str,
user_from: UserFrom,
invoke_from: InvokeFrom,
trace_session_id: str | None = None,
extra_context: Mapping[str, Any] | None = None,
) -> dict[str, Any]:
"""
@@ -78,6 +80,7 @@ def build_dify_run_context(
user_id=user_id,
user_from=user_from,
invoke_from=invoke_from,
trace_session_id=trace_session_id,
)
return run_context
@@ -413,7 +413,10 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline):
if trace_manager:
trace_manager.add_trace_task(
TraceTask(
TraceTaskName.MESSAGE_TRACE, conversation_id=self._conversation_id, message_id=self._message_id
TraceTaskName.MESSAGE_TRACE,
conversation_id=self._conversation_id,
message_id=self._message_id,
trace_session_id=self._application_generate_entity.extras.get("trace_session_id"),
)
)
@@ -417,10 +417,12 @@ class WorkflowPersistenceLayer(GraphEngineLayer):
conversation_id = self._system_variables().get(SystemVariableKey.CONVERSATION_ID.value)
external_trace_id = None
trace_session_id = None
parent_trace_context = None
if isinstance(self._application_generate_entity, (WorkflowAppGenerateEntity, AdvancedChatAppGenerateEntity)):
extras = self._application_generate_entity.extras
external_trace_id = extras.get("external_trace_id")
trace_session_id = extras.get("trace_session_id")
parent_trace_context = extras.get("parent_trace_context")
if isinstance(parent_trace_context, ParentTraceContext):
parent_trace_context = parent_trace_context.model_dump(exclude_none=True)
@@ -431,6 +433,7 @@ class WorkflowPersistenceLayer(GraphEngineLayer):
conversation_id=conversation_id,
user_id=self._trace_manager.user_id,
external_trace_id=external_trace_id,
trace_session_id=trace_session_id,
parent_trace_context=parent_trace_context,
)
self._trace_manager.add_trace_task(trace_task)
+64
View File
@@ -4,6 +4,7 @@ from collections.abc import Mapping
from typing import Any
from pydantic import BaseModel, ConfigDict, StrictStr, ValidationError
from werkzeug.exceptions import BadRequest
class ParentTraceContext(BaseModel):
@@ -72,6 +73,69 @@ def extract_external_trace_id_from_args(args: Mapping[str, Any]):
return {}
TRACE_SESSION_ID_HEADER = "X-Trace-Session-Id"
TRACE_SESSION_ID_ARG = "trace_session_id"
TRACE_SESSION_ID_MAX_LENGTH = 200
def _validate_trace_session_id(value: Any) -> str:
if not isinstance(value, str):
raise BadRequest("trace_session_id must be a string.")
normalized = value.strip()
if not normalized:
raise BadRequest("trace_session_id must be 1 to 200 characters after trimming.")
if len(normalized) > TRACE_SESSION_ID_MAX_LENGTH:
raise BadRequest("trace_session_id must be 1 to 200 characters after trimming.")
return normalized
def get_trace_session_id(request: Any) -> str | None:
"""
Resolve the Service API trace session ID from explicit request inputs.
Priority is ``X-Trace-Session-Id`` header, then ``trace_session_id`` query
parameter, then ``trace_session_id`` JSON body field. Only the resolved
highest-priority input is validated; lower-priority values are ignored.
"""
if TRACE_SESSION_ID_HEADER in request.headers:
return _validate_trace_session_id(request.headers.get(TRACE_SESSION_ID_HEADER))
if TRACE_SESSION_ID_ARG in request.args:
return _validate_trace_session_id(request.args.get(TRACE_SESSION_ID_ARG))
if getattr(request, "is_json", False):
json_data = getattr(request, "json", None)
if isinstance(json_data, Mapping) and TRACE_SESSION_ID_ARG in json_data:
return _validate_trace_session_id(json_data.get(TRACE_SESSION_ID_ARG))
return None
def extract_trace_session_id_from_args(args: Mapping[str, Any]) -> dict[str, str]:
"""
Extract normalized ``trace_session_id`` from generation args for entity extras.
"""
trace_session_id = args.get(TRACE_SESSION_ID_ARG)
if isinstance(trace_session_id, str):
normalized = trace_session_id.strip()
if normalized:
return {TRACE_SESSION_ID_ARG: normalized}
return {}
def omit_trace_session_id_from_payload(payload: Any) -> Any:
"""
Return a payload copy without transport-level ``trace_session_id``.
Controllers validate this field through :func:`get_trace_session_id` so lower-priority
body values cannot fail DTO validation before header/query priority is applied.
"""
if isinstance(payload, Mapping) and TRACE_SESSION_ID_ARG in payload:
return {key: value for key, value in payload.items() if key != TRACE_SESSION_ID_ARG}
return payload
def extract_parent_trace_context_from_args(args: Mapping[str, Any]) -> dict[str, ParentTraceContext]:
"""
Extract 'parent_trace_context' from args.
+14
View File
@@ -5,6 +5,7 @@ import os
import queue
import threading
import time
from collections.abc import Mapping
from datetime import timedelta
from typing import TYPE_CHECKING, Any, TypedDict
from uuid import UUID, uuid4
@@ -64,6 +65,11 @@ def _dump_parent_trace_context(parent_trace_context: Any) -> dict[str, str] | No
return None
def _get_trace_session_id(kwargs: Mapping[str, Any]) -> str | None:
value = kwargs.get("trace_session_id")
return value if isinstance(value, str) and value else None
class _AppTracingConfig(TypedDict, total=False):
enabled: bool
tracing_provider: str | None
@@ -873,6 +879,10 @@ class TraceTask:
if dumped_parent_trace_context:
metadata["parent_trace_context"] = dumped_parent_trace_context
trace_session_id = _get_trace_session_id(self.kwargs)
if trace_session_id:
metadata["trace_session_id"] = trace_session_id
workflow_trace_info = WorkflowTraceInfo(
trace_id=self.trace_id,
workflow_data=workflow_run.to_dict(),
@@ -956,6 +966,10 @@ class TraceTask:
if node_execution_id := kwargs.get("node_execution_id"):
metadata["node_execution_id"] = node_execution_id
trace_session_id = _get_trace_session_id(kwargs)
if trace_session_id:
metadata["trace_session_id"] = trace_session_id
message_tokens = message_data.message_tokens
message_trace_info = MessageTraceInfo(
+18 -1
View File
@@ -9,7 +9,11 @@ from sqlalchemy import select
from core.app.file_access import DatabaseFileAccessController
from core.db.session_factory import session_factory
from core.helper.trace_id_helper import ParentTraceContext, extract_parent_trace_context_from_args
from core.helper.trace_id_helper import (
ParentTraceContext,
extract_parent_trace_context_from_args,
extract_trace_session_id_from_args,
)
from core.tools.__base.tool import Tool
from core.tools.__base.tool_runtime import ToolRuntime
from core.tools.entities.tool_entities import (
@@ -38,6 +42,7 @@ class WorkflowTool(Tool):
"""
_parent_trace_context: ParentTraceContext | None
_trace_session_id: str | None
def __init__(
self,
@@ -58,6 +63,7 @@ class WorkflowTool(Tool):
self.label = label
self._latest_usage = LLMUsage.empty_usage()
self._parent_trace_context = None
self._trace_session_id = None
super().__init__(entity=entity, runtime=runtime)
@@ -103,6 +109,8 @@ class WorkflowTool(Tool):
generator_args.update(
extract_parent_trace_context_from_args({"parent_trace_context": self._parent_trace_context})
)
if self._trace_session_id:
generator_args.update(extract_trace_session_id_from_args({"trace_session_id": self._trace_session_id}))
result = generator.generate(
app_model=app,
@@ -215,6 +223,7 @@ class WorkflowTool(Tool):
label=self.label,
)
forked._parent_trace_context = self._parent_trace_context.model_copy() if self._parent_trace_context else None
forked._trace_session_id = self._trace_session_id
return forked
def set_parent_trace_context(
@@ -233,6 +242,14 @@ class WorkflowTool(Tool):
"""Remove parent trace context before invoking this tool outside a nested workflow."""
self._parent_trace_context = None
def set_trace_session_id(self, trace_session_id: str) -> None:
"""Attach parent trace session ID without exposing it as tool input."""
self._trace_session_id = trace_session_id
def clear_trace_session_id(self) -> None:
"""Remove trace session ID before invoking this tool outside a traced session."""
self._trace_session_id = None
def _resolve_user(self, user_id: str) -> Account | EndUser | None:
"""
Resolve user object in both HTTP and worker contexts.
+9
View File
@@ -382,6 +382,7 @@ class _WorkflowToolRuntimeBinding:
tool: Tool
conversation_id: str | None = None
parent_trace_context: ParentTraceContext | None = None
trace_session_id: str | None = None
class DifyToolNodeRuntime(ToolNodeRuntimeProtocol):
@@ -423,6 +424,7 @@ class DifyToolNodeRuntime(ToolNodeRuntimeProtocol):
None if variable_pool is None else get_system_text(variable_pool, SystemVariableKey.CONVERSATION_ID)
)
parent_trace_context: ParentTraceContext | None = None
trace_session_id: str | None = None
if self._is_workflow_tool_provider(node_data):
outer_workflow_run_id = (
None
@@ -434,11 +436,14 @@ class DifyToolNodeRuntime(ToolNodeRuntimeProtocol):
parent_workflow_run_id=outer_workflow_run_id,
parent_node_execution_id=node_execution_id,
)
if isinstance(self._run_context.trace_session_id, str) and self._run_context.trace_session_id:
trace_session_id = self._run_context.trace_session_id
return ToolRuntimeHandle(
raw=_WorkflowToolRuntimeBinding(
tool=tool_runtime,
conversation_id=conversation_id,
parent_trace_context=parent_trace_context,
trace_session_id=trace_session_id,
)
)
@@ -471,6 +476,10 @@ class DifyToolNodeRuntime(ToolNodeRuntimeProtocol):
)
elif hasattr(tool, "clear_parent_trace_context"):
tool.clear_parent_trace_context()
if runtime_binding.trace_session_id and hasattr(tool, "set_trace_session_id"):
tool.set_trace_session_id(runtime_binding.trace_session_id)
elif hasattr(tool, "clear_trace_session_id"):
tool.clear_trace_session_id()
try:
messages = ToolEngine.generic_invoke(