Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 23 additions & 0 deletions src/claude_agent_sdk/_internal/_trace_helpers.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,23 @@
"""Helper utilities for OpenTelemetry trace context propagation."""

from typing import Any


def inject_trace_into_message(message: dict[str, Any]) -> None:
"""Inject ambient OTel trace context into a message dict for the CLI.

Best-effort: no-op if opentelemetry-api is not installed or there's no
active span. The CLI reads these optional fields to stamp outbound MCP/tool
spans under the caller's current trace rather than the spawn-time trace.
"""
try:
from opentelemetry import propagate

carrier: dict[str, str] = {}
propagate.inject(carrier)
if "traceparent" in carrier:
message["traceparent"] = carrier["traceparent"]
if "tracestate" in carrier:
message["tracestate"] = carrier["tracestate"]
except Exception:
pass
2 changes: 2 additions & 0 deletions src/claude_agent_sdk/_internal/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
Message,
_warn_if_can_use_tool_shadowed,
)
from ._trace_helpers import inject_trace_into_message
from .message_parser import parse_message
from .query import Query
from .session_resume import (
Expand Down Expand Up @@ -217,6 +218,7 @@ async def _on_mirror_error(key: Any, error: str) -> None:
"message": {"role": "user", "content": prompt},
"parent_tool_use_id": None,
}
inject_trace_into_message(user_message)
await chosen_transport.write(json.dumps(user_message) + "\n")
query.spawn_task(query.wait_for_result_and_end_input())
elif isinstance(prompt, AsyncIterable):
Expand Down
28 changes: 28 additions & 0 deletions src/claude_agent_sdk/_internal/query.py
Original file line number Diff line number Diff line change
Expand Up @@ -806,6 +806,31 @@ async def stop_task(self, task_id: str) -> None:
}
)

async def set_trace_context(self) -> None:
"""Update the CLI's trace context from the ambient OTel span.

Best-effort: no-op if opentelemetry-api is not installed or there's no
active span. This lets the CLI stamp outbound MCP/tool calls under the
caller's current trace when ``CLAUDE_CODE_PROPAGATE_TRACEPARENT=1``,
rather than the spawn-time trace that was frozen at ``connect()``.
"""
try:
from opentelemetry import propagate

carrier: dict[str, str] = {}
propagate.inject(carrier)
if "traceparent" not in carrier:
return
request: dict[str, Any] = {
"subtype": "set_trace_context",
"traceparent": carrier["traceparent"],
}
if "tracestate" in carrier:
request["tracestate"] = carrier["tracestate"]
await self._send_control_request(request)
except Exception: # noqa: BLE001 - best-effort tracing must never break query()
logger.debug("set_trace_context control request failed", exc_info=True)

async def wait_for_result_and_end_input(self) -> None:
"""Wait for the first result (if needed) then close stdin.

Expand All @@ -832,10 +857,13 @@ async def stream_input(self, stream: AsyncIterable[dict[str, Any]]) -> None:
If SDK MCP servers or hooks are present, waits for the first result
before closing stdin to allow bidirectional control protocol communication.
"""
from ._trace_helpers import inject_trace_into_message

try:
async for message in stream:
if self._closed:
break
inject_trace_into_message(message)
await self.transport.write(json.dumps(message) + "\n")

await self.wait_for_result_and_end_input()
Expand Down
4 changes: 4 additions & 0 deletions src/claude_agent_sdk/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -297,6 +297,8 @@ async def query(
if not self._query or not self._transport:
raise CLIConnectionError("Not connected. Call connect() first.")

from ._internal._trace_helpers import inject_trace_into_message

# Handle string prompts
if isinstance(prompt, str):
message = {
Expand All @@ -305,13 +307,15 @@ async def query(
"parent_tool_use_id": None,
"session_id": session_id,
}
inject_trace_into_message(message)
await self._transport.write(json.dumps(message) + "\n")
else:
# Handle AsyncIterable prompts - stream them
async for msg in prompt:
# Ensure session_id is set on each message
if "session_id" not in msg:
msg["session_id"] = session_id
inject_trace_into_message(msg)
await self._transport.write(json.dumps(msg) + "\n")

async def interrupt(self) -> None:
Expand Down
Loading