mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
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:
co-authored by
autofix-ci[bot]
parent
f9320b2c91
commit
c8abb11bf0
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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()),
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user