Team Ai
Datasetpublic

codekingpro/portable-devtools

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