mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-07-27 00:17:53 +00:00
fix(gateway): prefer X-Trace-Id over metadata.deerflow_trace_id when header is set (#4283)
When TraceMiddleware binds a valid inbound X-Trace-Id, worker resolution prefers that id over config.metadata.deerflow_trace_id so logs, response headers, Langfuse, and runtime context stay aligned. Also add trace-context marker reset coverage for the inbound-header flag.
This commit is contained in:
parent
283cea567e
commit
0f088033fe
@ -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: <Token> 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`.
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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."""
|
||||
|
||||
@ -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"
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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()
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user