codekingpro/portable-devtools
114k
1"""Hook-based tool tracing for Claude Agent SDK.2 3Correlation state is scoped **per client session** via a4:class:`contextvars.ContextVar`. Each instrumented ``ClaudeSDKClient`` owns a5:class:`SessionState`; ``receive_response()`` binds it while processing the6stream so helper functions can look up the right state regardless of how many7clients are concurrently active in the process.8 9Hooks injected by ``_client.py`` are also bound to their owning10``SessionState`` so hook callbacks use the correct state even if the SDK runs11them in an async context that did not inherit ``receive_response``'s12ContextVar.13 14When no ContextVar is active, hooks use a module-level default session. This is15primarily for direct unit tests; real traffic under ``receive_response`` uses a16client-bound session.17"""18 19import logging20import threading21import time22import weakref23from contextvars import ContextVar24from dataclasses import dataclass, field25from datetime import datetime, timezone26from typing import TYPE_CHECKING, Any, Optional27 28from langsmith.run_helpers import get_current_run_tree29from langsmith.run_trees import RunTree30 31from ._tools import get_parent_run_tree32 33if TYPE_CHECKING:34 from claude_agent_sdk import (35 HookContext,36 HookInput,37 HookJSONOutput,38 )39 40logger = logging.getLogger(__name__)41 42 43# ── Per-session state ─────────────────────────────────────────────────────────44 45 46@dataclass47class SessionState:48 """All mutable correlation state for a single conversation.49 50 One instance is created per instrumented ``ClaudeSDKClient`` and bound to51 the ``_current_session`` ContextVar while that client is active.52 """53 54 # Key: tool_use_id → (run_tree, start_time)55 active_tool_runs: dict[str, tuple[Any, float]] = field(default_factory=dict)56 57 # Key: agent_id → RunTree for the subagent chain.58 # Populated by SubagentStart, consumed by SubagentStop.59 subagent_runs: dict[str, RunTree] = field(default_factory=dict)60 61 # Key: tool_use_id → tool_input dict.62 # When PreToolUse fires for an "Agent" tool, it stashes here.63 # SubagentStart pops it to find the matching Agent tool run.64 pending_agent_tools: dict[str, dict[str, Any]] = field(default_factory=dict)65 66 # Key: agent_id → Agent tool_use_id.67 # Maps a subagent back to the Agent tool that spawned it.68 agent_to_tool_mapping: dict[str, str] = field(default_factory=dict)69 70 # Key: Agent tool_use_id → RunTree.71 # SubagentStop moves the run here; PostToolUse sets outputs on it;72 # clear_active_tool_runs() ends + patches it.73 ended_subagent_runs: dict[str, RunTree] = field(default_factory=dict)74 75 # (transcript_path, subagent_RunTree) captured from SubagentStop.76 # Used for usage extraction and creating missing LLM runs.77 subagent_transcript_paths: list[tuple[str, RunTree]] = field(default_factory=list)78 79 # Main session transcript path, captured from BaseHookInput.transcript_path80 # on the first hook that fires (every hook inherits this field).81 main_transcript_path: Optional[str] = None82 83 # Root LangSmith run used for parenting root-level hook spans.84 root_run: Optional[RunTree] = None85 86 87# Module-level *default* session. Used when no ContextVar is set (e.g. tests88# that poke hooks directly, or hooks firing outside a traced conversation).89_default_session: SessionState = SessionState()90 91# ContextVar holding the active session for a conversation. Injected hook92# callables bind this explicitly before calling the shared hook function.93_current_session: ContextVar[Optional[SessionState]] = ContextVar(94 "langsmith_claude_agent_session", default=None95)96 97# Live sessions are only used by SDK MCP tool handlers when the SDK invokes the98# handler in a detached async context that did not inherit _current_session.99# Store weak values so this fallback registry never owns session lifetime.100_live_sessions_lock = threading.Lock()101_live_sessions: weakref.WeakValueDictionary[int, SessionState] = (102 weakref.WeakValueDictionary()103)104 105 106def _current_session_or_default() -> SessionState:107 """Return the session bound to the current context, or the default."""108 session = _current_session.get()109 if session is not None:110 return session111 return _default_session112 113 114def _session_for_hook() -> SessionState:115 """Resolve the session that owns the current hook invocation.116 117 Real Claude SDK hook invocations are wrapped by ``_bind_hook_to_session``118 in ``_client.py``, so the ContextVar should be set. The default session is119 only for tests or direct, unbound hook calls.120 """121 return _current_session_or_default()122 123 124def _register_session(session: SessionState) -> object:125 """Bind *session* to the ContextVar and return a reset token.126 127 The caller must pass the returned token to ``_unregister_session`` when128 the conversation ends.129 """130 with _live_sessions_lock:131 _live_sessions[id(session)] = session132 return _current_session.set(session)133 134 135def _set_session_root(session: SessionState, run_tree: RunTree) -> None:136 """Store the root LangSmith run for *session*."""137 session.root_run = run_tree138 139 140def _unregister_session(session: SessionState, token: Any) -> None:141 """Reset the ContextVar for the current session and drop the live entry."""142 try:143 _current_session.reset(token)144 except ValueError:145 # Token was created in a different context. Don't clobber an unrelated146 # current value — just log and continue. The live-sessions registry147 # below is still cleaned up so matching won't find a stale session.148 logger.debug("Could not reset _current_session with token from another context")149 finally:150 with _live_sessions_lock:151 _live_sessions.pop(id(session), None)152 153 154def _registered_sessions() -> list[SessionState]:155 """Return currently active client sessions."""156 with _live_sessions_lock:157 return list(_live_sessions.values())158 159 160# ── Public helpers (used by _client.py) ───────────────────────────────────────161 162 163def get_subagent_run_by_tool_id(tool_use_id: str) -> Optional[RunTree]:164 """Get a subagent run by the Agent tool's tool_use_id.165 166 Checks both active subagent runs and ended-but-not-finalised runs,167 because the SDK fires ``SubagentStop`` before the subagent's messages168 reach the client.169 """170 session = _current_session_or_default()171 # Check active subagents first172 for aid, tid in session.agent_to_tool_mapping.items():173 if tid == tool_use_id:174 return session.subagent_runs.get(aid)175 # Fall back to ended-but-not-finalised subagents176 return session.ended_subagent_runs.get(tool_use_id)177 178 179# ── Hook functions ────────────────────────────────────────────────────────────180 181 182async def pre_tool_use_hook(183 input_data: "HookInput",184 tool_use_id: Optional[str],185 context: "HookContext",186) -> "HookJSONOutput":187 """Trace tool execution before it starts.188 189 Args:190 input_data: Contains `tool_name`, `tool_input`, `session_id`, `agent_id`191 tool_use_id: Unique identifier for this tool invocation192 context: Hook context (currently contains only signal)193 194 Returns:195 Hook output (empty dict allows execution to proceed)196 """197 if not tool_use_id:198 return {}199 200 data: dict[str, Any] = dict(input_data) # flatten TypedDict union201 tool_name: str = str(data.get("tool_name", "unknown_tool"))202 tool_input: dict[str, Any] = dict(data.get("tool_input") or {})203 agent_id: Optional[str] = str(data["agent_id"]) if data.get("agent_id") else None204 session = _session_for_hook()205 206 # Capture main session transcript path from BaseHookInput207 if session.main_transcript_path is None and data.get("transcript_path"):208 session.main_transcript_path = str(data["transcript_path"])209 210 # If this is an Agent tool call, record it so SubagentStart can find it211 if tool_name == "Agent":212 session.pending_agent_tools[tool_use_id] = tool_input213 214 try:215 # Determine parent: subagent chain > root chain.216 # Tool runs are siblings of LLM runs, not children.217 parent: Optional[RunTree] = None218 if agent_id and agent_id in session.subagent_runs:219 parent = session.subagent_runs[agent_id]220 else:221 parent = (222 session.root_run223 if session is not _default_session224 else get_parent_run_tree()225 ) or get_current_run_tree()226 227 if not parent:228 return {}229 230 start_time = time.time()231 tool_run = parent.create_child(232 name=tool_name,233 run_type="tool",234 inputs={"input": tool_input} if tool_input else {},235 start_time=datetime.fromtimestamp(start_time, tz=timezone.utc),236 )237 238 try:239 tool_run.post()240 except Exception as e:241 logger.warning(f"Failed to post tool run for {tool_name}: {e}")242 243 session.active_tool_runs[tool_use_id] = (tool_run, start_time)244 245 except Exception as e:246 logger.warning(f"Error in PreToolUse hook for {tool_name}: {e}", exc_info=True)247 248 return {}249 250 251async def post_tool_use_hook(252 input_data: "HookInput",253 tool_use_id: Optional[str],254 context: "HookContext",255) -> "HookJSONOutput":256 """Trace tool execution after it completes.257 258 Args:259 input_data: Contains `tool_name`, `tool_input`, `tool_response`,260 `session_id`, etc.261 tool_use_id: Unique identifier for this tool invocation262 context: Hook context (currently contains only signal)263 264 Returns:265 Hook output (empty `dict` by default)266 """267 if not tool_use_id:268 return {}269 270 tool_name: str = str(input_data.get("tool_name", "unknown_tool"))271 tool_response = input_data.get("tool_response")272 session = _session_for_hook()273 274 try:275 run_info = session.active_tool_runs.pop(tool_use_id, None)276 if not run_info:277 return {}278 279 tool_run, _ = run_info280 281 if isinstance(tool_response, dict):282 outputs = tool_response283 elif isinstance(tool_response, list):284 outputs = {"content": tool_response}285 else:286 outputs = {"output": str(tool_response)} if tool_response else {}287 288 # Check if the tool execution was an error289 is_error = False290 if isinstance(tool_response, dict):291 is_error = tool_response.get("is_error", False)292 293 tool_run.end(294 outputs=outputs,295 error=outputs.get("output") if is_error else None,296 )297 298 try:299 tool_run.patch()300 except Exception as e:301 logger.warning(f"Failed to patch tool run for {tool_name}: {e}")302 303 # If this is an Agent tool, also set outputs on the stashed304 # subagent run. We don't end/patch the subagent here because305 # its AssistantMessages may not have been yielded to306 # receive_response() yet. clear_active_tool_runs() will307 # finalise it at the end of the conversation.308 subagent_run = session.ended_subagent_runs.get(tool_use_id)309 if subagent_run:310 try:311 subagent_run.outputs = outputs312 except Exception as e:313 logger.warning(f"Failed to set subagent run outputs: {e}")314 315 except Exception as e:316 logger.warning(317 f"Error in PostToolUse hook for {tool_name}: {e}",318 exc_info=True,319 )320 321 return {}322 323 324async def post_tool_use_failure_hook(325 input_data: "HookInput",326 tool_use_id: Optional[str],327 context: "HookContext",328) -> "HookJSONOutput":329 """Trace tool execution when it fails.330 331 This hook fires for built-in tool failures (Bash, Read, Write, etc.)332 and is mutually exclusive with :func:`post_tool_use_hook` — when a333 built-in tool fails, only ``PostToolUseFailure`` fires.334 335 Args:336 input_data: Contains ``tool_name``, ``tool_input``, ``error``,337 and optionally ``is_interrupt``.338 tool_use_id: Unique identifier for this tool invocation339 context: Hook context (currently contains only signal)340 341 Returns:342 Hook output (empty dict)343 """344 if not tool_use_id:345 return {}346 347 tool_name: str = str(input_data.get("tool_name", "unknown_tool"))348 error: str = str(input_data.get("error", "Unknown error"))349 session = _session_for_hook()350 351 try:352 run_info = session.active_tool_runs.pop(tool_use_id, None)353 if not run_info:354 return {}355 356 tool_run, _ = run_info357 358 tool_run.end(359 outputs={"error": error},360 error=error,361 )362 363 try:364 tool_run.patch()365 except Exception as e:366 logger.warning(f"Failed to patch failed tool run for {tool_name}: {e}")367 368 except Exception as e:369 logger.warning(370 f"Error in PostToolUseFailure hook for {tool_name}: {e}",371 exc_info=True,372 )373 374 return {}375 376 377async def subagent_start_hook(378 input_data: "HookInput",379 tool_use_id: Optional[str],380 context: "HookContext",381) -> "HookJSONOutput":382 """Create a chain run when a subagent starts.383 384 The subagent chain is nested under the Agent tool run that spawned it.385 Since the SDK passes a different ``tool_use_id`` to this hook than the386 one from ``PreToolUse`` for the Agent tool, we match them via the387 ``_pending_agent_tools`` queue.388 389 Args:390 input_data: Contains ``agent_id``, ``agent_type``, ``session_id``391 tool_use_id: SDK-internal session id (not the Agent tool's392 tool_use_id)393 context: Hook context394 395 Returns:396 Hook output (empty dict)397 """398 data: dict[str, Any] = dict(input_data)399 agent_id: Optional[str] = str(data["agent_id"]) if data.get("agent_id") else None400 agent_type: str = str(data.get("agent_type") or "subagent")401 session = _session_for_hook()402 403 if not agent_id:404 return {}405 406 try:407 # Find the Agent tool run that triggered this subagent.408 # pending_agent_tools is populated by pre_tool_use_hook when409 # tool_name == "Agent". Pop the most recent one.410 agent_tool_use_id: Optional[str] = None411 agent_tool_input: dict[str, Any] = {}412 parent: Optional[RunTree] = None413 414 if session.pending_agent_tools:415 agent_tool_use_id, agent_tool_input = session.pending_agent_tools.popitem()416 417 if agent_tool_use_id in session.active_tool_runs:418 agent_tool_run, _ = session.active_tool_runs[agent_tool_use_id]419 parent = agent_tool_run420 421 if parent is None:422 parent = (423 session.root_run424 if session is not _default_session425 else get_parent_run_tree()426 ) or get_current_run_tree()427 428 if not parent:429 return {}430 431 start_time = time.time()432 subagent_run = parent.create_child(433 name=agent_type,434 run_type="chain",435 inputs=agent_tool_input if agent_tool_input else {},436 start_time=datetime.fromtimestamp(start_time, tz=timezone.utc),437 )438 subagent_run.extra["metadata"] = {439 **subagent_run.extra.get("metadata", {}),440 "ls_agent_type": "subagent",441 }442 443 try:444 subagent_run.post()445 except Exception as e:446 logger.warning(f"Failed to post subagent run: {e}")447 448 # Store by agent_id so tool hooks and LLM run lookup can find it449 session.subagent_runs[agent_id] = subagent_run450 451 # Remember which Agent tool_use_id spawned this agent_id452 if agent_tool_use_id:453 session.agent_to_tool_mapping[agent_id] = agent_tool_use_id454 455 except Exception as e:456 logger.warning(f"Error in SubagentStart hook: {e}", exc_info=True)457 458 return {}459 460 461async def subagent_stop_hook(462 input_data: "HookInput",463 tool_use_id: Optional[str],464 context: "HookContext",465) -> "HookJSONOutput":466 """Move the subagent run to ended state when it finishes.467 468 Does NOT end/patch the run — ``PostToolUse`` for the Agent tool will469 set outputs, and ``clear_active_tool_runs()`` will finalise it at the470 end of the conversation.471 472 Args:473 input_data: Contains ``agent_id``, ``agent_type``, ``session_id``,474 ``agent_transcript_path``475 tool_use_id: SDK-internal session id476 context: Hook context477 478 Returns:479 Hook output (empty dict)480 """481 data: dict[str, Any] = dict(input_data)482 agent_id: Optional[str] = str(data["agent_id"]) if data.get("agent_id") else None483 transcript_path: Optional[str] = (484 str(data["agent_transcript_path"])485 if data.get("agent_transcript_path")486 else None487 )488 session = _session_for_hook()489 490 if not agent_id:491 return {}492 493 try:494 subagent_run = session.subagent_runs.pop(agent_id, None)495 if not subagent_run:496 return {}497 498 if transcript_path:499 session.subagent_transcript_paths.append((transcript_path, subagent_run))500 501 # Move to ended state so PostToolUse can set outputs.502 agent_tool_id = session.agent_to_tool_mapping.pop(agent_id, None)503 if agent_tool_id:504 session.ended_subagent_runs[agent_tool_id] = subagent_run505 else:506 # No matching Agent tool — just end it now507 subagent_run.end()508 try:509 subagent_run.patch()510 except Exception as e:511 logger.warning(f"Failed to patch subagent run: {e}")512 513 except Exception as e:514 logger.warning(f"Error in SubagentStop hook: {e}", exc_info=True)515 516 return {}517 518 519# ── Cleanup ───────────────────────────────────────────────────────────────────520 521 522def clear_active_tool_runs(session: Optional[SessionState] = None) -> None:523 """Finalise all runs and clear state for *session*.524 525 If *session* is omitted the current ContextVar-bound session is used526 (falling back to the module-level default session). ``receive_response``527 passes the per-call session explicitly.528 """529 if session is None:530 session = _current_session_or_default()531 532 # 1. End orphaned subagent runs (SubagentStop never fired)533 for agent_id, subagent_run in session.subagent_runs.items():534 try:535 subagent_run.end(error="Subagent run not completed (conversation ended)")536 subagent_run.patch()537 except Exception as e:538 logger.debug(f"Failed to clean up orphaned subagent run {agent_id}: {e}")539 540 # 2. Finalise ended subagent runs (outputs already set by PostToolUse)541 for tool_use_id, subagent_run in session.ended_subagent_runs.items():542 try:543 subagent_run.end()544 subagent_run.patch()545 except Exception as e:546 logger.debug(f"Failed to finalise ended subagent run {tool_use_id}: {e}")547 548 # 3. End orphaned tool runs549 for tool_use_id, (tool_run, _) in session.active_tool_runs.items():550 try:551 tool_run.end(error="Tool run not completed (conversation ended)")552 tool_run.patch()553 except Exception as e:554 logger.debug(f"Failed to clean up orphaned tool run {tool_use_id}: {e}")555 556 # 4. Reset session state557 session.active_tool_runs.clear()558 session.subagent_runs.clear()559 session.pending_agent_tools.clear()560 session.agent_to_tool_mapping.clear()561 session.ended_subagent_runs.clear()562 session.subagent_transcript_paths.clear()563 session.main_transcript_path = None564 session.root_run = None565 