Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
_hooks.py565 linesDownload Raw Back to claude_agent_sdk
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 
codekingpro/portable-devtools · Team Ai