diff --git a/docs/acp.md b/docs/acp.md index 5adad50e7..3d4d3522b 100644 --- a/docs/acp.md +++ b/docs/acp.md @@ -114,7 +114,11 @@ failure fails that room turn visibly instead of falling back. `OMP_MODEL` env variable. Set `api_key` with it and the adapter passes the key in the env variable that model's provider needs. - **Cursor:** question, plan and permission decisions default to `manual`, resolved by a - room participant with `/cursor `. Cursor omits the session id on its + room participant with `/cursor `. Band's own tools, including + `additional_tools`, are approved once per call in every `approval_mode` during + an active turn. If Cursor offers no allow-once option, the configured approval + policy applies instead. Late Band permission requests are refused after the + turn closes or is interrupted. Cursor omits the session id on its extension notifications, so the adapter holds a turn lock and binds them to that turn's session; Cursor turns are serialized. Decision prompts, timeout notices and `/cursor` replies post through `send_notice`, so they never count as the model's reply, and a diff --git a/src/band/adapters/cursor_acp.py b/src/band/adapters/cursor_acp.py index 4ff09c586..1e1741068 100644 --- a/src/band/adapters/cursor_acp.py +++ b/src/band/adapters/cursor_acp.py @@ -27,13 +27,7 @@ CursorQuestion, parse_cursor_questions, ) -from band.integrations.acp.client_runtime import ( - ALLOW_ALWAYS_KIND, - ACPRuntime, - option_id_of_kind, - permission_option_ids, - select_allow_option_id, -) +from band.integrations.acp.client_runtime import ACPRuntime from band.integrations.acp.client_types import ACPClientSessionState from band.integrations.acp.cursor import ( CURSOR_CLI_BINARY, @@ -45,6 +39,14 @@ PLAN_REQUESTED_TEMPLATE, ROOM_COMMAND, CursorCommandWord, + canonicalize_cursor_tool_name, +) +from band.integrations.acp.permissions import ( + ALLOW_ALWAYS_KIND, + ALLOW_ONCE_KIND, + option_id_of_kind, + permission_option_ids, + select_allow_option_id, ) from band.integrations.acp.session_config import SessionConfigResolver from band.runtime.custom_tools import CustomToolDef @@ -79,8 +81,10 @@ class CursorACPAdapterConfig(ACPClientAdapterConfig): api_key: Sets ``CURSOR_API_KEY`` unless ``env`` already does; exclusive with ``auth_token``. auth_token: Sets ``CURSOR_AUTH_TOKEN`` unless ``env`` already does. - approval_mode: How Cursor's permission requests are decided; - ``"manual"`` asks the room. + approval_mode: How Cursor's own tools are decided; ``"manual"`` asks + the room. Band tools, including additional_tools, are approved + once per call in every mode while the turn is active. If Cursor + offers no allow-once option, the configured policy applies. question_mode: How Cursor's questions are answered; ``"manual"`` asks the room. plan_mode: How Cursor's plans are settled; ``"manual"`` asks the room. @@ -188,6 +192,9 @@ def _credential_env(self) -> dict[str, str]: } return {name: value for name, value in credentials.items() if value} + def _canonical_tool_name(self, name: str) -> str: + return canonicalize_cursor_tool_name(name, self._own_tool_names) + async def on_message( self, msg: PlatformMessage, @@ -318,6 +325,9 @@ async def on_cleanup( # Wakes any decision _run_turn is parked on; the task itself keeps # running detached and winds down on its own (via _on_background_task_done) # once the runtime this stops out from under it closes the connection. + turn = self._active_turn + if turn is not None and turn.room_id == room_id: + turn.session_id = None self._cancel_room_decisions(room_id) await super().on_cleanup(room_id, expected_runtime=expected_runtime) @@ -334,6 +344,7 @@ async def on_interrupt(self, room_id: str, mode: ControlMode) -> None: missing reply.""" turn = self._active_turn if turn is not None and turn.room_id == room_id: + turn.session_id = None turn.tools.turn.settle() self._cancel_room_decisions(room_id) @@ -345,6 +356,12 @@ async def cleanup_all(self, *, final: bool = True) -> None: async def _resolve_cursor_permission( self, request: ACPPermissionRequest ) -> str | None: + if request.tool_call.name in self._own_tool_names: + if self._active_turn_for(request.room_id, request.session_id) is None: + return None + once = option_id_of_kind(request.options, ALLOW_ONCE_KIND) + if once is not None: + return once match self.config.approval_mode: case "auto_accept": return select_allow_option_id(request.options) diff --git a/src/band/adapters/omp_acp.py b/src/band/adapters/omp_acp.py index b5f89d954..bfd99d476 100644 --- a/src/band/adapters/omp_acp.py +++ b/src/band/adapters/omp_acp.py @@ -25,8 +25,8 @@ PermissionResolver, SpawnProcess, ) -from band.integrations.acp.client_runtime import ( - ACPCollectingClient, +from band.integrations.acp.collecting import ACPCollectingClient +from band.integrations.acp.permissions import ( ElicitationHandler, ElicitationNarrator, elicitation_requested_schema, diff --git a/src/band/integrations/acp/client_adapter.py b/src/band/integrations/acp/client_adapter.py index 909a37fd1..843189043 100644 --- a/src/band/integrations/acp/client_adapter.py +++ b/src/band/integrations/acp/client_adapter.py @@ -57,28 +57,27 @@ PlatformMessage, ) from band.integrations.acp.client_profiles import ACPClientProfile -from band.integrations.acp.client_runtime import ( - ACPCollectingClient, - ACPConnectionProtocol, - ACPRuntime, - ElicitationHandler, - ElicitationNarrator, - PermissionHandler, - PermissionNarrator, - allow_permission, - cancel_permission, - permission_option_ids, - select_allow_option_id, -) +from band.integrations.acp.client_runtime import ACPRuntime from band.integrations.acp.client_types import ( ACPClientSessionState, BandACPClient, ) +from band.integrations.acp.collecting import ACPCollectingClient from band.integrations.acp.model_selection import ( ACPModelOptions, apply_model_selection, locate_model_options, ) +from band.integrations.acp.permissions import ( + ElicitationHandler, + ElicitationNarrator, + PermissionHandler, + PermissionNarrator, + allow_permission, + cancel_permission, + permission_option_ids, + select_allow_option_id, +) from band.integrations.acp.room_emitter import RoomTurnEmitter from band.integrations.acp.session_config import ( CONFIG_FAILURE_PREFIX, @@ -91,6 +90,7 @@ SessionConfigSetter, apply_session_config_selections, ) +from band.integrations.acp.transport import ACPConnectionProtocol from band.integrations.acp.types import ACPToolCall from band.integrations.mcp import ( BandMCPBackend, diff --git a/src/band/integrations/acp/client_runtime.py b/src/band/integrations/acp/client_runtime.py index cb18c5f1d..516badefe 100644 --- a/src/band/integrations/acp/client_runtime.py +++ b/src/band/integrations/acp/client_runtime.py @@ -1,860 +1,80 @@ -"""Generic ACP subprocess runtime for outbound ACP bridges.""" +"""Connection lifecycle and session operations for outbound ACP clients.""" from __future__ import annotations import asyncio -import json import logging -from collections.abc import AsyncIterator, Awaitable, Callable, Mapping, Sequence -from contextlib import AbstractAsyncContextManager, asynccontextmanager -from typing import Any, Protocol, cast +from collections.abc import Callable +from contextlib import AbstractAsyncContextManager +from typing import Any, cast -from acp import connect_to_agent, spawn_agent_process, text_block -from acp.exceptions import RequestError -from acp.interfaces import Client +from acp import spawn_agent_process from acp.schema import ( ClientCapabilities, - ConfigOptionUpdate, - DeclineElicitationResponse, - LoadSessionResponse, - NewSessionResponse, - SetSessionConfigOptionResponse, ) -from band.integrations.acp.client_profiles import ( - ACPClientProfile, - NoopACPClientProfile, +from band.integrations.acp.collecting import ACPCollectingClient +from band.integrations.acp.permissions import ( + ALLOW_ALWAYS_KIND, + ElicitationHandler, + ElicitationNarrator, + PermissionHandler, + PermissionNarrator, + allow_permission, + cancel_permission, + elicitation_requested_schema, + elicitation_session_id, + option_id_of_kind, + permission_option_ids, + select_allow_option_id, ) -from band.integrations.acp.session_config import ( - SessionConfigOption, - session_config_options, +from band.integrations.acp.sessions import ( + ACP_SESSION_LOAD_TIMEOUT_SECONDS, + ACPSessionOperations, ) -from band.integrations.acp.types import ( - ACPToolCall, - ACPToolResult, - ChunkType, - CollectedChunk, - ToolStatus, +from band.integrations.acp.stderr import ( + STDERR_DRAIN_TIMEOUT_S, + STDERR_TAIL_LINES, + ACPStderrDrain, +) +from band.integrations.acp.stream import ChunkSink +from band.integrations.acp.transport import ( + ACP_STDIO_LIMIT_BYTES, + ACPConnectionProtocol, + ACPSpawnContextProtocol, + tcp_spawn_process, ) from band.integrations.mcp.backends import BandMCPTransport -logger = logging.getLogger(__name__) - -ACP_STDIO_LIMIT_BYTES = 16 * 1024 * 1024 -ACP_SESSION_LOAD_TIMEOUT_SECONDS = 5.0 -PermissionHandler = Callable[..., Awaitable[dict[str, object]]] -PermissionNarrator = Callable[[Awaitable[None]], Awaitable[None]] -ElicitationHandler = Callable[..., Awaitable[object]] -ElicitationNarrator = Callable[[Awaitable[None]], Awaitable[None]] -ChunkSink = Callable[[CollectedChunk], Awaitable[None]] - -# ACP grants a tool-call permission by *selecting one of the options the agent -# offered* (each carries an ``optionId`` and a ``kind``); the on-wire response is -# ``{"outcome": {"outcome": "selected", "optionId": ...}}`` or -# ``{"outcome": {"outcome": "cancelled"}}`` (see ``acp.schema`` AllowedOutcome / -# DeniedOutcome). There is no ``"allowed"`` literal — emitting one makes a -# spec-strict agent (e.g. codex-acp) fail to parse the response and abort the turn. -ALLOW_ALWAYS_KIND = "allow_always" -_ALLOW_OPTION_KINDS = ("allow_once", ALLOW_ALWAYS_KIND) - - -def _resolve_option_id(option: object) -> str | None: - """One offered option's wire id, preferring ``optionId`` over ``option_id``. - - Coalesces the camelCase (wire/JSON) and snake_case spellings on - *absence*, not falsiness — an explicit (if empty) id must not fall - through to the alias and get dropped. Accepts the ACP ``PermissionOption`` - objects or plain dicts/Mappings. - """ - if isinstance(option, Mapping): - option_id = option.get("optionId") - if option_id is None: - option_id = option.get("option_id") - else: - option_id = getattr(option, "option_id", None) - if option_id is None: - option_id = getattr(option, "optionId", None) - return str(option_id) if option_id is not None else None - - -def permission_option_ids(options: object) -> tuple[str, ...]: - """The wire option ids offered by a permission/tool-call request, in order. - - Accepts the ACP ``PermissionOption`` objects or plain dicts/Mappings. - """ - if not isinstance(options, (list, tuple)): - return () - return tuple( - option_id - for option in options - if (option_id := _resolve_option_id(option)) is not None - ) - - -def option_id_of_kind(options: object, kind: str) -> str | None: - """The ``optionId`` of the first offered option of ``kind``, else None.""" - if not isinstance(options, (list, tuple)): - return None - for option in options: - option_kind = ( - option.get("kind") - if isinstance(option, Mapping) - else getattr(option, "kind", None) - ) - option_id = _resolve_option_id(option) - if option_kind == kind and option_id is not None: - return option_id - return None - - -def select_allow_option_id(options: object) -> str | None: - """The ``optionId`` of an allow option offered in a permission request, else None. - - Prefers the least-privilege ``allow_once`` over ``allow_always``. Returns None - when the agent offered no allow option, so the caller cancels rather than - guessing (selecting a reject option would silently deny). - """ - for kind in _ALLOW_OPTION_KINDS: - if (option_id := option_id_of_kind(options, kind)) is not None: - return option_id - return None - - -def allow_permission(option_id: str) -> dict[str, object]: - """An ACP ``RequestPermissionResponse`` selecting (granting) ``option_id``.""" - return {"outcome": {"outcome": "selected", "optionId": option_id}} - - -def cancel_permission() -> dict[str, object]: - """An ACP ``RequestPermissionResponse`` cancelling the request.""" - return {"outcome": {"outcome": "cancelled"}} - - -def elicitation_session_id(mode: object, kwargs: dict[str, object]) -> str: - """Session id from ACP ``mode`` or test kwargs (camelCase / snake_case).""" - return str( - getattr(mode, "session_id", None) - or kwargs.get("session_id") - or kwargs.get("sessionId") - or "" - ) - - -def elicitation_requested_schema( - mode: object, kwargs: dict[str, object] -) -> object | None: - """Form schema from ACP ``mode`` or test kwargs (camelCase / snake_case).""" - return ( - getattr(mode, "requested_schema", None) - or kwargs.get("requested_schema") - or kwargs.get("requestedSchema") - ) - - -def _strict_json_equal(a: object, b: object) -> bool: - """JSON equality without Python's cross-type coercions. - - ``==`` treats ``True == 1`` and ``1 == 1.0`` as equal, which would let two - genuinely different JSON payloads pass a duplicate-echo proof. Two values - are equal here only if their JSON types match too, recursively. - """ - if type(a) is not type(b): - return False - if isinstance(a, dict) and isinstance(b, dict): - return a.keys() == b.keys() and all( - _strict_json_equal(value, b[key]) for key, value in a.items() - ) - if isinstance(a, list) and isinstance(b, list): - return len(a) == len(b) and all(_strict_json_equal(x, y) for x, y in zip(a, b)) - return a == b - - -def _is_echo_of(content: str, readable: str, echo: dict[str, object]) -> bool: - """True when ``content`` is ``readable`` followed by one JSON re-encoding of - exactly ``echo`` -- the one proven duplicated-echo shape, never a guessed - one. The trailing segment is parsed, not string-compared, so the - re-encoding's separators/escaping don't matter. - """ - if not content.startswith(readable): - return False - trailing = content[len(readable) :].strip() - try: - return _strict_json_equal(json.loads(trailing), echo) - except json.JSONDecodeError: - return False - - -def _readable_rendering(content: str, structured: dict[str, object]) -> str | None: - """The prefix of ``content`` that renders ``structured``, else ``None``. - - A FastMCP primitive wrap (``{"result": }``) renders as the wrapped - string verbatim -- required non-empty, or the "rendering" is vacuous and - proves nothing; any other object renders as a leading JSON document that - parses equal to ``structured``. - """ - result = structured.get("result") - if set(structured) == {"result"} and isinstance(result, str): - return result if result and content.startswith(result) else None - try: - leading_value, end = json.JSONDecoder().raw_decode(content) - except json.JSONDecodeError: - return None - return content[:end] if _strict_json_equal(leading_value, structured) else None - - -def _unwrap_structured_result( - content: str, raw_output: object -) -> tuple[str, dict[str, object]] | None: - """Recover a tool result's readable value from a duplicated structured echo. - - An MCP bridge (observed: Copilot) can forward both a tool result's readable - text and its ``structuredContent`` companion into one text block. The echo - invariant, stated once: ``content`` is a readable rendering of - ``structuredContent`` (see ``_readable_rendering``) followed by exactly one - JSON re-encoding of ``structuredContent`` (see ``_is_echo_of``). On that - full proof, returns ``(readable, echo)`` -- the cleaned value plus the - payload proven appended, which the chunk records so a later frame - re-reporting the same duplicate is recognized by exactly that shape. - - Returns ``None`` (leave ``content`` untouched) for anything less: a bridge - may synthesize a distinct human-facing ``content`` -- a summary, or prose - that merely quotes the structured value -- alongside a tool's real - structured result, and that legitimate text must never be clobbered just - because a shape matches or the value appears somewhere within it. - """ - if not isinstance(raw_output, dict): - return None - structured = raw_output.get("structuredContent") - if not isinstance(structured, dict): - return None - readable = _readable_rendering(content, structured) - if readable is not None and _is_echo_of(content, readable, structured): - return readable, structured - return None - - -def tcp_spawn_process( - host: str, - port: int, - *, - limit: int = ACP_STDIO_LIMIT_BYTES, -) -> Callable[..., AbstractAsyncContextManager[tuple[object, object]]]: - """Build a ``spawn_process`` callable that connects to an ACP server over TCP. - - Drop-in for the stdio ``spawn_agent_process`` seam in :class:`ACPRuntime`: the - runtime dials *into* an already-running ACP server (e.g. ``copilot --acp --port - N`` in a container) instead of spawning a subprocess. The returned callable - accepts and ignores the subprocess-shaped args the runtime forwards (the - command executable/args and ``transport_kwargs``) — host/port are captured - here — so no core change to ``ACPRuntime.start`` is needed. - """ - - @asynccontextmanager - async def _connect( - client: Client, - *_command: object, - env: dict[str, str] | None = None, - transport_kwargs: dict[str, object] | None = None, - ) -> AsyncIterator[tuple[object, object]]: - del _command, env, transport_kwargs # subprocess-only; unused for TCP - reader, writer = await asyncio.open_connection(host, port, limit=limit) - # connect_to_agent argument order is (client, input_stream=writer, - # output_stream=reader) and it type-guards writer: StreamWriter / - # reader: StreamReader. Unlike spawn_agent_process it does no cleanup, - # so we close the connection and transport ourselves. - conn = connect_to_agent(client, writer, reader) - try: - yield conn, writer - finally: - try: - await conn.close() - finally: - writer.close() - try: - await writer.wait_closed() - except Exception: - logger.debug("Error awaiting TCP writer close", exc_info=True) - - return _connect - - -class ACPConnectionProtocol(Protocol): - """Protocol for the ACP agent connection returned by spawn_agent_process.""" - - async def initialize(self, *, protocol_version: int) -> object: ... - - async def authenticate(self, *, method_id: str) -> object: ... - - async def new_session( - self, *, cwd: str, mcp_servers: list[object] - ) -> NewSessionResponse: ... - - async def load_session( - self, - *, - cwd: str, - session_id: str, - mcp_servers: list[object], - ) -> LoadSessionResponse | None: ... - - async def prompt(self, *, session_id: str, prompt: list[object]) -> object: ... - - async def set_config_option( - self, - *, - config_id: str, - session_id: str, - value: str, - ) -> SetSessionConfigOptionResponse | None: ... - - async def close_session(self, session_id: str) -> object: ... - - async def cancel(self, session_id: str) -> None: ... - - -class ACPSpawnContextProtocol(Protocol): - """Protocol for the spawn_agent_process async context manager.""" - - async def __aenter__(self) -> tuple[ACPConnectionProtocol, object]: ... - - async def __aexit__(self, exc_type: object, exc: object, tb: object) -> object: ... - - -class ACPCollectingClient(Client): # type: ignore[misc] # ACP Client has optional methods treated as abstract by pyrefly - """Generic ACP client that buffers session updates by session_id. - - The ``acp`` transport runs each incoming notification as its own task, so - consecutive ``session_update``s execute concurrently. A per-session lock - serializes the ingest→sink path and the permission handler (which posts to - the same room mid-turn): lock waiters wake FIFO and the tasks start in - wire-arrival order, so room posts keep the stream's causal order. - """ - - def __init__( - self, - profile: ACPClientProfile | None = None, - canonicalize_tool_name: Callable[[str], str] | None = None, - ) -> None: - self._profile = profile or NoopACPClientProfile() - # Rewrites a runtime's MCP spelling of a tool name (e.g. Copilot's - # ``band-band_send_message``) back to the canonical band name at the - # single point where tool-call chunks are born, so every downstream - # consumer (room narration, reply suppression) sees one vocabulary. - self._canonicalize_tool_name = canonicalize_tool_name or (lambda name: name) - self._session_chunks: dict[str, list[CollectedChunk]] = {} - self._permission_handlers: dict[str, PermissionHandler] = {} - self._elicitation_handlers: dict[str, ElicitationHandler] = {} - # Per session, the canonical tool_result chunk for each tool_call_id, so a - # call's stream of tool_call_updates folds into one result, finalized once - # when the call reaches a terminal status (see _ingest_tool_result). Reset - # per turn in reset_session. - self._result_chunks: dict[str, dict[str, CollectedChunk]] = {} - # ACP reports a result as a separate frame containing only the call id. - # Keep the originating call here, where result folding already happens, so - # a finalized result remains a complete typed lifecycle event. - self._tool_calls: dict[str, dict[str, ACPToolCall]] = {} - self._emitted_results: dict[str, set[str]] = {} - # An open text/thought run being coalesced until the next boundary, and the - # per-session live sink that finalized chunks are posted to, in order. - self._open_runs: dict[str, CollectedChunk] = {} - self._sinks: dict[str, ChunkSink] = {} - # Never popped in reset_session: replacing a lock a straggler task still - # holds would let two tasks into the session's critical section. - self._session_locks: dict[str, asyncio.Lock] = {} - # Each session's latest advertised config catalog; like the locks, - # it outlives reset_session, which runs every turn. - self._config_options: dict[str, tuple[SessionConfigOption, ...]] = {} - - def _session_lock(self, session_id: str) -> asyncio.Lock: - return self._session_locks.setdefault(session_id, asyncio.Lock()) - - def _session_narrator( - self, session_id: str - ) -> Callable[[Awaitable[None]], Awaitable[None]]: - """Serialize a narration awaitable under the session's chunk lock.""" - - async def narrate(action: Awaitable[None]) -> None: - async with self._session_lock(session_id): - await self._close_open_run(session_id) - await action - - return narrate - - async def session_update( - self, session_id: str, update: object, **kwargs: object - ) -> None: - del kwargs - if isinstance(update, ConfigOptionUpdate): - logger.debug("ACP session %s pushed new config options", session_id) - self.record_config_options(session_id, update.config_options) - return - async with self._session_lock(session_id): - chunk = self._chunk_from_update(update) - if chunk is not None: - await self._ingest(session_id, chunk) - - def _chunk_from_update(self, update: object) -> CollectedChunk | None: - """Parse one ACP session update without mutating the chunk buffer.""" - match getattr(update, "session_update", None): - case "agent_message_chunk": - return self._text_chunk(update, ChunkType.TEXT) - case "agent_thought_chunk": - return self._text_chunk(update, ChunkType.THOUGHT) - case "tool_call": - return self._tool_call_chunk(update) - case "tool_call_update": - return self._tool_result_chunk(update) - case "plan": - entries = getattr(update, "entries", []) - plan_text = "\n".join( - getattr(entry, "content", str(entry)) for entry in entries - ) - return CollectedChunk(chunk_type=ChunkType.PLAN, content=plan_text) - case _: - text = self._extract_text_from_content(update) - return ( - CollectedChunk(chunk_type=ChunkType.TEXT, content=text) - if text - else None - ) - - def _text_chunk(self, update: object, chunk_type: str) -> CollectedChunk: - return CollectedChunk( - chunk_type=chunk_type, - content=self._extract_text_from_content(update), - ) - - def _tool_call_chunk(self, update: object) -> CollectedChunk: - raw_input = getattr(update, "raw_input", None) - call = ACPToolCall.from_acp(update, canonicalize=self._canonicalize_tool_name) - metadata = { - "tool_call_id": call.tool_call_id, - "raw_input": raw_input, - "status": getattr(update, "status", ToolStatus.IN_PROGRESS), - } - return CollectedChunk( - chunk_type=ChunkType.TOOL_CALL, - content=call.name, - metadata=metadata, - tool=call, - ) - - def _tool_result_chunk(self, update: object) -> CollectedChunk: - tool_call_id = getattr(update, "tool_call_id", "") - status = getattr(update, "status", ToolStatus.COMPLETED) - metadata = { - "tool_call_id": tool_call_id, - "status": status, - } - # Prefer the human-readable content blocks over ``rawOutput``. An agent's - # terminal update often carries the structured result object in - # ``rawOutput`` (e.g. Copilot's ``{'content': ..., 'contents': [...]}``), - # which stringifies into an unreadable dict; the content blocks hold the - # same output as plain text. Fall back to ``rawOutput`` only when there are - # no content blocks, so an agent that reports output *only* via - # ``rawOutput`` still surfaces it (a blank result keeps the placeholder - # guard). ``from_raw`` records the fallback so _merge_tool_result never - # overwrites clean text with a later raw-only frame. - content = self._extract_text_from_tool_content(getattr(update, "content", None)) - from_raw = not content - echo: dict[str, object] | None = None - if from_raw: - raw_output = getattr(update, "raw_output", "") - content = str(raw_output) if raw_output else "" - else: - # ``rawOutput`` may still carry the least-processed copy of the same - # value under MCP's own ``structuredContent`` field (see - # _unwrap_structured_result); prefer it over the content blocks when - # recognizable, since a bridge that forwards both can duplicate a - # JSON-serialized string across them. - unwrapped = _unwrap_structured_result( - content, getattr(update, "raw_output", None) - ) - if unwrapped is not None: - content, echo = unwrapped - return CollectedChunk( - chunk_type=ChunkType.TOOL_RESULT, - content=content, - metadata=metadata, - from_raw=from_raw, - echo=echo, - tool=self._call_revision(update), - ) - - def _call_revision(self, update: object) -> ACPToolCall | None: - """The call identity a ``tool_call_update`` revises, when it reports one.""" - if not (getattr(update, "title", None) or getattr(update, "raw_input", None)): - return None - return ACPToolCall.from_acp(update, canonicalize=self._canonicalize_tool_name) - - # Chunk kinds that arrive as a stream of deltas for one logical message, so a - # run of them is coalesced into a single chunk (agents emit one delta per token - # or phrase). tool_call/tool_result/plan are discrete and never merged. - _COALESCED_CHUNK_TYPES = (ChunkType.TEXT, ChunkType.THOUGHT) - - async def _ingest(self, session_id: str, chunk: CollectedChunk) -> None: - """Route one parsed chunk through coalescing/collapse to the live sink. - - Text/thought deltas coalesce into an open run, finalized at the next - boundary (a different chunk type, or turn end). tool_call and plan are - discrete and finalize at once. A tool_result folds its call's frames and - finalizes when the call reaches a terminal status. Finalizing a chunk both - buffers it (for get_collected_chunks) and posts it to the sink, in order. - """ - if chunk.chunk_type == ChunkType.TOOL_RESULT: - revision, chunk.tool = chunk.tool, None - if ( - isinstance(revision, ACPToolCall) - and self._revise_held_call(session_id, revision) - and _carries_no_result(chunk) - ): - return - if isinstance(chunk.tool, ACPToolCall) and chunk.tool.tool_call_id: - self._tool_calls.setdefault(session_id, {})[chunk.tool.tool_call_id] = ( - chunk.tool - ) - if chunk.chunk_type in self._COALESCED_CHUNK_TYPES: - open_run = self._open_runs.get(session_id) - if open_run is not None and open_run.chunk_type == chunk.chunk_type: - open_run.content += chunk.content # merge the streamed delta - return - await self._close_open_run(session_id) - self._open_runs[session_id] = chunk - return - await self._close_open_run(session_id) - match chunk.chunk_type: - case ChunkType.TOOL_RESULT: - await self._ingest_tool_result(session_id, chunk) - case ChunkType.TOOL_CALL if _awaits_input(chunk): - # Named by a later tool_call_update (Cursor's "MCP: tool"), so it - # is held like an open run until then. - self._open_runs[session_id] = chunk - case _: - await self._finalize(session_id, chunk) - - def _revise_held_call(self, session_id: str, revision: ACPToolCall) -> bool: - """Apply ``revision`` to the held input-less call it names, if any.""" - held = self._open_runs.get(session_id) - if ( - held is None - or not isinstance(held.tool, ACPToolCall) - or held.tool.tool_call_id != revision.tool_call_id - ): - return False - held.tool, held.content = revision, revision.name - held.metadata["raw_input"] = revision.arguments - self._tool_calls.setdefault(session_id, {})[revision.tool_call_id] = revision - return True - - async def _close_open_run(self, session_id: str) -> None: - """Finalize the open text/thought run, if any — a boundary was reached.""" - run = self._open_runs.pop(session_id, None) - if run is not None: - await self._finalize(session_id, run) - - async def _ingest_tool_result(self, session_id: str, chunk: CollectedChunk) -> None: - """Fold a tool_result frame into its call and finalize once, at terminal. - - A call reports its result over several frames sharing a ``tool_call_id`` - (partial content blocks, then a terminal frame often carrying only the - structured ``rawOutput``). They fold into one canonical result, finalized - the first time the call reports a terminal status. A frame with no id can't - be correlated, so it stands alone. - """ - call_id = str(chunk.metadata.get("tool_call_id", "")) - call = self._tool_calls.get(session_id, {}).get(call_id) - if call is None: - call = ACPToolCall(tool_call_id=call_id, name="unknown", arguments={}) - chunk.tool = ACPToolResult( - call=call, - output=chunk.content, - status=chunk.metadata.get("status"), - ) - if not call_id: - await self._finalize(session_id, chunk) - return - results = self._result_chunks.setdefault(session_id, {}) - canonical = results.get(call_id) - if canonical is None: - results[call_id] = chunk - canonical = chunk - else: - self._fold_result(canonical, chunk) - emitted = self._emitted_results.setdefault(session_id, set()) - terminal = canonical.metadata.get("status") in ( - ToolStatus.COMPLETED, - ToolStatus.FAILED, - ) - # Finalize exactly once, at the first terminal frame — the earliest point a - # live, causally-ordered post is correct (waiting for the true last frame - # would defer every tool_result to turn-end, out of order). A later frame - # still folds into ``canonical`` (so get_collected_chunks reflects it), but - # the room event was already posted: the events API is append-only, so we - # can neither edit it nor re-post without duplicating the narration. A - # post-terminal content revision therefore stays in the buffer only — an - # accepted trade-off of live emission, not a bug to "fix" by re-emitting. - if terminal and call_id not in emitted: - emitted.add(call_id) - await self._finalize(session_id, canonical) - - def _fold_result(self, canonical: CollectedChunk, chunk: CollectedChunk) -> None: - """Fold a later frame into a call's canonical result. - - The last *reported* status wins — ACP status is optional, so a frame that - omits it (status is None) must not regress a recorded "completed". A frame - that merely re-reports a cleaned canonical's proven duplicate — its - content is the cleaned value plus a re-encoding of exactly the recorded - ``CollectedChunk.echo`` payload — carries no new information and must not - regress the cleaned value; anything else, including a genuinely new - readable result, still replaces it. Otherwise ACP ``content`` replaces the - preceding content collection, so the latest readable frame wins even when - it is shorter (e.g. a long-running command's streamed "still running..." - progress text superseded by a short "OK"). A raw-only or empty frame, - which carries no such replacement semantics, falls back to ranking by - completeness (see _result_key). - """ - incoming_status = chunk.metadata.get("status") - if incoming_status is not None: - canonical.metadata["status"] = incoming_status - if canonical.echo is not None and _is_echo_of( - chunk.content, canonical.content, canonical.echo - ): - best = canonical - elif chunk.content and not chunk.from_raw: - best = chunk - elif canonical.content and not canonical.from_raw: - best = canonical - else: - best = max(canonical, chunk, key=self._result_key) - canonical.content, canonical.from_raw, canonical.echo = ( - best.content, - best.from_raw, - best.echo, - ) - if isinstance(canonical.tool, ACPToolResult): - canonical.tool.output = canonical.content - canonical.tool.status = canonical.metadata.get("status") - - async def _finalize(self, session_id: str, chunk: CollectedChunk) -> None: - """Buffer a finalized chunk and post it to the session's live sink, if any. - - A sink failure is logged, not raised: the ``acp`` transport suppresses - notification-handler exceptions without a trace, so raising would lose - the failure. Narration is best-effort — the turn's reply still posts - (and fails loudly) from ``on_message``'s own task. - """ - self._session_chunks.setdefault(session_id, []).append(chunk) - sink = self._sinks.get(session_id) - if sink is None: - return - try: - await sink(chunk) - except Exception: - logger.exception( - "Failed to post %s chunk for ACP session %s to the room; " - "narration for this turn may be incomplete", - chunk.chunk_type, - session_id, - ) - - def set_sink(self, session_id: str, sink: ChunkSink | None) -> None: - if sink is None: - self._sinks.pop(session_id, None) - else: - self._sinks[session_id] = sink - - async def flush(self, session_id: str) -> None: - """Finalize anything still open at turn end: the coalesced run, then any - tool result whose call never reported a terminal status.""" - async with self._session_lock(session_id): - await self._close_open_run(session_id) - emitted = self._emitted_results.setdefault(session_id, set()) - for call_id, canonical in self._result_chunks.get(session_id, {}).items(): - if call_id not in emitted: - emitted.add(call_id) - await self._finalize(session_id, canonical) - - @staticmethod - def _result_key(chunk: CollectedChunk) -> tuple[bool, bool, int]: - """Rank the raw-only/empty frames ``_fold_result`` falls back to (neither - side is a readable-content replacement, so completeness decides): non-empty - beats empty, then the longer (more complete) frame. - """ - has_text = bool(chunk.content) - return has_text, has_text and not chunk.from_raw, len(chunk.content) - - async def request_permission( # type: ignore[override] # ACP Client uses specific types; we widen to object - self, - options: object, - session_id: str, - tool_call: object, - **kwargs: object, - ) -> dict[str, object]: - handler = self._permission_handlers.get(session_id) - if handler: - # A manual handler can wait for room input. Holding the ingestion lock - # for that wait stalls every live update, so the handler receives a - # narrow narrator for the denied tool-call/tool-result pair instead. - return await handler( - options=options, - session_id=session_id, - tool_call=tool_call, - narrate_permission=self._session_narrator(session_id), - **kwargs, - ) - - logger.debug("Auto-cancelling permission request for session %s", session_id) - return cancel_permission() - - async def create_elicitation( # type: ignore[override] - self, - message: str, - mode: object, - **kwargs: object, - ) -> object: - # The modern ACP client API packs session scope into ``mode`` - # (``ElicitationFormSessionMode.session_id``); kwargs only carry - # ``_meta``. Prefer the mode field so form handlers bind to the - # same room session that registered them. - session_id = elicitation_session_id(mode, kwargs) - handler = self._elicitation_handlers.get(session_id) - if handler is not None: - return await handler( - message=message, - mode=mode, - narrate_elicitation=self._session_narrator(session_id), - **kwargs, - ) - logger.debug( - "Auto-declining elicitation for session %s (no handler)", session_id - ) - return DeclineElicitationResponse(action="decline") - - def set_elicitation_handler( - self, - session_id: str, - handler: ElicitationHandler | None, - ) -> None: - if handler is None: - self._elicitation_handlers.pop(session_id, None) - else: - self._elicitation_handlers[session_id] = handler - - def set_permission_handler( - self, - session_id: str, - handler: PermissionHandler | None, - ) -> None: - if handler is None: - self._permission_handlers.pop(session_id, None) - else: - self._permission_handlers[session_id] = handler - - def record_config_options( - self, session_id: str, options: Sequence[SessionConfigOption] - ) -> None: - self._config_options[session_id] = tuple(options) - - def config_options(self, session_id: str) -> tuple[SessionConfigOption, ...]: - return self._config_options.get(session_id, ()) - - def forget_config_options(self, session_id: str) -> None: - self._config_options.pop(session_id, None) - - def reset_session(self, session_id: str) -> None: - self._session_chunks.pop(session_id, None) - self._permission_handlers.pop(session_id, None) - self._elicitation_handlers.pop(session_id, None) - self._result_chunks.pop(session_id, None) - self._tool_calls.pop(session_id, None) - self._emitted_results.pop(session_id, None) - self._open_runs.pop(session_id, None) - self._sinks.pop(session_id, None) - - def get_collected_text(self, session_id: str | None = None) -> str: - if session_id is not None: - chunks = self._session_chunks.get(session_id, []) - else: - chunks = [ - chunk - for session_chunks in self._session_chunks.values() - for chunk in session_chunks - ] - return "".join( - chunk.content for chunk in chunks if chunk.chunk_type == ChunkType.TEXT - ) - - def get_collected_chunks( - self, session_id: str | None = None - ) -> list[CollectedChunk]: - if session_id is not None: - return list(self._session_chunks.get(session_id, [])) - return [ - chunk - for session_chunks in self._session_chunks.values() - for chunk in session_chunks - ] +__all__ = [ + "ACP_SESSION_LOAD_TIMEOUT_SECONDS", + "ACP_STDIO_LIMIT_BYTES", + "ALLOW_ALWAYS_KIND", + "STDERR_DRAIN_TIMEOUT_S", + "STDERR_TAIL_LINES", + "ACPCollectingClient", + "ACPConnectionProtocol", + "ACPRuntime", + "ACPSpawnContextProtocol", + "ChunkSink", + "ElicitationHandler", + "ElicitationNarrator", + "PermissionHandler", + "PermissionNarrator", + "allow_permission", + "cancel_permission", + "elicitation_requested_schema", + "elicitation_session_id", + "option_id_of_kind", + "permission_option_ids", + "select_allow_option_id", + "tcp_spawn_process", +] - async def ext_method( - self, - method: str, - params: dict[str, object], - ) -> dict[str, object]: - return await self._profile.ext_method(method, params) - - async def ext_notification(self, method: str, params: dict[str, object]) -> None: - session_id = str(params.get("sessionId") or params.get("session_id") or "") - if not session_id: - # A profile written against the pre-extension_session_id - # ACPClientProfile protocol has no such attribute at all. - session_id = getattr(self._profile, "extension_session_id", None) or "" - if not session_id: - return - - chunks = await self._profile.ext_notification(method, params) - if chunks: - async with self._session_lock(session_id): - await self._close_open_run(session_id) - for chunk in chunks: - await self._finalize(session_id, chunk) - - @staticmethod - def _block_text(block: object) -> str: - """The ``text`` field of a single ACP content block, else ``""``. - - Accepts either the parsed pydantic model (``.text``) or a raw dict. - """ - text = getattr(block, "text", None) - if text is None and isinstance(block, dict): - text = block.get("text") - return str(text) if text else "" - - @staticmethod - def _extract_text_from_content(update: object) -> str: - return ACPCollectingClient._block_text(getattr(update, "content", None)) - - @staticmethod - def _extract_text_from_tool_content(content: object) -> str: - """Join the inline text blocks of a ``tool_call_update``'s ``content`` list. - - Unlike the single-block ``content`` on message/thought updates, a - ``ToolCallUpdate.content`` is a tagged-union list (``ContentToolCallContent`` - | ``FileEditToolCallContent`` | ``TerminalToolCallContent``, discriminated - by ``type``); only ``"content"`` entries wrap a text block, so entries of - another ``type`` (file-edit diffs, terminal references) are skipped by - their explicit tag rather than by happening to lack a ``.content`` field. - """ - if not isinstance(content, list): - return "" - texts = [ - ACPCollectingClient._block_text(getattr(item, "content", None)) - for item in content - if getattr(item, "type", None) == "content" - ] - return "\n".join(text for text in texts if text) +logger = logging.getLogger(__name__) -class ACPRuntime: +class ACPRuntime(ACPSessionOperations): """Generic ACP subprocess runtime shared by outbound ACP bridges.""" def __init__( @@ -870,6 +90,7 @@ def __init__( use_unstable_protocol: bool = False, pass_builtin_transport_options: bool = True, ) -> None: + super().__init__() self._command = list(command) self._env = env self._cwd = cwd @@ -881,75 +102,83 @@ def __init__( self._pass_builtin_transport_options = pass_builtin_transport_options self._conn: ACPConnectionProtocol | None = None - self._client: ACPCollectingClient | None = None self._ctx: ( AbstractAsyncContextManager[tuple[ACPConnectionProtocol, object]] | None ) = None self._stop_lock = asyncio.Lock() - self._agent_mcp_transport = BandMCPTransport.HTTP - self._agent_supports_session_load = False - self._agent_supports_session_close = False - self._config_lock = asyncio.Lock() + self._stderr = ACPStderrDrain(logger) async def start(self, *, respawn: bool = False) -> None: """Spawn or respawn the ACP agent subprocess.""" + async with self._stop_lock: + await self._start(respawn=respawn) + + async def _start(self, *, respawn: bool) -> None: + self._connection_failed = False logger.info( "%s ACP agent subprocess", "Respawning" if respawn else "Spawning", ) + ctx = self._spawn_context() + self._ctx = ctx + try: + self._conn, transport = await ctx.__aenter__() + if isinstance(transport, asyncio.subprocess.Process): + self._stderr.start(transport) + await self._initialize_connection(self._conn) + except (asyncio.CancelledError, KeyboardInterrupt): + self._stderr.expect_exit(connection_failed=self._connection_failed) + await self._cleanup_failed_start(ctx, "init cancel") + raise + except Exception: + await self._cleanup_failed_start(ctx, "init failure") + raise + # A connect-only transport (e.g. TCP) carries no command; describe it + # rather than logging a blank suffix. + logger.info( + "Connected to ACP agent: %s", + " ".join(self._command) or "", + ) + + def _spawn_context( + self, + ) -> AbstractAsyncContextManager[tuple[ACPConnectionProtocol, object]]: self._client = self._client_factory() # type: ignore[abstract] # ACP client protocol defines optional hooks as abstract spawn_kwargs: dict[str, Any] = {} if self._pass_builtin_transport_options: spawn_kwargs["transport_kwargs"] = {"limit": ACP_STDIO_LIMIT_BYTES} if self._use_unstable_protocol: spawn_kwargs["use_unstable_protocol"] = True - ctx = cast( + return cast( AbstractAsyncContextManager[tuple[ACPConnectionProtocol, object]], self._spawn_process( self._client, - # Splat the whole command: stdio forwards executable + args, while - # a TCP transport passes an empty command (host/port live in the - # injected spawn_process closure) and receives no positional args. + # TCP transports carry no command; stdio carries executable and args. *self._command, env=self._env, cwd=self._cwd, **spawn_kwargs, ), ) - self._ctx = ctx - try: - self._conn, _ = await ctx.__aenter__() - init_kwargs: dict[str, Any] = {"protocol_version": 1} - if self._client_capabilities is not None: - init_kwargs["client_capabilities"] = self._client_capabilities - init_response = await cast(Any, self._conn).initialize(**init_kwargs) - self._agent_mcp_transport = self._select_mcp_transport(init_response) - self._agent_supports_session_load = self._select_session_load(init_response) - self._agent_supports_session_close = self._select_session_close( - init_response - ) - if self._auth_method: - await self._conn.authenticate(method_id=self._auth_method) - logger.info("Authenticated with method: %s", self._auth_method) - except (asyncio.CancelledError, KeyboardInterrupt): - await self._cleanup_failed_start(ctx, "init cancel") - raise - except Exception: - await self._cleanup_failed_start(ctx, "init failure") - raise - # A connect-only transport (e.g. TCP) carries no command; describe it - # rather than logging a blank suffix. - logger.info( - "Connected to ACP agent: %s", - " ".join(self._command) or "", - ) + + async def _initialize_connection(self, conn: ACPConnectionProtocol) -> None: + init_kwargs: dict[str, Any] = {"protocol_version": 1} + if self._client_capabilities is not None: + init_kwargs["client_capabilities"] = self._client_capabilities + init_response = await cast(Any, conn).initialize(**init_kwargs) + self._agent_mcp_transport = self._select_mcp_transport(init_response) + self._agent_supports_session_load = self._select_session_load(init_response) + self._agent_supports_session_close = self._select_session_close(init_response) + if self._auth_method: + await conn.authenticate(method_id=self._auth_method) + logger.info("Authenticated with method: %s", self._auth_method) async def ensure_connection(self, *, can_respawn: bool) -> ACPConnectionProtocol: async with self._stop_lock: if self._conn is None: if self._ctx is None and can_respawn: - await self.start(respawn=False) + await self._start(respawn=False) else: raise RuntimeError( "ACP client not initialized. Call on_started first." @@ -961,215 +190,23 @@ async def ensure_connection(self, *, can_respawn: bool) -> ACPConnectionProtocol raise RuntimeError("ACP client connection dropped before prompt") return conn - async def create_session(self, *, cwd: str, mcp_servers: list[object]) -> str: - session = await self.create_session_response(cwd=cwd, mcp_servers=mcp_servers) - return session.session_id - - async def create_session_response( - self, *, cwd: str, mcp_servers: list[object] - ) -> NewSessionResponse: - conn = await self.ensure_connection(can_respawn=False) - response = cast( - NewSessionResponse, - await conn.new_session(cwd=cwd, mcp_servers=mcp_servers), - ) - self._record_config_options(response.session_id, response) - return response - - async def load_session( - self, - *, - cwd: str, - session_id: str, - mcp_servers: list[object], - ) -> bool: - return ( - await self.load_session_response( - cwd=cwd, - session_id=session_id, - mcp_servers=mcp_servers, - ) - ) is not None - - async def load_session_response( - self, - *, - cwd: str, - session_id: str, - mcp_servers: list[object], - ) -> LoadSessionResponse | None: - """Load a persisted ACP session when the connected agent supports it. - - ACP session IDs are meaningful only to the agent process that owns them. - A successful ``session/load`` is therefore the boundary where a persisted ID - becomes usable on this connection. An unsupported, unavailable, slow, or - erroring load returns ``None`` so callers can create a fresh session - without blocking a turn. - """ - if not self._agent_supports_session_load: - return None - - conn = await self.ensure_connection(can_respawn=False) - try: - response = await asyncio.wait_for( - conn.load_session( - cwd=cwd, - session_id=session_id, - mcp_servers=mcp_servers, - ), - timeout=ACP_SESSION_LOAD_TIMEOUT_SECONDS, - ) - except TimeoutError: - logger.warning( - "ACP session %s did not load within %s seconds", - session_id, - ACP_SESSION_LOAD_TIMEOUT_SECONDS, - ) - return None - except RequestError as error: - # Any load failure is equally recoverable: the caller falls back to - # a fresh session (with history replay) rather than letting a remote - # protocol error kill the bootstrap turn. - if self._is_missing_session_error(error): - logger.info("ACP session %s is no longer available", session_id) - else: - logger.warning( - "ACP session/load for %s failed (%s); using a new session", - session_id, - error, - ) - return None - self._record_config_options(session_id, response) - return response - - async def set_config_option( - self, - *, - session_id: str, - config_id: str, - value: str, - ) -> SetSessionConfigOptionResponse | None: - """Set one advertised select option and return the refreshed catalog.""" - conn = await self.ensure_connection(can_respawn=False) - response = await conn.set_config_option( - session_id=session_id, - config_id=config_id, - value=value, - ) - self._record_config_options(session_id, response) - return response - - def config_options(self, session_id: str) -> tuple[SessionConfigOption, ...]: - """The session's live catalog: its setup response, then every change.""" - if self._client is None: - return () - return self._client.config_options(session_id) - - @property - def config_lock(self) -> asyncio.Lock: - """Held across a runtime switch. - - Each step is checked against the catalog the previous one returned, so - an interleaved switch would validate against a model no longer current. - Session setup needs no lock: its room is not yet published to switch. - """ - return self._config_lock - - def _record_config_options(self, session_id: str, response: object) -> None: - options = session_config_options(response) - if options is not None and self._client is not None: - self._client.record_config_options(session_id, options) - - async def close_session(self, session_id: str) -> None: - """Close a session when the agent advertised lifecycle support.""" - if self._client is not None: - self._client.forget_config_options(session_id) - if not self._agent_supports_session_close: - return - conn = await self.ensure_connection(can_respawn=False) - await conn.close_session(session_id) - - async def prompt( - self, - *, - session_id: str, - prompt_text: str, - on_chunk: ChunkSink | None = None, - ) -> list[CollectedChunk]: - conn = await self.ensure_connection(can_respawn=False) - if on_chunk is not None and self._client is not None: - self._client.set_sink(session_id, on_chunk) - try: - await conn.prompt(session_id=session_id, prompt=[text_block(prompt_text)]) - if self._client is not None: - await self._client.flush(session_id) - finally: - if self._client is not None: - self._client.set_sink(session_id, None) - return self.get_collected_chunks(session_id) - - async def cancel_turn(self, session_id: str) -> None: - """Tell the agent to stop a room's in-flight prompt.""" - conn = await self.ensure_connection(can_respawn=False) - await conn.cancel(session_id) - - def reset_session(self, session_id: str) -> None: - if self._client is not None: - self._client.reset_session(session_id) - - def set_permission_handler( - self, - session_id: str, - handler: PermissionHandler | None, - ) -> None: - if self._client is not None: - self._client.set_permission_handler(session_id, handler) - - def set_elicitation_handler( - self, - session_id: str, - handler: ElicitationHandler | None, - ) -> None: - if self._client is not None: - self._client.set_elicitation_handler(session_id, handler) - - def get_collected_chunks(self, session_id: str) -> list[CollectedChunk]: - if self._client is None: - return [] - return self._client.get_collected_chunks(session_id) - - @property - def client(self) -> ACPCollectingClient | None: - """The active collecting client, once started. - - ``ACPRuntime`` only forwards the handful of client methods a real turn - needs (``get_collected_chunks``, ``reset_session``, ...); this is the - read-only escape hatch for callers that need something else off the - buffer directly (e.g. ``get_collected_text``) rather than growing - ``ACPRuntime`` a passthrough per client method. - """ - return self._client - - @property - def agent_mcp_transport(self) -> BandMCPTransport: - """The MCP transport the connected agent negotiated during ``start()``.""" - return self._agent_mcp_transport - async def stop(self) -> None: ctx: AbstractAsyncContextManager[tuple[ACPConnectionProtocol, object]] | None async with self._stop_lock: + self._stderr.expect_exit(connection_failed=self._connection_failed) ctx = self._ctx self._ctx = None self._conn = None self._client = None self._agent_supports_session_load = False self._agent_supports_session_close = False - if ctx is None: - return - try: - await ctx.__aexit__(None, None, None) - except Exception: - logger.exception("Error during ACP runtime shutdown") + try: + if ctx is not None: + await ctx.__aexit__(None, None, None) + except Exception: + logger.exception("Error during ACP runtime shutdown") + finally: + await self._stderr.finish() async def _cleanup_failed_start( self, @@ -1180,6 +217,8 @@ async def _cleanup_failed_start( await ctx.__aexit__(None, None, None) except Exception: logger.exception("Error cleaning up ACP subprocess after %s", reason) + finally: + await self._stderr.finish() self._ctx = None self._conn = None self._agent_supports_session_load = False @@ -1207,22 +246,3 @@ def _select_session_close(init_response: object) -> bool: capabilities = getattr(init_response, "agent_capabilities", None) session_capabilities = getattr(capabilities, "session_capabilities", None) return getattr(session_capabilities, "close", None) is not None - - @staticmethod - def _is_missing_session_error(error: RequestError) -> bool: - """Whether an ACP ``session/load`` failure means the session is absent.""" - return error.code == -32002 or ( - "session" in str(error).lower() and "not found" in str(error).lower() - ) - - -def _awaits_input(chunk: CollectedChunk) -> bool: - """True for a pending tool_call reported before its input.""" - return chunk.metadata.get( - "status" - ) == ToolStatus.PENDING and not chunk.metadata.get("raw_input") - - -def _carries_no_result(chunk: CollectedChunk) -> bool: - """True for a tool_call_update frame that only revised its call's identity.""" - return chunk.metadata.get("status") is None and not chunk.content diff --git a/src/band/integrations/acp/client_types.py b/src/band/integrations/acp/client_types.py index 4021b9078..6fda9fa3d 100644 --- a/src/band/integrations/acp/client_types.py +++ b/src/band/integrations/acp/client_types.py @@ -4,7 +4,7 @@ from dataclasses import dataclass, field -from band.integrations.acp.client_runtime import ACPCollectingClient +from band.integrations.acp.collecting import ACPCollectingClient @dataclass diff --git a/src/band/integrations/acp/collecting.py b/src/band/integrations/acp/collecting.py new file mode 100644 index 000000000..f8bc3ddd8 --- /dev/null +++ b/src/band/integrations/acp/collecting.py @@ -0,0 +1,200 @@ +"""ACP client callbacks with serialized session output and decisions.""" + +from __future__ import annotations + +import asyncio +import logging +from collections.abc import Awaitable, Callable, Sequence + +from acp.interfaces import Client +from acp.schema import ConfigOptionUpdate, DeclineElicitationResponse + +from band.integrations.acp.client_profiles import ACPClientProfile, NoopACPClientProfile +from band.integrations.acp.permissions import ( + ElicitationHandler, + PermissionHandler, + cancel_permission, + elicitation_session_id, +) +from band.integrations.acp.session_config import SessionConfigOption +from band.integrations.acp.stream import ACPChunkStream, ChunkSink +from band.integrations.acp.types import CollectedChunk +from band.integrations.acp.updates import ACPUpdateParser + +logger = logging.getLogger(__name__) + + +class ACPCollectingClient(ACPUpdateParser, Client): # type: ignore[misc] # ACP Client has optional methods treated as abstract by pyrefly + """Generic ACP client that buffers session updates by session_id. + + The ``acp`` transport runs each incoming notification as its own task, so + consecutive ``session_update``s execute concurrently. A per-session lock + serializes the ingest→sink path and the permission handler (which posts to + the same room mid-turn): lock waiters wake FIFO and the tasks start in + wire-arrival order, so room posts keep the stream's causal order. + """ + + def __init__( + self, + profile: ACPClientProfile | None = None, + canonicalize_tool_name: Callable[[str], str] | None = None, + ) -> None: + self._profile = profile or NoopACPClientProfile() + super().__init__(canonicalize_tool_name) + self._stream = ACPChunkStream() + self._permission_handlers: dict[str, PermissionHandler] = {} + self._elicitation_handlers: dict[str, ElicitationHandler] = {} + # Retain locks while a straggler can still hold one. + self._session_locks: dict[str, asyncio.Lock] = {} + # Each session's latest advertised config catalog; like the locks, + # it outlives reset_session, which runs every turn. + self._config_options: dict[str, tuple[SessionConfigOption, ...]] = {} + + def _session_lock(self, session_id: str) -> asyncio.Lock: + return self._session_locks.setdefault(session_id, asyncio.Lock()) + + def _session_narrator( + self, session_id: str + ) -> Callable[[Awaitable[None]], Awaitable[None]]: + """Serialize a narration awaitable under the session's chunk lock.""" + + async def narrate(action: Awaitable[None]) -> None: + async with self._session_lock(session_id): + await self._stream.close_open_run(session_id) + await action + + return narrate + + async def session_update( + self, session_id: str, update: object, **kwargs: object + ) -> None: + del kwargs + if isinstance(update, ConfigOptionUpdate): + logger.debug("ACP session %s pushed new config options", session_id) + self.record_config_options(session_id, update.config_options) + return + async with self._session_lock(session_id): + chunk = self._chunk_from_update(update) + if chunk is not None: + await self._stream.ingest(session_id, chunk) + + async def request_permission( # type: ignore[override] # ACP Client uses specific types; we widen to object + self, + options: object, + session_id: str, + tool_call: object, + **kwargs: object, + ) -> dict[str, object]: + handler = self._permission_handlers.get(session_id) + if handler: + # A manual handler can wait for room input. Holding the ingestion lock + # for that wait stalls every live update, so the handler receives a + # narrow narrator for the denied tool-call/tool-result pair instead. + return await handler( + options=options, + session_id=session_id, + tool_call=tool_call, + narrate_permission=self._session_narrator(session_id), + **kwargs, + ) + + logger.debug("Auto-cancelling permission request for session %s", session_id) + return cancel_permission() + + async def create_elicitation( # type: ignore[override] + self, + message: str, + mode: object, + **kwargs: object, + ) -> object: + # The modern ACP client API packs session scope into ``mode`` + # (``ElicitationFormSessionMode.session_id``); kwargs only carry + # ``_meta``. Prefer the mode field so form handlers bind to the + # same room session that registered them. + session_id = elicitation_session_id(mode, kwargs) + handler = self._elicitation_handlers.get(session_id) + if handler is not None: + return await handler( + message=message, + mode=mode, + narrate_elicitation=self._session_narrator(session_id), + **kwargs, + ) + logger.debug( + "Auto-declining elicitation for session %s (no handler)", session_id + ) + return DeclineElicitationResponse(action="decline") + + def set_elicitation_handler( + self, + session_id: str, + handler: ElicitationHandler | None, + ) -> None: + if handler is None: + self._elicitation_handlers.pop(session_id, None) + else: + self._elicitation_handlers[session_id] = handler + + def set_permission_handler( + self, + session_id: str, + handler: PermissionHandler | None, + ) -> None: + if handler is None: + self._permission_handlers.pop(session_id, None) + else: + self._permission_handlers[session_id] = handler + + def record_config_options( + self, session_id: str, options: Sequence[SessionConfigOption] + ) -> None: + self._config_options[session_id] = tuple(options) + + def config_options(self, session_id: str) -> tuple[SessionConfigOption, ...]: + return self._config_options.get(session_id, ()) + + def forget_config_options(self, session_id: str) -> None: + self._config_options.pop(session_id, None) + + def reset_session(self, session_id: str) -> None: + self._stream.reset_session(session_id) + self._permission_handlers.pop(session_id, None) + self._elicitation_handlers.pop(session_id, None) + + def set_sink(self, session_id: str, sink: ChunkSink | None) -> None: + self._stream.set_sink(session_id, sink) + + async def flush(self, session_id: str) -> None: + async with self._session_lock(session_id): + await self._stream.flush(session_id) + + def get_collected_text(self, session_id: str | None = None) -> str: + return self._stream.get_collected_text(session_id) + + def get_collected_chunks( + self, session_id: str | None = None + ) -> list[CollectedChunk]: + return self._stream.get_collected_chunks(session_id) + + async def ext_method( + self, + method: str, + params: dict[str, object], + ) -> dict[str, object]: + return await self._profile.ext_method(method, params) + + async def ext_notification(self, method: str, params: dict[str, object]) -> None: + session_id = str(params.get("sessionId") or params.get("session_id") or "") + if not session_id: + # A profile written against the pre-extension_session_id + # ACPClientProfile protocol has no such attribute at all. + session_id = getattr(self._profile, "extension_session_id", None) or "" + if not session_id: + return + + chunks = await self._profile.ext_notification(method, params) + if chunks: + async with self._session_lock(session_id): + await self._stream.close_open_run(session_id) + for chunk in chunks: + await self._stream.finalize(session_id, chunk) diff --git a/src/band/integrations/acp/cursor.py b/src/band/integrations/acp/cursor.py index 7498a3511..d870534ae 100644 --- a/src/band/integrations/acp/cursor.py +++ b/src/band/integrations/acp/cursor.py @@ -1,14 +1,50 @@ -"""Dependency-free room decision vocabulary for the Cursor ACP adapter.""" +"""Cursor MCP tool-title identity and room decision vocabulary.""" from __future__ import annotations +from collections.abc import Collection from enum import StrEnum +from band.runtime.tools.registry import ( + BAND_MCP_SERVER_NAME, + canonicalize_mcp_tool_name, + mcp_tool_spelling, +) + DECISION_NOT_PENDING_TEMPLATE = "Cursor decision `{token}` is not pending." DECISION_UNAUTHORIZED_MESSAGE = "You are not authorized to resolve Cursor decisions." ROOM_COMMAND = "/cursor" CURSOR_CLI_BINARY = "agent" +CURSOR_TITLE_SEPARATOR = ": " + + +def cursor_mcp_title(title: str) -> tuple[str, str] | None: + """Parse Cursor's ``provider-tool: tool`` display title, failing closed.""" + spelling, separator, tool = title.partition(CURSOR_TITLE_SEPARATOR) + if not separator or not spelling or not tool or spelling == tool: + return None + if any(char.isspace() for char in spelling + tool) or ":" in spelling + tool: + return None + return spelling, tool + + +def is_cursor_band_tool(title: str, own_names: Collection[str]) -> bool: + """Whether both title halves identify the same registered Band MCP tool.""" + parsed = cursor_mcp_title(title) + if parsed is None: + return False + spelling, tool = parsed + return tool in own_names and spelling == mcp_tool_spelling( + BAND_MCP_SERVER_NAME, tool + ) + + +def canonicalize_cursor_tool_name(name: str, own_names: Collection[str]) -> str: + """Decode a registered tool's exact Cursor title before generic MCP names.""" + if is_cursor_band_tool(name, own_names): + return name.partition(CURSOR_TITLE_SEPARATOR)[2] + return canonicalize_mcp_tool_name(name, own_names) class CursorCommandWord(StrEnum): diff --git a/src/band/integrations/acp/permissions.py b/src/band/integrations/acp/permissions.py new file mode 100644 index 000000000..b80b53a50 --- /dev/null +++ b/src/band/integrations/acp/permissions.py @@ -0,0 +1,94 @@ +"""ACP permission options and elicitation request fields.""" + +from __future__ import annotations + +from collections.abc import Awaitable, Callable, Mapping + +PermissionHandler = Callable[..., Awaitable[dict[str, object]]] +PermissionNarrator = Callable[[Awaitable[None]], Awaitable[None]] +ElicitationHandler = Callable[..., Awaitable[object]] +ElicitationNarrator = Callable[[Awaitable[None]], Awaitable[None]] + +# ACP requires selecting an offered option id, never a synthesized grant. +ALLOW_ONCE_KIND = "allow_once" +ALLOW_ALWAYS_KIND = "allow_always" +_ALLOW_OPTION_KINDS = (ALLOW_ONCE_KIND, ALLOW_ALWAYS_KIND) + + +def _resolve_option_id(option: object) -> str | None: + """One offered option's wire id, preferring ``optionId`` over ``option_id``.""" + if isinstance(option, Mapping): + option_id = option.get("optionId") + if option_id is None: + option_id = option.get("option_id") + else: + option_id = getattr(option, "option_id", None) + if option_id is None: + option_id = getattr(option, "optionId", None) + return str(option_id) if option_id is not None else None + + +def permission_option_ids(options: object) -> tuple[str, ...]: + """The wire option ids offered by a permission/tool-call request, in order.""" + if not isinstance(options, (list, tuple)): + return () + return tuple( + option_id + for option in options + if (option_id := _resolve_option_id(option)) is not None + ) + + +def option_id_of_kind(options: object, kind: str) -> str | None: + """The ``optionId`` of the first offered option of ``kind``, else None.""" + if not isinstance(options, (list, tuple)): + return None + for option in options: + option_kind = ( + option.get("kind") + if isinstance(option, Mapping) + else getattr(option, "kind", None) + ) + option_id = _resolve_option_id(option) + if option_kind == kind and option_id is not None: + return option_id + return None + + +def select_allow_option_id(options: object) -> str | None: + """The ``optionId`` of an allow option offered in a permission request, else None.""" + for kind in _ALLOW_OPTION_KINDS: + if (option_id := option_id_of_kind(options, kind)) is not None: + return option_id + return None + + +def allow_permission(option_id: str) -> dict[str, object]: + """An ACP ``RequestPermissionResponse`` selecting (granting) ``option_id``.""" + return {"outcome": {"outcome": "selected", "optionId": option_id}} + + +def cancel_permission() -> dict[str, object]: + """An ACP ``RequestPermissionResponse`` cancelling the request.""" + return {"outcome": {"outcome": "cancelled"}} + + +def elicitation_session_id(mode: object, kwargs: dict[str, object]) -> str: + """Session id from ACP ``mode`` or test kwargs (camelCase / snake_case).""" + return str( + getattr(mode, "session_id", None) + or kwargs.get("session_id") + or kwargs.get("sessionId") + or "" + ) + + +def elicitation_requested_schema( + mode: object, kwargs: dict[str, object] +) -> object | None: + """Form schema from ACP ``mode`` or test kwargs (camelCase / snake_case).""" + return ( + getattr(mode, "requested_schema", None) + or kwargs.get("requested_schema") + or kwargs.get("requestedSchema") + ) diff --git a/src/band/integrations/acp/results.py b/src/band/integrations/acp/results.py new file mode 100644 index 000000000..9632dba82 --- /dev/null +++ b/src/band/integrations/acp/results.py @@ -0,0 +1,89 @@ +"""Readable ACP tool results and streamed result replacement.""" + +from __future__ import annotations + +import json + +from band.integrations.acp.types import ACPToolResult, CollectedChunk + + +def _strict_json_equal(a: object, b: object) -> bool: + """Compare JSON types strictly; Python equates booleans and numbers.""" + if type(a) is not type(b): + return False + if isinstance(a, dict) and isinstance(b, dict): + return a.keys() == b.keys() and all( + _strict_json_equal(value, b[key]) for key, value in a.items() + ) + if isinstance(a, list) and isinstance(b, list): + return len(a) == len(b) and all(_strict_json_equal(x, y) for x, y in zip(a, b)) + return a == b + + +def _is_echo_of(content: str, readable: str, echo: dict[str, object]) -> bool: + """Recognize a readable value followed by an exact JSON echo.""" + if not content.startswith(readable): + return False + trailing = content[len(readable) :].strip() + try: + return _strict_json_equal(json.loads(trailing), echo) + except json.JSONDecodeError: + return False + + +def _readable_rendering(content: str, structured: dict[str, object]) -> str | None: + """Find the non-empty rendering of MCP structured content.""" + result = structured.get("result") + if set(structured) == {"result"} and isinstance(result, str): + return result if result and content.startswith(result) else None + try: + leading_value, end = json.JSONDecoder().raw_decode(content) + except json.JSONDecodeError: + return None + return content[:end] if _strict_json_equal(leading_value, structured) else None + + +def unwrap_structured_result( + content: str, raw_output: object +) -> tuple[str, dict[str, object]] | None: + """Remove only a proven duplicate of MCP structured content.""" + if not isinstance(raw_output, dict): + return None + structured = raw_output.get("structuredContent") + if not isinstance(structured, dict): + return None + readable = _readable_rendering(content, structured) + if readable is not None and _is_echo_of(content, readable, structured): + return readable, structured + return None + + +def fold_result(canonical: CollectedChunk, chunk: CollectedChunk) -> None: + """Preserve reported status and readable output across result frames.""" + incoming_status = chunk.metadata.get("status") + if incoming_status is not None: + canonical.metadata["status"] = incoming_status + if canonical.echo is not None and _is_echo_of( + chunk.content, canonical.content, canonical.echo + ): + best = canonical + elif chunk.content and not chunk.from_raw: + best = chunk + elif canonical.content and not canonical.from_raw: + best = canonical + else: + best = max(canonical, chunk, key=_result_key) + canonical.content, canonical.from_raw, canonical.echo = ( + best.content, + best.from_raw, + best.echo, + ) + if isinstance(canonical.tool, ACPToolResult): + canonical.tool.output = canonical.content + canonical.tool.status = canonical.metadata.get("status") + + +def _result_key(chunk: CollectedChunk) -> tuple[bool, bool, int]: + """Rank fallback frames by readable content and completeness.""" + has_text = bool(chunk.content) + return has_text, has_text and not chunk.from_raw, len(chunk.content) diff --git a/src/band/integrations/acp/sessions.py b/src/band/integrations/acp/sessions.py new file mode 100644 index 000000000..d09278a32 --- /dev/null +++ b/src/band/integrations/acp/sessions.py @@ -0,0 +1,232 @@ +"""Session operations on a connected ACP agent.""" + +from __future__ import annotations + +import asyncio +import logging +from abc import ABC, abstractmethod +from typing import cast + +from acp import text_block +from acp.exceptions import RequestError +from acp.schema import ( + LoadSessionResponse, + NewSessionResponse, + SetSessionConfigOptionResponse, +) + +from band.integrations.acp.collecting import ACPCollectingClient +from band.integrations.acp.permissions import ElicitationHandler, PermissionHandler +from band.integrations.acp.session_config import ( + SessionConfigOption, + session_config_options, +) +from band.integrations.acp.stream import ChunkSink +from band.integrations.acp.transport import ACPConnectionProtocol +from band.integrations.acp.types import CollectedChunk +from band.integrations.mcp.backends import BandMCPTransport + +logger = logging.getLogger(__name__) +ACP_SESSION_LOAD_TIMEOUT_SECONDS = 5.0 + + +class ACPSessionOperations(ABC): + """Control sessions through a lifecycle-owned connection.""" + + def __init__(self) -> None: + self._client: ACPCollectingClient | None = None + self._connection_failed = False + self._agent_mcp_transport = BandMCPTransport.HTTP + self._agent_supports_session_load = False + self._agent_supports_session_close = False + self._config_lock = asyncio.Lock() + + @abstractmethod + async def ensure_connection(self, *, can_respawn: bool) -> ACPConnectionProtocol: + """Return the lifecycle owner's active connection.""" + raise NotImplementedError + + async def create_session(self, *, cwd: str, mcp_servers: list[object]) -> str: + session = await self.create_session_response(cwd=cwd, mcp_servers=mcp_servers) + return session.session_id + + async def create_session_response( + self, *, cwd: str, mcp_servers: list[object] + ) -> NewSessionResponse: + conn = await self.ensure_connection(can_respawn=False) + response = cast( + NewSessionResponse, + await conn.new_session(cwd=cwd, mcp_servers=mcp_servers), + ) + self._record_config_options(response.session_id, response) + return response + + async def load_session( + self, + *, + cwd: str, + session_id: str, + mcp_servers: list[object], + ) -> bool: + return ( + await self.load_session_response( + cwd=cwd, + session_id=session_id, + mcp_servers=mcp_servers, + ) + ) is not None + + async def load_session_response( + self, + *, + cwd: str, + session_id: str, + mcp_servers: list[object], + ) -> LoadSessionResponse | None: + """Load a persisted session, falling back on an unavailable or failed load.""" + if not self._agent_supports_session_load: + return None + + conn = await self.ensure_connection(can_respawn=False) + try: + response = await asyncio.wait_for( + conn.load_session( + cwd=cwd, + session_id=session_id, + mcp_servers=mcp_servers, + ), + timeout=ACP_SESSION_LOAD_TIMEOUT_SECONDS, + ) + except TimeoutError: + logger.warning( + "ACP session %s did not load within %s seconds", + session_id, + ACP_SESSION_LOAD_TIMEOUT_SECONDS, + ) + return None + except RequestError as error: + # Any load failure is equally recoverable: the caller falls back to + # a fresh session (with history replay) rather than letting a remote + # protocol error kill the bootstrap turn. + if self._is_missing_session_error(error): + logger.info("ACP session %s is no longer available", session_id) + else: + logger.warning( + "ACP session/load for %s failed (%s); using a new session", + session_id, + error, + ) + return None + self._record_config_options(session_id, response) + return response + + async def set_config_option( + self, + *, + session_id: str, + config_id: str, + value: str, + ) -> SetSessionConfigOptionResponse | None: + """Set one advertised select option and return the refreshed catalog.""" + conn = await self.ensure_connection(can_respawn=False) + response = await conn.set_config_option( + session_id=session_id, + config_id=config_id, + value=value, + ) + self._record_config_options(session_id, response) + return response + + def config_options(self, session_id: str) -> tuple[SessionConfigOption, ...]: + """The session's live catalog: its setup response, then every change.""" + if self._client is None: + return () + return self._client.config_options(session_id) + + @property + def config_lock(self) -> asyncio.Lock: + """Held across a runtime switch.""" + return self._config_lock + + def _record_config_options(self, session_id: str, response: object) -> None: + options = session_config_options(response) + if options is not None and self._client is not None: + self._client.record_config_options(session_id, options) + + async def close_session(self, session_id: str) -> None: + """Close a session when the agent advertised lifecycle support.""" + if self._client is not None: + self._client.forget_config_options(session_id) + if not self._agent_supports_session_close: + return + conn = await self.ensure_connection(can_respawn=False) + await conn.close_session(session_id) + + async def prompt( + self, + *, + session_id: str, + prompt_text: str, + on_chunk: ChunkSink | None = None, + ) -> list[CollectedChunk]: + conn = await self.ensure_connection(can_respawn=False) + if on_chunk is not None and self._client is not None: + self._client.set_sink(session_id, on_chunk) + try: + await conn.prompt(session_id=session_id, prompt=[text_block(prompt_text)]) + if self._client is not None: + await self._client.flush(session_id) + except ConnectionError: + self._connection_failed = True + raise + finally: + if self._client is not None: + self._client.set_sink(session_id, None) + return self.get_collected_chunks(session_id) + + async def cancel_turn(self, session_id: str) -> None: + """Tell the agent to stop a room's in-flight prompt.""" + conn = await self.ensure_connection(can_respawn=False) + await conn.cancel(session_id) + + def reset_session(self, session_id: str) -> None: + if self._client is not None: + self._client.reset_session(session_id) + + def set_permission_handler( + self, + session_id: str, + handler: PermissionHandler | None, + ) -> None: + if self._client is not None: + self._client.set_permission_handler(session_id, handler) + + def set_elicitation_handler( + self, + session_id: str, + handler: ElicitationHandler | None, + ) -> None: + if self._client is not None: + self._client.set_elicitation_handler(session_id, handler) + + def get_collected_chunks(self, session_id: str) -> list[CollectedChunk]: + if self._client is None: + return [] + return self._client.get_collected_chunks(session_id) + + @property + def client(self) -> ACPCollectingClient | None: + """The active client for operations beyond the runtime session interface.""" + return self._client + + @property + def agent_mcp_transport(self) -> BandMCPTransport: + """The MCP transport the connected agent negotiated during ``start()``.""" + return self._agent_mcp_transport + + @staticmethod + def _is_missing_session_error(error: RequestError) -> bool: + """Whether an ACP ``session/load`` failure means the session is absent.""" + return error.code == -32002 or ( + "session" in str(error).lower() and "not found" in str(error).lower() + ) diff --git a/src/band/integrations/acp/stderr.py b/src/band/integrations/acp/stderr.py new file mode 100644 index 000000000..f44b5e5ed --- /dev/null +++ b/src/band/integrations/acp/stderr.py @@ -0,0 +1,93 @@ +"""Bounded stderr diagnostics for an ACP subprocess.""" + +from __future__ import annotations + +import asyncio +import logging +from collections import deque + +STDERR_TAIL_LINES = 20 +STDERR_DRAIN_TIMEOUT_S = 5.0 +STDERR_EXIT_POLL_INTERVAL_S = 0.05 +STDERR_EXIT_FLUSH_TIMEOUT_S = 0.1 +STDERR_LINE_LOG_TEMPLATE = "ACP agent stderr: %s" + + +class ACPStderrDrain: + """Drain one process and distinguish a crash from an intentional exit.""" + + def __init__(self, logger: logging.Logger) -> None: + self._logger = logger + self._stopping = False + self._stdout: asyncio.StreamReader | None = None + self._process: asyncio.subprocess.Process | None = None + self._task: asyncio.Task[None] | None = None + + def start(self, process: asyncio.subprocess.Process) -> None: + self._stopping = False + self._stdout = process.stdout + self._process = process + self._task = asyncio.create_task(self._drain(process)) + + def expect_exit(self, *, connection_failed: bool = False) -> None: + if connection_failed: + return + if self._process is not None and self._process.returncode is not None: + return + # Cleanup after stdout EOF did not cause the connection to close. + if self._stdout is not None and self._stdout.at_eof(): + return + self._stopping = True + + async def _drain(self, process: asyncio.subprocess.Process) -> None: + if process.stderr is None: + return + tail: deque[str] = deque(maxlen=STDERR_TAIL_LINES) + reader = asyncio.create_task(self._read_lines(process.stderr, tail)) + try: + # An existing Process.wait() can await inherited pipes after exit. + while process.returncode is None: + await asyncio.sleep(STDERR_EXIT_POLL_INTERVAL_S) + unexpected = reader.result() if reader.done() else not self._stopping + try: + await asyncio.wait_for(reader, timeout=STDERR_EXIT_FLUSH_TIMEOUT_S) + except TimeoutError: + pass + if unexpected and process.returncode != 0: + self._logger.warning( + "ACP agent exited with code %s; stderr tail:\n%s", + process.returncode, + "\n".join(tail), + ) + finally: + reader.cancel() + await asyncio.gather(reader, return_exceptions=True) + + async def _read_lines(self, stderr: asyncio.StreamReader, tail: deque[str]) -> bool: + while True: + try: + line = await stderr.readline() + except (ValueError, asyncio.LimitOverrunError): + self._logger.debug("Skipped ACP stderr line exceeding the stream limit") + continue + except OSError: + self._logger.debug("Error reading ACP agent stderr", exc_info=True) + break + if not line: + break + text = line.decode(errors="replace").rstrip("\r\n") + tail.append(text) + self._logger.debug(STDERR_LINE_LOG_TEMPLATE, text) + # A later stop cannot reclassify an EOF already observed. + return not self._stopping + + async def finish(self) -> None: + task, self._task = self._task, None + if task is None: + return + try: + await asyncio.wait_for(task, timeout=STDERR_DRAIN_TIMEOUT_S) + except TimeoutError: + self._logger.debug("Timed out draining ACP agent stderr") + except Exception: + self._logger.debug("Error draining ACP agent stderr", exc_info=True) diff --git a/src/band/integrations/acp/stream.py b/src/band/integrations/acp/stream.py new file mode 100644 index 000000000..0653ec1ec --- /dev/null +++ b/src/band/integrations/acp/stream.py @@ -0,0 +1,192 @@ +"""Coalesce ACP deltas and publish completed tool lifecycles in order.""" + +from __future__ import annotations + +import logging +from collections.abc import Awaitable, Callable + +from band.integrations.acp.results import fold_result +from band.integrations.acp.types import ( + ACPToolCall, + ACPToolResult, + ChunkType, + CollectedChunk, + ToolStatus, +) + +logger = logging.getLogger(__name__) +ChunkSink = Callable[[CollectedChunk], Awaitable[None]] + + +class ACPChunkStream: + """Buffer and publish output; the caller serializes operations per session.""" + + _COALESCED_CHUNK_TYPES = (ChunkType.TEXT, ChunkType.THOUGHT) + + def __init__(self) -> None: + self._session_chunks: dict[str, list[CollectedChunk]] = {} + self._result_chunks: dict[str, dict[str, CollectedChunk]] = {} + self._tool_calls: dict[str, dict[str, ACPToolCall]] = {} + self._emitted_results: dict[str, set[str]] = {} + self._open_runs: dict[str, CollectedChunk] = {} + self._sinks: dict[str, ChunkSink] = {} + + async def ingest(self, session_id: str, chunk: CollectedChunk) -> None: + """Coalesce text deltas and publish discrete tool lifecycle boundaries.""" + if chunk.chunk_type == ChunkType.TOOL_RESULT: + revision, chunk.tool = chunk.tool, None + if ( + isinstance(revision, ACPToolCall) + and self._revise_held_call(session_id, revision) + and _carries_no_result(chunk) + ): + return + if isinstance(chunk.tool, ACPToolCall) and chunk.tool.tool_call_id: + self._tool_calls.setdefault(session_id, {})[chunk.tool.tool_call_id] = ( + chunk.tool + ) + if chunk.chunk_type in self._COALESCED_CHUNK_TYPES: + open_run = self._open_runs.get(session_id) + if open_run is not None and open_run.chunk_type == chunk.chunk_type: + open_run.content += chunk.content # merge the streamed delta + return + await self.close_open_run(session_id) + self._open_runs[session_id] = chunk + return + await self.close_open_run(session_id) + match chunk.chunk_type: + case ChunkType.TOOL_RESULT: + await self._ingest_tool_result(session_id, chunk) + case ChunkType.TOOL_CALL if _awaits_input(chunk): + # Named by a later tool_call_update (Cursor's "MCP: tool"), so it + # is held like an open run until then. + self._open_runs[session_id] = chunk + case _: + await self.finalize(session_id, chunk) + + def _revise_held_call(self, session_id: str, revision: ACPToolCall) -> bool: + """Apply ``revision`` to the held input-less call it names, if any.""" + held = self._open_runs.get(session_id) + if ( + held is None + or not isinstance(held.tool, ACPToolCall) + or held.tool.tool_call_id != revision.tool_call_id + ): + return False + held.tool, held.content = revision, revision.name + held.metadata["raw_input"] = revision.arguments + self._tool_calls.setdefault(session_id, {})[revision.tool_call_id] = revision + return True + + async def close_open_run(self, session_id: str) -> None: + """Finalize the open text/thought run, if any — a boundary was reached.""" + run = self._open_runs.pop(session_id, None) + if run is not None: + await self.finalize(session_id, run) + + async def _ingest_tool_result(self, session_id: str, chunk: CollectedChunk) -> None: + """Publish a tool result once, when its call first becomes terminal.""" + call_id = str(chunk.metadata.get("tool_call_id", "")) + call = self._tool_calls.get(session_id, {}).get(call_id) + if call is None: + call = ACPToolCall(tool_call_id=call_id, name="unknown", arguments={}) + chunk.tool = ACPToolResult( + call=call, + output=chunk.content, + status=chunk.metadata.get("status"), + ) + if not call_id: + await self.finalize(session_id, chunk) + return + results = self._result_chunks.setdefault(session_id, {}) + canonical = results.get(call_id) + if canonical is None: + results[call_id] = chunk + canonical = chunk + else: + fold_result(canonical, chunk) + emitted = self._emitted_results.setdefault(session_id, set()) + terminal = canonical.metadata.get("status") in ( + ToolStatus.COMPLETED, + ToolStatus.FAILED, + ) + # Room events are append-only; later revisions stay in the buffer. + if terminal and call_id not in emitted: + emitted.add(call_id) + await self.finalize(session_id, canonical) + + async def finalize(self, session_id: str, chunk: CollectedChunk) -> None: + """Buffer and publish a chunk; sink errors must survive ACP exception suppression.""" + self._session_chunks.setdefault(session_id, []).append(chunk) + sink = self._sinks.get(session_id) + if sink is None: + return + try: + await sink(chunk) + except Exception: + logger.exception( + "Failed to post %s chunk for ACP session %s to the room; " + "narration for this turn may be incomplete", + chunk.chunk_type, + session_id, + ) + + def set_sink(self, session_id: str, sink: ChunkSink | None) -> None: + if sink is None: + self._sinks.pop(session_id, None) + else: + self._sinks[session_id] = sink + + async def flush(self, session_id: str) -> None: + """Finalize anything still open at turn end: the coalesced run, then any + tool result whose call never reported a terminal status.""" + await self.close_open_run(session_id) + emitted = self._emitted_results.setdefault(session_id, set()) + for call_id, canonical in self._result_chunks.get(session_id, {}).items(): + if call_id not in emitted: + emitted.add(call_id) + await self.finalize(session_id, canonical) + + def get_collected_text(self, session_id: str | None = None) -> str: + if session_id is not None: + chunks = self._session_chunks.get(session_id, []) + else: + chunks = [ + chunk + for session_chunks in self._session_chunks.values() + for chunk in session_chunks + ] + return "".join( + chunk.content for chunk in chunks if chunk.chunk_type == ChunkType.TEXT + ) + + def get_collected_chunks( + self, session_id: str | None = None + ) -> list[CollectedChunk]: + if session_id is not None: + return list(self._session_chunks.get(session_id, [])) + return [ + chunk + for session_chunks in self._session_chunks.values() + for chunk in session_chunks + ] + + def reset_session(self, session_id: str) -> None: + self._session_chunks.pop(session_id, None) + self._result_chunks.pop(session_id, None) + self._tool_calls.pop(session_id, None) + self._emitted_results.pop(session_id, None) + self._open_runs.pop(session_id, None) + self._sinks.pop(session_id, None) + + +def _awaits_input(chunk: CollectedChunk) -> bool: + """True for a pending tool_call reported before its input.""" + return chunk.metadata.get( + "status" + ) == ToolStatus.PENDING and not chunk.metadata.get("raw_input") + + +def _carries_no_result(chunk: CollectedChunk) -> bool: + """True for a tool_call_update frame that only revised its call's identity.""" + return chunk.metadata.get("status") is None and not chunk.content diff --git a/src/band/integrations/acp/transport.py b/src/band/integrations/acp/transport.py new file mode 100644 index 000000000..8e4730c07 --- /dev/null +++ b/src/band/integrations/acp/transport.py @@ -0,0 +1,99 @@ +"""ACP connection contracts and stdio/TCP transport configuration.""" + +from __future__ import annotations + +import asyncio +import logging +from collections.abc import AsyncIterator, Callable +from contextlib import AbstractAsyncContextManager, asynccontextmanager +from typing import Protocol + +from acp import connect_to_agent +from acp.interfaces import Client +from acp.schema import ( + LoadSessionResponse, + NewSessionResponse, + SetSessionConfigOptionResponse, +) + +logger = logging.getLogger(__name__) +ACP_STDIO_LIMIT_BYTES = 16 * 1024 * 1024 + + +def tcp_spawn_process( + host: str, + port: int, + *, + limit: int = ACP_STDIO_LIMIT_BYTES, +) -> Callable[..., AbstractAsyncContextManager[tuple[object, object]]]: + """Connect to ACP over TCP using the stdio-shaped spawn interface.""" + + @asynccontextmanager + async def _connect( + client: Client, + *_command: object, + env: dict[str, str] | None = None, + transport_kwargs: dict[str, object] | None = None, + ) -> AsyncIterator[tuple[object, object]]: + del _command, env, transport_kwargs # subprocess-only; unused for TCP + reader, writer = await asyncio.open_connection(host, port, limit=limit) + # connect_to_agent argument order is (client, input_stream=writer, + # output_stream=reader) and it type-guards writer: StreamWriter / + # reader: StreamReader. Unlike spawn_agent_process it does no cleanup, + # so we close the connection and transport ourselves. + conn = connect_to_agent(client, writer, reader) + try: + yield conn, writer + finally: + try: + await conn.close() + finally: + writer.close() + try: + await writer.wait_closed() + except Exception: + logger.debug("Error awaiting TCP writer close", exc_info=True) + + return _connect + + +class ACPConnectionProtocol(Protocol): + """Protocol for the ACP agent connection returned by spawn_agent_process.""" + + async def initialize(self, *, protocol_version: int) -> object: ... + + async def authenticate(self, *, method_id: str) -> object: ... + + async def new_session( + self, *, cwd: str, mcp_servers: list[object] + ) -> NewSessionResponse: ... + + async def load_session( + self, + *, + cwd: str, + session_id: str, + mcp_servers: list[object], + ) -> LoadSessionResponse | None: ... + + async def prompt(self, *, session_id: str, prompt: list[object]) -> object: ... + + async def set_config_option( + self, + *, + config_id: str, + session_id: str, + value: str, + ) -> SetSessionConfigOptionResponse | None: ... + + async def close_session(self, session_id: str) -> object: ... + + async def cancel(self, session_id: str) -> None: ... + + +class ACPSpawnContextProtocol(Protocol): + """Protocol for the spawn_agent_process async context manager.""" + + async def __aenter__(self) -> tuple[ACPConnectionProtocol, object]: ... + + async def __aexit__(self, exc_type: object, exc: object, tb: object) -> object: ... diff --git a/src/band/integrations/acp/updates.py b/src/band/integrations/acp/updates.py new file mode 100644 index 000000000..7d5c283c9 --- /dev/null +++ b/src/band/integrations/acp/updates.py @@ -0,0 +1,128 @@ +"""Project ACP session-update frames into typed chunks.""" + +from __future__ import annotations + +from collections.abc import Callable + +from band.integrations.acp.results import unwrap_structured_result +from band.integrations.acp.types import ( + ACPToolCall, + ChunkType, + CollectedChunk, + ToolStatus, +) + + +class ACPUpdateParser: + """Parse updates while preserving backend-specific tool-call overrides.""" + + def __init__( + self, canonicalize_tool_name: Callable[[str], str] | None = None + ) -> None: + self._canonicalize_tool_name = canonicalize_tool_name or (lambda name: name) + + def _chunk_from_update(self, update: object) -> CollectedChunk | None: + """Parse one ACP session update without mutating the chunk buffer.""" + match getattr(update, "session_update", None): + case "agent_message_chunk": + return self._text_chunk(update, ChunkType.TEXT) + case "agent_thought_chunk": + return self._text_chunk(update, ChunkType.THOUGHT) + case "tool_call": + return self._tool_call_chunk(update) + case "tool_call_update": + return self._tool_result_chunk(update) + case "plan": + entries = getattr(update, "entries", []) + plan_text = "\n".join( + getattr(entry, "content", str(entry)) for entry in entries + ) + return CollectedChunk(chunk_type=ChunkType.PLAN, content=plan_text) + case _: + text = self._extract_text_from_content(update) + return ( + CollectedChunk(chunk_type=ChunkType.TEXT, content=text) + if text + else None + ) + + def _text_chunk(self, update: object, chunk_type: str) -> CollectedChunk: + return CollectedChunk( + chunk_type=chunk_type, + content=self._extract_text_from_content(update), + ) + + def _tool_call_chunk(self, update: object) -> CollectedChunk: + raw_input = getattr(update, "raw_input", None) + call = ACPToolCall.from_acp(update, canonicalize=self._canonicalize_tool_name) + metadata = { + "tool_call_id": call.tool_call_id, + "raw_input": raw_input, + "status": getattr(update, "status", ToolStatus.IN_PROGRESS), + } + return CollectedChunk( + chunk_type=ChunkType.TOOL_CALL, + content=call.name, + metadata=metadata, + tool=call, + ) + + def _tool_result_chunk(self, update: object) -> CollectedChunk: + tool_call_id = getattr(update, "tool_call_id", "") + status = getattr(update, "status", ToolStatus.COMPLETED) + metadata = { + "tool_call_id": tool_call_id, + "status": status, + } + # Readable blocks take precedence over raw output. + content = self._extract_text_from_tool_content(getattr(update, "content", None)) + from_raw = not content + echo: dict[str, object] | None = None + if from_raw: + raw_output = getattr(update, "raw_output", "") + content = str(raw_output) if raw_output else "" + else: + # MCP can repeat structured content after the readable result. + unwrapped = unwrap_structured_result( + content, getattr(update, "raw_output", None) + ) + if unwrapped is not None: + content, echo = unwrapped + return CollectedChunk( + chunk_type=ChunkType.TOOL_RESULT, + content=content, + metadata=metadata, + from_raw=from_raw, + echo=echo, + tool=self._call_revision(update), + ) + + def _call_revision(self, update: object) -> ACPToolCall | None: + """The call identity a ``tool_call_update`` revises, when it reports one.""" + if not (getattr(update, "title", None) or getattr(update, "raw_input", None)): + return None + return ACPToolCall.from_acp(update, canonicalize=self._canonicalize_tool_name) + + @staticmethod + def _block_text(block: object) -> str: + """The ``text`` field of a single ACP content block, else ``""``.""" + text = getattr(block, "text", None) + if text is None and isinstance(block, dict): + text = block.get("text") + return str(text) if text else "" + + @staticmethod + def _extract_text_from_content(update: object) -> str: + return ACPUpdateParser._block_text(getattr(update, "content", None)) + + @staticmethod + def _extract_text_from_tool_content(content: object) -> str: + """Read inline text only; file edits and terminal references have distinct tags.""" + if not isinstance(content, list): + return "" + texts = [ + ACPUpdateParser._block_text(getattr(item, "content", None)) + for item in content + if getattr(item, "type", None) == "content" + ] + return "\n".join(text for text in texts if text) diff --git a/tests/adapters/test_cursor_acp_adapter.py b/tests/adapters/test_cursor_acp_adapter.py index 944ccaa2f..bd3bcb619 100644 --- a/tests/adapters/test_cursor_acp_adapter.py +++ b/tests/adapters/test_cursor_acp_adapter.py @@ -4,6 +4,7 @@ import asyncio import re +import sys from collections.abc import AsyncIterator, Awaitable, Callable from datetime import UTC, datetime from pathlib import Path @@ -13,7 +14,9 @@ import pytest import pytest_asyncio -from acp.schema import PermissionOption +from acp.helpers import start_tool_call, update_tool_call +from acp.schema import AllowedOutcome, PermissionOption, ToolCallUpdate +from pydantic import BaseModel, create_model from band.adapters.cursor_acp import ( DECISION_UNAUTHORIZED_MESSAGE, @@ -24,10 +27,15 @@ ) from band.client.streaming import ControlMode from band.core.protocols import AgentToolsProtocol -from band.core.types import AgentInput, HistoryProvider, PlatformMessage +from band.core.types import AgentInput, Capability, HistoryProvider, PlatformMessage from band.integrations.acp.client_adapter import ACPPermissionRequest +from band.integrations.acp.client_runtime import ACPRuntime from band.integrations.acp.cursor import PLAN_REQUESTED_TEMPLATE +from band.integrations.acp.permissions import ALLOW_ALWAYS_KIND, ALLOW_ONCE_KIND from band.integrations.acp.types import ACPToolCall +from band.runtime.custom_tools import CustomToolDef +from band.runtime.tools.registry import BAND_MCP_SERVER_NAME, mcp_tool_spelling +from band.runtime.tools.types import BandTool from band.testing import MISSING_REPLY_FAILURE, FakeAgentTools, failure_reports from tests.integrations.acp.acp_toolkit.agent import FakeACPAgent from tests.integrations.acp.acp_toolkit.harness import ( @@ -36,7 +44,7 @@ pair_in_process, started_acp_adapter, ) -from tests.mcpclient import crash_backend +from tests.mcpclient import STORE_MEMORY_ARGS, crash_backend class DecisionTools(FakeAgentTools): @@ -119,6 +127,10 @@ def said(tools: FakeAgentTools) -> list[str]: return [cast(str, sent["content"]) for sent in tools.messages_sent] +def permission_requests(tools: FakeAgentTools) -> list[str]: + return [message for message in said(tools) if "needs permission" in message] + + class CursorRoom: """A Cursor room whose ``agent acp`` peer is a scripted in-process ACP agent: each message goes through ``on_event`` with its own tools, as the @@ -157,8 +169,18 @@ async def turns_finished(self) -> None: async def cursor_room() -> AsyncIterator[Callable[..., Awaitable[CursorRoom]]]: adapters: list[CursorACPAdapter] = [] - async def open_room(agent: FakeACPAgent, **config: Any) -> CursorRoom: - adapter = CursorACPAdapter(CursorACPAdapterConfig(**config)) + async def open_room( + agent: FakeACPAgent, + *, + capabilities: set[Capability] | None = None, + additional_tools: list[CustomToolDef] | None = None, + **config: Any, + ) -> CursorRoom: + adapter = CursorACPAdapter( + CursorACPAdapterConfig(**config), + capabilities=capabilities, + additional_tools=additional_tools, + ) pair_in_process(adapter, agent) await adapter.on_started("Cursor", "Cursor agent under test") adapters.append(adapter) @@ -178,6 +200,204 @@ def cursor_in( ) +class EchoInput(BaseModel): + text: str + + +def echo(arguments: EchoInput) -> str: + return arguments.text + + +def custom_reply() -> str: + return "custom reply" + + +def cursor_custom_tools( + echo_handler: Callable[[EchoInput], str], +) -> list[CustomToolDef]: + return [ + (EchoInput, echo_handler), + ( + create_model( + f"{mcp_tool_spelling(BAND_MCP_SERVER_NAME, BandTool.SEND_MESSAGE)}Input" + ), + custom_reply, + ), + ] + + +@pytest.fixture +def cursor_approval_room( + cursor_room: Callable[..., Awaitable[CursorRoom]], +) -> Callable[[str, str], Awaitable[CursorRoom]]: + async def open_room(title: str, approval_mode: str) -> CursorRoom: + return await cursor_room( + FakeACPAgent() + .will_ask_permission(title=title) + .will_call_mcp_tool( + "reply", + BandTool.SEND_MESSAGE, + arguments={"content": "finished", "mentions": ["@user-1"]}, + ), + approval_mode=approval_mode, + capabilities={Capability.MEMORY}, + additional_tools=cursor_custom_tools(echo), + decision_timeout_s=0.1, + ) + + return open_room + + +@pytest.mark.asyncio +@pytest.mark.parametrize("approval_mode", ["manual", "auto_accept", "auto_decline"]) +@pytest.mark.parametrize( + "title", + [ + pytest.param("band-band_store_memory: band_store_memory", id="memory"), + pytest.param( + "band-band_send_message: band_send_message", + id="reply-with-custom-name-collision", + ), + pytest.param("band-echo: echo", id="custom-tool"), + ], +) +async def test_registered_band_tools_are_approved_without_asking_the_room( + cursor_approval_room: Callable[[str, str], Awaitable[CursorRoom]], + approval_mode: str, + title: str, +) -> None: + room = await cursor_approval_room(title, approval_mode) + tools = await room.send("run the tool") + await room.turns_finished() + + assert room.agent.approved is True + assert permission_requests(tools) == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("approval_mode", "approved", "room_requests"), + [ + ("manual", False, 1), + ("auto_accept", True, 0), + ("auto_decline", False, 0), + ], +) +@pytest.mark.parametrize( + "title", + [ + pytest.param("other-band_send_message: band_send_message", id="other-server"), + pytest.param( + "band-band_store_memory: band_send_message", id="mismatched-identity" + ), + pytest.param("shell", id="native-tool"), + pytest.param("shell: shell", id="native-display-title"), + pytest.param("band-unknown: unknown", id="unregistered-tool"), + ], +) +async def test_other_tools_follow_cursor_approval_policy( + cursor_approval_room: Callable[[str, str], Awaitable[CursorRoom]], + approval_mode: str, + title: str, + approved: bool, + room_requests: int, +) -> None: + room = await cursor_approval_room(title, approval_mode) + tools = await room.send("run the tool") + await room.turns_finished() + + assert room.agent.approved is approved + assert len(permission_requests(tools)) == room_requests + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("approval_mode", "first_edit", "second_edit"), + [ + ("manual", False, True), + ("auto_accept", True, True), + ("auto_decline", False, False), + ], +) +async def test_project_workflow_keeps_band_tools_available_across_human_gates_and_rejoin( + cursor_room: Callable[..., Awaitable[CursorRoom]], + tmp_path: Path, + approval_mode: str, + first_edit: bool, + second_edit: bool, +) -> None: + project = tmp_path / "project.txt" + project.write_text("original") + audit = tmp_path / "audit.txt" + marker = "project workflow completed" + + def record_work(arguments: EchoInput) -> str: + with audit.open("a") as stream: + stream.write(arguments.text + "\n") + return arguments.text + + async def edit_project(agent: FakeACPAgent, session_id: str) -> None: + await agent.emit(session_id, start_tool_call("native-edit", "shell")) + process = await asyncio.create_subprocess_exec( + sys.executable, + "-c", + "from pathlib import Path; import sys; Path(sys.argv[1]).write_text(sys.argv[2])", + str(project), + "repaired", + ) + if await process.wait() != 0: + raise RuntimeError("project edit failed") + await agent.emit( + session_id, update_tool_call("native-edit", raw_output="project repaired") + ) + + agent = ( + FakeACPAgent() + .will_ask_permission(title="band-echo: echo") + .will_call_mcp_tool( + "custom-work", "echo", arguments={"text": marker}, requires_approval=True + ) + ) + agent.will_ask_permission(title="shell").will_execute_if_approved(edit_project) + for tool, arguments in [ + (BandTool.STORE_MEMORY, {**STORE_MEMORY_ARGS, "content": marker}), + (BandTool.SEND_MESSAGE, {"content": marker, "mentions": ["@user-1"]}), + ]: + agent.will_ask_permission( + title=f"{mcp_tool_spelling(BAND_MCP_SERVER_NAME, tool)}: {tool}" + ).will_call_mcp_tool(tool, tool, arguments=arguments, requires_approval=True) + agent.will_say("I already posted the result.") + room = await cursor_room( + agent, + approval_mode=approval_mode, + capabilities={Capability.MEMORY}, + additional_tools=cursor_custom_tools(record_work), + ) + + for attempt, edited in enumerate((first_edit, second_edit)): + project.write_text("original") + turn = await room.send("repair the project and record the result") + if approval_mode == "manual": + assert project.read_text() == "original" + [request] = permission_requests(turn) + assert "shell" in request + command = "deny" if attempt == 0 else "select" + option = "" if attempt == 0 else " allow-1" + decision = await room.send( + f"/cursor {command} {decision_token(request)}{option}" + ) + assert failure_reports(decision) == [] + await room.turns_finished() + + assert project.read_text() == ("repaired" if edited else "original") + assert audit.read_text().splitlines() == [marker] * (attempt + 1) + assert turn.memory_contents == [marker] + assert said(turn) == [*permission_requests(turn), marker] + assert turn.turn.replied + assert failure_reports(turn) == [] + await room.adapter.on_cleanup("room-1") + + class TestCursorACPAdapterConfig: @pytest.mark.parametrize( ("settings", "error"), @@ -288,7 +508,139 @@ async def test_cwd_becomes_a_room_workspace_root(self, tmp_path: Path) -> None: assert launch.cwd == str(tmp_path / "room-a") +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("approval_mode", "persistent", "approved", "room_requests"), + [ + ("manual", False, True, 0), + ("auto_accept", False, True, 0), + ("auto_decline", False, True, 0), + ("manual", True, False, 1), + ("auto_accept", True, True, 0), + ("auto_decline", True, False, 0), + ], +) +async def test_raw_input_band_permission_obeys_grant_duration( + cursor_room: Callable[..., Awaitable[CursorRoom]], + approval_mode: str, + persistent: bool, + approved: bool, + room_requests: int, +) -> None: + agent = FakeACPAgent() + + @agent.on_prompt + async def ask_band_permission(peer: FakeACPAgent, session_id: str) -> None: + response = await peer.ask_permission( + session_id, + ToolCallUpdate( + toolCallId="reply", + title="band-band_store_memory: band_store_memory" + if persistent + else "Reply to room", + rawInput=None + if persistent + else { + "providerIdentifier": BAND_MCP_SERVER_NAME, + "toolName": BandTool.STORE_MEMORY, + "args": {}, + }, + ), + [ + PermissionOption( + optionId="grant", + name="Allow", + kind=ALLOW_ALWAYS_KIND if persistent else ALLOW_ONCE_KIND, + ) + ], + ) + peer.approved = isinstance(response.outcome, AllowedOutcome) + await peer.call_mcp_tool( + session_id=session_id, + server=BAND_MCP_SERVER_NAME, + tool_name=BandTool.SEND_MESSAGE, + arguments={"content": "finished", "mentions": ["@user-1"]}, + ) + + room = await cursor_room( + agent, + approval_mode=approval_mode, + capabilities={Capability.MEMORY}, + decision_timeout_s=0.1, + ) + tools = await room.send("reply") + await room.turns_finished() + + assert agent.approved is approved + assert len(permission_requests(tools)) == room_requests + + class TestCursorACPAdapterDecisions: + @pytest.mark.asyncio + async def test_stale_runtime_cleanup_preserves_live_turn_permission(self) -> None: + tools = DecisionTools() + adapter = CursorACPAdapter() + adapter._runtimes["room-1"] = ACPRuntime(command=[]) + adapter._active_turn = _turn("room-1", tools, "user-1", "session-1") + stale_runtime = ACPRuntime(command=[]) + + await adapter.on_cleanup("room-1", expected_runtime=stale_runtime) + request = ACPPermissionRequest( + room_id="room-1", + session_id="session-1", + tool_call=ACPToolCall("reply", BandTool.SEND_MESSAGE, {}), + options=( + PermissionOption( + optionId="once", name="Allow once", kind=ALLOW_ONCE_KIND + ), + ), + ) + + assert await adapter._resolve_cursor_permission(request) == "once" + assert tools.messages == [] + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "inactive", + ["absent", "other-room", "other-session", "interrupted", "cleaned-up"], + ) + @pytest.mark.parametrize("approval_mode", ["manual", "auto_accept", "auto_decline"]) + async def test_band_permission_requires_a_matching_live_turn( + self, inactive: str, approval_mode: str + ) -> None: + tools = DecisionTools() + adapter = CursorACPAdapter(CursorACPAdapterConfig(approval_mode=approval_mode)) + if inactive != "absent": + adapter._active_turn = _turn( + "other-room" if inactive == "other-room" else "room-1", + tools, + "user-1", + "other-session" if inactive == "other-session" else "session-1", + ) + if inactive == "interrupted": + await adapter.on_interrupt("room-1", ControlMode.STOP) + if inactive == "cleaned-up": + await adapter.on_cleanup("room-1") + request = ACPPermissionRequest( + room_id="room-1", + session_id="session-1", + tool_call=ACPToolCall( + "reply", + adapter._canonical_tool_name( + "band-band_send_message: band_send_message" + ), + {}, + ), + options=( + PermissionOption( + optionId="once", name="Allow once", kind=ALLOW_ONCE_KIND + ), + ), + ) + + assert await adapter._resolve_cursor_permission(request) is None + assert tools.messages == [] + @pytest.mark.asyncio async def test_manual_question_requires_a_valid_room_answer(self) -> None: tools = DecisionTools() @@ -660,8 +1012,7 @@ def request(call_id: str) -> ACPPermissionRequest: ) assert await asyncio.gather(first, repeat) == ["allow-always"] * 2 - asks = [message for message in tools.messages if "needs permission" in message] - assert len(asks) == 1 + assert len(permission_requests(tools)) == 1 @pytest.mark.asyncio @pytest.mark.parametrize( diff --git a/tests/e2e/baseline/smoke/behavior/test_cursor_workflows.py b/tests/e2e/baseline/smoke/behavior/test_cursor_workflows.py index 2fcadf9a3..3ee5ae931 100644 --- a/tests/e2e/baseline/smoke/behavior/test_cursor_workflows.py +++ b/tests/e2e/baseline/smoke/behavior/test_cursor_workflows.py @@ -23,9 +23,9 @@ MemoryType, ) from band.core.types import Capability +from band.integrations.acp.cursor import is_cursor_band_tool from band.integrations.acp.room_emitter import ACP_SESSION_CLOSED_EVENT -from band.runtime.tools.effects import turn_effect -from band.runtime.tools.types import TurnEffect +from band.runtime.tools.types import BandTool from tests.e2e.baseline.agents import Adapter, per_adapter from tests.e2e.baseline.settings import BaselineSettings from tests.e2e.baseline.smoke.samples.approvalroom import ( @@ -60,7 +60,6 @@ TURN_BUDGET_S = BaselineSettings().e2e_timeout WORKFLOW_BUDGET = slow_turn_budget(TURN_BUDGET_S, barriers=6) RECOVERY_BUDGET = slow_turn_budget(TURN_BUDGET_S, barriers=4) -MAX_PERMISSION_REQUESTS = 8 PROJECT_TEST_TIMEOUT_S = 30 SOURCE_FILE = "calculator.py" TEST_FILE = "test_calculator.py" @@ -276,7 +275,7 @@ async def _permission_request( return request -async def _decide_permissions_until_closed( +async def _decide_permission_and_expect_close( room: ApprovalRoom, request: re.Match[str], *, @@ -284,55 +283,39 @@ async def _decide_permissions_until_closed( reply_marker: str, deny_first_tool: str | None = None, ) -> None: - """Answer every permission request until the turn closes, then expect the reply. - - A request can follow the reply (a memory write, a read-back), and an unanswered - one holds the turn open, so the reply alone does not end the decisions. - """ - for attempt in range(MAX_PERMISSION_REQUESTS): - outcome = Outcome.APPROVE - if deny_first_tool is not None and attempt == 0: - assert deny_first_tool in request["tool"], ( - f"Cursor did not request the expected project action: {request['tool']}" - ) - outcome = Outcome.DECLINE - after_decision = await room.decide(outcome, request) - await _within_turn( - room, - checkpoint, - f"Cursor did not confirm the {outcome} decision", - room.shown( - room.dialect.notice(outcome, request).text, since=after_decision - ), - ) - pending = await _within_turn( - room, - checkpoint, - f"Cursor did not finish after {outcome}", - _next_request_or_close(room, checkpoint), + """Decide the single native call; Band calls must not need room approval.""" + assert not is_cursor_band_tool(request["tool"], tuple(BandTool)), ( + f"Cursor asked the room to approve a Band tool: {request['tool']}" + ) + outcome = Outcome.APPROVE + if deny_first_tool is not None: + assert deny_first_tool in request["tool"], ( + f"Cursor did not request the expected project action: {request['tool']}" ) - if pending is None: - await _within_turn( - room, - checkpoint, - "Cursor closed the turn without its reply", - room.shown(reply_marker, since=checkpoint.cursor), - ) - return - requested_tool = pending["tool"].partition(":")[0] - if ( - deny_first_tool is not None - and turn_effect(requested_tool) is not TurnEffect.REPLY - ): - pytest.fail( - "Cursor requested another permission after the denied action: " - f"{room.said_since(checkpoint.cursor)}" - ) - request = pending - pytest.fail( - f"Cursor requested permission more than {MAX_PERMISSION_REQUESTS} times: " + outcome = Outcome.DECLINE + after_decision = await room.decide(outcome, request) + await _within_turn( + room, + checkpoint, + f"Cursor did not confirm the {outcome} decision", + room.shown(room.dialect.notice(outcome, request).text, since=after_decision), + ) + pending = await _within_turn( + room, + checkpoint, + f"Cursor did not finish after {outcome}", + _next_request_or_close(room, checkpoint), + ) + assert pending is None, ( + "Cursor requested another permission after the single native action: " f"{room.said_since(checkpoint.cursor)}" ) + await _within_turn( + room, + checkpoint, + "Cursor closed the turn without its reply", + room.shown(reply_marker, since=checkpoint.cursor), + ) @per_adapter(Adapter.CURSOR_ACP) @@ -379,7 +362,7 @@ async def test_repairs_a_failing_project_after_a_human_gate( f"deny the copy, do not retry; reply with {denied_reply}.", ) request = await _permission_request(room, checkpoint) - await _decide_permissions_until_closed( + await _decide_permission_and_expect_close( room, request, checkpoint=checkpoint, @@ -400,7 +383,7 @@ async def test_repairs_a_failing_project_after_a_human_gate( request = await _permission_request(room, checkpoint) # Nothing may change while the human gate is pending. _assert_project_unchanged(root, original_state) - await _decide_permissions_until_closed( + await _decide_permission_and_expect_close( room, request, checkpoint=checkpoint, reply_marker=report ) assert source.read_text() != original diff --git a/tests/e2e/baseline/smoke/samples/approvals.py b/tests/e2e/baseline/smoke/samples/approvals.py index c2a227d2e..3d9963357 100644 --- a/tests/e2e/baseline/smoke/samples/approvals.py +++ b/tests/e2e/baseline/smoke/samples/approvals.py @@ -455,7 +455,6 @@ def cursor_test_adapter( workspace_for_room: WorkspaceResolver | None = None, custom_section: str = SHELL_PROMPT, capabilities: set[Capability] | None = None, - inject_band_tools: bool = True, ) -> CursorACPAdapter: from band.adapters.cursor_acp import ( # noqa: PLC0415 CursorACPAdapter, @@ -468,7 +467,6 @@ def cursor_test_adapter( "plan_mode": "auto_accept", "decision_timeout_s": setup.wait_timeout_s, "decision_authorized_senders": setup.approvers, - "inject_band_tools": inject_band_tools, # Cursor saves an "allow always" grant to its config dir, which would # let later cells run that command unasked. "env": { @@ -484,13 +482,6 @@ def cursor_test_adapter( ) -def cursor_approval_adapter( - settings: BaselineSettings, setup: AgentSetup -) -> CursorACPAdapter: - # Band MCP replies need a separate Cursor permission after the shell decision. - return cursor_test_adapter(settings, setup, inject_band_tools=False) - - def _allow_option(options: str, *, lasting: bool) -> str: """The allow option among a permission request's option ids: the one-shot one, or with ``lasting`` the one that also covers repeats.""" @@ -625,7 +616,7 @@ def _opencode_handled(reply: str) -> Callable[[re.Match[str]], Notice]: workdir_root=lambda settings: settings.backends.codex_cwd, ), Adapter.CURSOR_ACP: ApprovalDialect( - build=cursor_approval_adapter, + build=cursor_test_adapter, request=template_pattern(PERMISSION_REQUESTED_TEMPLATE), reply=_cursor_reply, notice=_cursor_notice, diff --git a/tests/integrations/acp/acp_toolkit/agent.py b/tests/integrations/acp/acp_toolkit/agent.py index ccf96e682..1aa7a5eef 100644 --- a/tests/integrations/acp/acp_toolkit/agent.py +++ b/tests/integrations/acp/acp_toolkit/agent.py @@ -22,6 +22,7 @@ ) from acp.schema import ( AgentCapabilities, + AllowedOutcome, ConfigOptionUpdate, InitializeResponse, LoadSessionResponse, @@ -365,6 +366,7 @@ def will_call_mcp_tool( arguments: dict[str, Any], server: str = "band", title: str | None = None, + requires_approval: bool = False, ) -> FakeACPAgent: """Call an advertised MCP tool between ACP call and result updates. @@ -388,6 +390,18 @@ async def _action(a: FakeACPAgent, sid: str) -> None: ) await a.emit(sid, update_tool_call(tool_call_id, raw_output=result)) + if requires_approval: + return self.will_execute_if_approved(_action) + self._script.append(_action) + return self + + def will_execute_if_approved(self, action: PromptHandler) -> FakeACPAgent: + """Execute an action only when the preceding permission was granted.""" + + async def _action(a: FakeACPAgent, sid: str) -> None: + if a.approved is True: + await action(a, sid) + self._script.append(_action) return self @@ -468,7 +482,10 @@ async def _action(a: FakeACPAgent, sid: str) -> None: ), ], ) - a.approved = allow_option_id in str(resp) + a.approved = ( + isinstance(resp.outcome, AllowedOutcome) + and resp.outcome.option_id == allow_option_id + ) self._script.append(_action) return self diff --git a/tests/integrations/acp/acp_toolkit/harness.py b/tests/integrations/acp/acp_toolkit/harness.py index c6af9e139..fe0efefce 100644 --- a/tests/integrations/acp/acp_toolkit/harness.py +++ b/tests/integrations/acp/acp_toolkit/harness.py @@ -63,6 +63,15 @@ def __init__(self) -> None: super().__init__() self.transcript: list[RoomActivity] = [] + @property + def reply(self) -> Reply: + return Reply( + messages=self.messages_sent, + events=self.events_sent, + transcript=self.transcript, + memories=self.memories, + ) + async def send_message( self, content: str, mentions: list[str] | list[dict[str, str]] | None = None ) -> dict[str, Any]: @@ -305,12 +314,7 @@ def last_reply(self) -> Reply: value — it reads this instead. """ assert self._last_tools is not None, "send() has not been called yet" - return Reply( - messages=self._last_tools.messages_sent, - events=self._last_tools.events_sent, - transcript=self._last_tools.transcript, - memories=self._last_tools.memories, - ) + return self._last_tools.reply async def send( self, @@ -345,12 +349,7 @@ async def send( is_session_bootstrap=bootstrap, room_id=room, ) - return Reply( - messages=tools.messages_sent, - events=tools.events_sent, - transcript=tools.transcript, - memories=tools.memories, - ) + return tools.reply def session_id(self, room: str) -> str: return self.adapter._room_to_session[room].session_id diff --git a/tests/integrations/acp/acp_toolkit/peer.py b/tests/integrations/acp/acp_toolkit/peer.py new file mode 100644 index 000000000..f34115a5c --- /dev/null +++ b/tests/integrations/acp/acp_toolkit/peer.py @@ -0,0 +1,125 @@ +"""A real stdio ACP peer with scripted process exits and stderr.""" + +from __future__ import annotations + +import argparse +import asyncio +import os +import sys +from enum import StrEnum +from pathlib import Path +from typing import Any + +from acp import run_agent +from acp.schema import InitializeResponse + +from band.integrations.acp.client_adapter import ACPClientAdapterConfig +from band.runtime.tools import BandTool +from band.runtime.tools.registry import BAND_MCP_SERVER_NAME +from tests.integrations.acp.acp_toolkit.agent import FakeACPAgent + +INITIALIZE_PENDING_LINE = "initialization pending" +PEER_EXIT_CODE = 3 +CRASH_MARKER_ARG = "--crash-marker" +RECOVERED_REPLY = "Recovered after the connection failed." + + +class ExitStage(StrEnum): + INITIALIZE = "initialize" + INITIALIZE_WAIT = "initialize-wait" + PROMPT = "prompt" + PROTOCOL_ERROR = "protocol-error" + STDOUT_EOF = "stdout-eof" + EOF = "eof" + + +def stdio_peer_config( + stage: ExitStage, lines: list[str], *, crash_marker: Path | None = None +) -> ACPClientAdapterConfig: + # Windows' venv redirector keeps stdout open until the peer exits. + command = [ + sys._base_executable, + "-m", + "tests.integrations.acp.acp_toolkit.peer", + stage, + str(PEER_EXIT_CODE), + *lines, + ] + if crash_marker is not None: + command.extend((CRASH_MARKER_ARG, str(crash_marker))) + return ACPClientAdapterConfig( + command=command, + env={"__PYVENV_LAUNCHER__": sys.executable}, + inject_band_tools=crash_marker is not None, + ) + + +class StdioPeer(FakeACPAgent): + def __init__( + self, stage: ExitStage, code: int, lines: list[str], crash_marker: Path | None + ) -> None: + super().__init__() + self.stage = stage + self.code = code + self.lines = lines + self.crash_marker = crash_marker + self.on_prompt(self.exit_on_prompt) + + def exit_process(self) -> None: + sys.stderr.write("\n".join(self.lines) + "\n") + sys.stderr.flush() + # Exit inside a request without letting ACP turn the failure into a reply. + os._exit(self.code) + + async def initialize( + self, protocol_version: int, client_capabilities: Any = None, **kwargs: Any + ) -> InitializeResponse: + match self.stage: + case ExitStage.INITIALIZE: + self.exit_process() + case ExitStage.INITIALIZE_WAIT: + sys.stderr.write(INITIALIZE_PENDING_LINE + "\n") + sys.stderr.flush() + await asyncio.Future[None]() + return await super().initialize(protocol_version, client_capabilities, **kwargs) + + async def exit_on_prompt(self, agent: FakeACPAgent, session_id: str) -> None: + del agent + if self.crash_marker is not None: + if self.crash_marker.exists(): + await self.call_mcp_tool( + session_id=session_id, + server=BAND_MCP_SERVER_NAME, + tool_name=BandTool.SEND_MESSAGE, + arguments={"content": RECOVERED_REPLY, "mentions": ["@user-456"]}, + ) + return + self.crash_marker.touch() + match self.stage: + case ExitStage.PROMPT: + self.exit_process() + case ExitStage.PROTOCOL_ERROR: + sys.stdout.write("[]\n") + sys.stdout.flush() + case ExitStage.STDOUT_EOF: + os.close(sys.stdout.fileno()) + case _: + return + # Keep stderr alive until runtime cleanup closes the agent's stdin. + await asyncio.Future[None]() + + +async def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("stage", type=ExitStage, choices=list(ExitStage)) + parser.add_argument("code", type=int) + parser.add_argument("lines", nargs="*") + parser.add_argument(CRASH_MARKER_ARG, type=Path) + args = parser.parse_args() + peer = StdioPeer(args.stage, args.code, args.lines, args.crash_marker) + await run_agent(peer) + peer.exit_process() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/tests/integrations/acp/test_client_runtime.py b/tests/integrations/acp/test_client_runtime.py index 14803b695..ca474dabf 100644 --- a/tests/integrations/acp/test_client_runtime.py +++ b/tests/integrations/acp/test_client_runtime.py @@ -5,6 +5,9 @@ import asyncio import json import logging +from collections.abc import AsyncIterator, Callable, Iterator +from contextlib import AbstractAsyncContextManager, asynccontextmanager +from pathlib import Path from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -16,21 +19,219 @@ NewSessionResponse, ) +from band.core.types import AgentInput, HistoryProvider, PlatformMessage +from band.integrations.acp.client_adapter import ACPClientAdapter from band.integrations.acp.client_profiles import ( CursorACPClientProfile, NoopACPClientProfile, ) from band.integrations.acp.client_runtime import ( ACP_STDIO_LIMIT_BYTES, + STDERR_TAIL_LINES, ACPCollectingClient, ACPRuntime, select_allow_option_id, tcp_spawn_process, ) from band.integrations.acp.session_config import select_ids +from band.integrations.acp.stderr import STDERR_LINE_LOG_TEMPLATE from band.integrations.acp.types import ChunkType, CollectedChunk from band.integrations.mcp import BandMCPTransport from tests.integrations.acp.acp_toolkit import FakeSpawn, select_option +from tests.integrations.acp.acp_toolkit.harness import TranscriptTools +from tests.integrations.acp.acp_toolkit.peer import ( + INITIALIZE_PENDING_LINE, + PEER_EXIT_CODE, + RECOVERED_REPLY, + ExitStage, + stdio_peer_config, +) +from tests.paths import REPO_ROOT + +RUNTIME_LOGGER = "band.integrations.acp.client_runtime" +StdioRuntimeFactory = Callable[ + [ExitStage, list[str]], AbstractAsyncContextManager[ACPRuntime] +] + + +@pytest.fixture +def stdio_runtime() -> StdioRuntimeFactory: + @asynccontextmanager + async def build(stage: ExitStage, lines: list[str]) -> AsyncIterator[ACPRuntime]: + config = stdio_peer_config(stage, lines) + runtime = ACPRuntime( + command=list(config.command), + cwd=str(REPO_ROOT), + env=config.env, + ) + try: + yield runtime + finally: + await runtime.stop() + + return build + + +def runtime_warnings(caplog: pytest.LogCaptureFixture) -> list[logging.LogRecord]: + return [ + record + for record in caplog.records + if record.name == RUNTIME_LOGGER and record.levelno == logging.WARNING + ] + + +def assert_crash_report(caplog: pytest.LogCaptureFixture, lines: list[str]) -> None: + [warning] = runtime_warnings(caplog) + assert warning.getMessage() == ( + f"ACP agent exited with code {PEER_EXIT_CODE}; stderr tail:\n" + + "\n".join(lines) + ) + + +async def deliver_stdio_turn( + adapter: ACPClientAdapter, message: PlatformMessage, tools: TranscriptTools +) -> None: + await adapter.on_event( + AgentInput( + msg=message, + tools=tools, + history=HistoryProvider(raw=[]), + participants_msg=None, + contacts_msg=None, + is_session_bootstrap=True, + room_id=message.room_id, + ) + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "stage", + [ + ExitStage.INITIALIZE, + ExitStage.PROMPT, + ExitStage.STDOUT_EOF, + ExitStage.PROTOCOL_ERROR, + ], +) +async def test_crashed_stdio_agent_reports_exit_and_stderr( + stdio_runtime: StdioRuntimeFactory, + caplog: pytest.LogCaptureFixture, + tmp_path: Path, + stage: ExitStage, +) -> None: + lines = ["authentication failed", "peer crashed"] + with caplog.at_level(logging.WARNING, logger=RUNTIME_LOGGER): + async with stdio_runtime(stage, lines) as runtime: + with pytest.raises(ConnectionError, match="Connection closed"): + await runtime.start() + session_id = await runtime.create_session( + cwd=str(tmp_path), mcp_servers=[] + ) + await runtime.prompt(session_id=session_id, prompt_text="crash") + + assert_crash_report(caplog, lines) + + +@pytest.mark.asyncio +async def test_crashed_stdio_agent_reports_only_the_stderr_tail( + stdio_runtime: StdioRuntimeFactory, caplog: pytest.LogCaptureFixture +) -> None: + lines = [f"diagnostic {index}" for index in range(STDERR_TAIL_LINES + 5)] + with caplog.at_level(logging.WARNING, logger=RUNTIME_LOGGER): + async with stdio_runtime(ExitStage.INITIALIZE, lines) as runtime: + with pytest.raises(ConnectionError, match="Connection closed"): + await runtime.start() + + [warning] = runtime_warnings(caplog) + assert warning.getMessage().splitlines()[1:] == lines[-STDERR_TAIL_LINES:] + + +@pytest.mark.asyncio +async def test_deliberate_stdio_stop_does_not_warn_about_nonzero_exit( + stdio_runtime: StdioRuntimeFactory, caplog: pytest.LogCaptureFixture +) -> None: + with caplog.at_level(logging.WARNING, logger=RUNTIME_LOGGER): + async with stdio_runtime(ExitStage.EOF, ["deliberately stopped"]) as runtime: + await runtime.start() + + assert runtime_warnings(caplog) == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stage", [ExitStage.STDOUT_EOF, ExitStage.PROTOCOL_ERROR]) +async def test_failed_stdio_turn_reports_crash_and_next_turn_recovers( + caplog: pytest.LogCaptureFixture, + tmp_path: Path, + sample_platform_message: PlatformMessage, + stage: ExitStage, +) -> None: + lines = ["fatal peer diagnostic"] + adapter = ACPClientAdapter( + stdio_peer_config(stage, lines, crash_marker=tmp_path / "crashed"), + workspace_for_room=lambda room_id: str(REPO_ROOT), + ) + failed = TranscriptTools() + recovered = TranscriptTools() + with caplog.at_level(logging.WARNING, logger=RUNTIME_LOGGER): + await adapter.on_started("Stdio peer", "Real subprocess") + try: + with pytest.raises(ConnectionError, match="Connection closed"): + await deliver_stdio_turn(adapter, sample_platform_message, failed) + await deliver_stdio_turn(adapter, sample_platform_message, recovered) + finally: + await adapter.stop() + + assert_crash_report(caplog, lines) + assert failed.reply.texts == [] + assert len(failed.reply.errors) == 1 + assert recovered.reply.texts == [RECOVERED_REPLY] + assert recovered.reply.errors == [] + assert recovered.turn.replied + + +class InitializationObserver(logging.Handler): + def __init__(self) -> None: + super().__init__() + self.started = asyncio.Event() + + def emit(self, record: logging.LogRecord) -> None: + if record.getMessage() == STDERR_LINE_LOG_TEMPLATE % INITIALIZE_PENDING_LINE: + self.started.set() + + +@pytest.fixture +def initialization_observer() -> Iterator[InitializationObserver]: + observer = InitializationObserver() + logger = logging.getLogger(RUNTIME_LOGGER) + logger.addHandler(observer) + try: + yield observer + finally: + logger.removeHandler(observer) + + +@pytest.mark.asyncio +async def test_cancelled_stdio_start_does_not_warn_about_nonzero_exit( + stdio_runtime: StdioRuntimeFactory, + initialization_observer: InitializationObserver, + caplog: pytest.LogCaptureFixture, +) -> None: + with caplog.at_level(logging.DEBUG, logger=RUNTIME_LOGGER): + async with stdio_runtime( + ExitStage.INITIALIZE_WAIT, ["cancelled initialization"] + ) as runtime: + startup = asyncio.create_task(runtime.start()) + try: + await initialization_observer.started.wait() + startup.cancel() + with pytest.raises(asyncio.CancelledError): + await startup + finally: + startup.cancel() + await asyncio.gather(startup, return_exceptions=True) + + assert runtime_warnings(caplog) == [] class TestSelectAllowOptionId: @@ -832,7 +1033,6 @@ async def test_create_session_and_prompt_use_active_connection(self) -> None: runtime = ACPRuntime(command=["codex"]) runtime._conn = mock_conn runtime._client = ACPCollectingClient() - runtime._client._session_chunks["sess-1"] = [] session_id = await runtime.create_session(cwd="/tmp", mcp_servers=[]) chunks = await runtime.prompt(session_id=session_id, prompt_text="hello") @@ -934,7 +1134,7 @@ async def test_load_session_timeout_is_treated_as_unavailable(self) -> None: runtime._agent_supports_session_load = True with patch( - "band.integrations.acp.client_runtime.ACP_SESSION_LOAD_TIMEOUT_SECONDS", + "band.integrations.acp.sessions.ACP_SESSION_LOAD_TIMEOUT_SECONDS", 0.01, ): assert not await runtime.load_session( diff --git a/tests/integrations/acp/test_cursor_identity.py b/tests/integrations/acp/test_cursor_identity.py new file mode 100644 index 000000000..7c5668e85 --- /dev/null +++ b/tests/integrations/acp/test_cursor_identity.py @@ -0,0 +1,88 @@ +"""Cursor MCP identity survives permissions and room narration.""" + +from __future__ import annotations + +import pytest +from acp.helpers import update_agent_message_text +from acp.schema import ToolCallProgress, ToolCallStart +from pydantic import BaseModel, create_model + +from band.adapters.cursor_acp import CursorACPAdapter +from band.integrations.acp.room_emitter import RoomTurnEmitter +from band.integrations.acp.types import ToolCallRoomEvent, ToolResultRoomEvent +from band.runtime.tools import BAND_MCP_SERVER_NAME, BandTool, mcp_tool_spelling +from band.testing import FakeAgentTools, events_of_type + + +class AliasArguments(BaseModel): + """An accepted custom name that overlaps the MCP reply spelling.""" + + content: str + + +async def alias_reply(arguments: AliasArguments) -> str: + return arguments.content + + +@pytest.mark.asyncio +@pytest.mark.parametrize("with_alias", [False, True], ids=["plain", "custom-alias"]) +@pytest.mark.parametrize( + ("title", "is_reply"), + [ + ("band-band_send_message: band_send_message", True), + ("other-band_send_message: band_send_message", False), + ("band-band_send_message: band_no_reply", False), + ], + ids=["band-title", "foreign-server", "mismatched-halves"], +) +async def test_cursor_title_preserves_reply_identity_in_room_events( + with_alias: bool, + title: str, + is_reply: bool, +) -> None: + alias_model = create_model( + f"{mcp_tool_spelling(BAND_MCP_SERVER_NAME, BandTool.SEND_MESSAGE)}Input", + __base__=AliasArguments, + ) + adapter = CursorACPAdapter( + additional_tools=[(alias_model, alias_reply)] if with_alias else None, + ) + client = adapter._runtime_client_factory() + tools = FakeAgentTools() + async with RoomTurnEmitter( + tools, session_id="session", room_id="room", records_tool_effects=True + ) as emitter: + client.set_sink("session", emitter.emit) + await client.session_update( + "session", + ToolCallStart( + session_update="tool_call", + tool_call_id="reply", + title=title, + status="in_progress", + raw_input={"content": "Reply"}, + ), + ) + await client.session_update( + "session", + ToolCallProgress( + session_update="tool_call_update", + tool_call_id="reply", + status="completed", + raw_output={"id": "message"}, + ), + ) + await client.session_update("session", update_agent_message_text("Posted it.")) + await client.flush("session") + + call = ToolCallRoomEvent.model_validate_json( + events_of_type(tools, "tool_call")[0]["content"] + ) + result = ToolResultRoomEvent.model_validate_json( + events_of_type(tools, "tool_result")[0]["content"] + ) + assert tools.turn.replied is is_reply + assert call.name == result.name == (BandTool.SEND_MESSAGE if is_reply else title) + assert [event["content"] for event in events_of_type(tools, "thought")] == ( + [] if is_reply else ["Posted it."] + ) diff --git a/tests/integrations/acp/test_stderr.py b/tests/integrations/acp/test_stderr.py new file mode 100644 index 000000000..a227e36fe --- /dev/null +++ b/tests/integrations/acp/test_stderr.py @@ -0,0 +1,126 @@ +"""Subprocess-boundary regressions for ACP stderr diagnostics.""" + +from __future__ import annotations + +import asyncio +import logging +import sys +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager +from pathlib import Path + +import pytest + +from band.integrations.acp.stderr import ACPStderrDrain + +LOGGER = logging.getLogger(__name__) +EXIT_CODE = 3 +FATAL_LINE = "fatal agent diagnostic" +# The Windows venv redirector keeps inherited stdout open until its child exits. +PYTHON_EXECUTABLE = getattr(sys, "_base_executable", sys.executable) + + +async def wait_for_exit(process: asyncio.subprocess.Process) -> None: + async with asyncio.timeout(3): + while process.returncode is None: + await asyncio.sleep(0.01) + + +@asynccontextmanager +async def inherited_stderr_peer( + release_file: Path, *, buffered_stdout: bool +) -> AsyncIterator[asyncio.subprocess.Process]: + helper = ( + "import pathlib, sys, time\n" + "release = pathlib.Path(sys.argv[1])\n" + "while not release.exists(): time.sleep(0.01)\n" + ) + source = ( + "import os, subprocess, sys\n" + "subprocess.Popen([sys.executable, '-c', sys.argv[1], sys.argv[2]], " + "stdin=subprocess.DEVNULL, stdout=subprocess.DEVNULL)\n" + "sys.stderr.write(sys.argv[3] + '\\n'); sys.stderr.flush()\n" + "sys.stdout.write(sys.argv[4]); sys.stdout.flush()\n" + "os._exit(int(sys.argv[5]))\n" + ) + process = await asyncio.create_subprocess_exec( + PYTHON_EXECUTABLE, + "-c", + source, + helper, + str(release_file), + FATAL_LINE, + "partial ACP message" if buffered_stdout else "", + str(EXIT_CODE), + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + ) + try: + yield process + finally: + release_file.touch() + await process.communicate() + + +def crash_reports(caplog: pytest.LogCaptureFixture) -> list[str]: + return [ + record.getMessage() + for record in caplog.records + if record.name == LOGGER.name and record.levelno == logging.WARNING + ] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("buffered_stdout", [False, True]) +async def test_crash_is_reported_while_helper_keeps_stderr_open( + tmp_path: Path, caplog: pytest.LogCaptureFixture, buffered_stdout: bool +) -> None: + release_file = tmp_path / "release-helper" + with caplog.at_level(logging.WARNING, logger=LOGGER.name): + async with inherited_stderr_peer( + release_file, buffered_stdout=buffered_stdout + ) as process: + drain = ACPStderrDrain(LOGGER) + drain.start(process) + await wait_for_exit(process) + if buffered_stdout: + assert process.stdout is not None and not process.stdout.at_eof() + drain.expect_exit() + async with asyncio.timeout(1): + await drain.finish() + assert not release_file.exists() + assert crash_reports(caplog) == [ + f"ACP agent exited with code {EXIT_CODE}; stderr tail:\n{FATAL_LINE}" + ] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("error", [ConnectionResetError("reset"), OSError("pipe")]) +async def test_stderr_read_failure_does_not_escape_cleanup( + caplog: pytest.LogCaptureFixture, error: OSError +) -> None: + process = await asyncio.create_subprocess_exec( + PYTHON_EXECUTABLE, + "-c", + "import sys; sys.exit(3)", + stderr=asyncio.subprocess.PIPE, + ) + assert process.stderr is not None + process.stderr.set_exception(error) + drain = ACPStderrDrain(LOGGER) + with caplog.at_level(logging.DEBUG, logger=LOGGER.name): + drain.start(process) + try: + await drain.finish() + finally: + await wait_for_exit(process) + await process.wait() + assert [ + record.levelno + for record in caplog.records + if record.name == LOGGER.name + and record.getMessage() == "Error reading ACP agent stderr" + ] == [logging.DEBUG] + assert crash_reports(caplog) == [ + f"ACP agent exited with code {EXIT_CODE}; stderr tail:\n" + ]