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
@@ -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()
|
||||
|
||||
+3
@@ -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()
|
||||
|
||||
+6
-2
@@ -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"
|
||||
|
||||
+10
-2
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user