diff --git a/examples/otel_tracing.py b/examples/otel_tracing.py new file mode 100644 index 000000000..6eda8114d --- /dev/null +++ b/examples/otel_tracing.py @@ -0,0 +1,89 @@ +#!/usr/bin/env python3 +"""Example: OpenTelemetry tracing with the Claude Agent SDK. + +This example shows how to wire up distributed tracing so every SDK +call -- session start, message, tool invocation -- appears as a span +in your observability backend (Jaeger, Zipkin, OTLP-compatible, ...). + +Prerequisites +------------- +Install the SDK with the ``[otel]`` extra and a span exporter:: + + pip install claude-agent-sdk[otel] \ + opentelemetry-sdk \ + opentelemetry-exporter-otlp-proto-grpc + +Then run a local Jaeger instance (the all-in-one Docker image is the +fastest way to get a collector + UI):: + + docker run -d --name jaeger \ + -p 16686:16686 \ + -p 4317:4317 \ + jaegertracing/all-in-one:latest + +Finally, run this script:: + + python examples/otel_tracing.py + +Open http://localhost:16686 to browse the traces in the Jaeger UI. +""" + +import anyio + +from claude_agent_sdk import ( + AssistantMessage, + ClaudeAgentOptions, + ResultMessage, + TextBlock, + enable_tracing, + query, +) + + +def setup_otel() -> None: + """Configure an OpenTelemetry TracerProvider with an OTLP exporter.""" + from opentelemetry import trace + from opentelemetry.exporter.otlp.proto.grpc.trace_exporter import OTLPSpanExporter + from opentelemetry.sdk.resources import Resource + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.export import BatchSpanProcessor + + resource = Resource.create({"service.name": "claude-agent-sdk-example"}) + provider = TracerProvider(resource=resource) + exporter = OTLPSpanExporter(endpoint="http://localhost:4317", insecure=True) + provider.add_span_processor(BatchSpanProcessor(exporter)) + trace.set_tracer_provider(provider) + + +async def traced_query() -> None: + """Run a simple query with tracing enabled.""" + print("Running a traced query...") + + options = ClaudeAgentOptions(max_turns=1) + async for message in query( + prompt="What is the square root of 144? Answer in one sentence.", + options=options, + ): + if isinstance(message, AssistantMessage): + for block in message.content: + if isinstance(block, TextBlock): + print(f"Claude: {block.text}") + elif isinstance(message, ResultMessage): + print(f"Turns: {message.num_turns}, Cost: ${message.total_cost_usd:.4f}") + + print("\nDone. Check Jaeger UI at http://localhost:16686") + + +def main() -> None: + # 1. Configure the OTel SDK (provider, exporter, processor). + setup_otel() + + # 2. Tell the Claude Agent SDK to start emitting spans. + enable_tracing() + + # 3. Run the traced query. + anyio.run(traced_query) + + +if __name__ == "__main__": + main() diff --git a/src/claude_agent_sdk/__init__.py b/src/claude_agent_sdk/__init__.py index 50d5b9fe6..bc7dc9ed9 100644 --- a/src/claude_agent_sdk/__init__.py +++ b/src/claude_agent_sdk/__init__.py @@ -51,6 +51,7 @@ list_subagents, list_subagents_from_store, ) +from ._internal.tracing import disable_tracing, enable_tracing from ._internal.transport import Transport from ._version import __version__ from .client import ClaudeSDKClient @@ -528,6 +529,8 @@ async def call_tool(name: str, arguments: dict[str, Any]) -> Any: __all__ = [ # Main exports "query", + "enable_tracing", + "disable_tracing", "__version__", # Transport "Transport", diff --git a/src/claude_agent_sdk/_internal/client.py b/src/claude_agent_sdk/_internal/client.py index 7b0a59387..33e8efd2a 100644 --- a/src/claude_agent_sdk/_internal/client.py +++ b/src/claude_agent_sdk/_internal/client.py @@ -11,6 +11,7 @@ HookEvent, HookMatcher, Message, + ResultMessage, _warn_if_can_use_tool_shadowed, ) from .message_parser import parse_message @@ -22,6 +23,7 @@ materialize_resume_session, ) from .session_store_validation import validate_session_store_options +from .tracing import record_span_event, set_span_attributes, start_span from .transport import Transport from .transport.subprocess_cli import SubprocessCLITransport @@ -95,6 +97,27 @@ async def _process_query_inner( options: ClaudeAgentOptions, transport: Transport | None, materialized: MaterializedResume | None, + ) -> AsyncGenerator[Message, None]: + # Build span attributes for the query-level tracing span. + span_attrs: dict[str, Any] = {} + if options.model: + span_attrs["claude_agent_sdk.model"] = options.model + if options.max_turns is not None: + span_attrs["claude_agent_sdk.max_turns"] = options.max_turns + + with start_span("claude_agent_sdk.query", attributes=span_attrs) as query_span: + async for msg in self._process_query_inner_traced( + prompt, options, transport, materialized, query_span + ): + yield msg + + async def _process_query_inner_traced( + self, + prompt: str | AsyncIterable[dict[str, Any]], + options: ClaudeAgentOptions, + transport: Transport | None, + materialized: MaterializedResume | None, + query_span: Any, ) -> AsyncGenerator[Message, None]: # Validate and configure permission settings (matching TypeScript SDK logic) configured_options = options @@ -227,6 +250,25 @@ async def _on_mirror_error(key: Any, error: str) -> None: async for data in query.receive_messages(): message = parse_message(data) if message is not None: + record_span_event( + query_span, + f"message.{type(message).__name__}", + ) + if isinstance(message, ResultMessage): + attrs: dict[str, Any] = {} + if message.num_turns is not None: + attrs["claude_agent_sdk.num_turns"] = message.num_turns + if message.is_error is not None: + attrs["claude_agent_sdk.is_error"] = message.is_error + if message.duration_ms is not None: + attrs["claude_agent_sdk.duration_ms"] = message.duration_ms + if message.session_id: + attrs["claude_agent_sdk.session_id"] = message.session_id + if message.total_cost_usd is not None: + attrs["claude_agent_sdk.total_cost_usd"] = ( + message.total_cost_usd + ) + set_span_attributes(query_span, attrs) yield message finally: diff --git a/src/claude_agent_sdk/_internal/query.py b/src/claude_agent_sdk/_internal/query.py index 66f10f06e..14f8bd4c3 100644 --- a/src/claude_agent_sdk/_internal/query.py +++ b/src/claude_agent_sdk/_internal/query.py @@ -29,6 +29,7 @@ ToolPermissionContext, ) from ._task_compat import TaskHandle, spawn_detached +from .tracing import start_span from .transport import Transport if TYPE_CHECKING: @@ -428,33 +429,38 @@ async def _handle_control_request(self, request: SDKControlRequest) -> None: if subtype == "can_use_tool": permission_request: SDKControlPermissionRequest = request_data # type: ignore[assignment] + tool_name = permission_request["tool_name"] original_input = permission_request["input"] # Handle tool permission request if not self.can_use_tool: raise Exception("canUseTool callback is not provided") - context = ToolPermissionContext( - signal=None, # TODO: Add abort signal support - suggestions=[ - PermissionUpdate.from_dict(s) - for s in ( - permission_request.get("permission_suggestions") or [] - ) - ], - tool_use_id=permission_request.get("tool_use_id"), - agent_id=permission_request.get("agent_id"), - blocked_path=permission_request.get("blocked_path"), - decision_reason=permission_request.get("decision_reason"), - title=permission_request.get("title"), - display_name=permission_request.get("display_name"), - description=permission_request.get("description"), - ) + with start_span( + "claude_agent_sdk.tool_permission", + attributes={"tool.name": tool_name}, + ): + context = ToolPermissionContext( + signal=None, # TODO: Add abort signal support + suggestions=[ + PermissionUpdate.from_dict(s) + for s in ( + permission_request.get("permission_suggestions") or [] + ) + ], + tool_use_id=permission_request.get("tool_use_id"), + agent_id=permission_request.get("agent_id"), + blocked_path=permission_request.get("blocked_path"), + decision_reason=permission_request.get("decision_reason"), + title=permission_request.get("title"), + display_name=permission_request.get("display_name"), + description=permission_request.get("description"), + ) - response = await self.can_use_tool( - permission_request["tool_name"], - permission_request["input"], - context, - ) + response = await self.can_use_tool( + tool_name, + permission_request["input"], + context, + ) # Convert PermissionResult to expected dict format if isinstance(response, PermissionResultAllow): @@ -507,9 +513,23 @@ async def _handle_control_request(self, request: SDKControlRequest) -> None: # Type narrowing - we've verified these are not None above assert isinstance(server_name, str) assert isinstance(mcp_message, dict) - mcp_response = await self._handle_sdk_mcp_request( - server_name, mcp_message - ) + mcp_method = mcp_message.get("method", "") + mcp_tool_name = "" + if mcp_method == "tools/call": + params = mcp_message.get("params", {}) + if isinstance(params, dict): + mcp_tool_name = params.get("name", "") + with start_span( + "claude_agent_sdk.tool_call", + attributes={ + "mcp.server": server_name, + "mcp.method": mcp_method, + "tool.name": mcp_tool_name, + }, + ): + mcp_response = await self._handle_sdk_mcp_request( + server_name, mcp_message + ) # Wrap the MCP response as expected by the control protocol response_data = {"mcp_response": mcp_response} diff --git a/src/claude_agent_sdk/_internal/tracing.py b/src/claude_agent_sdk/_internal/tracing.py new file mode 100644 index 000000000..3ea4b0c0c --- /dev/null +++ b/src/claude_agent_sdk/_internal/tracing.py @@ -0,0 +1,208 @@ +"""OpenTelemetry tracing integration for Claude Agent SDK. + +This module provides opt-in distributed tracing for the SDK's key lifecycle +events. When ``opentelemetry-api`` is installed and tracing has been enabled +(via :func:`enable_tracing` or by passing ``trace=True`` to +:class:`~claude_agent_sdk.ClaudeAgentOptions`), the SDK emits spans for: + +* **claude_agent_sdk.session** -- the top-level span covering an entire + ``ClaudeSDKClient`` session (connect -> disconnect). +* **claude_agent_sdk.query** -- one-shot ``query()`` calls. +* **claude_agent_sdk.tool_call** -- SDK MCP tool invocations routed through + the control protocol. +* **claude_agent_sdk.tool_permission** -- ``can_use_tool`` permission callback + invocations. +* **claude_agent_sdk.message** -- each message received from the CLI + subprocess (assistant, result, system, ...). + +All tracing is best-effort: failures in the tracing layer are logged at +DEBUG level and never propagate to the caller. When ``opentelemetry-api`` +is not installed the public helpers are harmless no-ops. +""" + +from __future__ import annotations + +import logging +from collections.abc import Iterator +from contextlib import contextmanager +from typing import Any + +logger = logging.getLogger(__name__) + +# --------------------------------------------------------------------------- +# Optional opentelemetry-api import +# --------------------------------------------------------------------------- + +try: + from opentelemetry import trace as otel_trace + from opentelemetry.trace import ( + Span, + StatusCode, + Tracer, + ) + + _HAS_OTEL = True +except ImportError: # pragma: no cover – tested via mock + _HAS_OTEL = False + + # Minimal stand-ins so the rest of this module type-checks without + # opentelemetry-api installed. + class Span: # type: ignore[no-redef] + """No-op span stub.""" + + class StatusCode: # type: ignore[no-redef] + OK = "OK" + ERROR = "ERROR" + UNSET = "UNSET" + + class Tracer: # type: ignore[no-redef] + """No-op tracer stub.""" + +# --------------------------------------------------------------------------- +# Module-level state +# --------------------------------------------------------------------------- + +_TRACER_NAME = "claude_agent_sdk" +_enabled: bool = False +_tracer: Tracer | None = None + + +# --------------------------------------------------------------------------- +# Public API +# --------------------------------------------------------------------------- + + +def enable_tracing( + *, + tracer_name: str = "claude_agent_sdk", +) -> None: + """Enable OpenTelemetry tracing for the Claude Agent SDK. + + Call this once during application startup, **after** you have configured + the OpenTelemetry SDK (e.g. set up a ``TracerProvider`` with an exporter). + + Args: + tracer_name: Name passed to ``opentelemetry.trace.get_tracer()``. + Defaults to ``"claude_agent_sdk"``. + + Raises: + RuntimeError: If ``opentelemetry-api`` is not installed. Install + the SDK with the ``[otel]`` extra to pull it in:: + + pip install claude-agent-sdk[otel] + """ + if not _HAS_OTEL: + raise RuntimeError( + "opentelemetry-api is not installed. " + "Install claude-agent-sdk with the [otel] extra: " + "pip install claude-agent-sdk[otel]" + ) + global _enabled, _tracer, _TRACER_NAME # noqa: PLW0603 + _TRACER_NAME = tracer_name + _tracer = otel_trace.get_tracer(tracer_name) + _enabled = True + + +def disable_tracing() -> None: + """Disable OpenTelemetry tracing. + + Useful in tests or when you want to stop emitting spans at runtime. + """ + global _enabled, _tracer # noqa: PLW0603 + _enabled = False + _tracer = None + + +def is_tracing_enabled() -> bool: + """Return whether tracing is currently active.""" + return _enabled and _tracer is not None + + +# --------------------------------------------------------------------------- +# Internal helpers -- called from client / query code +# --------------------------------------------------------------------------- + + +def _get_tracer() -> Tracer | None: + """Return the active tracer, or *None* when tracing is off.""" + if not _enabled: + return None + return _tracer + + +@contextmanager +def start_span( + name: str, + attributes: dict[str, Any] | None = None, +) -> Iterator[Span | None]: + """Context manager that starts an OTel span when tracing is enabled. + + When tracing is disabled (or opentelemetry-api is absent) the context + manager yields ``None`` and does nothing. + + The span is automatically ended when the block exits. If the block + raises, the span records the exception and sets ``StatusCode.ERROR``. + + Args: + name: Span name -- conventionally ``"claude_agent_sdk."``. + attributes: Optional dict of span attributes set at creation time. + """ + tracer = _get_tracer() + if tracer is None or not _HAS_OTEL: + yield None + return + + try: + span = tracer.start_span(name, attributes=attributes or {}) + except Exception: + logger.debug("Failed to start span %r", name, exc_info=True) + yield None + return + + try: + yield span + except BaseException as exc: + try: + span.set_status(StatusCode.ERROR, str(exc)) + span.record_exception(exc) + except Exception: + logger.debug("Failed to record exception on span", exc_info=True) + raise + else: + try: + span.set_status(StatusCode.OK) + except Exception: + logger.debug("Failed to set span status", exc_info=True) + finally: + try: + span.end() + except Exception: + logger.debug("Failed to end span", exc_info=True) + + +def record_span_event( + span: Span | None, + name: str, + attributes: dict[str, Any] | None = None, +) -> None: + """Add an event to an active span (no-op when span is None).""" + if span is None or not _HAS_OTEL: + return + try: + span.add_event(name, attributes=attributes or {}) + except Exception: + logger.debug("Failed to add event to span", exc_info=True) + + +def set_span_attributes( + span: Span | None, + attributes: dict[str, Any], +) -> None: + """Set attributes on an active span (no-op when span is None).""" + if span is None or not _HAS_OTEL: + return + try: + for key, value in attributes.items(): + span.set_attribute(key, value) + except Exception: + logger.debug("Failed to set span attributes", exc_info=True) diff --git a/src/claude_agent_sdk/client.py b/src/claude_agent_sdk/client.py index 03705f085..b356706e1 100644 --- a/src/claude_agent_sdk/client.py +++ b/src/claude_agent_sdk/client.py @@ -1,5 +1,6 @@ """Claude SDK Client for interacting with Claude Code.""" +import contextlib import json import os from collections.abc import AsyncIterable, AsyncIterator @@ -8,6 +9,7 @@ from . import Transport from ._errors import CLIConnectionError +from ._internal.tracing import record_span_event, set_span_attributes, start_span if TYPE_CHECKING: from ._internal.session_resume import MaterializedResume @@ -78,6 +80,8 @@ def __init__( self._transport: Transport | None = None self._query: Any | None = None self._materialized: MaterializedResume | None = None + self._session_span: Any | None = None + self._session_span_ctx: Any | None = None def _convert_hooks_to_internal_format( self, hooks: dict[HookEvent, list[HookMatcher]] @@ -134,8 +138,25 @@ async def _empty_stream() -> AsyncIterator[dict[str, Any]]: if self._custom_transport is None else None ) + + # Start a session-level tracing span (best-effort). + span_attrs: dict[str, Any] = {} + if self.options.model: + span_attrs["claude_agent_sdk.model"] = self.options.model + if self.options.max_turns is not None: + span_attrs["claude_agent_sdk.max_turns"] = self.options.max_turns + span_ctx = start_span("claude_agent_sdk.session", attributes=span_attrs) + try: + span = span_ctx.__enter__() + except Exception: # noqa: BLE001 + span = None + span_ctx = None # type: ignore[assignment] + self._session_span = span + self._session_span_ctx = span_ctx + try: await self._connect_inner(prompt, actual_prompt) + record_span_event(self._session_span, "connected") except BaseException: # If connect fails after the subprocess has spawned (e.g. at # query.initialize()), close the subprocess/read task *before* @@ -282,6 +303,25 @@ async def receive_messages(self) -> AsyncIterator[Message]: async for data in self._query.receive_messages(): message = parse_message(data) if message is not None: + record_span_event( + self._session_span, + f"message.{type(message).__name__}", + ) + if isinstance(message, ResultMessage): + attrs: dict[str, Any] = {} + if message.num_turns is not None: + attrs["claude_agent_sdk.num_turns"] = message.num_turns + if message.is_error is not None: + attrs["claude_agent_sdk.is_error"] = message.is_error + if message.duration_ms is not None: + attrs["claude_agent_sdk.duration_ms"] = message.duration_ms + if message.session_id: + attrs["claude_agent_sdk.session_id"] = message.session_id + if message.total_cost_usd is not None: + attrs["claude_agent_sdk.total_cost_usd"] = ( + message.total_cost_usd + ) + set_span_attributes(self._session_span, attrs) yield message async def query( @@ -619,6 +659,12 @@ async def disconnect(self) -> None: if self._materialized is not None: await self._materialized.cleanup() self._materialized = None + # End the session tracing span (best-effort). + if self._session_span_ctx is not None: + with contextlib.suppress(Exception): + self._session_span_ctx.__exit__(None, None, None) + self._session_span = None + self._session_span_ctx = None async def __aenter__(self) -> "ClaudeSDKClient": """Enter async context - automatically connects with empty stream for interactive use.""" diff --git a/tests/test_tracing.py b/tests/test_tracing.py new file mode 100644 index 000000000..33d963e44 --- /dev/null +++ b/tests/test_tracing.py @@ -0,0 +1,320 @@ +"""Tests for OpenTelemetry tracing integration.""" + +from __future__ import annotations + +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import anyio +import pytest + +from claude_agent_sdk._internal import tracing as tracing_mod +from claude_agent_sdk._internal.tracing import ( + disable_tracing, + enable_tracing, + is_tracing_enabled, + record_span_event, + set_span_attributes, + start_span, +) + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + + +@pytest.fixture(autouse=True) +def _reset_tracing_state() -> Any: + """Ensure each test starts with tracing disabled.""" + disable_tracing() + yield + disable_tracing() + + +def _make_mock_tracer() -> MagicMock: + mock_tracer = MagicMock() + mock_span = MagicMock() + mock_tracer.start_span.return_value = mock_span + return mock_tracer + + +def _force_enable_with_mock_tracer(mock_tracer: MagicMock) -> None: + """Directly inject a mock tracer into the tracing module state. + + This bypasses the ``enable_tracing()`` path (which tries to call + ``otel_trace.get_tracer()``), so tests do not need opentelemetry-api + installed or patched. + """ + tracing_mod._enabled = True + tracing_mod._tracer = mock_tracer + + +# --------------------------------------------------------------------------- +# enable_tracing / disable_tracing +# --------------------------------------------------------------------------- + + +class TestEnableDisable: + def test_enable_tracing_without_otel_raises(self) -> None: + with ( + patch.object(tracing_mod, "_HAS_OTEL", False), + pytest.raises(RuntimeError, match="opentelemetry-api is not installed"), + ): + enable_tracing() + + def test_enable_tracing_with_otel(self) -> None: + mock_tracer = _make_mock_tracer() + mock_otel = MagicMock() + mock_otel.get_tracer.return_value = mock_tracer + with ( + patch.object(tracing_mod, "_HAS_OTEL", True), + patch.object(tracing_mod, "otel_trace", mock_otel, create=True), + ): + enable_tracing(tracer_name="test-sdk") + + assert is_tracing_enabled() + mock_otel.get_tracer.assert_called_once_with("test-sdk") + + def test_disable_tracing(self) -> None: + mock_tracer = _make_mock_tracer() + _force_enable_with_mock_tracer(mock_tracer) + assert is_tracing_enabled() + + disable_tracing() + assert not is_tracing_enabled() + + def test_enable_tracing_default_name(self) -> None: + mock_tracer = _make_mock_tracer() + mock_otel = MagicMock() + mock_otel.get_tracer.return_value = mock_tracer + with ( + patch.object(tracing_mod, "_HAS_OTEL", True), + patch.object(tracing_mod, "otel_trace", mock_otel, create=True), + ): + enable_tracing() + mock_otel.get_tracer.assert_called_once_with("claude_agent_sdk") + + +# --------------------------------------------------------------------------- +# start_span +# --------------------------------------------------------------------------- + + +class TestStartSpan: + def test_noop_when_disabled(self) -> None: + """start_span yields None when tracing is not enabled.""" + with start_span("test.span") as span: + assert span is None + + def test_creates_span_when_enabled(self) -> None: + mock_tracer = _make_mock_tracer() + with patch.object(tracing_mod, "_HAS_OTEL", True): + _force_enable_with_mock_tracer(mock_tracer) + + with start_span("test.span", attributes={"key": "val"}) as span: + assert span is not None + assert span is mock_tracer.start_span.return_value + + mock_tracer.start_span.assert_called_once_with( + "test.span", attributes={"key": "val"} + ) + span.set_status.assert_called() # type: ignore[union-attr] + span.end.assert_called_once() # type: ignore[union-attr] + + def test_records_exception_on_error(self) -> None: + mock_tracer = _make_mock_tracer() + mock_span = mock_tracer.start_span.return_value + with patch.object(tracing_mod, "_HAS_OTEL", True): + _force_enable_with_mock_tracer(mock_tracer) + + with ( + pytest.raises(ValueError, match="boom"), + start_span("test.error"), + ): + raise ValueError("boom") + + mock_span.record_exception.assert_called_once() + mock_span.end.assert_called_once() + + def test_span_attributes_default_empty(self) -> None: + mock_tracer = _make_mock_tracer() + with patch.object(tracing_mod, "_HAS_OTEL", True): + _force_enable_with_mock_tracer(mock_tracer) + + with start_span("test.default_attrs"): + pass + + mock_tracer.start_span.assert_called_once_with( + "test.default_attrs", attributes={} + ) + + def test_tracer_start_span_failure_yields_none(self) -> None: + """If tracer.start_span raises, we yield None and don't crash.""" + mock_tracer = _make_mock_tracer() + mock_tracer.start_span.side_effect = RuntimeError("tracer broken") + with patch.object(tracing_mod, "_HAS_OTEL", True): + _force_enable_with_mock_tracer(mock_tracer) + + with start_span("test.broken") as span: + assert span is None + + +# --------------------------------------------------------------------------- +# record_span_event / set_span_attributes +# --------------------------------------------------------------------------- + + +class TestSpanHelpers: + def test_record_event_noop_on_none(self) -> None: + # Should not raise + record_span_event(None, "test_event", {"key": "val"}) + + def test_record_event_on_span(self) -> None: + mock_span = MagicMock() + with patch.object(tracing_mod, "_HAS_OTEL", True): + record_span_event(mock_span, "test_event", {"k": "v"}) + mock_span.add_event.assert_called_once_with( + "test_event", attributes={"k": "v"} + ) + + def test_record_event_default_attrs(self) -> None: + mock_span = MagicMock() + with patch.object(tracing_mod, "_HAS_OTEL", True): + record_span_event(mock_span, "evt") + mock_span.add_event.assert_called_once_with("evt", attributes={}) + + def test_set_attributes_noop_on_none(self) -> None: + set_span_attributes(None, {"key": "val"}) + + def test_set_attributes_on_span(self) -> None: + mock_span = MagicMock() + with patch.object(tracing_mod, "_HAS_OTEL", True): + set_span_attributes(mock_span, {"a": 1, "b": "two"}) + mock_span.set_attribute.assert_any_call("a", 1) + mock_span.set_attribute.assert_any_call("b", "two") + + def test_record_event_tolerates_span_error(self) -> None: + mock_span = MagicMock() + mock_span.add_event.side_effect = RuntimeError("span broken") + with patch.object(tracing_mod, "_HAS_OTEL", True): + record_span_event(mock_span, "test_event") # should not raise + + def test_set_attributes_tolerates_span_error(self) -> None: + mock_span = MagicMock() + mock_span.set_attribute.side_effect = RuntimeError("span broken") + with patch.object(tracing_mod, "_HAS_OTEL", True): + set_span_attributes(mock_span, {"k": "v"}) # should not raise + + +# --------------------------------------------------------------------------- +# Integration: spans in query path +# --------------------------------------------------------------------------- + + +class TestQueryTracing: + """Verify that the internal client emits spans when tracing is on.""" + + def test_query_creates_span(self) -> None: + """process_query wraps execution in a claude_agent_sdk.query span.""" + from claude_agent_sdk import query + from claude_agent_sdk.types import AssistantMessage, ResultMessage, TextBlock + + mock_tracer = _make_mock_tracer() + + async def _test() -> None: + with patch.object(tracing_mod, "_HAS_OTEL", True): + _force_enable_with_mock_tracer(mock_tracer) + + with patch( + "claude_agent_sdk._internal.client.InternalClient.process_query" + ) as mock_pq: + + async def mock_gen() -> Any: + yield AssistantMessage( + content=[TextBlock(text="12")], + model="claude-opus-4-1-20250805", + ) + yield ResultMessage( + subtype="success", + duration_ms=100, + duration_api_ms=80, + is_error=False, + num_turns=1, + session_id="sess-1", + total_cost_usd=0.01, + ) + + mock_pq.return_value = mock_gen() + + messages = [] + async for msg in query(prompt="test"): + messages.append(msg) + + assert len(messages) == 2 + assert isinstance(messages[0], AssistantMessage) + assert isinstance(messages[1], ResultMessage) + + anyio.run(_test) + + def test_client_session_span_lifecycle(self) -> None: + """ClaudeSDKClient creates a session span on connect, ends on disconnect.""" + from claude_agent_sdk import ClaudeAgentOptions + from claude_agent_sdk.client import ClaudeSDKClient + + mock_tracer = _make_mock_tracer() + + async def _test() -> None: + with patch.object(tracing_mod, "_HAS_OTEL", True): + _force_enable_with_mock_tracer(mock_tracer) + + client = ClaudeSDKClient(options=ClaudeAgentOptions(model="test-model")) + + # Patch connect_inner and disconnect internals + with patch.object(client, "_connect_inner", new_callable=AsyncMock): + await client.connect("test prompt") + + # Session span should have been created + assert client._session_span is not None + mock_tracer.start_span.assert_called() + + # Check span was started with session name + call_args = mock_tracer.start_span.call_args + assert call_args[0][0] == "claude_agent_sdk.session" + assert ( + call_args[1]["attributes"]["claude_agent_sdk.model"] + == "test-model" + ) + + # Now disconnect + await client.disconnect() + assert client._session_span is None + + anyio.run(_test) + + def test_session_span_ends_on_failed_connect(self) -> None: + """If connect fails, disconnect still cleans up the span.""" + from claude_agent_sdk import ClaudeAgentOptions + from claude_agent_sdk.client import ClaudeSDKClient + + mock_tracer = _make_mock_tracer() + + async def _test() -> None: + with patch.object(tracing_mod, "_HAS_OTEL", True): + _force_enable_with_mock_tracer(mock_tracer) + + client = ClaudeSDKClient(options=ClaudeAgentOptions(model="test-model")) + + async def fail_connect(*a: Any, **kw: Any) -> None: + raise RuntimeError("connect failed") + + with ( + patch.object(client, "_connect_inner", side_effect=fail_connect), + pytest.raises(RuntimeError, match="connect failed"), + ): + await client.connect("test") + + # After the failed connect + disconnect, span should be cleaned up + assert client._session_span is None + assert client._session_span_ctx is None + + anyio.run(_test)