diff --git a/backend/AGENTS.md b/backend/AGENTS.md index 31cd0b12e..31a462bc7 100644 --- a/backend/AGENTS.md +++ b/backend/AGENTS.md @@ -694,7 +694,7 @@ A terminal-native UI over the embedded harness, exposed as the `deerflow` consol Request trace correlation is controlled by `logging.enhance.enabled` at **both** entry points, gated through the shared helper `deerflow.config.app_config.is_trace_correlation_enabled` so the Gateway and embedded paths cannot drift: -- **Gateway HTTP**: `app.gateway.trace_middleware.TraceMiddleware` binds one request-level trace id per HTTP request, inheriting inbound `X-Trace-Id` when present or generating a new id otherwise. The middleware writes the final value to every HTTP response at `http.response.start`, which covers SSE / streaming responses without consuming the body. +- **Gateway HTTP**: `app.gateway.trace_middleware.TraceMiddleware` binds one request-level trace id per HTTP request, inheriting inbound `X-Trace-Id` when present or generating a new id otherwise. A **valid** inbound header also marks the request so `runtime/runs/worker.py` prefers that id over `config.metadata.deerflow_trace_id`, keeping logs, response headers, Langfuse, and runtime context aligned when callers send both. The middleware writes the final value to every HTTP response at `http.response.start`, which covers SSE / streaming responses without consuming the body. - **Embedded / TUI / CLI**: `DeerFlowClient.stream()` mints (or inherits) a request-level trace id per turn only when the flag is on. When it is off, no fresh id is minted — a caller that explicitly wraps `stream()` in `request_trace_context(...)` still opts in, because the downstream `get_current_trace_id()` read propagates that value into Langfuse metadata regardless of the flag. Because `stream()` is a sync generator (which shares the caller's context), the id binding is set/reset around each `next()` step rather than around `yield from`: this keeps LangGraph node execution and its log records inside the binding, while returning control to the caller with the ContextVar restored — avoids cross-request leak between yields and `ValueError: was created in a different Context` on GC-driven close of an abandoned generator (regression pinned by `tests/test_client_langfuse_metadata.py::test_stream_does_not_leak_trace_id_to_caller_context_between_yields` and `::test_stream_abandoned_generator_close_does_not_raise_cross_context`). The same ContextVar value is injected into enhanced log records as `trace_id` and into Langfuse metadata as `deerflow_trace_id`. diff --git a/backend/app/gateway/trace_middleware.py b/backend/app/gateway/trace_middleware.py index 014cebbe7..057b65e82 100644 --- a/backend/app/gateway/trace_middleware.py +++ b/backend/app/gateway/trace_middleware.py @@ -9,7 +9,13 @@ from starlette.datastructures import Headers, MutableHeaders from starlette.types import ASGIApp, Message, Receive, Scope, Send from deerflow.config.app_config import is_trace_correlation_enabled -from deerflow.trace_context import TRACE_ID_HEADER, request_trace_context +from deerflow.trace_context import ( + TRACE_ID_HEADER, + mark_trace_id_from_request_header, + normalize_trace_id, + request_trace_context, + reset_trace_id_from_request_header, +) logger = logging.getLogger(__name__) @@ -39,16 +45,21 @@ class TraceMiddleware: headers = Headers(scope=scope) incoming_trace_id = headers.get(TRACE_ID_HEADER) + header_provided = normalize_trace_id(incoming_trace_id) is not None with request_trace_context(incoming_trace_id) as trace_id: + header_token = mark_trace_id_from_request_header(from_header=header_provided) + try: - async def send_with_trace(message: Message) -> None: - if message["type"] == "http.response.start": - response_headers = MutableHeaders(scope=message) - response_headers[TRACE_ID_HEADER] = trace_id - await send(message) + async def send_with_trace(message: Message) -> None: + if message["type"] == "http.response.start": + response_headers = MutableHeaders(scope=message) + response_headers[TRACE_ID_HEADER] = trace_id + await send(message) - await self.app(scope, receive, send_with_trace) + await self.app(scope, receive, send_with_trace) + finally: + reset_trace_id_from_request_header(header_token) def resolve_trace_enabled(config: Any) -> bool: diff --git a/backend/packages/harness/deerflow/runtime/runs/worker.py b/backend/packages/harness/deerflow/runtime/runs/worker.py index f8dc08a53..0427fe1f8 100644 --- a/backend/packages/harness/deerflow/runtime/runs/worker.py +++ b/backend/packages/harness/deerflow/runtime/runs/worker.py @@ -56,7 +56,11 @@ from deerflow.runtime.goal import ( from deerflow.runtime.serialization import serialize from deerflow.runtime.stream_bridge import StreamBridge from deerflow.runtime.user_context import get_effective_user_id, resolve_runtime_user_id -from deerflow.trace_context import DEERFLOW_TRACE_METADATA_KEY, get_current_trace_id, normalize_trace_id +from deerflow.trace_context import ( + DEERFLOW_TRACE_METADATA_KEY, + is_trace_id_from_request_header, + resolve_deerflow_trace_id, +) from deerflow.tracing import inject_langfuse_metadata from deerflow.utils.messages import message_to_text from deerflow.workspace_changes import capture_workspace_snapshot, record_workspace_changes @@ -363,9 +367,13 @@ async def run_agent( runtime_ctx = _build_runtime_context(thread_id, run_id, config.get("context"), ctx.app_config) runtime_ctx[CURRENT_RUN_PRE_EXISTING_MESSAGE_IDS_KEY] = frozenset(pre_existing_message_ids) incoming_metadata = config.get("metadata") if isinstance(config.get("metadata"), dict) else {} - deerflow_trace_id = normalize_trace_id(incoming_metadata.get(DEERFLOW_TRACE_METADATA_KEY)) or get_current_trace_id() + deerflow_trace_id = resolve_deerflow_trace_id(incoming_metadata.get(DEERFLOW_TRACE_METADATA_KEY)) if deerflow_trace_id: runtime_ctx[DEERFLOW_TRACE_METADATA_KEY] = deerflow_trace_id + if is_trace_id_from_request_header(): + merged_metadata = dict(incoming_metadata) + merged_metadata[DEERFLOW_TRACE_METADATA_KEY] = deerflow_trace_id + config["metadata"] = merged_metadata # Expose the run-scoped journal under a sentinel key so middleware can # write audit events (e.g. SafetyFinishReasonMiddleware recording # suppressed tool calls). Double-underscore prefix marks it as a diff --git a/backend/packages/harness/deerflow/trace_context.py b/backend/packages/harness/deerflow/trace_context.py index 1a7661798..5f9726996 100644 --- a/backend/packages/harness/deerflow/trace_context.py +++ b/backend/packages/harness/deerflow/trace_context.py @@ -17,6 +17,10 @@ DEERFLOW_TRACE_METADATA_KEY: Final[str] = "deerflow_trace_id" _MAX_TRACE_ID_LENGTH: Final[int] = 512 _current_trace_id: Final[ContextVar[str | None]] = ContextVar("deerflow_current_trace_id", default=None) +_trace_id_from_request_header: Final[ContextVar[bool]] = ContextVar( + "deerflow_trace_id_from_request_header", + default=False, +) def generate_trace_id() -> str: @@ -65,6 +69,34 @@ def get_current_trace_id() -> str | None: return _current_trace_id.get() +def mark_trace_id_from_request_header(*, from_header: bool) -> Token[bool]: + """Record whether the current trace id came from a valid inbound header.""" + return _trace_id_from_request_header.set(from_header) + + +def reset_trace_id_from_request_header(token: Token[bool]) -> None: + """Restore the inbound-header flag captured by *token*.""" + _trace_id_from_request_header.reset(token) + + +def is_trace_id_from_request_header() -> bool: + """Return ``True`` when a valid ``X-Trace-Id`` header bound the request.""" + return _trace_id_from_request_header.get() + + +def resolve_deerflow_trace_id(metadata_trace_id: object) -> str | None: + """Resolve the effective ``deerflow_trace_id`` for a run. + + When Gateway ``TraceMiddleware`` bound a valid inbound ``X-Trace-Id``, + that value wins over ``config.metadata.deerflow_trace_id`` so logs, + response headers, Langfuse, and runtime context stay aligned. Otherwise + caller metadata wins, then the ambient request trace context. + """ + if is_trace_id_from_request_header(): + return get_current_trace_id() + return normalize_trace_id(metadata_trace_id) or get_current_trace_id() + + @contextmanager def request_trace_context(trace_id: str | None = None) -> Iterator[str]: """Bind a request trace id for the duration of a request or entry point.""" diff --git a/backend/tests/test_trace_context.py b/backend/tests/test_trace_context.py index 9244e841f..29de503a2 100644 --- a/backend/tests/test_trace_context.py +++ b/backend/tests/test_trace_context.py @@ -9,7 +9,15 @@ from __future__ import annotations import pytest -from deerflow.trace_context import _MAX_TRACE_ID_LENGTH, normalize_trace_id +from deerflow.trace_context import ( + _MAX_TRACE_ID_LENGTH, + is_trace_id_from_request_header, + mark_trace_id_from_request_header, + normalize_trace_id, + request_trace_context, + reset_trace_id_from_request_header, + resolve_deerflow_trace_id, +) class TestNormalizeTraceIdAcceptsPrintableAscii: @@ -84,3 +92,30 @@ class TestNormalizeTraceIdRejectsUnsafeInput: def test_rejects_surrogate_pair_pieces(self) -> None: assert normalize_trace_id("trace-\ud83d") is None + + +class TestResolveDeerflowTraceId: + def test_header_marker_defaults_false_and_resets(self) -> None: + assert is_trace_id_from_request_header() is False + token = mark_trace_id_from_request_header(from_header=True) + try: + assert is_trace_id_from_request_header() is True + finally: + reset_trace_id_from_request_header(token) + assert is_trace_id_from_request_header() is False + + def test_metadata_wins_without_inbound_header(self) -> None: + with request_trace_context("ambient-trace"): + assert resolve_deerflow_trace_id("metadata-trace") == "metadata-trace" + + def test_inbound_header_overrides_metadata(self) -> None: + with request_trace_context("header-trace"): + token = mark_trace_id_from_request_header(from_header=True) + try: + assert resolve_deerflow_trace_id("metadata-trace") == "header-trace" + finally: + reset_trace_id_from_request_header(token) + + def test_falls_back_to_ambient_context(self) -> None: + with request_trace_context("ambient-only"): + assert resolve_deerflow_trace_id(None) == "ambient-only" diff --git a/backend/tests/test_trace_middleware.py b/backend/tests/test_trace_middleware.py index cca3a9a84..cd10eb91e 100644 --- a/backend/tests/test_trace_middleware.py +++ b/backend/tests/test_trace_middleware.py @@ -5,7 +5,11 @@ from fastapi.responses import Response, StreamingResponse from starlette.testclient import TestClient from app.gateway.trace_middleware import TraceMiddleware, resolve_trace_enabled -from deerflow.trace_context import TRACE_ID_HEADER, get_current_trace_id +from deerflow.trace_context import ( + TRACE_ID_HEADER, + get_current_trace_id, + is_trace_id_from_request_header, +) def _make_app(*, enabled: bool) -> FastAPI: @@ -16,6 +20,10 @@ def _make_app(*, enabled: bool) -> FastAPI: async def plain() -> dict[str, str | None]: return {"trace_id": get_current_trace_id()} + @app.get("/header-flag") + async def header_flag() -> dict[str, bool]: + return {"from_header": is_trace_id_from_request_header()} + @app.get("/stream") async def stream() -> StreamingResponse: async def body(): @@ -76,6 +84,16 @@ def test_trace_header_overwrites_duplicate_downstream_value() -> None: assert response.headers.get_list(TRACE_ID_HEADER) == ["canonical-trace"] +def test_trace_header_marks_inbound_header_flag() -> None: + client = TestClient(_make_app(enabled=True)) + + with_header = client.get("/header-flag", headers={TRACE_ID_HEADER: "trace-from-upstream"}) + without_header = client.get("/header-flag") + + assert with_header.json() == {"from_header": True} + assert without_header.json() == {"from_header": False} + + def test_trace_header_rejects_crafted_non_ascii_and_generates_fresh_id() -> None: """A caller-crafted ``X-Trace-Id`` containing codepoints > 0x7E must not reach the response header. Prior to tightening ``normalize_trace_id`` such diff --git a/backend/tests/test_worker_langfuse_metadata.py b/backend/tests/test_worker_langfuse_metadata.py index 697debb1d..ee820cb1d 100644 --- a/backend/tests/test_worker_langfuse_metadata.py +++ b/backend/tests/test_worker_langfuse_metadata.py @@ -14,7 +14,12 @@ import pytest from deerflow.runtime.runs.manager import RunRecord from deerflow.runtime.runs.schemas import DisconnectMode, RunStatus from deerflow.runtime.runs.worker import RunContext, run_agent -from deerflow.trace_context import DEERFLOW_TRACE_METADATA_KEY, request_trace_context +from deerflow.trace_context import ( + DEERFLOW_TRACE_METADATA_KEY, + mark_trace_id_from_request_header, + request_trace_context, + reset_trace_id_from_request_header, +) class _FakeAgent: @@ -287,6 +292,56 @@ async def test_run_agent_preserves_caller_metadata_overrides(monkeypatch): assert metadata["langfuse_trace_name"] == "lead-agent" +@pytest.mark.asyncio +async def test_run_agent_inbound_header_trace_overrides_metadata(monkeypatch): + """A valid inbound ``X-Trace-Id`` wins over ``config.metadata.deerflow_trace_id``.""" + monkeypatch.setenv("LANGFUSE_TRACING", "true") + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "pk-lf-test") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "sk-lf-test") + from deerflow.config.tracing_config import reset_tracing_config + + reset_tracing_config() + + fake_agent = _FakeAgent() + + def agent_factory(config): + return fake_agent + + record = RunRecord( + run_id="run-header-override", + thread_id="thread-header", + assistant_id="lead-agent", + status=RunStatus.pending, + on_disconnect=DisconnectMode.cancel, + ) + record.abort_event = asyncio.Event() + ctx = RunContext(checkpointer=None) + + with request_trace_context("header-trace-1"): + header_token = mark_trace_id_from_request_header(from_header=True) + try: + await run_agent( + _FakeBridge(), + _FakeRunManager(), + record, + ctx=ctx, + agent_factory=agent_factory, + graph_input={"messages": []}, + config={ + "configurable": {"thread_id": "thread-header"}, + "metadata": { + DEERFLOW_TRACE_METADATA_KEY: "metadata-trace-ignored", + }, + }, + ) + finally: + reset_trace_id_from_request_header(header_token) + + metadata = fake_agent.captured_config.get("metadata") or {} + assert metadata[DEERFLOW_TRACE_METADATA_KEY] == "header-trace-1" + assert fake_agent.captured_config.get("context", {}).get(DEERFLOW_TRACE_METADATA_KEY) == "header-trace-1" + + @pytest.mark.asyncio async def test_run_agent_skips_metadata_when_langfuse_disabled(monkeypatch): fake_agent = _FakeAgent()