"""Claude Agent SDK backend adapter. One :class:`ClaudeSdkBackend` per :class:`ClaudeAgent`. A live session is one ``ClaudeSDKClient`` (one claude subprocess) and runs one turn at a time; the sessions of every agent live in one shared :class:`~beaver_gateway.core.sessions.SessionPool` that owns TTL and memory-pressure eviction. Sessions are keyed by ``conversation_id`` when the caller passes one or by a text-only fingerprint of ``messages[:-1]`` for stateless callers (``/v1/messages``). Without a live session the adapter resumes ``session_id`` from the session store (after closing any ``tool_use`` left open by a crash), or seeds the incoming history into the store via ``core/transcript`` and resumes that. Per-turn ``**options`` beyond the protocol's: ``kind`` (conversation kind, picks the prompt assembly and the pool TTL), ``pinned`` (never evicted), ``tools=False`` (no MCP at all - forks and jobs), ``observer`` (callback receiving every raw SDK message, subagent ones included), ``turn_id``. Events on the wire are the Anthropic ``MessageStreamEvent`` family: one ``message_start``/``message_stop`` envelope per turn, block indices rebased across the API calls claude makes inside the turn. Process isolation: claude is spawned through a small exec wrapper that drops every inherited environment variable outside a whitelist and, when ``RunnerConfig.user`` is set, switches to that uid before exec (done in the wrapper rather than via ``subprocess(user=...)``, which uvloop rejects). """ from __future__ import annotations import asyncio import contextlib import fnmatch import hashlib import json import logging import os import pwd import shutil import sys import tempfile import time import uuid from collections.abc import Mapping from dataclasses import dataclass, field from pathlib import Path from typing import TYPE_CHECKING, Any, Self, cast import claude_agent_sdk from claude_agent_sdk import ( AssistantMessage, ClaudeAgentOptions, ClaudeSDKClient, MirrorErrorMessage, ResultMessage, StreamEvent, TextBlock, ThinkingBlock, ToolResultBlock, ToolUseBlock, UserMessage, project_key_for_directory, ) from beaver_gateway.core import prompt as prompt_assembly from beaver_gateway.core.events import ( StopReason, build_content_block_stop, build_input_json_delta, build_message_delta, build_message_start, build_message_stop, build_signature_delta, build_text_block_start, build_text_delta, build_thinking_block_start, build_thinking_delta, build_tool_use_block_start, ) from beaver_gateway.core.sessions import Session, SessionClient, SessionPool from beaver_gateway.core.transcript import ( build_entries, close_open_tool_uses, fingerprint, text_of, ) from beaver_gateway.core.turn_capture import TurnCapture, TurnUsage if TYPE_CHECKING: from collections.abc import AsyncIterator, Awaitable, Callable, Iterable, Sequence from anthropic.types import MessageParam from claude_agent_sdk import McpSdkServerConfig, SessionStore from beaver_gateway.agents.base import BaseAgent from beaver_gateway.agents.claude import ClaudeAgent from beaver_gateway.core.events import MessageStreamEvent _log = logging.getLogger("beaver_gateway.backends.claude_sdk") __all__ = [ "ClaudeSdkBackend", "RunnerConfig", "SessionClient", "ToolServerFactory", "UsageSink", "fingerprint", ] ENV_KEEP: tuple[str, ...] = ( "PATH", "HOME", "LANG", "LANGUAGE", "LC_ALL", "LC_CTYPE", "TZ", "TERM", "USER", "LOGNAME", "SHELL", "TMPDIR", "PWD", "SSL_CERT_FILE", "SSL_CERT_DIR", "NODE_EXTRA_CA_CERTS", "NODE_OPTIONS", "IS_SANDBOX", "HTTP_PROXY", "HTTPS_PROXY", "NO_PROXY", "http_proxy", "https_proxy", "no_proxy", ) ENV_KEEP_PREFIXES: tuple[str, ...] = ("CLAUDE_", "ANTHROPIC_", "DISABLE_") _STOP_REASONS: dict[str, StopReason] = { "end_turn": "end_turn", "tool_use": "tool_use", "max_tokens": "max_tokens", "stop_sequence": "stop_sequence", "refusal": "refusal", "pause_turn": "pause_turn", } ClientFactory = "Callable[[ClaudeAgentOptions], SessionClient]" UsageSink = "Callable[[UsageEvent], Awaitable[None]]" ToolServerFactory = "Callable[[str, str], McpSdkServerConfig | None]" """``(conversation_key, kind) -> in-process MCP server config`` or ``None``.""" @dataclass(frozen=True, slots=True) class RunnerConfig: user: str | None = None home: Path | None = None @dataclass(frozen=True, slots=True) class UsageEvent: agent_name: str model: str effort: str | None conversation_id: str | None session_id: str | None usage: TurnUsage class _RunnerClient(ClaudeSDKClient): """``ClaudeSDKClient`` that hands the materialized resume dir to the runner uid.""" def __init__(self, options: ClaudeAgentOptions, *, uid: int | None) -> None: super().__init__(options=options) self._runner_uid = uid async def _connect_inner(self, prompt: Any, actual_prompt: Any) -> None: materialized = self._materialized if materialized is not None and self._runner_uid is not None: _chown_tree(materialized.config_dir, self._runner_uid) await super()._connect_inner(prompt, actual_prompt) class ClaudeSdkBackend: def __init__( self, *, agent: ClaudeAgent, mcp_internal_urls: Mapping[str, str], session_store: SessionStore, mcp_tool_names: Mapping[str, Sequence[str]] | None = None, runner: RunnerConfig | None = None, usage_sink: Callable[[UsageEvent], Awaitable[None]] | None = None, client_factory: Callable[[ClaudeAgentOptions], SessionClient] | None = None, work_dir: Path | None = None, pool: SessionPool | None = None, tool_server: Callable[[str, str], McpSdkServerConfig | None] | None = None, ) -> None: self._agent = agent self._store = session_store self._runner = runner or RunnerConfig() self._usage_sink = usage_sink self._factory = client_factory or self._default_factory self._work_dir = work_dir or Path(tempfile.gettempdir()) / "beaver-claude" self._servers = _mcp_servers(agent, mcp_internal_urls) self._mcp_disallowed = _mcp_disallowed(agent, mcp_tool_names or {}) self._pool = pool if pool is not None else SessionPool() self._tool_server = tool_server self._uid, self._gid = _resolve_ids(self._runner.user) self._wrapper: Path | None = None @property def agent(self) -> ClaudeAgent: return self._agent @property def pool(self) -> SessionPool: return self._pool @property def sessions(self) -> dict[str, dict[str, Any]]: return { row["key"]: row for row in self._pool.snapshot() if row["agent"] == self._agent.name } async def __aenter__(self) -> Self: return self async def __aexit__(self, exc_type: object, exc: object, tb: object) -> None: await self.aclose() async def aclose(self) -> None: await self._pool.close_all(agent=self._agent.name) async def close(self, key: str) -> None: await self._pool.close(key) async def interrupt(self, key: str) -> bool: live = self._pool.get(key) if live is None or not live.busy: return False live.interrupt_requested = True await live.client.interrupt() return True def live(self, key: str) -> Session | None: return self._pool.get(key) async def repair_session(self, session_id: str) -> int: """Close ``tool_use`` blocks a crash left without a result; count added.""" key = self._store_key(session_id) entries = await self._store.load(cast("Any", key)) if not entries: return 0 fixes = close_open_tool_uses(cast("list[Mapping[str, Any]]", entries)) if fixes: await self._store.append(cast("Any", key), cast("Any", fixes)) _log.warning( "session %s: closed %d open tool_use with synthetic results", session_id, len(fixes), ) return len(fixes) async def complete( self, *, agent: BaseAgent, messages: Iterable[MessageParam], system: str | None = None, # noqa: ARG002 - the agent owns its prompt conversation_id: str | None = None, session_id: str | None = None, capture: TurnCapture | None = None, kind: str = "deep", pinned: bool = False, tools: bool = True, observer: Callable[[Any], None] | None = None, turn_id: str | None = None, **options: Any, # noqa: ARG002 - per-request knobs are not supported ) -> AsyncIterator[MessageStreamEvent]: if agent.name != self._agent.name: msg = f"backend bound to {self._agent.name!r}, got {agent.name!r}" raise ValueError(msg) history = [dict(m) for m in messages] if not history or history[-1].get("role") != "user": msg = "the last message must be a user turn" raise ValueError(msg) prompt = _prompt_text(history[-1].get("content")) prior = history[:-1] key = conversation_id or fingerprint(prior) spec = _SessionSpec(kind=kind, pinned=pinned, tools=tools) live = await self._acquire(key, session_id=session_id, history=prior, spec=spec) message_id = f"msg_{uuid.uuid4().hex}" yield build_message_start(message_id=message_id, model=self._agent.model) async with live.lock: live.running_turn = turn_id or message_id live.last_used = time.monotonic() try: turn = await self._run_turn(live, prompt, observer=observer) except Exception: live.running_turn = None await self._pool.close(key) if not (live.resumed and live.turns == 0): raise _log.exception( "resume of %s failed, reseeding from history", live.session_id ) live = await self._acquire( key, session_id=None, history=prior, spec=spec ) async with live.lock: live.running_turn = turn_id or message_id turn = await self._run_turn(live, prompt, observer=observer) for event in turn.events: yield event live.turns += 1 live.last_used = time.monotonic() live.running_turn = None usage, interrupted = await self._after_turn( live, turn, conversation_id=conversation_id, history=history, capture=capture, ) if turn.result is not None and turn.result.is_error and not interrupted: msg = f"claude: {turn.result.result or turn.result.subtype}" raise RuntimeError(msg) yield build_message_delta( stop_reason=turn.stop_reason, usage=_wire_usage(usage) ) yield build_message_stop() async def _after_turn( self, live: Session, turn: _Turn, *, conversation_id: str | None, history: list[dict[str, Any]], capture: TurnCapture | None, ) -> tuple[TurnUsage, bool]: if turn.result is not None and turn.result.session_id: live.session_id = turn.result.session_id if conversation_id is None: self._rekey(live.key, fingerprint([*history, *turn.synthesized])) usage = _usage_of(turn.result) interrupted = live.interrupt_requested live.interrupt_requested = False if capture is not None: capture.synthesized_messages = turn.synthesized capture.session_id = live.session_id capture.usage = usage capture.interrupted = interrupted if self._usage_sink is not None: await self._usage_sink( UsageEvent( agent_name=self._agent.name, model=self._agent.model, effort=self._agent.options.effort, conversation_id=conversation_id, session_id=live.session_id, usage=usage, ) ) return usage, interrupted async def _run_turn( self, live: Session, prompt: str, *, observer: Callable[[Any], None] | None = None, ) -> _Turn: streaming = self._agent.options.include_partial_messages turn = _Turn() raw: list[Any] = [] next_index = 0 offset = 0 await live.client.query(prompt) async for message in live.client.receive_response(): if observer is not None: observer(message) if isinstance(message, MirrorErrorMessage): live.dirty = True _log.error( "session %s: mirror error, marked dirty: %s", live.session_id, message.error, ) continue if getattr(message, "parent_tool_use_id", None) is not None: continue if isinstance(message, StreamEvent): event = message.event if event.get("type") == "message_start": offset = next_index continue index = event.get("index") if isinstance(index, int): next_index = max(next_index, offset + index + 1) if streaming: turn.events.extend(_emit_stream_event(event, offset + index)) elif isinstance(message, AssistantMessage): raw.append(message) if not streaming: for block in message.content: turn.events.extend(_emit_block(block, next_index)) next_index += 1 elif isinstance(message, UserMessage): raw.append(message) elif isinstance(message, ResultMessage): turn.result = message turn.stop_reason = _STOP_REASONS.get( message.stop_reason or "", "end_turn" ) turn.synthesized = synthesize_turn_messages(raw) _log.info( "turn: agent=%s session=%s events=%d synthesized=%d stop=%s", self._agent.name, live.session_id, len(turn.events), len(turn.synthesized), turn.stop_reason, ) return turn async def _acquire( self, key: str, *, session_id: str | None, history: list[dict[str, Any]], spec: _SessionSpec, ) -> Session: live = self._pool.get(key) if live is not None: return live resume = session_id if resume is not None: await self.repair_session(resume) elif history: resume = await self._seed(history) await self._pool.make_room() live = await self._spawn(resume, key=key, spec=spec) return self._pool.add(live) def _store_key(self, session_id: str) -> dict[str, str]: return { "project_key": project_key_for_directory(str(self._agent.cwd)), "session_id": session_id, } async def _seed(self, history: list[dict[str, Any]]) -> str: session_id = str(uuid.uuid4()) entries = build_entries( history, session_id=session_id, cwd=str(self._agent.cwd), model=self._agent.model, permission_mode=self._agent.options.permission_mode, ) await self._store.append( cast("Any", self._store_key(session_id)), cast("Any", entries) ) _log.info( "seeded session %s with %d entries from %d messages", session_id, len(entries), len(history), ) return session_id async def _spawn( self, resume: str | None, *, key: str, spec: _SessionSpec ) -> Session: options = self._build_options(resume, key=key, spec=spec) client = self._factory(options) await client.connect() _log.info( "spawned claude: agent=%s kind=%s resume=%s tools=%s user=%s", self._agent.name, spec.kind, resume, spec.tools, self._runner.user, ) return Session( key=key, agent=self._agent.name, kind=spec.kind, client=client, session_id=resume, resumed=resume is not None, pinned=spec.pinned, ) def _default_factory(self, options: ClaudeAgentOptions) -> SessionClient: return _RunnerClient(options, uid=self._uid) def _build_options( self, resume: str | None, *, key: str, spec: _SessionSpec ) -> ClaudeAgentOptions: agent = self._agent opt = agent.options env = dict(opt.env) if self._runner.home is not None: env["HOME"] = str(self._runner.home) env.setdefault("CLAUDE_CONFIG_DIR", str(self._runner.home / ".claude")) plugins = self._plugins() sources = agent.prompt_for(spec.kind) system_prompt = ( prompt_assembly.assemble(sources) if sources else agent.system_prompt ) servers: dict[str, Any] = dict(self._servers) if spec.tools else {} gateway = ( self._tool_server(key, spec.kind) if spec.tools and self._tool_server is not None and agent.gateway_tools else None ) if gateway is not None: servers[str(gateway["name"])] = gateway return ClaudeAgentOptions( model=agent.model or None, effort=cast("Any", opt.effort), system_prompt=system_prompt, setting_sources=[], strict_mcp_config=True, mcp_servers=cast("Any", servers), permission_mode=cast("Any", opt.permission_mode), tools=list(opt.tools) if opt.tools is not None else None, disallowed_tools=[*opt.disallowed_tools, *self._mcp_disallowed], cwd=str(agent.cwd), add_dirs=list(opt.add_dirs), env=env, cli_path=str(self._exec_wrapper(extra_keep=tuple(env))), include_partial_messages=opt.include_partial_messages, session_store=self._store, session_store_flush=cast("Any", opt.session_store_flush), resume=resume, plugins=cast("Any", plugins), skills="all" if plugins else None, max_turns=opt.max_turns, stderr=lambda line: _log.warning( "claude[%s]: %s", agent.name, line.rstrip() ), ) def _plugins(self) -> list[dict[str, str]]: plugins: list[dict[str, str]] = [] root = self._work_dir / "plugins" / self._agent.name for raw in sorted(self._agent.skill_sets, key=str): source = Path(str(raw)) name = source.name target = root / name if target.exists(): shutil.rmtree(target) shutil.copytree(source, target / "skills") (target / ".claude-plugin").mkdir(parents=True, exist_ok=True) (target / ".claude-plugin" / "plugin.json").write_text( json.dumps({"name": name, "version": "0.0.0"}), encoding="utf-8" ) _chmod_tree(target) plugins.append({"type": "local", "path": str(target)}) return plugins def _exec_wrapper(self, *, extra_keep: tuple[str, ...]) -> Path: if self._wrapper is not None: return self._wrapper keep = sorted({*ENV_KEEP, *self._agent.options.env_keep, *extra_keep}) target = _claude_binary() script = _WRAPPER.format( python=sys.executable, target=json.dumps(target), keep=json.dumps(keep), prefixes=json.dumps(list(ENV_KEEP_PREFIXES)), uid=repr(self._uid), gid=repr(self._gid), ) digest = hashlib.sha256(script.encode("utf-8")).hexdigest()[:12] path = self._work_dir / f"claude-exec-{digest}.py" path.parent.mkdir(parents=True, exist_ok=True) if not path.exists() or path.read_text(encoding="utf-8") != script: path.write_text(script, encoding="utf-8") path.chmod(0o755) path.parent.chmod(0o755) self._wrapper = path return path def _rekey(self, old: str, new: str) -> None: stale = self._pool.rekey(old, new) if stale is not None: asyncio.get_running_loop().create_task(_disconnect(stale)) @dataclass(frozen=True, slots=True) class _SessionSpec: kind: str pinned: bool tools: bool @dataclass class _Turn: events: list[Any] = field(default_factory=list) synthesized: list[dict[str, Any]] = field(default_factory=list) result: ResultMessage | None = None stop_reason: StopReason = "end_turn" _WRAPPER = """#!{python} import os import sys TARGET = {target} KEEP = set({keep}) PREFIXES = tuple({prefixes}) UID = {uid} GID = {gid} env = {{ k: v for k, v in os.environ.items() if k in KEEP or k.startswith(PREFIXES) }} if UID is not None and os.getuid() != UID: os.setgroups([]) os.setgid(GID) os.setuid(UID) os.execve(TARGET, [TARGET, *sys.argv[1:]], env) """ async def _disconnect(live: Session) -> None: try: await live.client.disconnect() except Exception: # noqa: BLE001 _log.exception("disconnect failed for session %s", live.session_id) def _claude_binary() -> str: bundled = Path(claude_agent_sdk.__file__).parent / "_bundled" / "claude" if bundled.is_file(): return str(bundled) found = shutil.which("claude") if found is None: msg = "claude CLI not found: neither bundled in claude_agent_sdk nor on PATH" raise FileNotFoundError(msg) return found def _resolve_ids(user: str | None) -> tuple[int | None, int | None]: if user is None: return None, None record = pwd.getpwuid(int(user)) if user.isdigit() else pwd.getpwnam(user) return record.pw_uid, record.pw_gid def _chown_tree(root: Path, uid: int) -> None: for path in [root, *root.rglob("*")]: with contextlib.suppress(OSError): os.chown(path, uid, -1) def _chmod_tree(root: Path) -> None: for path in [root, *root.rglob("*")]: with contextlib.suppress(OSError): path.chmod(0o755 if path.is_dir() else 0o644) def _mcp_servers( agent: ClaudeAgent, urls: Mapping[str, str] ) -> dict[str, dict[str, Any]]: servers: dict[str, dict[str, Any]] = {} for exposed in agent.expose_mcps: url = urls.get(exposed.name) if url is None: msg = f"agent {agent.name!r} exposes MCP {exposed.name!r} without a URL" raise ValueError(msg) servers[exposed.name] = {"type": "http", "url": url} return servers def _mcp_disallowed( agent: ClaudeAgent, catalog: Mapping[str, Sequence[str]] ) -> list[str]: out: list[str] = [] for exposed in agent.expose_mcps: if exposed.tools is None and not exposed.deny: continue names = list(catalog.get(exposed.name, ())) if not names: _log.warning( "MCP %r has no tool catalog; disallowing the whole server for %s", exposed.name, agent.name, ) out.append(f"mcp__{exposed.name}") continue for name in names: allowed = exposed.tools is None or name in exposed.tools denied = any(fnmatch.fnmatchcase(name, pat) for pat in exposed.deny) if not allowed or denied: out.append(f"mcp__{exposed.name}__{name}") return out def _prompt_text(content: Any) -> str: text = text_of(content) if not text: msg = "user message has no text content" raise ValueError(msg) return text def synthesize_turn_messages(raw: Iterable[Any]) -> list[dict[str, Any]]: out: list[dict[str, Any]] = [] for message in raw: if isinstance(message, AssistantMessage): out.append( { "role": "assistant", "content": [_block_to_dict(b) for b in message.content], } ) elif isinstance(message, UserMessage): content = message.content if isinstance(content, list) and content: out.append( {"role": "user", "content": [_block_to_dict(b) for b in content]} ) return out def _block_to_dict(block: Any) -> dict[str, Any]: if isinstance(block, TextBlock): return {"type": "text", "text": block.text} if isinstance(block, ToolUseBlock): return { "type": "tool_use", "id": block.id, "name": block.name, "input": block.input, } if isinstance(block, ToolResultBlock): result: dict[str, Any] = { "type": "tool_result", "tool_use_id": block.tool_use_id, "content": block.content, } if block.is_error is not None: result["is_error"] = block.is_error return result if isinstance(block, ThinkingBlock): return { "type": "thinking", "thinking": block.thinking, "signature": block.signature, } msg = f"unknown content block type: {type(block).__name__}" raise TypeError(msg) def _usage_of(result: ResultMessage | None) -> TurnUsage: if result is None: return TurnUsage() usage = result.usage or {} return TurnUsage( input_tokens=_int(usage.get("input_tokens")), output_tokens=_int(usage.get("output_tokens")), cache_read_tokens=_int(usage.get("cache_read_input_tokens")), cache_creation_tokens=_int(usage.get("cache_creation_input_tokens")), cost_usd=result.total_cost_usd, duration_ms=result.duration_ms, num_turns=result.num_turns, ) def _wire_usage(usage: TurnUsage) -> dict[str, int]: return { "input_tokens": usage.input_tokens, "output_tokens": usage.output_tokens, "cache_read_input_tokens": usage.cache_read_tokens, "cache_creation_input_tokens": usage.cache_creation_tokens, } def _int(value: Any) -> int: return value if isinstance(value, int) else 0 def _emit_stream_event( event: Mapping[str, Any], index: int ) -> Iterable[MessageStreamEvent]: etype = event.get("type") if etype == "content_block_start": block = event.get("content_block") if not isinstance(block, Mapping): return () btype = block.get("type") if btype == "text": return (build_text_block_start(index),) if btype == "thinking": return (build_thinking_block_start(index),) if btype == "tool_use": return ( build_tool_use_block_start( index, tool_use_id=str(block.get("id", "")), name=str(block.get("name", "")), ), ) return () if etype == "content_block_delta": delta = event.get("delta") if not isinstance(delta, Mapping): return () dtype = delta.get("type") if dtype == "text_delta": return (build_text_delta(index, str(delta.get("text", ""))),) if dtype == "thinking_delta": return (build_thinking_delta(index, str(delta.get("thinking", ""))),) if dtype == "signature_delta": return (build_signature_delta(index, str(delta.get("signature", ""))),) if dtype == "input_json_delta": return (build_input_json_delta(index, str(delta.get("partial_json", ""))),) return () if etype == "content_block_stop": return (build_content_block_stop(index),) return () def _emit_block(block: Any, index: int) -> Iterable[MessageStreamEvent]: if isinstance(block, TextBlock): return ( build_text_block_start(index), build_text_delta(index, block.text), build_content_block_stop(index), ) if isinstance(block, ThinkingBlock): return ( build_thinking_block_start(index), build_thinking_delta(index, block.thinking), build_signature_delta(index, block.signature), build_content_block_stop(index), ) if isinstance(block, ToolUseBlock): partial = json.dumps(block.input, separators=(",", ":"), ensure_ascii=False) return ( build_tool_use_block_start(index, tool_use_id=block.id, name=block.name), build_input_json_delta(index, partial), build_content_block_stop(index), ) return ()