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
@@ -0,0 +1,180 @@
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from flask import Flask
from werkzeug.exceptions import BadRequest
from controllers.service_api.app import completion as completion_module
from controllers.service_api.app import workflow as workflow_module
from core.helper.trace_id_helper import get_trace_session_id
from models.model import AppMode
class _Request:
def __init__(self, *, headers=None, args=None, json=None, is_json=True):
self.headers = headers or {}
self.args = args or {}
self.json = json
self.is_json = is_json
def test_trace_session_id_header_query_body_priority_matches_service_api_contract():
req = _Request(
headers={"X-Trace-Session-Id": "header"},
args={"trace_session_id": "query"},
json={"trace_session_id": "body"},
)
assert get_trace_session_id(req) == "header"
def test_trace_session_id_invalid_highest_priority_raises_bad_request():
req = _Request(
headers={"X-Trace-Session-Id": " "},
args={"trace_session_id": "query"},
json={"trace_session_id": "body"},
)
with pytest.raises(BadRequest):
get_trace_session_id(req)
def _app(mode: AppMode) -> SimpleNamespace:
return SimpleNamespace(id="app-1", mode=mode, tenant_id="tenant-1")
def _end_user() -> SimpleNamespace:
return SimpleNamespace(id="user-1")
def _assert_generate_trace_session_id(mock_generate_service: MagicMock, expected: str) -> None:
_, kwargs = mock_generate_service.generate.call_args
assert kwargs["args"]["trace_session_id"] == expected
@patch("controllers.service_api.app.completion.AppGenerateService")
@patch("controllers.service_api.app.completion.service_api_ns")
def test_chat_api_rejects_invalid_highest_priority_query_trace_session_id_without_generating(
mock_service_api_ns: MagicMock,
mock_generate_service: MagicMock,
app: Flask,
):
payload = {"inputs": {}, "query": "hello", "trace_session_id": "body-session"}
mock_service_api_ns.payload = payload
with app.test_request_context(
"/chat-messages?trace_session_id=%20%20%20",
method="POST",
json=payload,
):
with pytest.raises(BadRequest):
completion_module.ChatApi().post.__wrapped__(
completion_module.ChatApi(),
_app(AppMode.CHAT),
_end_user(),
)
mock_generate_service.generate.assert_not_called()
@patch("controllers.service_api.app.workflow.AppGenerateService")
@patch("controllers.service_api.app.workflow.service_api_ns")
def test_workflow_run_api_rejects_invalid_highest_priority_body_trace_session_id_without_generating(
mock_service_api_ns: MagicMock,
mock_generate_service: MagicMock,
app: Flask,
):
payload = {"inputs": {}, "trace_session_id": 123}
mock_service_api_ns.payload = payload
with app.test_request_context("/workflows/run", method="POST", json=payload):
with pytest.raises(BadRequest):
workflow_module.WorkflowRunApi().post.__wrapped__(
workflow_module.WorkflowRunApi(),
_app(AppMode.WORKFLOW),
_end_user(),
)
mock_generate_service.generate.assert_not_called()
@patch("controllers.service_api.app.completion.helper.compact_generate_response", return_value={"answer": "ok"})
@patch("controllers.service_api.app.completion.AppGenerateService")
@patch("controllers.service_api.app.completion.service_api_ns")
def test_completion_api_passes_header_trace_session_id_when_body_value_is_invalid_lower_priority(
mock_service_api_ns: MagicMock,
mock_generate_service: MagicMock,
mock_compact: MagicMock,
app: Flask,
):
payload = {"inputs": {}, "trace_session_id": 123}
mock_service_api_ns.payload = payload
mock_generate_service.generate.return_value = "response"
with app.test_request_context(
"/completion-messages",
method="POST",
json=payload,
headers={"X-Trace-Session-Id": " header-session "},
):
response = completion_module.CompletionApi().post.__wrapped__(
completion_module.CompletionApi(),
_app(AppMode.COMPLETION),
_end_user(),
)
assert response == {"answer": "ok"}
_assert_generate_trace_session_id(mock_generate_service, "header-session")
@patch("controllers.service_api.app.completion.helper.compact_generate_response", return_value={"answer": "ok"})
@patch("controllers.service_api.app.completion.AppGenerateService")
@patch("controllers.service_api.app.completion.service_api_ns")
def test_chat_api_passes_query_trace_session_id_when_body_value_is_invalid_lower_priority(
mock_service_api_ns: MagicMock,
mock_generate_service: MagicMock,
mock_compact: MagicMock,
app: Flask,
):
payload = {"inputs": {}, "query": "hello", "trace_session_id": 123}
mock_service_api_ns.payload = payload
mock_generate_service.generate.return_value = "response"
with app.test_request_context(
"/chat-messages?trace_session_id=query-session",
method="POST",
json=payload,
):
response = completion_module.ChatApi().post.__wrapped__(
completion_module.ChatApi(),
_app(AppMode.CHAT),
_end_user(),
)
assert response == {"answer": "ok"}
_assert_generate_trace_session_id(mock_generate_service, "query-session")
@patch("controllers.service_api.app.workflow.helper.compact_generate_response", return_value={"result": "ok"})
@patch("controllers.service_api.app.workflow.AppGenerateService")
@patch("controllers.service_api.app.workflow.service_api_ns")
def test_workflow_run_api_passes_body_trace_session_id(
mock_service_api_ns: MagicMock,
mock_generate_service: MagicMock,
mock_compact: MagicMock,
app: Flask,
):
payload = {"inputs": {}, "trace_session_id": " body-session "}
mock_service_api_ns.payload = payload
mock_generate_service.generate.return_value = "response"
with app.test_request_context("/workflows/run", method="POST", json=payload):
response = workflow_module.WorkflowRunApi().post.__wrapped__(
workflow_module.WorkflowRunApi(),
_app(AppMode.WORKFLOW),
_end_user(),
)
assert response == {"result": "ok"}
_assert_generate_trace_session_id(mock_generate_service, "body-session")
@@ -290,7 +290,7 @@ class TestAdvancedChatAppGeneratorInternals:
workflow=workflow,
node_id="node-1",
user=SimpleNamespace(id="user-id"),
args={"inputs": {"foo": "bar"}},
args={"inputs": {"foo": "bar"}, "trace_session_id": "session-1"},
streaming=False,
)
@@ -298,6 +298,7 @@ class TestAdvancedChatAppGeneratorInternals:
assert prefill_calls == [(workflow, "user-id")]
assert captured["variable_loader"] is var_loader
assert captured["application_generate_entity"].single_iteration_run.node_id == "node-1"
assert captured["application_generate_entity"].extras["trace_session_id"] == "session-1"
def test_single_loop_generate_builds_debug_task(self, monkeypatch: pytest.MonkeyPatch):
generator = AdvancedChatAppGenerator()
@@ -348,7 +349,7 @@ class TestAdvancedChatAppGeneratorInternals:
workflow=workflow,
node_id="node-2",
user=SimpleNamespace(id="user-id"),
args=SimpleNamespace(inputs={"foo": "bar"}),
args=SimpleNamespace(inputs={"foo": "bar"}, trace_session_id="session-1"),
streaming=False,
)
@@ -356,6 +357,7 @@ class TestAdvancedChatAppGeneratorInternals:
assert prefill_calls == [(workflow, "user-id")]
assert captured["variable_loader"] is var_loader
assert captured["application_generate_entity"].single_loop_run.node_id == "node-2"
assert captured["application_generate_entity"].extras["trace_session_id"] == "session-1"
def test_generate_internal_flow_initial_conversation_with_pause_layer(self, monkeypatch: pytest.MonkeyPatch):
generator = AdvancedChatAppGenerator()
@@ -99,6 +99,7 @@ class TestAdvancedChatAppRunnerConversationVariables:
mock_app_generate_entity.call_depth = 0
mock_app_generate_entity.single_iteration_run = None
mock_app_generate_entity.single_loop_run = None
mock_app_generate_entity.extras = {}
mock_app_generate_entity.trace_manager = None
# Create runner
@@ -244,6 +245,7 @@ class TestAdvancedChatAppRunnerConversationVariables:
mock_app_generate_entity.call_depth = 0
mock_app_generate_entity.single_iteration_run = None
mock_app_generate_entity.single_loop_run = None
mock_app_generate_entity.extras = {}
mock_app_generate_entity.trace_manager = None
# Create runner
@@ -404,6 +406,7 @@ class TestAdvancedChatAppRunnerConversationVariables:
mock_app_generate_entity.call_depth = 0
mock_app_generate_entity.single_iteration_run = None
mock_app_generate_entity.single_loop_run = None
mock_app_generate_entity.extras = {}
mock_app_generate_entity.trace_manager = None
# Create runner
@@ -63,6 +63,7 @@ def build_runner():
gen.call_depth = 0
gen.single_iteration_run = None
gen.single_loop_run = None
gen.extras = {}
gen.trace_manager = None
runner = AdvancedChatAppRunner(
@@ -134,6 +134,42 @@ class TestGenerateSuccess:
get_conv.assert_called_once()
def test_generate_does_not_include_trace_session_id_in_extras(self, generator, mocker: MockerFixture):
app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent")
user = DummyAccount("user")
generator._resolve_agent = mocker.MagicMock(
return_value=(mocker.MagicMock(id="agent1"), mocker.MagicMock(id="snap1"), mocker.MagicMock())
)
generator._prepare_user_inputs = mocker.MagicMock(return_value={})
generator._init_generate_records = mocker.MagicMock(
return_value=(mocker.MagicMock(id="conv", mode="agent"), mocker.MagicMock(id="msg"))
)
generator._handle_response = mocker.MagicMock(return_value="raw-response")
mocker.patch(
f"{MODULE}.AgentAppConfigManager.get_app_config",
return_value=mocker.MagicMock(variables=[], tenant_id="tenant", app_id="app1"),
)
mocker.patch(f"{MODULE}.ModelConfigConverter.convert", return_value=mocker.MagicMock(model="gpt-4o-mini"))
mocker.patch(f"{MODULE}.TraceQueueManager", return_value=mocker.MagicMock())
generate_entity = mocker.patch(
f"{MODULE}.AgentAppGenerateEntity", return_value=mocker.MagicMock(task_id="t", user_id="user")
)
mocker.patch(f"{MODULE}.MessageBasedAppQueueManager", return_value=mocker.MagicMock())
mocker.patch(f"{MODULE}.threading.Thread", return_value=mocker.MagicMock())
mocker.patch(f"{MODULE}.AgentAppGenerateResponseConverter.convert", return_value={"result": "ok"})
generator.generate(
app_model=app_model,
user=user,
args={"query": "hello", "inputs": {}, "trace_session_id": "session-1"},
invoke_from=InvokeFrom.WEB_APP,
streaming=True,
)
assert generate_entity.call_args.kwargs["extras"] == {"auto_generate_conversation_name": True}
class TestGenerateWorker:
@pytest.fixture(autouse=True)
@@ -125,7 +125,7 @@ class TestAgentChatAppGeneratorGenerate:
return_value={"result": "ok"},
)
app_entity = mocker.MagicMock(task_id="task", user_id="user", invoke_from=invoke_from)
mocker.patch(
generate_entity = mocker.patch(
"core.app.apps.agent_chat.app_generator.AgentChatAppGenerateEntity",
return_value=app_entity,
)
@@ -136,11 +136,13 @@ class TestAgentChatAppGeneratorGenerate:
"conversation_id": "conv",
"model_config": {"model": {"provider": "p"}},
"files": [{"id": "f1"}],
"trace_session_id": "session-1",
}
result = generator.generate(app_model=app_model, user=user, args=args, invoke_from=invoke_from, streaming=True)
assert result == {"result": "ok"}
assert generate_entity.call_args.kwargs["extras"]["trace_session_id"] == "session-1"
thread_obj.start.assert_called_once()
def test_generate_without_file_config(self, generator, mocker: MockerFixture):
@@ -56,7 +56,7 @@ class TestChatAppGenerator:
generator = ChatAppGenerator()
app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1")
user = SimpleNamespace(id="user-1", session_id="session-1")
args = {"query": "hi", "inputs": {}, "model_config": {"foo": "bar"}}
args = {"query": "hi", "inputs": {}, "model_config": {"foo": "bar"}, "trace_session_id": "session-1"}
with (
patch("core.app.apps.chat.app_generator.ConversationService.get_conversation", return_value=None),
@@ -70,7 +70,10 @@ class TestChatAppGenerator:
patch("core.app.apps.chat.app_generator.ModelConfigConverter.convert", return_value=SimpleNamespace()),
patch("core.app.apps.chat.app_generator.FileUploadConfigManager.convert", return_value=None),
patch("core.app.apps.chat.app_generator.file_factory.build_from_mappings", return_value=[]),
patch("core.app.apps.chat.app_generator.ChatAppGenerateEntity", DummyGenerateEntity),
patch(
"core.app.apps.chat.app_generator.ChatAppGenerateEntity",
Mock(side_effect=DummyGenerateEntity),
) as generate_entity,
patch("core.app.apps.chat.app_generator.TraceQueueManager", return_value=SimpleNamespace()),
patch("core.app.apps.chat.app_generator.MessageBasedAppQueueManager", DummyQueueManager),
patch(
@@ -91,6 +94,7 @@ class TestChatAppGenerator:
result = generator.generate(app_model, user, args, InvokeFrom.DEBUGGER, streaming=False)
assert result == {"ok": True}
assert generate_entity.call_args.kwargs["extras"]["trace_session_id"] == "session-1"
def test_generate_rejects_model_config_override_for_non_debugger(self):
generator = ChatAppGenerator()
@@ -30,7 +30,10 @@ def generator(mocker: MockerFixture):
mocker.patch.object(module, "MessageBasedAppQueueManager", return_value=MagicMock())
mocker.patch.object(module, "TraceQueueManager", return_value=MagicMock())
mocker.patch.object(module, "CompletionAppGenerateEntity", side_effect=lambda **kwargs: SimpleNamespace(**kwargs))
generate_entity = mocker.patch.object(
module, "CompletionAppGenerateEntity", side_effect=lambda **kwargs: SimpleNamespace(**kwargs)
)
gen.generate_entity = generate_entity
return gen
@@ -92,12 +95,13 @@ class TestCompletionAppGenerator:
result = generator.generate(
app_model=_build_app_model(),
user=_build_user(),
args={"query": "q", "inputs": {"a": 1}, "files": []},
args={"query": "q", "inputs": {"a": 1}, "files": [], "trace_session_id": "session-1"},
invoke_from=InvokeFrom.WEB_APP,
streaming=True,
)
assert result == "converted"
assert generator.generate_entity.call_args.kwargs["extras"]["trace_session_id"] == "session-1"
module.file_factory.build_from_mappings.assert_not_called()
def test_generate_success_with_files(self, generator, mocker: MockerFixture):
@@ -0,0 +1,11 @@
from core.helper.trace_id_helper import extract_trace_session_id_from_args
def test_extract_trace_session_id_from_args_for_generator_extras():
assert extract_trace_session_id_from_args({"trace_session_id": "session-1"}) == {
"trace_session_id": "session-1",
}
def test_extract_trace_session_id_from_args_missing_value_keeps_extras_clean():
assert extract_trace_session_id_from_args({"inputs": {}}) == {}
@@ -80,6 +80,7 @@ def test_generate_includes_parent_trace_context_in_extras(monkeypatch):
"parent_workflow_run_id": "outer-workflow-run-1",
"parent_node_execution_id": "outer-node-execution-1",
},
"trace_session_id": "session-1",
},
invoke_from="service-api",
streaming=False,
@@ -93,6 +94,7 @@ def test_generate_includes_parent_trace_context_in_extras(monkeypatch):
"parent_workflow_run_id": "outer-workflow-run-1",
"parent_node_execution_id": "outer-node-execution-1",
}
assert extras["trace_session_id"] == "session-1"
def test_resume_delegates_to_generate(mocker: MockerFixture):
@@ -6,7 +6,7 @@ from types import SimpleNamespace
import pytest
from core.app.apps.workflow_app_runner import WorkflowBasedAppRunner
from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom
from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY, InvokeFrom, UserFrom
from core.app.entities.queue_entities import (
QueueAgentLogEvent,
QueueHumanInputFormFilledEvent,
@@ -85,6 +85,35 @@ class TestWorkflowBasedAppRunner:
invoke_from=InvokeFrom.DEBUGGER,
)
def test_init_graph_includes_trace_session_id_in_run_context(self, monkeypatch: pytest.MonkeyPatch):
runner = WorkflowBasedAppRunner(queue_manager=SimpleNamespace(), app_id="app")
runtime_state = GraphRuntimeState(
variable_pool=VariablePool.from_bootstrap(system_variables=default_system_variables()),
start_at=0.0,
)
captured = {}
def fake_from_graph_init_context(**kwargs):
captured["run_context"] = kwargs["graph_init_context"].run_context
return SimpleNamespace()
monkeypatch.setattr(
"core.app.apps.workflow_app_runner.DifyNodeFactory.from_graph_init_context",
fake_from_graph_init_context,
)
monkeypatch.setattr("core.app.apps.workflow_app_runner.Graph.init", lambda **_kwargs: SimpleNamespace())
runner._init_graph(
graph_config={"nodes": [], "edges": []},
graph_runtime_state=runtime_state,
user_from=UserFrom.ACCOUNT,
invoke_from=InvokeFrom.DEBUGGER,
root_node_id="root",
trace_session_id="session-1",
)
assert captured["run_context"][DIFY_RUN_CONTEXT_KEY].trace_session_id == "session-1"
def test_prepare_single_node_execution_requires_run(self):
runner = WorkflowBasedAppRunner(queue_manager=SimpleNamespace(), app_id="app")
@@ -145,6 +174,57 @@ class TestWorkflowBasedAppRunner:
assert graph is not None
assert variable_pool is graph_runtime_state.variable_pool
def test_get_graph_and_variable_pool_for_single_node_run_includes_trace_session_id(
self, monkeypatch: pytest.MonkeyPatch
):
runner = WorkflowBasedAppRunner(queue_manager=SimpleNamespace(), app_id="app")
graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool.from_bootstrap(system_variables=default_system_variables()),
start_at=0.0,
)
graph_config = {
"nodes": [{"id": "node-1", "data": {"type": "start", "version": "1"}}],
"edges": [],
}
workflow = SimpleNamespace(tenant_id="tenant", id="workflow", graph_dict=graph_config)
captured = {}
def fake_from_graph_init_context(**kwargs):
captured["run_context"] = kwargs["graph_init_context"].run_context
return SimpleNamespace()
class _NodeCls:
@staticmethod
def extract_variable_selector_to_variable_mapping(graph_config, config):
return {}
from core.app.apps import workflow_app_runner
monkeypatch.setattr(
"core.app.apps.workflow_app_runner.DifyNodeFactory.from_graph_init_context",
fake_from_graph_init_context,
)
monkeypatch.setattr("core.app.apps.workflow_app_runner.Graph.init", lambda **kwargs: SimpleNamespace())
monkeypatch.setattr(workflow_app_runner, "resolve_workflow_node_class", lambda **_kwargs: _NodeCls)
monkeypatch.setattr("core.app.apps.workflow_app_runner.load_into_variable_pool", lambda **kwargs: None)
monkeypatch.setattr(
"core.app.apps.workflow_app_runner.WorkflowEntry.mapping_user_inputs_to_variable_pool",
lambda **kwargs: None,
)
runner._get_graph_and_variable_pool_for_single_node_run(
workflow=workflow,
node_id="node-1",
user_inputs={},
graph_runtime_state=graph_runtime_state,
node_type_filter_key="iteration_id",
node_type_label="iteration",
user_id="00000000-0000-0000-0000-000000000001",
trace_session_id="session-1",
)
assert captured["run_context"][DIFY_RUN_CONTEXT_KEY].trace_session_id == "session-1"
def test_get_graph_and_variable_pool_preloads_constructor_variables_before_graph_init(
self, monkeypatch: pytest.MonkeyPatch
):
@@ -51,6 +51,7 @@ def test_run_uses_single_node_execution_branch(
app_generate_entity.task_id = "task-id"
app_generate_entity.call_depth = 0
app_generate_entity.trace_manager = None
app_generate_entity.extras = {"trace_session_id": "session-1"}
app_generate_entity.single_iteration_run = single_iteration_run
app_generate_entity.single_loop_run = single_loop_run
@@ -101,6 +102,7 @@ def test_run_uses_single_node_execution_branch(
single_iteration_run=single_iteration_run,
single_loop_run=single_loop_run,
user_id="user",
trace_session_id="session-1",
)
init_graph.assert_not_called()
@@ -55,6 +55,100 @@ class TestWorkflowAppGeneratorValidation:
streaming=False,
)
def test_single_iteration_generate_includes_trace_session_id_in_extras(self, monkeypatch: pytest.MonkeyPatch):
generator = WorkflowAppGenerator()
app_config = WorkflowUIBasedAppConfig(
tenant_id="tenant",
app_id="app",
app_mode=AppMode.WORKFLOW,
additional_features=AppAdditionalFeatures(),
variables=[],
workflow_id="workflow-id",
)
captured: dict[str, object] = {}
monkeypatch.setattr(
"core.app.apps.workflow.app_generator.WorkflowAppConfigManager.get_app_config",
lambda **kwargs: app_config,
)
monkeypatch.setattr(
"core.app.apps.workflow.app_generator.DifyCoreRepositoryFactory.create_workflow_execution_repository",
lambda **kwargs: SimpleNamespace(),
)
monkeypatch.setattr(
"core.app.apps.workflow.app_generator.DifyCoreRepositoryFactory.create_workflow_node_execution_repository",
lambda **kwargs: SimpleNamespace(),
)
monkeypatch.setattr("core.app.apps.workflow.app_generator.DraftVarLoader", lambda **kwargs: SimpleNamespace())
monkeypatch.setattr("core.app.apps.workflow.app_generator.sessionmaker", lambda **kwargs: SimpleNamespace())
monkeypatch.setattr(
"core.app.apps.workflow.app_generator.db",
SimpleNamespace(engine=object(), session=lambda: SimpleNamespace()),
)
monkeypatch.setattr(
"core.app.apps.workflow.app_generator.WorkflowDraftVariableService",
lambda session: SimpleNamespace(prefill_conversation_variable_default_values=lambda *args, **kwargs: None),
)
monkeypatch.setattr(generator, "_generate", lambda **kwargs: captured.update(kwargs) or {"ok": True})
generator.single_iteration_generate(
app_model=SimpleNamespace(id="app", tenant_id="tenant"),
workflow=SimpleNamespace(id="workflow-id"),
node_id="node-1",
user=SimpleNamespace(id="user-id"),
args={"inputs": {"foo": "bar"}, "trace_session_id": "session-1"},
streaming=False,
)
assert captured["application_generate_entity"].extras["trace_session_id"] == "session-1"
def test_single_loop_generate_includes_trace_session_id_in_extras(self, monkeypatch: pytest.MonkeyPatch):
generator = WorkflowAppGenerator()
app_config = WorkflowUIBasedAppConfig(
tenant_id="tenant",
app_id="app",
app_mode=AppMode.WORKFLOW,
additional_features=AppAdditionalFeatures(),
variables=[],
workflow_id="workflow-id",
)
captured: dict[str, object] = {}
monkeypatch.setattr(
"core.app.apps.workflow.app_generator.WorkflowAppConfigManager.get_app_config",
lambda **kwargs: app_config,
)
monkeypatch.setattr(
"core.app.apps.workflow.app_generator.DifyCoreRepositoryFactory.create_workflow_execution_repository",
lambda **kwargs: SimpleNamespace(),
)
monkeypatch.setattr(
"core.app.apps.workflow.app_generator.DifyCoreRepositoryFactory.create_workflow_node_execution_repository",
lambda **kwargs: SimpleNamespace(),
)
monkeypatch.setattr("core.app.apps.workflow.app_generator.DraftVarLoader", lambda **kwargs: SimpleNamespace())
monkeypatch.setattr("core.app.apps.workflow.app_generator.sessionmaker", lambda **kwargs: SimpleNamespace())
monkeypatch.setattr(
"core.app.apps.workflow.app_generator.db",
SimpleNamespace(engine=object(), session=lambda: SimpleNamespace()),
)
monkeypatch.setattr(
"core.app.apps.workflow.app_generator.WorkflowDraftVariableService",
lambda session: SimpleNamespace(prefill_conversation_variable_default_values=lambda *args, **kwargs: None),
)
monkeypatch.setattr(generator, "_generate", lambda **kwargs: captured.update(kwargs) or {"ok": True})
generator.single_loop_generate(
app_model=SimpleNamespace(id="app", tenant_id="tenant"),
workflow=SimpleNamespace(id="workflow-id"),
node_id="node-2",
user=SimpleNamespace(id="user-id"),
args=SimpleNamespace(inputs={"foo": "bar"}, trace_session_id="session-1"),
streaming=False,
)
assert captured["application_generate_entity"].extras["trace_session_id"] == "session-1"
with pytest.raises(ValueError, match="inputs is required"):
generator.single_loop_generate(
app_model=SimpleNamespace(),
@@ -351,6 +351,7 @@ def _build_workflow_generate_entity_for_roundtrip() -> WorkflowResumptionContext
stream=False,
invoke_from=InvokeFrom.DEBUGGER,
workflow_execution_id="workflow-exec-roundtrip",
extras={"trace_session_id": "session-1"},
)
),
)
@@ -379,6 +380,7 @@ def _build_advanced_chat_generate_entity_for_roundtrip() -> WorkflowResumptionCo
invoke_from=InvokeFrom.DEBUGGER,
workflow_run_id="advanced-run-id",
query="Explain serialization behavior",
extras={"trace_session_id": "session-1"},
)
),
)
@@ -406,3 +408,4 @@ def test_workflow_resumption_context_dumps_loads_roundtrip(state: WorkflowResump
assert loaded.serialized_graph_runtime_state == state.serialized_graph_runtime_state
restored_entity = loaded.get_generate_entity()
assert isinstance(restored_entity, type(state.generate_entity.entity))
assert restored_entity.extras["trace_session_id"] == "session-1"
@@ -38,6 +38,7 @@ from core.app.entities.task_entities import (
)
from core.app.task_pipeline.easy_ui_based_generate_task_pipeline import EasyUIBasedGenerateTaskPipeline
from core.base.tts import AudioTrunk
from core.ops.entities.trace_entity import TraceTaskName
from graphon.file import FileTransferMethod
from graphon.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta, LLMUsage
from graphon.model_runtime.entities.message_entities import AssistantPromptMessage, TextPromptMessageContent
@@ -899,8 +900,10 @@ class TestEasyUiBasedGenerateTaskPipeline:
def test_save_message_persists_fields_and_emits_trace(self, monkeypatch: pytest.MonkeyPatch):
conversation = SimpleNamespace(id="conv", mode=AppMode.CHAT)
message = SimpleNamespace(id="msg", created_at=datetime.now(UTC))
application_generate_entity = _make_entity(ChatAppGenerateEntity, AppMode.CHAT)
application_generate_entity.extras = {"trace_session_id": "session-1"}
pipeline = EasyUIBasedGenerateTaskPipeline(
application_generate_entity=_make_entity(ChatAppGenerateEntity, AppMode.CHAT),
application_generate_entity=application_generate_entity,
queue_manager=SimpleNamespace(),
conversation=conversation,
message=message,
@@ -946,7 +949,12 @@ class TestEasyUiBasedGenerateTaskPipeline:
assert message_obj.message == "serialized-prompt"
assert message_obj.answer == "hello"
assert message_obj.provider_response_latency == 5.0
assert trace_manager.add_trace_task.called
trace_manager.add_trace_task.assert_called_once()
trace_task = trace_manager.add_trace_task.call_args.args[0]
assert trace_task.trace_type == TraceTaskName.MESSAGE_TRACE
assert trace_task.conversation_id == "conv"
assert trace_task.message_id == "msg"
assert trace_task.kwargs["trace_session_id"] == "session-1"
assert len(sent_payloads) == 1
def test_save_message_raises_when_message_not_found(self):
@@ -224,6 +224,7 @@ class TestWorkflowPersistenceLayer:
layer, _, _, _ = _make_layer(
extras={
"external_trace_id": "trace",
"trace_session_id": "session-1",
"parent_trace_context": {
"parent_workflow_run_id": "outer-workflow-run-1",
"parent_node_execution_id": "outer-node-execution-1",
@@ -245,6 +246,7 @@ class TestWorkflowPersistenceLayer:
):
captured["trace_type"] = self.trace_type
captured["external_trace_id"] = self.kwargs.get("external_trace_id")
captured["trace_session_id"] = self.kwargs.get("trace_session_id")
captured["parent_trace_context"] = self.kwargs.get("parent_trace_context")
captured["workflow_run_id"] = workflow_run_id
return {"ok": True}
@@ -257,6 +259,7 @@ class TestWorkflowPersistenceLayer:
trace_task = trace_tasks[0]
assert trace_task.trace_type == TraceTaskName.WORKFLOW_TRACE
assert trace_task.kwargs["external_trace_id"] == "trace"
assert trace_task.kwargs["trace_session_id"] == "session-1"
assert trace_task.kwargs["parent_trace_context"] == {
"parent_workflow_run_id": "outer-workflow-run-1",
"parent_node_execution_id": "outer-node-execution-1",
@@ -266,6 +269,7 @@ class TestWorkflowPersistenceLayer:
assert captured["trace_type"] == TraceTaskName.WORKFLOW_TRACE
assert captured["external_trace_id"] == "trace"
assert captured["trace_session_id"] == "session-1"
assert captured["parent_trace_context"] == {
"parent_workflow_run_id": "outer-workflow-run-1",
"parent_node_execution_id": "outer-node-execution-1",
@@ -1,10 +1,13 @@
import pytest
from werkzeug.exceptions import BadRequest
from core.helper.trace_id_helper import (
ParentTraceContext,
extract_external_trace_id_from_args,
extract_parent_trace_context_from_args,
extract_trace_session_id_from_args,
get_external_trace_id,
get_trace_session_id,
is_valid_trace_id,
)
@@ -17,6 +20,90 @@ class DummyRequest:
self.is_json = is_json
class _Request:
def __init__(self, *, headers=None, args=None, json=None, is_json=True):
self.headers = headers or {}
self.args = args or {}
self.json = json
self.is_json = is_json
def test_get_trace_session_id_prefers_header_over_query_and_body():
request = _Request(
headers={"X-Trace-Session-Id": " header-session "},
args={"trace_session_id": "query-session"},
json={"trace_session_id": "body-session"},
)
assert get_trace_session_id(request) == "header-session"
def test_get_trace_session_id_prefers_query_over_body():
request = _Request(
args={"trace_session_id": " query-session "},
json={"trace_session_id": "body-session"},
)
assert get_trace_session_id(request) == "query-session"
def test_get_trace_session_id_reads_body_when_no_higher_priority_input():
request = _Request(json={"trace_session_id": " body/session:123 "})
assert get_trace_session_id(request) == "body/session:123"
def test_get_trace_session_id_ignores_invalid_lower_priority_value():
request = _Request(
headers={"X-Trace-Session-Id": "header-session"},
json={"trace_session_id": " "},
)
assert get_trace_session_id(request) == "header-session"
@pytest.mark.parametrize(
"trace_session_request",
[
_Request(headers={"X-Trace-Session-Id": " "}, json={"trace_session_id": "body-session"}),
_Request(headers={"X-Trace-Session-Id": 123}),
_Request(headers={"X-Trace-Session-Id": "x" * 201}),
],
)
def test_get_trace_session_id_rejects_invalid_highest_priority_input(trace_session_request):
with pytest.raises(BadRequest) as exc_info:
get_trace_session_id(trace_session_request)
assert "trace_session_id" in str(exc_info.value)
def test_get_trace_session_id_does_not_read_trace_id_or_traceparent():
request = _Request(
headers={
"X-Trace-Id": "trace-id",
"traceparent": "00-5b8aa5a2d2c872e8321cf37308d69df2-051581bf3bb55c45-01",
},
args={"trace_id": "query-trace-id"},
json={"trace_id": "body-trace-id"},
)
assert get_trace_session_id(request) is None
def test_extract_trace_session_id_from_args_returns_trimmed_value():
args = {"trace_session_id": " session-1 "}
assert extract_trace_session_id_from_args(args) == {"trace_session_id": "session-1"}
def test_extract_trace_session_id_from_args_returns_empty_dict_when_missing():
assert extract_trace_session_id_from_args({}) == {}
def test_extract_trace_session_id_from_args_returns_empty_dict_when_blank_after_trim():
assert extract_trace_session_id_from_args({"trace_session_id": " "}) == {}
class TestTraceIdHelper:
"""Test cases for trace_id_helper.py"""
@@ -0,0 +1,140 @@
import json
from datetime import datetime, timedelta
from types import SimpleNamespace
from unittest.mock import MagicMock
from core.ops.entities.trace_entity import TraceTaskName
from core.ops.ops_trace_manager import TraceTask
class _DummySession:
scalar_values: list[object | None] = []
def __init__(self, engine):
self._values = list(self.scalar_values)
self._index = 0
def __enter__(self):
return self
def __exit__(self, exc_type, exc_val, exc_tb):
return False
def execute(self, *args, **kwargs):
return self
def scalar(self, *args, **kwargs):
if self._index >= len(self._values):
return None
value = self._values[self._index]
self._index += 1
return value
def scalars(self, *args, **kwargs):
return self
def all(self):
return []
def _make_workflow_run():
return SimpleNamespace(
workflow_id="wf-1",
tenant_id="tenant-1",
id="run-1",
elapsed_time=1,
status="succeeded",
inputs_dict={},
outputs_dict={},
version="1",
error=None,
total_tokens=0,
created_at=datetime(2026, 1, 1, 0, 0, 0),
finished_at=datetime(2026, 1, 1, 0, 0, 1),
triggered_from="user",
app_id="app-1",
to_dict=lambda self=None: {"id": "run-1"},
)
def _make_message_data():
created_at = datetime(2026, 1, 1, 0, 0, 0)
data = {
"id": "message-1",
"app_id": "app-1",
"conversation_id": "conv-1",
"created_at": created_at,
"updated_at": created_at + timedelta(seconds=1),
"message": "hello",
"provider_response_latency": 1,
"message_tokens": 0,
"answer_tokens": 0,
"answer": "world",
"error": "",
"status": "normal",
"model_provider": "provider",
"model_id": "model",
"from_end_user_id": "end-user-1",
"from_account_id": None,
"agent_based": False,
"workflow_run_id": None,
"from_source": "api",
"message_metadata": json.dumps({"usage": {}}),
}
class _MessageData:
def __init__(self, values):
self.__dict__.update(values)
def to_dict(self):
return dict(self.__dict__)
return _MessageData(data)
def test_workflow_trace_metadata_includes_trace_session_id(monkeypatch):
repo = MagicMock()
repo.get_workflow_run_by_id_without_tenant.return_value = _make_workflow_run()
monkeypatch.setattr(TraceTask, "_get_workflow_run_repo", classmethod(lambda cls: repo))
monkeypatch.setattr("core.ops.ops_trace_manager.Session", _DummySession)
monkeypatch.setattr("core.ops.ops_trace_manager.db", SimpleNamespace(engine=MagicMock()))
monkeypatch.setattr("core.telemetry.gateway.is_enterprise_telemetry_enabled", lambda: False)
_DummySession.scalar_values = [None, None]
task = TraceTask(
TraceTaskName.WORKFLOW_TRACE,
workflow_execution=SimpleNamespace(id_="run-1", total_tokens=0),
conversation_id="conv-1",
user_id="user-1",
trace_session_id="session-1",
)
trace_info = task.workflow_trace(workflow_run_id="run-1", conversation_id="conv-1", user_id="user-1")
assert task.kwargs["trace_session_id"] == "session-1"
assert trace_info.metadata["trace_session_id"] == "session-1"
def test_message_trace_metadata_includes_trace_session_id(monkeypatch):
db_session = MagicMock()
db_session.scalars.return_value.all.return_value = ["chat"]
db_session.scalar.return_value = None
monkeypatch.setattr(
"core.ops.ops_trace_manager.db",
SimpleNamespace(engine=MagicMock(), session=db_session),
)
monkeypatch.setattr("core.ops.ops_trace_manager.Session", _DummySession)
monkeypatch.setattr("core.ops.ops_trace_manager.get_message_data", lambda message_id: _make_message_data())
monkeypatch.setattr("core.telemetry.gateway.is_enterprise_telemetry_enabled", lambda: False)
_DummySession.scalar_values = ["tenant-1"]
task = TraceTask(
TraceTaskName.MESSAGE_TRACE,
message_id="message-1",
trace_session_id="session-1",
)
trace_info = task.message_trace(message_id="message-1", **task.kwargs)
assert task.kwargs["trace_session_id"] == "session-1"
assert trace_info.metadata["trace_session_id"] == "session-1"
@@ -174,6 +174,36 @@ def test_workflow_tool_passes_parent_trace_context_from_runtime(monkeypatch: pyt
}
def test_workflow_tool_passes_parent_trace_session_id(monkeypatch: pytest.MonkeyPatch):
"""Ensure nested workflows inherit the parent observability session ID."""
tool = _build_tool()
tool.entity.parameters = [
ToolParameter.get_simple_instance(
name="trace_session_id",
llm_description="User workflow input",
typ=ToolParameter.ToolParameterType.STRING,
required=False,
),
]
tool.set_trace_session_id("session-1")
monkeypatch.setattr(tool, "_get_app", lambda *args, **kwargs: None)
monkeypatch.setattr(tool, "_get_workflow", lambda *args, **kwargs: None)
mock_user = Mock()
monkeypatch.setattr(tool, "_resolve_user", lambda *args, **kwargs: mock_user)
generate_mock = MagicMock(return_value={"data": {}})
monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock)
monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None)
list(tool.invoke("test_user", {"trace_session_id": "user-input-session"}))
call_kwargs = generate_mock.call_args.kwargs
assert call_kwargs["args"]["inputs"]["trace_session_id"] == "user-input-session"
assert call_kwargs["args"]["trace_session_id"] == "session-1"
def test_workflow_tool_keeps_user_inputs_named_like_trace_runtime_keys(monkeypatch: pytest.MonkeyPatch):
"""Ensure private trace context does not overwrite same-named workflow inputs."""
tool = _build_tool()
@@ -250,6 +280,28 @@ def test_workflow_tool_can_clear_parent_trace_context(monkeypatch: pytest.Monkey
assert "parent_trace_context" not in call_kwargs["args"]
def test_workflow_tool_can_clear_trace_session_id(monkeypatch: pytest.MonkeyPatch):
"""Ensure reused WorkflowTool instances do not keep stale trace session IDs."""
tool = _build_tool()
tool.set_trace_session_id("session-1")
tool.clear_trace_session_id()
monkeypatch.setattr(tool, "_get_app", lambda *args, **kwargs: None)
monkeypatch.setattr(tool, "_get_workflow", lambda *args, **kwargs: None)
mock_user = Mock()
monkeypatch.setattr(tool, "_resolve_user", lambda *args, **kwargs: mock_user)
generate_mock = MagicMock(return_value={"data": {}})
monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock)
monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None)
list(tool.invoke("test_user", {}))
call_kwargs = generate_mock.call_args.kwargs
assert "trace_session_id" not in call_kwargs["args"]
@pytest.mark.parametrize(
"runtime_parameters",
[
@@ -187,6 +187,44 @@ def test_get_runtime_stores_parent_trace_context_for_workflow_tools(
assert workflow_runtime.runtime.runtime_parameters == {}
def test_get_runtime_stores_trace_session_id_for_workflow_tools(
runtime: DifyToolNodeRuntime,
) -> None:
variable_pool: VariablePool = build_test_variable_pool(
variables=build_system_variables(
conversation_id="conversation-id",
workflow_execution_id="workflow-run-id",
)
)
workflow_runtime = MagicMock()
workflow_runtime.runtime.runtime_parameters = {}
runtime._run_context.trace_session_id = "session-1"
node_data = ToolNodeData.model_validate(
{
"type": "tool",
"title": "Tool",
"provider_id": "provider",
"provider_type": ToolProviderType.WORKFLOW,
"provider_name": "provider",
"tool_name": "lookup",
"tool_label": "Lookup",
"tool_configurations": {},
"tool_parameters": {},
}
)
with patch.object(ToolManager, "get_workflow_tool_runtime", return_value=workflow_runtime):
tool_runtime = runtime.get_runtime(
node_id="node-id",
node_data=node_data,
variable_pool=variable_pool,
node_execution_id="node-execution-id",
)
assert tool_runtime.raw.trace_session_id == "session-1"
assert workflow_runtime.runtime.runtime_parameters == {}
def test_get_runtime_leaves_non_workflow_tool_runtime_parameters_unchanged(
runtime: DifyToolNodeRuntime,
) -> None:
@@ -470,6 +470,46 @@ def test_dify_tool_node_runtime_injects_outer_workflow_run_id_for_workflow_tools
get_runtime.assert_called_once()
def test_dify_tool_node_runtime_stores_trace_session_id_for_workflow_tools(
monkeypatch: pytest.MonkeyPatch,
) -> None:
runtime_tool = SimpleNamespace(runtime=SimpleNamespace(runtime_parameters={}))
get_runtime = MagicMock(return_value=runtime_tool)
monkeypatch.setattr(node_runtime.ToolManager, "get_workflow_tool_runtime", get_runtime)
monkeypatch.setattr(
node_runtime,
"get_system_text",
lambda _pool, key: (
"outer-workflow-run-id" if key == node_runtime.SystemVariableKey.WORKFLOW_EXECUTION_ID else None
),
)
run_context = _build_run_context()
run_context[DIFY_RUN_CONTEXT_KEY].trace_session_id = "session-1"
runtime = node_runtime.DifyToolNodeRuntime(run_context)
node_data = ToolNodeData(
title="Workflow Tool Node",
desc=None,
provider_id="workflow-provider-id",
provider_type=ToolProviderType.WORKFLOW,
provider_name="workflow-provider",
tool_name="workflow-tool",
tool_label="Workflow Tool",
tool_configurations={},
tool_parameters={},
)
handle = runtime.get_runtime(
node_id="tool-node",
node_data=node_data,
variable_pool=object(),
node_execution_id="node-execution-id",
)
assert handle.raw.trace_session_id == "session-1"
assert runtime_tool.runtime.runtime_parameters == {}
def test_dify_tool_node_runtime_does_not_inject_outer_workflow_run_id_for_non_workflow_tools(
monkeypatch: pytest.MonkeyPatch,
) -> None: