codekingpro/portable-devtools
114k
1"""Client instrumentation for Claude Agent SDK."""2 3import logging4import time5import weakref6from collections.abc import AsyncGenerator, AsyncIterable7from datetime import datetime, timezone8from functools import cache9from typing import Any, Optional10 11from langsmith._internal import _context12from langsmith.run_helpers import get_current_run_tree, trace13 14from ._config import get_tracing_config15from ._hooks import (16 SessionState,17 _current_session,18 _register_session,19 _set_session_root,20 _unregister_session,21 clear_active_tool_runs,22 get_subagent_run_by_tool_id,23 post_tool_use_failure_hook,24 post_tool_use_hook,25 pre_tool_use_hook,26 subagent_start_hook,27 subagent_stop_hook,28)29from ._messages import (30 build_llm_input,31 flatten_content_blocks,32 unwrap_message_dicts,33)34from ._tools import (35 clear_parent_run_tree,36 get_parent_run_tree,37 set_parent_run_tree,38)39from ._transcripts import LLM_RUN_NAME, reconcile_from_transcripts40from ._usage import extract_usage_metadata41 42logger = logging.getLogger(__name__)43 44TRACE_CHAIN_NAME = "claude.conversation"45 46 47@cache48def _get_package_version(package_name: str) -> str | None:49 try:50 from importlib.metadata import version51 52 return version(package_name)53 except Exception:54 return None55 56 57class TurnLifecycle:58 """Track ongoing model runs so consecutive messages are recorded correctly.59 60 The Claude Agent SDK may deliver a single assistant turn as multiple61 ``AssistantMessage`` events (e.g. one with ``ThinkingBlock``, another62 with ``TextBlock``/``ToolUseBlock``). Messages that share the same63 ``message_id`` are accumulated into a single LLM run.64 """65 66 def __init__(self, query_start_time: Optional[float] = None):67 self.current_run: Optional[Any] = None68 self.current_message_id: Optional[str] = None69 self.next_start_time: Optional[float] = query_start_time70 # message_id → RunTree for all LLM runs created this conversation.71 # Used to retroactively set usage from transcripts.72 self.llm_runs_by_message_id: dict[str, Any] = {}73 # Runs that have been end()ed but not yet patch()ed.74 # Deferred so transcript usage can be set before the single patch().75 self._pending_patch: list[Any] = []76 77 def start_llm_run(78 self,79 message: Any,80 prompt: Any,81 history: list[dict[str, Any]],82 parent: Optional[Any] = None,83 ) -> Optional[dict[str, Any]]:84 """Begin or continue a model run for *message*.85 86 If *message* has the same ``message_id`` as the current run the87 output is appended; otherwise a new run is started (ending any88 previous one first).89 """90 message_id = getattr(message, "message_id", None)91 start = self.next_start_time or time.time()92 93 # Same turn – just accumulate the output blocks and update usage.94 # Return None so the caller does NOT append a duplicate history95 # entry; the original entry in ``history`` is updated in place.96 if message_id and message_id == self.current_message_id and self.current_run:97 content = flatten_content_blocks(getattr(message, "content", None))98 if content and self.current_run.outputs:99 prev = self.current_run.outputs.get("content", [])100 if isinstance(prev, list) and isinstance(content, list):101 merged = prev + content102 self.current_run.outputs["content"] = merged103 # Update the existing history entry in place so104 # subsequent LLM runs see a single merged message.105 for entry in reversed(history):106 if entry.get("role") == "assistant":107 entry["content"] = merged108 break109 elif isinstance(content, list):110 self.current_run.outputs["content"] = content111 self._set_usage_from_message(message, self.current_run)112 return None113 114 # Different turn – end previous but defer patch() until115 # transcript usage is available.116 if self.current_run:117 self.current_run.end()118 self._pending_patch.append(self.current_run)119 120 final_output, run = begin_llm_run_from_assistant_messages(121 [message], prompt, history, start_time=start, parent=parent122 )123 self.current_run = run124 self.current_message_id = message_id125 self.next_start_time = None126 127 if run:128 if message_id:129 self.llm_runs_by_message_id[message_id] = run130 self._set_usage_from_message(message, run)131 132 return final_output133 134 @staticmethod135 def _set_usage_from_message(message: Any, run: Any) -> None:136 """Set usage metadata on a run from a live AssistantMessage.137 138 Always overwrites — later chunks in the same turn have more139 accurate counts. Transcript-based usage will overwrite again140 if available.141 """142 raw_usage = getattr(message, "usage", None)143 if not raw_usage:144 return145 usage_meta = extract_usage_metadata(raw_usage)146 if usage_meta:147 meta = run.extra.setdefault("metadata", {})148 meta["usage_metadata"] = usage_meta149 150 def mark_next_start(self) -> None:151 """Mark when the next assistant message will start."""152 self.next_start_time = time.time()153 154 def close(self) -> None:155 """End any open run and add to pending patch list."""156 if self.current_run:157 self.current_run.end()158 self._pending_patch.append(self.current_run)159 self.current_run = None160 161 def flush(self) -> None:162 """Patch all deferred LLM runs. Call after usage has been set."""163 for run in self._pending_patch:164 try:165 run.patch()166 except Exception as e:167 logger.warning(f"Failed to patch LLM run: {e}")168 self._pending_patch.clear()169 170 171def begin_llm_run_from_assistant_messages(172 messages: list[Any],173 prompt: Any,174 history: list[dict[str, Any]],175 start_time: Optional[float] = None,176 parent: Optional[Any] = None,177) -> tuple[Optional[dict[str, Any]], Optional[Any]]:178 """Create a traced model run from assistant messages."""179 if not messages or type(messages[-1]).__name__ != "AssistantMessage":180 return None, None181 182 last_msg = messages[-1]183 model = getattr(last_msg, "model", None)184 if parent is None:185 parent = get_parent_run_tree() or get_current_run_tree()186 if not parent:187 return None, None188 189 inputs = build_llm_input(prompt, history)190 outputs = [191 {"content": flatten_content_blocks(m.content), "role": "assistant"}192 for m in messages193 if hasattr(m, "content")194 ]195 196 llm_metadata: dict[str, Any] = {"ls_provider": "anthropic"}197 if model:198 llm_metadata["ls_model_name"] = model199 200 llm_run = parent.create_child(201 name=LLM_RUN_NAME,202 run_type="llm",203 inputs={"messages": inputs} if inputs else {},204 extra={"metadata": llm_metadata},205 start_time=datetime.fromtimestamp(start_time, tz=timezone.utc)206 if start_time207 else None,208 )209 210 try:211 llm_run.post()212 except Exception as e:213 logger.warning(f"Failed to post LLM run: {e}")214 215 # Set outputs after posting so they are sent with end_time on the patch.216 llm_run.outputs = outputs[-1] if len(outputs) == 1 else {"content": outputs}217 218 final_content = (219 {"content": flatten_content_blocks(last_msg.content), "role": "assistant"}220 if hasattr(last_msg, "content")221 else None222 )223 return final_content, llm_run224 225 226def _bind_hook_to_session(hook: Any, session: Optional[SessionState]) -> Any:227 """Return a hook callable that runs with *session* bound, if provided."""228 if session is None:229 return hook230 231 async def _bound(input_data: Any, tool_use_id: Any, context: Any) -> Any:232 token = _current_session.set(session)233 try:234 return await hook(input_data, tool_use_id, context)235 finally:236 _current_session.reset(token)237 238 return _bound239 240 241def _inject_tracing_hooks(options: Any, session: Optional[SessionState] = None) -> None:242 """Inject LangSmith tracing hooks into ClaudeAgentOptions.243 244 If *session* is provided, injected hook callables bind that session around245 each hook invocation. This is important because the Claude SDK may execute246 hooks in async contexts that do not inherit the ``receive_response``247 ContextVar; binding at hook injection time keeps each client isolated.248 """249 if not hasattr(options, "hooks"):250 return251 252 # Initialize hooks dict if not present253 if options.hooks is None:254 options.hooks = {}255 256 for event in (257 "PreToolUse",258 "PostToolUse",259 "PostToolUseFailure",260 "SubagentStart",261 "SubagentStop",262 ):263 if event not in options.hooks:264 options.hooks[event] = []265 266 try:267 from claude_agent_sdk import HookMatcher # type: ignore[import-not-found]268 269 langsmith_pre_matcher = HookMatcher(270 matcher=None, hooks=[_bind_hook_to_session(pre_tool_use_hook, session)]271 )272 langsmith_post_matcher = HookMatcher(273 matcher=None, hooks=[_bind_hook_to_session(post_tool_use_hook, session)]274 )275 langsmith_failure_matcher = HookMatcher(276 matcher=None,277 hooks=[_bind_hook_to_session(post_tool_use_failure_hook, session)],278 )279 langsmith_subagent_start_matcher = HookMatcher(280 matcher=None, hooks=[_bind_hook_to_session(subagent_start_hook, session)]281 )282 langsmith_subagent_stop_matcher = HookMatcher(283 matcher=None, hooks=[_bind_hook_to_session(subagent_stop_hook, session)]284 )285 286 options.hooks["PreToolUse"].insert(0, langsmith_pre_matcher)287 options.hooks["PostToolUse"].insert(0, langsmith_post_matcher)288 options.hooks["PostToolUseFailure"].insert(0, langsmith_failure_matcher)289 options.hooks["SubagentStart"].insert(0, langsmith_subagent_start_matcher)290 options.hooks["SubagentStop"].insert(0, langsmith_subagent_stop_matcher)291 292 logger.debug("Injected LangSmith tracing hooks into ClaudeAgentOptions")293 except ImportError:294 logger.warning("Failed to import HookMatcher from claude_agent_sdk")295 except Exception as e:296 logger.warning(f"Failed to inject tracing hooks: {e}")297 298 299def _wrap_tool_handler(300 original_handler: Any,301 session: Optional[SessionState] = None,302 tool_name: Optional[str] = None,303) -> Any:304 """Wrap an MCP tool handler to propagate LangSmith run context.305 306 The Claude SDK runs hooks and tool handlers in different async task307 contexts, so contextvars set in ``PreToolUse`` are invisible to the308 handler. This wrapper copies the active tool run into the contextvar before309 calling the original handler, so ``@traceable`` calls inside the handler310 nest correctly.311 """312 313 async def _wrapped(args: Any) -> Any:314 # The most recently added active tool run is the one PreToolUse just315 # created for this invocation. Prefer an explicitly bound client316 # session because tool handlers may run in an async context that did317 # not inherit _current_session.318 tool_run = _get_last_active_tool_run(session, args=args, tool_name=tool_name)319 if tool_run:320 token = _context._PARENT_RUN_TREE_REF.set(weakref.ref(tool_run))321 session_token = (322 _current_session.set(session) if session is not None else None323 )324 try:325 return await original_handler(args)326 finally:327 if session_token is not None:328 _current_session.reset(session_token)329 _context._PARENT_RUN_TREE_REF.reset(token)330 return await original_handler(args)331 332 _wrapped._langsmith_wrapped = True # type: ignore[attr-defined]333 _wrapped._langsmith_original_handler = original_handler # type: ignore[attr-defined]334 _wrapped._langsmith_session = session # type: ignore[attr-defined]335 _wrapped._langsmith_tool_name = tool_name # type: ignore[attr-defined]336 return _wrapped337 338 339def _tool_run_matches(run: Any, args: Any, tool_name: Optional[str]) -> bool:340 """Return whether *run* appears to be for this SDK MCP handler call.341 342 Matching is intentionally strict: we require both the tool name and the343 handler args to line up with what the ``PreToolUse`` hook recorded. This344 avoids cross-attributing a handler invocation to the wrong client's active345 tool run under concurrency.346 """347 if not tool_name:348 return False349 run_name = str(getattr(run, "name", ""))350 # SDK MCP tools show up in hook data as e.g. ``mcp__weather__get_weather``351 # while the handler only knows its short name ``get_weather``.352 name_matches = (353 tool_name == run_name or tool_name in run_name or run_name in tool_name354 )355 if not name_matches:356 return False357 inputs = getattr(run, "inputs", None)358 if not isinstance(inputs, dict):359 return False360 # PreToolUse stores {} when the tool had no inputs, otherwise361 # {"input": <tool_input>}. Normalise both sides before comparing.362 recorded = inputs.get("input", {}) if inputs else {}363 return recorded == (args or {})364 365 366def _newest_matching_tool_run(367 sessions: list[SessionState], args: Any, tool_name: Optional[str]368) -> Any:369 """Return the most recently created active tool run that matches."""370 candidates: list[tuple[float, Any]] = []371 for candidate_session in sessions:372 for run, start_time in candidate_session.active_tool_runs.values():373 if _tool_run_matches(run, args, tool_name):374 candidates.append((start_time, run))375 if not candidates:376 return None377 return max(candidates, key=lambda item: item[0])[1]378 379 380def _get_last_active_tool_run(381 session: Optional[SessionState] = None,382 *,383 args: Any = None,384 tool_name: Optional[str] = None,385) -> Any:386 """Return the active tool run for an SDK MCP handler, or None.387 388 Lookup order:389 390 1. The session explicitly bound to the handler (if any).391 2. The session bound to the current ContextVar.392 3. The module-level default session (unit tests / unbound callers).393 4. Any live client session — only used when the handler is unbound and the394 SDK invoked it in a detached async context. Requires an exact tool395 name + args match to avoid cross-attribution across clients.396 """397 from ._hooks import (398 _current_session,399 _current_session_or_default,400 _registered_sessions,401 )402 403 # If we have a specific session (explicitly bound, current-context, or the404 # test default), just return its newest active tool run. There is no405 # cross-client ambiguity at that point.406 def _newest_in(s: SessionState) -> Any:407 if not s.active_tool_runs:408 return None409 latest_id = max(410 s.active_tool_runs,411 key=lambda tid: s.active_tool_runs[tid][1],412 )413 return s.active_tool_runs[latest_id][0]414 415 if session is not None:416 return _newest_in(session)417 418 current_session = _current_session.get()419 if current_session is not None:420 run = _newest_in(current_session)421 if run is not None:422 return run423 424 default_session = _current_session_or_default()425 if default_session is not current_session:426 run = _newest_in(default_session)427 if run is not None:428 return run429 430 # Last resort: the SDK invoked this handler in a detached async context and431 # the handler object wasn't bound to a session. Require strict tool-name +432 # args match so concurrent clients can't steal each other's attribution.433 return _newest_matching_tool_run(_registered_sessions(), args, tool_name)434 435 436def instrument_claude_client(original_class: Any) -> None:437 """Patch ``ClaudeSDKClient`` **in place** to trace calls.438 439 In-place patching (rather than subclassing + reference replacement)440 ensures that callers who imported ``ClaudeSDKClient`` *before*441 ``configure_claude_agent_sdk()`` was called still get instrumented.442 """443 if getattr(original_class, "_langsmith_instrumented", False):444 return # Already wrapped, avoid double-tracing445 446 # ── stash originals ──────────────────────────────────────────────447 _orig_init = original_class.__init__448 _orig_query = original_class.query449 _orig_receive_response = original_class.receive_response450 451 # ── patched __init__ ─────────────────────────────────────────────452 def _traced_init(self: Any, *args: Any, **kwargs: Any) -> None:453 options = kwargs.get("options") or (args[0] if args else None)454 self._ls_session = SessionState()455 if options:456 _inject_tracing_hooks(options, self._ls_session)457 _orig_init(self, *args, **kwargs)458 self._ls_prompt = None459 self._ls_start_time = None460 self._ls_streamed_input = None461 462 # ── patched query ────────────────────────────────────────────────463 async def _traced_query(self: Any, *args: Any, **kwargs: Any) -> Any:464 self._ls_start_time = time.time()465 self._ls_streamed_input = None466 prompt = args[0] if args else kwargs.get("prompt")467 468 if prompt is None:469 pass470 elif isinstance(prompt, str):471 self._ls_prompt = prompt472 elif isinstance(prompt, AsyncIterable):473 collector: list[dict[str, Any]] = []474 self._ls_streamed_input = collector475 self._ls_prompt = None476 477 async def _gen_wrapper() -> AsyncGenerator[dict[str, Any], None]:478 async for msg in prompt:479 collector.append(msg)480 yield msg481 482 if args:483 args = (_gen_wrapper(),) + args[1:]484 else:485 kwargs["prompt"] = _gen_wrapper()486 else:487 self._ls_prompt = str(prompt)488 489 return await _orig_query(self, *args, **kwargs)490 491 # ── patched receive_response ─────────────────────────────────────492 async def _traced_receive_response(self: Any) -> AsyncGenerator[Any, None]:493 messages = _orig_receive_response(self)494 495 trace_inputs: dict[str, Any] = {}496 trace_metadata: dict[str, Any] = {497 "ls_integration": "claude-agent-sdk",498 "ls_integration_version": _get_package_version("claude_agent_sdk"),499 }500 501 awaiting_streamed_input = self._ls_streamed_input is not None502 503 if self._ls_prompt:504 trace_inputs["prompt"] = self._ls_prompt505 506 if hasattr(self, "options") and self.options:507 if hasattr(self.options, "system_prompt") and self.options.system_prompt:508 system_prompt = self.options.system_prompt509 if isinstance(system_prompt, str):510 trace_inputs["system"] = system_prompt511 elif isinstance(system_prompt, dict):512 if system_prompt.get("type") == "preset":513 preset_text = (514 f"preset: {system_prompt.get('preset', 'claude_code')}"515 )516 if "append" in system_prompt:517 preset_text += f"\nappend: {system_prompt['append']}"518 trace_inputs["system"] = preset_text519 else:520 trace_inputs["system"] = system_prompt521 522 for attr in ["model", "permission_mode", "max_turns"]:523 if hasattr(self.options, attr):524 val = getattr(self.options, attr)525 if val is not None:526 trace_metadata[attr] = val527 528 config = get_tracing_config()529 user_metadata = config.get("metadata") or {}530 531 trace_kwargs: dict[str, Any] = {532 "name": config.get("name") or TRACE_CHAIN_NAME,533 "run_type": "chain",534 "inputs": trace_inputs,535 "metadata": {536 **trace_metadata,537 **user_metadata,538 "ls_agent_type": "root",539 },540 }541 if config.get("project_name"):542 trace_kwargs["project_name"] = config["project_name"]543 if config.get("tags"):544 trace_kwargs["tags"] = config["tags"]545 546 async with trace(**trace_kwargs) as run:547 # Bind this client's state container to the ContextVar so stream548 # helpers on this SDK event loop pick it up (see549 # _hooks.SessionState). This keeps concurrent ClaudeSDKClient550 # instances — eval runs, FastAPI handlers, Celery workers,551 # asyncio.gather — from corrupting each other's correlation state.552 session = getattr(self, "_ls_session", None)553 if session is None:554 session = self._ls_session = SessionState()555 session_token = _register_session(session)556 _set_session_root(session, run)557 parent_token = set_parent_run_tree(run)558 tracker = TurnLifecycle(self._ls_start_time)559 collected_by_ctx: dict[Optional[str], list[dict[str, Any]]] = {None: []}560 561 prompt_for_llm: Any = self._ls_prompt562 563 try:564 async for msg in messages:565 if awaiting_streamed_input and self._ls_streamed_input:566 unwrapped_messages = unwrap_message_dicts(567 self._ls_streamed_input568 )569 if unwrapped_messages:570 run.inputs["messages"] = unwrapped_messages571 prompt_for_llm = self._ls_streamed_input572 awaiting_streamed_input = False573 574 msg_type = type(msg).__name__575 576 if msg_type == "AssistantMessage":577 parent_tool_use_id = getattr(msg, "parent_tool_use_id", None)578 llm_parent = (579 get_subagent_run_by_tool_id(parent_tool_use_id)580 if parent_tool_use_id581 else None582 )583 584 ctx_key = parent_tool_use_id585 ctx_history = collected_by_ctx.setdefault(ctx_key, [])586 587 content = tracker.start_llm_run(588 msg,589 prompt_for_llm if parent_tool_use_id is None else None,590 ctx_history,591 parent=llm_parent,592 )593 if content:594 ctx_history.append(content)595 596 elif msg_type == "UserMessage":597 parent_tool_use_id = getattr(msg, "parent_tool_use_id", None)598 ctx_key = parent_tool_use_id599 ctx_history = collected_by_ctx.setdefault(ctx_key, [])600 601 if hasattr(msg, "content"):602 flattened = flatten_content_blocks(msg.content)603 if (604 isinstance(flattened, list)605 and flattened606 and isinstance(flattened[0], dict)607 and flattened[0].get("type") == "tool_result"608 ):609 for block in flattened:610 tool_use_id = block.get("tool_use_id")611 ctx_history.append(612 {613 "role": "tool",614 "content": block.get("content", ""),615 "tool_call_id": tool_use_id,616 }617 )618 if (619 tool_use_id620 and tool_use_id in session.active_tool_runs621 ):622 tool_run, _ = session.active_tool_runs.pop(623 tool_use_id624 )625 result_content = block.get("content", "")626 is_error = block.get("is_error", False)627 tool_run.end(628 outputs={"output": result_content},629 error=str(result_content)630 if is_error631 else None,632 )633 try:634 tool_run.patch()635 except Exception as e:636 logger.warning(637 "Failed to patch"638 f" orphaned tool run: {e}"639 )640 else:641 ctx_history.append(642 {643 "content": flattened,644 "role": "user",645 }646 )647 tracker.mark_next_start()648 elif msg_type == "ResultMessage":649 meta = {650 k: v651 for k, v in {652 "num_turns": getattr(msg, "num_turns", None),653 "session_id": getattr(msg, "session_id", None),654 "duration_ms": getattr(msg, "duration_ms", None),655 "duration_api_ms": getattr(656 msg, "duration_api_ms", None657 ),658 "is_error": getattr(msg, "is_error", None),659 }.items()660 if v is not None661 }662 if meta:663 run.metadata.update(meta)664 665 yield msg666 main_collected = collected_by_ctx.get(None, [])667 run.end(outputs=main_collected[-1] if main_collected else None)668 except Exception:669 logger.exception("Error while tracing Claude Agent stream")670 finally:671 tracker.close()672 reconcile_from_transcripts(tracker, session=session)673 tracker.flush()674 clear_parent_run_tree(parent_token)675 try:676 clear_active_tool_runs(session)677 finally:678 _unregister_session(session, session_token)679 680 # ── apply patches to the class itself ────────────────────────────681 original_class.__init__ = _traced_init682 original_class.query = _traced_query683 original_class.receive_response = _traced_receive_response684 original_class._langsmith_instrumented = True685 686 687def instrument_sdk_mcp_tool(tool_class: Any) -> None:688 """Patch ``SdkMcpTool.__init__`` to auto-wrap handlers.689 690 Wrapping happens at construction time so that any tool created691 *after* ``configure_claude_agent_sdk()`` automatically gets692 run-context propagation, regardless of how ``tool`` or693 ``create_sdk_mcp_server`` were imported.694 """695 if getattr(tool_class, "_langsmith_handler_patched", False):696 return697 698 _orig_init = tool_class.__init__699 700 def _patched_init(self: Any, *args: Any, **kwargs: Any) -> None:701 _orig_init(self, *args, **kwargs)702 handler = self.handler703 if callable(handler) and not getattr(handler, "_langsmith_wrapped", False):704 self.handler = _wrap_tool_handler(705 handler, tool_name=getattr(self, "name", None)706 )707 708 tool_class.__init__ = _patched_init709 tool_class._langsmith_handler_patched = True710 