Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -203,6 +203,12 @@ def _sanitize_url_string(value: str) -> str:
)


def _is_mcp_tool_call(message: Any) -> bool:
request = message.get("request") if isinstance(message, dict) else None
mcp_message = request.get("message") if isinstance(request, dict) else None
return isinstance(mcp_message, dict) and mcp_message.get("method") == "tools/call"


class ClaudeAgentSdkCassetteTransport(Transport):
"""Record or replay the SDK<->CLI JSON protocol at the transport layer."""

Expand All @@ -213,6 +219,8 @@ def __init__(
prompt: str | Any,
options: ClaudeAgentOptions,
record_mode: str | None = None,
pause_before_sdk_messages: bool = False,
pause_before_mcp_tool_calls: bool = False,
) -> None:
_require_sdk()
self._cassette_name = cassette_name
Expand All @@ -227,6 +235,10 @@ def __init__(
self._ready = False
self._cursor_lock = anyio.Lock()
self._cursor_changed = anyio.Event()
self._mcp_tool_call_read = anyio.Event()
self._mcp_tool_call_gate = anyio.Event() if pause_before_mcp_tool_calls else None
self._sdk_message_blocked = anyio.Event() if pause_before_sdk_messages else None
self._sdk_message_gate = anyio.Event() if pause_before_sdk_messages else None
self._control_request_ids: dict[str, str] = {}

async def connect(self) -> None:
Expand Down Expand Up @@ -273,14 +285,25 @@ async def _read_messages_impl(self):
assert self._delegate is not None
async for message in self._delegate.read_messages():
self._events.append({"op": "read", "payload": message})
if _is_mcp_tool_call(message):
self._mcp_tool_call_read.set()
if self._mcp_tool_call_gate is not None:
await self._mcp_tool_call_gate.wait()
if self._sdk_message_gate is not None:
await self._wait_before_sdk_message(message)
yield message
return

while True:
event = await self._wait_for_event("read", allow_eof=True)
if event is None:
return
yield self._remap_read_message(event["payload"])
message = self._remap_read_message(event["payload"])
if _is_mcp_tool_call(message) and self._mcp_tool_call_gate is not None:
await self._mcp_tool_call_gate.wait()
if self._sdk_message_gate is not None:
await self._wait_before_sdk_message(message)
yield message

async def close(self) -> None:
self._ready = False
Expand All @@ -300,11 +323,36 @@ async def close(self) -> None:
def is_ready(self) -> bool:
return self._ready

async def wait_for_mcp_tool_call(self) -> None:
"""Wait until replay delivers a local MCP ``tools/call`` request."""
await self._mcp_tool_call_read.wait()

def release_mcp_tool_calls(self) -> None:
assert self._mcp_tool_call_gate is not None
self._mcp_tool_call_gate.set()

async def wait_until_sdk_message_blocked(self) -> None:
assert self._sdk_message_blocked is not None
await self._sdk_message_blocked.wait()

def release_sdk_messages(self) -> None:
assert self._sdk_message_gate is not None
self._sdk_message_gate.set()

async def end_input(self) -> None:
if self._recording:
assert self._delegate is not None
await self._delegate.end_input()

async def _wait_before_sdk_message(self, message: Any) -> None:
assert self._sdk_message_gate is not None
assert self._sdk_message_blocked is not None
message_type = message.get("type") if isinstance(message, dict) else None
if message_type in {"control_response", "control_request", "control_cancel_request"}:
return
self._sdk_message_blocked.set()
await self._sdk_message_gate.wait()

def _should_replay(self) -> bool:
if self._record_mode == "all":
return False
Expand All @@ -326,6 +374,8 @@ async def _wait_for_event(self, op: str, *, allow_eof: bool = False) -> dict[str
old_event = self._cursor_changed
self._cursor_changed = anyio.Event()
old_event.set()
if op == "read" and _is_mcp_tool_call(event["payload"]):
self._mcp_tool_call_read.set()
return event

waiter = self._cursor_changed
Expand Down Expand Up @@ -381,10 +431,14 @@ def make_cassette_transport(
prompt: str | Any,
options: "ClaudeAgentOptions",
record_mode: str | None = None,
pause_before_sdk_messages: bool = False,
pause_before_mcp_tool_calls: bool = False,
) -> ClaudeAgentSdkCassetteTransport:
return ClaudeAgentSdkCassetteTransport(
cassette_name=cassette_name,
prompt=prompt,
options=options,
record_mode=record_mode,
pause_before_sdk_messages=pause_before_sdk_messages,
pause_before_mcp_tool_calls=pause_before_mcp_tool_calls,
)
14 changes: 12 additions & 2 deletions py/src/braintrust/integrations/claude_agent_sdk/integration.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,12 @@

from braintrust.integrations.base import BaseIntegration

from .patchers import ClaudeSDKClientPatcher, ClaudeSDKQueryPatcher, SdkMcpToolPatcher
from .patchers import (
ClaudeSDKClientPatcher,
ClaudeSDKMessageReaderPatcher,
ClaudeSDKQueryPatcher,
SdkMcpToolPatcher,
)


class ClaudeAgentSDKIntegration(BaseIntegration):
Expand All @@ -11,4 +16,9 @@ class ClaudeAgentSDKIntegration(BaseIntegration):
name = "claude_agent_sdk"
import_names = ("claude_agent_sdk",)
min_version = "0.1.10"
patchers = (ClaudeSDKClientPatcher, ClaudeSDKQueryPatcher, SdkMcpToolPatcher)
patchers = (
ClaudeSDKMessageReaderPatcher,
ClaudeSDKClientPatcher,
ClaudeSDKQueryPatcher,
SdkMcpToolPatcher,
)
21 changes: 18 additions & 3 deletions py/src/braintrust/integrations/claude_agent_sdk/patchers.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,23 @@
"""Claude Agent SDK patchers — replacement patchers for ClaudeSDKClient, query, and SdkMcpTool."""
"""Claude Agent SDK patchers for message reading, clients, queries, and SDK MCP tools."""

from braintrust.integrations.base import ClassReplacementPatcher
from braintrust.integrations.base import ClassReplacementPatcher, FunctionWrapperPatcher

from .tracing import _create_client_wrapper_class, _create_query_wrapper_function, _create_tool_wrapper_class
from .tracing import (
_create_client_wrapper_class,
_create_query_wrapper_function,
_create_tool_wrapper_class,
_wrap_query_read_messages,
)


class ClaudeSDKMessageReaderPatcher(FunctionWrapperPatcher):
"""Observe SDK messages before local MCP control requests can dispatch."""

name = "claude_agent_sdk.message_reader"
target_module = "claude_agent_sdk._internal.query"
target_path = "Query._read_messages"
wrapper = staticmethod(_wrap_query_read_messages)
priority = 50


class ClaudeSDKClientPatcher(ClassReplacementPatcher):
Expand Down
Loading