From 75540d5462e0ac4b4a1b9723c897dada90b08fc4 Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Thu, 3 Sep 2026 14:15:41 -0700 Subject: [PATCH 01/23] feat(runtime): add Orka harness v2 ACP mode Signed-off-by: Sertac Ozercan --- README.md | 13 +- docs/orka.md | 65 +- docs/runtime-adapters.md | 19 +- runtimes/common/README.md | 23 +- runtimes/common/agentkit_serve_common/acp.py | 687 ++++++++++++++++++ runtimes/common/agentkit_serve_common/cli.py | 27 +- .../common/agentkit_serve_common/runtime.py | 4 +- runtimes/common/tests/test_acp_protocol.py | 644 ++++++++++++++++ runtimes/common/tests/test_cli_protocol.py | 33 + .../langgraph/agentkit_serve/agent_factory.py | 4 + runtimes/langgraph/tests/test_guardrails.py | 2 + .../agentkit_serve/agent_factory.py | 30 +- .../tests/test_guardrails.py | 7 +- .../agentkit_serve/agent_factory.py | 7 + .../tests/test_unsupported_features.py | 2 + 15 files changed, 1546 insertions(+), 21 deletions(-) create mode 100644 runtimes/common/agentkit_serve_common/acp.py create mode 100644 runtimes/common/tests/test_acp_protocol.py diff --git a/README.md b/README.md index 7cecdea..fa73b7b 100644 --- a/README.md +++ b/README.md @@ -63,7 +63,7 @@ The image also exposes: ## Select a protocol surface -AgentKit builds one custom-agent image and selects the HTTP protocol at runtime: +AgentKit builds one custom-agent image and selects the protocol at runtime: ```sh # default standalone OpenAI-compatible surface @@ -78,6 +78,10 @@ docker run \ -e AGENTKIT_PROTOCOL=orka \ -e AGENTKIT_AUTH_TOKEN=dev-token \ ... image + +# ACP stdio child for an Orka harness v2 supervisor. The supervisor supplies +# the provider proxy, MCP broker, model, and image-bound configuration digest. +agentkit-serve --config /agent/agent.yaml --protocol acp ``` Protocol endpoints: @@ -87,6 +91,7 @@ Protocol endpoints: | `openai` | `/healthz`, `/v1/models`, `/v1/chat/completions` | Default, non-streaming Chat Completions. | | `foundry` | `/readiness`, `/invocations`, `/responses` | `/responses` is `foundry-responses-minimal`: synchronous/non-streaming only. | | `orka` | `/v1/health`, `/v1/capabilities`, `/v1/turns`, `/v1/turns/{turnID}/events`, `/v1/turns/{turnID}/continue`, `/v1/turns/{turnID}/cancel` | Observed-mode `orka.harness.v1` over HTTP+SSE by default. AgentKit reports frames; Orka enforces policy. Brokered read/write/coordination are feature-gated for conformance. | +| `acp` | stdin/stdout | ACP protocol v1 child mode for Orka `orka.harness.v2`. It opens no network listener and accepts only the supervisor's loopback provider proxy and prompt-scoped HTTP MCP server. | After deploying the image with `AGENTKIT_PROTOCOL=orka` and an `AGENTKIT_AUTH_TOKEN` sourced from the Orka client-auth Secret, render an Orka @@ -119,6 +124,12 @@ fields (`runtimeSessionID`, `turnID`, `correlationID`), `createdAt`, `contentText` for runtime output, and `completed` / `failed` terminal payloads. See [`docs/orka.md`](docs/orka.md) for complete request/response examples. +For harness v2, layer Orka's supervisor onto a digest-pinned AgentKit image and +register the resulting service as a strict-governed external `AgentRuntime`. +The supervisor starts `agentkit-serve --protocol acp` as an isolated child and +keeps provider, MCP, workspace, permission, and publication authority outside +AgentKit. See [`docs/orka.md`](docs/orka.md) for the composition contract. + Note: this repository's `agentkit-serve` runtime is distinct from OpenAI's public AgentKit/Agents SDK product surface unless a future adapter explicitly targets it. diff --git a/docs/orka.md b/docs/orka.md index 2af06f7..487507e 100644 --- a/docs/orka.md +++ b/docs/orka.md @@ -1,4 +1,60 @@ -# Register an AgentKit image with Orka +# Use an AgentKit image with Orka + +AgentKit supports two distinct Orka integrations. New BYO deployments should +use `orka.harness.v2`: Orka's hardened supervisor runs AgentKit as an ACP stdio +child. The older `orka.harness.v1` mode remains an observed HTTP+SSE adapter. + +## Harness v2 BYO runtime + +Build the AgentKit agent image first and address it by digest. In an Orka +checkout, layer the v2 supervisor onto that immutable image: + +```sh +make docker-build-acp-agentkit-runtime \ + AGENTKIT_RUNTIME_IMAGE=ghcr.io/acme/fibey@sha256: \ + ACP_AGENTKIT_RUNTIME_IMG=ghcr.io/acme/fibey-orka-v2:dev +``` + +The composed image runs `orka-acp-runtime`. For each RuntimeSession, the +supervisor starts this child under a private UID/GID and session tree: + +```sh +/opt/agentkit/bin/agentkit-serve \ + --config /agent/agent.yaml \ + --protocol acp +``` + +The v2 path is strict: + +- `/agent/agent.yaml` must not contain direct `tools`, `brokeredTools`, or + context providers; +- the registered model must equal `model.name` in the baked config; +- `agentConfigurationDigest` is `sha256:` plus the SHA-256 of the exact + `/agent/agent.yaml` bytes; +- the runtime advertises only the `agentkit-serve-acp` adapter digest; +- Orka sends `AgentConfiguration: null`; the image-bound config is authoritative; +- provider calls use the supervisor's loopback proxy, and tools use its one + prompt-scoped loopback HTTP MCP server; +- the child retains successful user/assistant history for Session continuation + and discards cancelled or failed prompt history. + +Deploy the composed image as an operator-owned v2 supervisor service, configure +the standard `ORKA_ACP_*` profile, fence, token-file, and runtime identity +settings, then register it with Orka's strict-governed `AgentRuntime` sample. +The profile model and `agentConfigurationDigest` must match the baked AgentKit +config, and the registration's `adapterName` must be `agentkit-serve-acp`. + +Orka freezes the AgentRuntime UID, generation, endpoint, profile, authentication +Secret versions, and observed instance into each Task binding. It revalidates +them before dispatch and recovery mutations. `Task.spec.execution.workspace` +is not supported for external runtimes; repository input still uses +`Task.spec.workspace`. + +See Orka's `website/docs/guides/bring-your-own-agent-runtime.md` and +`config/samples/core_v1alpha1_agentruntime.yaml` for the registration and +authentication contract. + +## Harness v1 observed mode AgentKit images can expose observed-mode `orka.harness.v1` without rebuilding the agent. Start the same image with Orka mode enabled: @@ -240,10 +296,11 @@ For deeper local validation, run the common Python Orka protocol tests: uv run --directory runtimes/common --extra dev pytest -q tests/test_orka_protocol.py ``` -## Render an AgentRuntime manifest +## Render a harness v1 AgentRuntime manifest -The current Orka `AgentRuntime` CRD supports external endpoints first. Deploy the -AgentKit image yourself (for example as a Kubernetes Deployment/Service) with: +The AgentKit renderer currently emits the harness v1 registration shape. Deploy +the AgentKit image yourself, for example as a Kubernetes Deployment/Service, +with: - `AGENTKIT_PROTOCOL=orka` - `AGENTKIT_BIND=0.0.0.0` so the Kubernetes Service can reach the harness outside diff --git a/docs/runtime-adapters.md b/docs/runtime-adapters.md index 9778674..f1d31fa 100644 --- a/docs/runtime-adapters.md +++ b/docs/runtime-adapters.md @@ -12,10 +12,11 @@ must be identical across adapters: | Module | Responsibility | |---|---| | `config.py` | Strict `/agent/agent.yaml` reader and ABI version check. | -| `cli.py` | `agentkit-serve --config ... --protocol openai\|foundry\|orka`, bind/port handling, auth startup gates. | +| `cli.py` | `agentkit-serve --config ... --protocol openai\|foundry\|orka\|acp`, bind/port handling, auth startup gates. | | `server.py` | FastAPI app and OpenAI-compatible response/error envelopes. | | `foundry.py` | Foundry `/readiness`, `/invocations`, and minimal `/responses` skin. | | `orka.py` | Observed-mode `orka.harness.v1` HTTP+SSE skin. | +| `acp.py` | Strict ACP stdio child for an Orka `orka.harness.v2` supervisor. | | `conversation.py` | Protocol request normalization into `RunRequest`. | | `runtime.py` | `RuntimeFactory`, `RuntimeSession`, `RunResult`, `AgentRunError`. | | `adapter_support.py` | API-key lookup, tool env projection, timeout parsing, error normalization. | @@ -26,7 +27,7 @@ The protocol app factories receive an adapter module that satisfies `RuntimeSession.run(request)`, so it never imports pydantic-ai, Microsoft Agent Framework, LangChain, OpenAI SDK types, Azure, Foundry SDKs, or Orka controllers. -## HTTP surface +## Protocol surfaces All adapters can serve the same selected protocol surface. `openai` is the default. `foundry` and `orka` are selected with `--protocol` or @@ -45,6 +46,13 @@ port; generated images expose both `8080` and `8088` in OCI metadata for that case. Orka mode exposes `orka.harness.v1` health, capabilities, turn acceptance, SSE replay, and cancel endpoints. +ACP mode opens no listener. It speaks newline-delimited ACP JSON-RPC on stdin +and stdout. The child verifies the configured model and SHA-256 digest of the +exact `/agent/agent.yaml` bytes before accepting a session. It rejects baked +direct tools, `brokeredTools`, and context providers. At session creation it +accepts at most one loopback HTTP MCP server with bearer authentication, which +is the prompt-scoped broker created by the Orka supervisor. + Request behavior is intentionally narrow: - `stream: true` returns HTTP 400 with code `stream_unsupported`. @@ -76,7 +84,7 @@ configured `baseURL` at runtime. ## Network posture -The generated image defaults to `AGENTKIT_BIND=127.0.0.1`. At runtime: +The generated image defaults to `AGENTKIT_BIND=127.0.0.1`. In HTTP modes: - loopback binds need no token except in Orka mode, - non-loopback binds such as `0.0.0.0` require `AGENTKIT_AUTH_TOKEN`, and @@ -87,6 +95,11 @@ OpenAI `/healthz` and Orka `/v1/health` and `/v1/capabilities` are intentionally unauthenticated so container platforms and orchestrators can probe/discover the service. Orka turn, event, cancel, and output endpoints always require a token. +ACP mode ignores bind and port settings because it uses stdio. The Orka +supervisor injects only `AGENTKIT_ACP_PROVIDER_BASE_URL`, +`AGENTKIT_ACP_PROVIDER_TOKEN`, `AGENTKIT_ACP_MODEL`, and +`AGENTKIT_ACP_AGENT_CONFIGURATION_DIGEST` into the child. + ## Tool lifecycle and env projection Tools are MCP servers declared in the ABI. Stdio tools use `name`, `command`, diff --git a/runtimes/common/README.md b/runtimes/common/README.md index 3c9816c..cfa53db 100644 --- a/runtimes/common/README.md +++ b/runtimes/common/README.md @@ -9,9 +9,11 @@ or LangGraph. - `config.py` — strict `/agent/agent.yaml` ABI loader, version check, env requirement validation, and provider-neutral schema validation. -- `cli.py` — `agentkit-serve --config ... --protocol openai|foundry|orka`, +- `cli.py` — `agentkit-serve --config ... --protocol acp|openai|foundry|orka`, bind/port handling, and startup auth gates for non-loopback binds and Orka protected endpoints. +- `acp.py` — ACP protocol v1 over newline-delimited JSON-RPC on stdio for the + Orka harness v2 supervisor. - `server.py` — FastAPI app for `/healthz`, `/v1/models`, and `/v1/chat/completions`. - `foundry.py` — reusable Foundry Hosted Agent protocol wrapper for @@ -38,6 +40,25 @@ packages or touches raw framework agent lifecycle. This keeps framework dependency lock-in inside each adapter's `agent_factory.py`. +## Orka ACP child mode + +Orka harness v2 starts the adapter as one child process per RuntimeSession: + +```sh +agentkit-serve --config /agent/agent.yaml --protocol acp +``` + +ACP mode requires `AGENTKIT_ACP_AGENT_CONFIGURATION_DIGEST` to equal the +`sha256:` digest of the exact config file bytes and `AGENTKIT_ACP_MODEL` to +equal `model.name`. It replaces the baked model endpoint and credential with +`AGENTKIT_ACP_PROVIDER_BASE_URL` and `AGENTKIT_ACP_PROVIDER_TOKEN`. + +The child accepts one ACP session, text prompts, cancellation, and at most one +loopback HTTP MCP server carrying a bearer Authorization header. The runtime +keeps successful user and assistant turns for later prompts. It rejects baked +`tools`, `brokeredTools`, and context providers. Orka owns process and workspace +isolation, prompt-scoped MCP authority, provider proxying, and cleanup proof. + ## Adding a runtime adapter A new single-agent adapter should provide: diff --git a/runtimes/common/agentkit_serve_common/acp.py b/runtimes/common/agentkit_serve_common/acp.py new file mode 100644 index 0000000..668d35e --- /dev/null +++ b/runtimes/common/agentkit_serve_common/acp.py @@ -0,0 +1,687 @@ +"""Agent Client Protocol stdio adapter for Orka ACP supervisors. + +The Orka supervisor owns process isolation, prompt leases, provider and MCP +brokers, workspace governance, and descendant cleanup. This module is only the +provider child. It translates the newline-delimited ACP JSON-RPC protocol into +the framework-neutral ``RuntimeFactory`` / ``RuntimeSession`` seam. +""" + +from __future__ import annotations + +import asyncio +import hashlib +import inspect +import ipaddress +import json +import os +import secrets +import sys +from contextlib import redirect_stdout +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, Awaitable, BinaryIO, Callable, Mapping +from urllib.parse import urlsplit + +from .config import AgentSpec +from .conversation import ConversationTurn, RunRequest +from .runtime import AgentRunError, RuntimeFactory, RuntimeSession + +ACP_PROTOCOL_VERSION = 1 +ACP_AGENT_CONFIGURATION_DIGEST_ENV = "AGENTKIT_ACP_AGENT_CONFIGURATION_DIGEST" +ACP_MODEL_ENV = "AGENTKIT_ACP_MODEL" +ACP_PROVIDER_BASE_URL_ENV = "AGENTKIT_ACP_PROVIDER_BASE_URL" +ACP_PROVIDER_TOKEN_ENV = "AGENTKIT_ACP_PROVIDER_TOKEN" + +_JSONRPC_VERSION = "2.0" +_METHOD_INITIALIZE = "initialize" +_METHOD_SESSION_NEW = "session/new" +_METHOD_SESSION_PROMPT = "session/prompt" +_METHOD_SESSION_CANCEL = "session/cancel" +_METHOD_SESSION_UPDATE = "session/update" +_METHOD_CANCEL_REQUEST = "$/cancel_request" +_MAX_MESSAGE_BYTES = 8 << 20 + +_PARSE_ERROR = -32700 +_INVALID_REQUEST = -32600 +_METHOD_NOT_FOUND = -32601 +_INVALID_PARAMS = -32602 +_INTERNAL_ERROR = -32603 +_RUNTIME_ERROR = -32000 + +_MessageSender = Callable[[Mapping[str, Any]], Awaitable[None]] + + +class ACPProtocolError(Exception): + """A JSON-RPC error safe to return to the ACP client.""" + + def __init__(self, code: int, message: str, data: Mapping[str, Any] | None = None) -> None: + super().__init__(message) + self.code = code + self.message = message + self.data = dict(data) if data is not None else None + + +class ACPConfigurationError(ValueError): + """An ACP startup binding or strict-mode configuration failure.""" + + +@dataclass +class _SessionState: + session_id: str + context: RuntimeSession + runtime: RuntimeSession + previous_environment: dict[str, str | None] + history: list[ConversationTurn] = field(default_factory=list) + active_request_key: str | None = None + active_run: asyncio.Task[Any] | None = None + cancel_requested: bool = False + + +def _factory_supports_http_mcp(factory: RuntimeFactory) -> bool: + capability = getattr(factory, "supports_acp_http_mcp", None) + return bool(capability()) if callable(capability) else False + + +def _request_key(value: Any) -> str: + if isinstance(value, bool) or not isinstance(value, (int, str)): + raise ACPProtocolError(_INVALID_REQUEST, "JSON-RPC id must be a string or integer") + return json.dumps(value, ensure_ascii=False, separators=(",", ":")) + + +def _required_object(value: Any, *, name: str = "params") -> dict[str, Any]: + if not isinstance(value, dict): + raise ACPProtocolError(_INVALID_PARAMS, f"{name} must be an object") + return value + + +def _required_string(value: Any, *, name: str) -> str: + if not isinstance(value, str) or not value.strip(): + raise ACPProtocolError(_INVALID_PARAMS, f"{name} must be a non-empty string") + if value != value.strip(): + raise ACPProtocolError(_INVALID_PARAMS, f"{name} must not contain surrounding whitespace") + return value + + +def _required_prompt_text(value: Any, *, name: str) -> str: + if not isinstance(value, str) or not value.strip(): + raise ACPProtocolError(_INVALID_PARAMS, f"{name} must be a non-empty string") + return value + + +async def _discard_runtime_session(runtime: RuntimeSession, session_id: str) -> None: + discard = getattr(runtime, "discard_session", None) + if not callable(discard): + return + result = discard(session_id) + if inspect.isawaitable(result): + await result + + +def _loopback_http_url(value: str, *, name: str) -> str: + try: + parsed = urlsplit(value) + parsed.port + except ValueError as exc: + raise ACPProtocolError(_INVALID_PARAMS, f"{name} must be a valid loopback HTTP URL") from exc + if ( + parsed.scheme not in {"http", "https"} + or not parsed.hostname + or parsed.username is not None + or parsed.password is not None + or parsed.query + or parsed.fragment + ): + raise ACPProtocolError(_INVALID_PARAMS, f"{name} must be an absolute loopback HTTP URL") + hostname = parsed.hostname.lower() + if hostname != "localhost": + try: + address = ipaddress.ip_address(hostname) + except ValueError as exc: + raise ACPProtocolError( + _INVALID_PARAMS, + f"{name} must use a loopback IP address", + ) from exc + if not address.is_loopback: + raise ACPProtocolError(_INVALID_PARAMS, f"{name} must use a loopback IP address") + return value + + +def _safe_environment_value(value: Any, *, name: str) -> str: + if not isinstance(value, str) or not value: + raise ACPProtocolError(_INVALID_PARAMS, f"{name} must be a non-empty string") + if "\r" in value or "\n" in value or "\x00" in value: + raise ACPProtocolError(_INVALID_PARAMS, f"{name} contains forbidden control characters") + return value + + +def _required_environment(name: str) -> str: + try: + return _safe_environment_value(os.environ.get(name), name=name) + except ACPProtocolError as exc: + raise ACPConfigurationError(exc.message) from exc + + +def validate_acp_runtime_binding(config_path: str | Path, spec: AgentSpec) -> None: + """Fail closed unless Orka's immutable profile matches the baked config.""" + + if spec.tools: + raise ACPConfigurationError("ACP strict mode rejects baked direct tools") + if spec.brokered_tools: + raise ACPConfigurationError("ACP strict mode rejects baked brokeredTools") + if spec.context.providers: + raise ACPConfigurationError("ACP strict mode rejects baked context providers") + + expected_digest = _required_environment(ACP_AGENT_CONFIGURATION_DIGEST_ENV) + prefix = "sha256:" + encoded_digest = expected_digest.removeprefix(prefix) + if ( + not expected_digest.startswith(prefix) + or len(encoded_digest) != 64 + or encoded_digest.lower() != encoded_digest + ): + raise ACPConfigurationError( + f"{ACP_AGENT_CONFIGURATION_DIGEST_ENV} must be a lowercase sha256 digest" + ) + try: + bytes.fromhex(encoded_digest) + except ValueError as exc: + raise ACPConfigurationError( + f"{ACP_AGENT_CONFIGURATION_DIGEST_ENV} must be a lowercase sha256 digest" + ) from exc + + path = Path(config_path) + try: + config_bytes = path.read_bytes() + except OSError as exc: + raise ACPConfigurationError(f"cannot read ACP agent config {path}: {exc}") from exc + actual_digest = prefix + hashlib.sha256(config_bytes).hexdigest() + if not secrets.compare_digest(expected_digest, actual_digest): + raise ACPConfigurationError( + f"{ACP_AGENT_CONFIGURATION_DIGEST_ENV} does not match the exact agent config bytes" + ) + + expected_model = _required_environment(ACP_MODEL_ENV) + if not secrets.compare_digest(expected_model, spec.model.name): + raise ACPConfigurationError(f"{ACP_MODEL_ENV} does not match model.name in the agent config") + + provider_base_url = _required_environment(ACP_PROVIDER_BASE_URL_ENV) + try: + _loopback_http_url(provider_base_url, name=ACP_PROVIDER_BASE_URL_ENV) + except ACPProtocolError as exc: + raise ACPConfigurationError(exc.message) from exc + _required_environment(ACP_PROVIDER_TOKEN_ENV) + + +class ACPStdioServer: + """Concurrent ACP JSON-RPC dispatcher over an injected message sender.""" + + def __init__(self, spec: AgentSpec, factory: RuntimeFactory, send: _MessageSender) -> None: + self.spec = spec + self.factory = factory + self.send = send + self.initialized = False + self.http_mcp = ( + _factory_supports_http_mcp(factory) + and not spec.tools + and not spec.brokered_tools + and not spec.context.providers + ) + self.sessions: dict[str, _SessionState] = {} + self.requests: dict[str, asyncio.Task[None]] = {} + self.prompt_requests: dict[str, _SessionState] = {} + self.session_creation_active = False + self.closed = False + + async def accept_line(self, line: bytes) -> None: + """Decode and schedule one newline-delimited JSON-RPC message.""" + + if len(line) > _MAX_MESSAGE_BYTES: + await self._send_error(None, _INVALID_REQUEST, "ACP message exceeds the 8 MiB limit") + return + try: + message = json.loads(line) + except (UnicodeDecodeError, json.JSONDecodeError): + await self._send_error(None, _PARSE_ERROR, "invalid JSON") + return + await self.accept(message) + + async def accept(self, message: Any) -> None: + if self.closed: + return + if not isinstance(message, dict) or message.get("jsonrpc") != _JSONRPC_VERSION: + response_id = message.get("id") if isinstance(message, dict) else None + await self._send_error(response_id, _INVALID_REQUEST, "invalid JSON-RPC request") + return + method = message.get("method") + if not isinstance(method, str) or not method: + await self._send_error(message.get("id"), _INVALID_REQUEST, "JSON-RPC method is required") + return + if "id" not in message: + await self._handle_notification(method, message.get("params")) + return + response_id = message["id"] + try: + key = _request_key(response_id) + except ACPProtocolError as exc: + await self._send_protocol_error(response_id, exc) + return + if key in self.requests: + await self._send_error(response_id, _INVALID_REQUEST, "duplicate active JSON-RPC id") + return + task = asyncio.create_task( + self._dispatch_request(response_id, key, method, message.get("params")), + name=f"agentkit-acp-{method}", + ) + self.requests[key] = task + task.add_done_callback(lambda completed, request_key=key: self._request_done(request_key, completed)) + # Give prompt dispatch a chance to register its cancellation target before + # the reader accepts an immediately following session/cancel notification. + await asyncio.sleep(0) + + async def wait_idle(self) -> None: + while self.requests: + active = tuple(self.requests.items()) + await asyncio.gather(*(task for _, task in active), return_exceptions=True) + for key, task in active: + if task.done() and self.requests.get(key) is task: + self.requests.pop(key, None) + + async def close(self) -> None: + if self.closed: + return + self.closed = True + for state in self.sessions.values(): + state.cancel_requested = True + if state.active_run is not None and not state.active_run.done(): + state.active_run.cancel() + await self.wait_idle() + errors: list[BaseException] = [] + for state in reversed(tuple(self.sessions.values())): + try: + await state.context.__aexit__(None, None, None) + except BaseException as exc: # noqa: BLE001 - close every runtime before surfacing failure. + errors.append(exc) + finally: + self._restore_environment(state.previous_environment) + self.sessions.clear() + if errors: + raise errors[0] + + def _request_done(self, key: str, task: asyncio.Task[None]) -> None: + if self.requests.get(key) is task: + self.requests.pop(key, None) + # Retrieve exceptions even if stdout failed after the caller disconnected. + if not task.cancelled(): + task.exception() + + async def _dispatch_request(self, response_id: Any, key: str, method: str, params: Any) -> None: + try: + if method == _METHOD_INITIALIZE: + result = self._initialize(params) + elif method == _METHOD_SESSION_NEW: + result = await self._new_session(params) + elif method == _METHOD_SESSION_PROMPT: + result = await self._prompt(key, params) + else: + raise ACPProtocolError(_METHOD_NOT_FOUND, f"unsupported ACP method {method!r}") + except ACPProtocolError as exc: + await self._send_protocol_error(response_id, exc) + return + except AgentRunError as exc: + data = {"code": exc.code or exc.__class__.__name__} + await self._send_error(response_id, _RUNTIME_ERROR, "AgentKit runtime prompt failed", data) + return + except BaseException as exc: # noqa: BLE001 - keep provider details off stdout. + if isinstance(exc, (KeyboardInterrupt, SystemExit)): + raise + await self._send_error( + response_id, + _INTERNAL_ERROR, + "AgentKit ACP request failed", + {"code": exc.__class__.__name__}, + ) + return + await self.send({"jsonrpc": _JSONRPC_VERSION, "id": response_id, "result": result}) + + async def _handle_notification(self, method: str, params: Any) -> None: + if method == _METHOD_SESSION_CANCEL: + if not isinstance(params, dict) or not isinstance(params.get("sessionId"), str): + return + state = self.sessions.get(params["sessionId"]) + self._cancel_state(state) + return + if method == _METHOD_CANCEL_REQUEST: + if not isinstance(params, dict) or "requestId" not in params: + return + try: + key = _request_key(params["requestId"]) + except ACPProtocolError: + return + self._cancel_state(self.prompt_requests.get(key)) + + def _initialize(self, params: Any) -> dict[str, Any]: + request = _required_object(params) + if request.get("protocolVersion") != ACP_PROTOCOL_VERSION: + raise ACPProtocolError( + _INVALID_PARAMS, + f"protocolVersion must be {ACP_PROTOCOL_VERSION}", + ) + self.initialized = True + mcp_capabilities: dict[str, bool] = {"http": True} if self.http_mcp else {} + return { + "protocolVersion": ACP_PROTOCOL_VERSION, + "agentCapabilities": { + "loadSession": False, + "promptCapabilities": { + "image": False, + "audio": False, + "embeddedContext": False, + }, + "mcpCapabilities": mcp_capabilities, + "sessionCapabilities": {}, + "auth": {}, + }, + "agentInfo": { + "name": "agentkit", + "title": "AgentKit ACP runtime", + "version": "0.0.0", + }, + } + + async def _new_session(self, params: Any) -> dict[str, Any]: + if not self.initialized: + raise ACPProtocolError(_INVALID_REQUEST, "initialize must complete before session/new") + if self.sessions or self.session_creation_active: + raise ACPProtocolError(_INVALID_REQUEST, "ACP child already owns a session") + self.session_creation_active = True + try: + return await self._create_session(params) + finally: + self.session_creation_active = False + + async def _create_session(self, params: Any) -> dict[str, Any]: + request = _required_object(params) + cwd = _required_string(request.get("cwd"), name="cwd") + if not os.path.isabs(cwd) or os.path.realpath(cwd) != os.path.realpath(os.getcwd()): + raise ACPProtocolError(_INVALID_PARAMS, "cwd must match the ACP child working directory") + additional = request.get("additionalDirectories", []) + if not isinstance(additional, list) or additional: + raise ACPProtocolError(_INVALID_PARAMS, "additionalDirectories are not supported") + if self.spec.tools: + raise ACPProtocolError(_INVALID_PARAMS, "ACP strict mode rejects baked direct tools") + if self.spec.brokered_tools: + raise ACPProtocolError(_INVALID_PARAMS, "ACP strict mode rejects baked brokeredTools") + if self.spec.context.providers: + raise ACPProtocolError(_INVALID_PARAMS, "ACP strict mode rejects baked context providers") + + mcp_servers = request.get("mcpServers", []) + if not isinstance(mcp_servers, list): + raise ACPProtocolError(_INVALID_PARAMS, "mcpServers must be an array") + if len(mcp_servers) > 1: + raise ACPProtocolError(_INVALID_PARAMS, "ACP strict mode accepts at most one MCP server") + if mcp_servers and not self.http_mcp: + raise ACPProtocolError(_INVALID_PARAMS, "this runtime cannot consume ACP HTTP MCP servers") + + session_id = "agentkit-" + secrets.token_hex(16) + projected, environment = self._project_spec(session_id, mcp_servers) + previous_environment = self._install_environment(environment) + context: RuntimeSession | None = None + try: + context = self.factory.build_runtime(projected) + runtime = await context.__aenter__() + except BaseException: + error_info = sys.exc_info() + try: + if context is not None: + try: + await context.__aexit__(*error_info) + except BaseException: + pass + finally: + self._restore_environment(previous_environment) + raise + self.sessions[session_id] = _SessionState( + session_id=session_id, + context=context, + runtime=runtime, + previous_environment=previous_environment, + ) + return {"sessionId": session_id} + + def _project_spec( + self, + session_id: str, + mcp_servers: list[Any], + ) -> tuple[AgentSpec, dict[str, str]]: + provider_base_url = os.environ.get(ACP_PROVIDER_BASE_URL_ENV, "") + provider_token = os.environ.get(ACP_PROVIDER_TOKEN_ENV, "") + provider_base_url = _loopback_http_url(provider_base_url, name=ACP_PROVIDER_BASE_URL_ENV) + _safe_environment_value(provider_token, name=ACP_PROVIDER_TOKEN_ENV) + + data = self.spec.model_dump(by_alias=True) + data["model"]["baseURL"] = provider_base_url + data["model"]["apiKeyEnv"] = ACP_PROVIDER_TOKEN_ENV + data["model"]["auth"] = None + data["tools"] = [] + environment: dict[str, str] = {} + prefix = "AGENTKIT_ACP_SESSION_" + session_id.removeprefix("agentkit-").upper() + seen_names: set[str] = set() + for index, raw_server in enumerate(mcp_servers): + server = _required_object(raw_server, name=f"mcpServers[{index}]") + if server.get("type") != "http": + raise ACPProtocolError(_INVALID_PARAMS, "ACP mode supports only HTTP MCP servers") + if server.get("command") or server.get("args") or server.get("env"): + raise ACPProtocolError(_INVALID_PARAMS, "HTTP MCP servers must not carry process fields") + name = _required_string(server.get("name"), name=f"mcpServers[{index}].name") + if name in seen_names: + raise ACPProtocolError(_INVALID_PARAMS, f"duplicate MCP server name {name!r}") + seen_names.add(name) + url = _required_string(server.get("url"), name=f"mcpServers[{index}].url") + url = _loopback_http_url(url, name=f"mcpServers[{index}].url") + url_env = f"{prefix}_MCP_{index}_URL" + environment[url_env] = url + + headers = server.get("headers", []) + if not isinstance(headers, list): + raise ACPProtocolError(_INVALID_PARAMS, f"mcpServers[{index}].headers must be an array") + projected_headers: list[dict[str, str]] = [] + seen_headers: set[str] = set() + for header_index, raw_header in enumerate(headers): + header = _required_object( + raw_header, + name=f"mcpServers[{index}].headers[{header_index}]", + ) + header_name = _required_string( + header.get("name"), + name=f"mcpServers[{index}].headers[{header_index}].name", + ) + canonical_header = header_name.lower() + if canonical_header in seen_headers: + raise ACPProtocolError(_INVALID_PARAMS, f"duplicate MCP header {header_name!r}") + seen_headers.add(canonical_header) + header_value = _safe_environment_value( + header.get("value"), + name=f"mcpServers[{index}].headers[{header_index}].value", + ) + value_env = f"{prefix}_MCP_{index}_HEADER_{header_index}" + environment[value_env] = header_value + projected_headers.append({"name": header_name, "valueEnv": value_env}) + authorization = next( + ( + header + for header in headers + if isinstance(header, dict) + and isinstance(header.get("name"), str) + and header["name"].lower() == "authorization" + ), + None, + ) + authorization_value = "" if authorization is None else str(authorization.get("value", "")) + if not authorization_value.startswith("Bearer ") or not authorization_value[7:]: + raise ACPProtocolError( + _INVALID_PARAMS, + "ACP HTTP MCP server must include a bearer Authorization header", + ) + data["tools"].append( + { + "name": name, + "type": "mcp", + "transport": "streamable-http", + "urlEnv": url_env, + "headers": projected_headers, + } + ) + try: + projected = AgentSpec.model_validate(data) + except ValueError as exc: + raise ACPProtocolError(_INVALID_PARAMS, "ACP MCP server configuration is invalid") from exc + return projected, environment + + async def _prompt(self, request_key: str, params: Any) -> dict[str, str]: + request = _required_object(params) + session_id = _required_string(request.get("sessionId"), name="sessionId") + state = self.sessions.get(session_id) + if state is None: + raise ACPProtocolError(_INVALID_PARAMS, "unknown ACP sessionId") + if state.active_run is not None and not state.active_run.done(): + raise ACPProtocolError(_RUNTIME_ERROR, "ACP session already has an active prompt") + prompt = request.get("prompt") + if not isinstance(prompt, list) or not prompt: + raise ACPProtocolError(_INVALID_PARAMS, "prompt must be a non-empty array") + text_blocks: list[str] = [] + for index, raw_block in enumerate(prompt): + block = _required_object(raw_block, name=f"prompt[{index}]") + if block.get("type") != "text": + raise ACPProtocolError(_INVALID_PARAMS, "ACP mode accepts only text prompt blocks") + text_blocks.append( + _required_prompt_text(block.get("text"), name=f"prompt[{index}].text") + ) + + run_request = RunRequest( + prompt="\n".join(text_blocks), + history=tuple(state.history), + session_id=session_id, + ) + state.cancel_requested = False + state.active_request_key = request_key + state.active_run = asyncio.create_task( + state.runtime.run(run_request), + name="agentkit-acp-runtime-prompt", + ) + self.prompt_requests[request_key] = state + cancelled = False + runtime_error: BaseException | None = None + try: + result = await state.active_run + except asyncio.CancelledError: + cancelled = True + result = None + except BaseException as exc: # noqa: BLE001 - cancellation wins over its fallout. + result = None + runtime_error = exc + finally: + self.prompt_requests.pop(request_key, None) + state.active_request_key = None + state.active_run = None + + if cancelled or state.cancel_requested or runtime_error is not None or result is None: + await _discard_runtime_session(state.runtime, session_id) + if cancelled or state.cancel_requested: + return {"stopReason": "cancelled"} + if runtime_error is not None: + raise runtime_error + if result is None: + raise ACPProtocolError(_INTERNAL_ERROR, "runtime returned no prompt result") + if result.text: + await self.send( + { + "jsonrpc": _JSONRPC_VERSION, + "method": _METHOD_SESSION_UPDATE, + "params": { + "sessionId": session_id, + "update": { + "sessionUpdate": "agent_message_chunk", + "content": {"type": "text", "text": result.text}, + }, + }, + } + ) + state.history.append(ConversationTurn(role="user", text=run_request.prompt)) + state.history.append(ConversationTurn(role="assistant", text=result.text)) + return {"stopReason": "end_turn"} + + def _cancel_state(self, state: _SessionState | None) -> None: + if state is None: + return + state.cancel_requested = True + if state.active_run is not None and not state.active_run.done(): + state.active_run.cancel() + + @staticmethod + def _install_environment(environment: Mapping[str, str]) -> dict[str, str | None]: + previous = {name: os.environ.get(name) for name in environment} + os.environ.update(environment) + return previous + + @staticmethod + def _restore_environment(environment: Mapping[str, str | None]) -> None: + for name, previous in environment.items(): + if previous is None: + os.environ.pop(name, None) + else: + os.environ[name] = previous + + async def _send_protocol_error(self, response_id: Any, error: ACPProtocolError) -> None: + await self._send_error(response_id, error.code, error.message, error.data) + + async def _send_error( + self, + response_id: Any, + code: int, + message: str, + data: Mapping[str, Any] | None = None, + ) -> None: + error: dict[str, Any] = {"code": code, "message": message} + if data is not None: + error["data"] = dict(data) + await self.send({"jsonrpc": _JSONRPC_VERSION, "id": response_id, "error": error}) + + +async def serve_acp_stdio( + spec: AgentSpec, + factory: RuntimeFactory, + *, + reader: BinaryIO | None = None, + writer: BinaryIO | None = None, +) -> None: + """Serve ACP until stdin closes, then close every runtime session.""" + + input_stream = reader or sys.stdin.buffer + output_stream = writer or sys.stdout.buffer + write_lock = asyncio.Lock() + + async def send(message: Mapping[str, Any]) -> None: + encoded = json.dumps(message, ensure_ascii=False, separators=(",", ":")).encode("utf-8") + b"\n" + if len(encoded) > _MAX_MESSAGE_BYTES: + raise RuntimeError("ACP response exceeds the 8 MiB limit") + async with write_lock: + output_stream.write(encoded) + output_stream.flush() + + server = ACPStdioServer(spec, factory, send) + try: + while True: + line = await asyncio.to_thread(input_stream.readline, _MAX_MESSAGE_BYTES + 2) + if not line: + break + await server.accept_line(line.rstrip(b"\r\n")) + finally: + await server.close() + + +def run_acp_stdio(spec: AgentSpec, factory: RuntimeFactory) -> None: + """Synchronous console entrypoint for ACP stdio mode.""" + + protocol_output = sys.stdout.buffer + with redirect_stdout(sys.stderr): + asyncio.run(serve_acp_stdio(spec, factory, writer=protocol_output)) diff --git a/runtimes/common/agentkit_serve_common/cli.py b/runtimes/common/agentkit_serve_common/cli.py index 6274353..c3b9d29 100644 --- a/runtimes/common/agentkit_serve_common/cli.py +++ b/runtimes/common/agentkit_serve_common/cli.py @@ -1,8 +1,9 @@ """Shared CLI / network-posture core for AgentKit runtime adapters. -Loads ``/agent/agent.yaml``, selects one protocol skin, applies the network -posture, and runs uvicorn. Each adapter's console script calls :func:`run` with -its own framework-specific ``agent_factory`` module — the only per-adapter input. +Loads ``/agent/agent.yaml`` and selects one protocol skin. HTTP modes apply the +network posture and run uvicorn; ACP uses stdio. Each adapter's console script +calls :func:`run` with its own framework-specific ``agent_factory`` module, the +only per-adapter input. Protocol modes: @@ -10,6 +11,7 @@ * ``foundry``: ``/readiness``, ``/invocations``, minimal non-streaming ``/responses``. * ``orka``: observed-mode ``orka.harness.v1`` over HTTP+SSE. +* ``acp``: Orka-owned ACP protocol v1 over newline-delimited JSON-RPC on stdio. Network posture: @@ -30,6 +32,12 @@ import uvicorn +from .acp import ( + ACPConfigurationError, + ACPProtocolError, + run_acp_stdio, + validate_acp_runtime_binding, +) from .config import ConfigError, load, load_or_exit from .foundry import create_foundry_app from .orka import create_orka_app @@ -38,7 +46,7 @@ # Hosts that mean "loopback only" — a bind to any of these needs no auth token. _LOOPBACK_HOSTS = frozenset({"127.0.0.1", "localhost", "::1", "::ffff:127.0.0.1"}) -_PROTOCOLS = frozenset({"openai", "foundry", "orka"}) +_PROTOCOLS = frozenset({"acp", "openai", "foundry", "orka"}) DEFAULT_CONFIG_PATH = "/agent/agent.yaml" DEFAULT_PORT = 8080 @@ -98,10 +106,8 @@ def _resolve_port(protocol: str, spec_port: int | None) -> int: return DEFAULT_FOUNDRY_PORT return spec_port or DEFAULT_PORT - - def _load_spec_or_exit(path: str, protocol: str): # noqa: ANN001 - if protocol != "orka": + if protocol not in {"acp", "orka"}: return load_or_exit(path) try: return load(path) @@ -131,6 +137,13 @@ def run(factory: RuntimeFactory, argv: list[str] | None = None) -> None: # builds a runtime session. os.environ["AGENTKIT_PROTOCOL"] = protocol spec = _load_spec_or_exit(args.config, protocol) + if protocol == "acp": + try: + validate_acp_runtime_binding(args.config, spec) + except (ACPConfigurationError, ACPProtocolError) as exc: + _fail(str(exc)) + run_acp_stdio(spec, factory) + return if spec.brokered_tools and protocol != "foundry": _fail( "brokeredTools require AGENTKIT_PROTOCOL=foundry (or --protocol foundry); " diff --git a/runtimes/common/agentkit_serve_common/runtime.py b/runtimes/common/agentkit_serve_common/runtime.py index 77ffc13..336a648 100644 --- a/runtimes/common/agentkit_serve_common/runtime.py +++ b/runtimes/common/agentkit_serve_common/runtime.py @@ -24,9 +24,9 @@ def offline_orka_echo_enabled() -> bool: - """Whether adapter factories should use the no-provider Orka echo runtime.""" + """Whether adapter factories should use the no-provider Orka/ACP echo runtime.""" - return os.environ.get("AGENTKIT_PROTOCOL", "").strip().lower() == "orka" and ( + return os.environ.get("AGENTKIT_PROTOCOL", "").strip().lower() in {"acp", "orka"} and ( os.environ.get(OFFLINE_ORKA_ECHO_ENV, "").strip().lower() in {"1", "true", "yes", "on"} ) diff --git a/runtimes/common/tests/test_acp_protocol.py b/runtimes/common/tests/test_acp_protocol.py new file mode 100644 index 0000000..6f90b0f --- /dev/null +++ b/runtimes/common/tests/test_acp_protocol.py @@ -0,0 +1,644 @@ +from __future__ import annotations + +import asyncio +import hashlib +import io +import json +import os +from collections.abc import Callable +from types import TracebackType +from typing import Any + +import pytest + +from agentkit_serve_common import acp +from agentkit_serve_common.acp import ACPConfigurationError, ACPStdioServer +from agentkit_serve_common.config import AgentSpec +from agentkit_serve_common.conversation import ConversationTurn, RunRequest +from agentkit_serve_common.runtime import ( + AgentRunError, + OfflineEchoRuntimeFactory, + RunResult, + RuntimeSession, + offline_orka_echo_enabled, +) + + +def _spec(**overrides: Any) -> AgentSpec: + data: dict[str, Any] = { + "abiVersion": "v0", + "metadata": {"name": "acp-test"}, + "model": { + "provider": "openai-compatible", + "baseURL": "https://baked.example.invalid/v1", + "name": "test-model", + "apiKeyEnv": "BAKED_MODEL_TOKEN", + }, + "instructions": "Be concise.", + "tools": [], + "expose": {"openai": True, "port": 8080}, + } + data.update(overrides) + return AgentSpec.model_validate(data) + + +def _set_provider_environment(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv(acp.ACP_PROVIDER_BASE_URL_ENV, "http://127.0.0.1:43123/v1") + monkeypatch.setenv(acp.ACP_PROVIDER_TOKEN_ENV, "provider-session-token") + + +def _initialize(request_id: int = 1) -> dict[str, Any]: + return { + "jsonrpc": "2.0", + "id": request_id, + "method": "initialize", + "params": { + "protocolVersion": 1, + "clientCapabilities": { + "fs": {"readTextFile": False, "writeTextFile": False}, + "terminal": False, + }, + "clientInfo": {"name": "orka", "version": "test"}, + }, + } + + +def _new_session(request_id: int = 2, *, mcp_servers: list[dict[str, Any]] | None = None) -> dict[str, Any]: + return { + "jsonrpc": "2.0", + "id": request_id, + "method": "session/new", + "params": { + "cwd": os.getcwd(), + "mcpServers": [] if mcp_servers is None else mcp_servers, + }, + } + + +def _prompt(request_id: int, session_id: str, text: str) -> dict[str, Any]: + return { + "jsonrpc": "2.0", + "id": request_id, + "method": "session/prompt", + "params": { + "sessionId": session_id, + "prompt": [{"type": "text", "text": text}], + }, + } + + +def _orka_mcp_server(url: str = "http://127.0.0.1:43124/session") -> dict[str, Any]: + return { + "type": "http", + "name": "orka", + "url": url, + "headers": [{"name": "Authorization", "value": "Bearer mcp-session-token"}], + } + + +def _response(messages: list[dict[str, Any]], request_id: int) -> dict[str, Any]: + matches = [message for message in messages if message.get("id") == request_id] + assert len(matches) == 1 + return matches[0] + + +async def _send_to( + server: ACPStdioServer, + messages: list[dict[str, Any]], + request: dict[str, Any], +) -> dict[str, Any]: + await server.accept(request) + await server.wait_idle() + return _response(messages, request["id"]) + + +class RecordingRuntime: + def __init__(self, outcomes: list[RunResult | BaseException]) -> None: + self.outcomes = list(outcomes) + self.requests: list[RunRequest] = [] + self.discarded_sessions: list[str] = [] + self.entered = 0 + self.exited = 0 + + async def __aenter__(self) -> RuntimeSession: + self.entered += 1 + return self + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc: BaseException | None, + tb: TracebackType | None, + ) -> bool | None: + self.exited += 1 + return None + + async def run(self, request: RunRequest) -> RunResult: + self.requests.append(request) + outcome = self.outcomes.pop(0) + if isinstance(outcome, BaseException): + raise outcome + return outcome + + async def discard_session(self, session_id: str) -> None: + self.discarded_sessions.append(session_id) + + +class RecordingFactory: + def __init__( + self, + runtime_builder: Callable[[], RuntimeSession], + *, + supports_http_mcp: bool = False, + ) -> None: + self.runtime_builder = runtime_builder + self.supports_http_mcp = supports_http_mcp + self.specs: list[AgentSpec] = [] + self.runtimes: list[RuntimeSession] = [] + self.environment_snapshots: list[dict[str, str | None]] = [] + + def supports_acp_http_mcp(self) -> bool: + return self.supports_http_mcp + + def build_runtime(self, spec: AgentSpec) -> RuntimeSession: + self.specs.append(spec) + environment_names = [ + name + for tool in spec.tools + for name in [tool.url_env, *(header.value_env for header in tool.headers)] + if name is not None + ] + self.environment_snapshots.append({name: os.environ.get(name) for name in environment_names}) + runtime = self.runtime_builder() + self.runtimes.append(runtime) + return runtime + + +def test_offline_echo_round_trip_uses_canonical_acp_shapes(monkeypatch): + _set_provider_environment(monkeypatch) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + + async def send(message): # noqa: ANN001 + messages.append(dict(message)) + + server = ACPStdioServer(_spec(), OfflineEchoRuntimeFactory(), send) + initialized = await _send_to(server, messages, _initialize()) + assert initialized["result"]["protocolVersion"] == 1 + assert initialized["result"]["agentCapabilities"]["mcpCapabilities"] == {} + + created = await _send_to(server, messages, _new_session()) + session_id = created["result"]["sessionId"] + completed = await _send_to(server, messages, _prompt(3, session_id, "hello")) + + assert completed["result"] == {"stopReason": "end_turn"} + assert { + "jsonrpc": "2.0", + "method": "session/update", + "params": { + "sessionId": session_id, + "update": { + "sessionUpdate": "agent_message_chunk", + "content": {"type": "text", "text": "offline echo: hello"}, + }, + }, + } in messages + assert server.sessions[session_id].history == [ + ConversationTurn(role="user", text="hello"), + ConversationTurn(role="assistant", text="offline echo: hello"), + ] + await server.close() + + asyncio.run(exercise()) + + +def test_stdio_server_frames_json_rpc_and_closes_on_eof(monkeypatch): + _set_provider_environment(monkeypatch) + monkeypatch.setattr(acp.secrets, "token_hex", lambda size: "b" * (size * 2)) + session_id = "agentkit-" + "b" * 32 + requests = [ + _initialize(), + _new_session(), + _prompt(3, session_id, "through stdio"), + ] + reader = io.BytesIO( + b"".join( + json.dumps(request, separators=(",", ":")).encode() + b"\n" + for request in requests + ) + ) + writer = io.BytesIO() + + asyncio.run( + acp.serve_acp_stdio( + _spec(), + OfflineEchoRuntimeFactory(), + reader=reader, + writer=writer, + ) + ) + + output = [json.loads(line) for line in writer.getvalue().splitlines()] + assert _response(output, 3)["result"] == {"stopReason": "end_turn"} + assert any(message.get("method") == "session/update" for message in output) + + +def test_offline_echo_gate_applies_to_acp(monkeypatch): + monkeypatch.setenv("AGENTKIT_PROTOCOL", "acp") + monkeypatch.setenv("AGENTKIT_ORKA_OFFLINE_ECHO", "true") + + assert offline_orka_echo_enabled() is True + + +def test_session_reuses_runtime_and_passes_full_successful_history(monkeypatch): + _set_provider_environment(monkeypatch) + runtime = RecordingRuntime([RunResult("first answer"), RunResult("second answer")]) + factory = RecordingFactory(lambda: runtime, supports_http_mcp=True) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + + async def send(message): # noqa: ANN001 + messages.append(dict(message)) + + server = ACPStdioServer(_spec(), factory, send) + initialized = await _send_to(server, messages, _initialize()) + assert initialized["result"]["agentCapabilities"]["mcpCapabilities"] == {"http": True} + created = await _send_to( + server, + messages, + _new_session(mcp_servers=[_orka_mcp_server()]), + ) + session_id = created["result"]["sessionId"] + + assert (await _send_to(server, messages, _prompt(3, session_id, " first question ")))[ + "result" + ] == {"stopReason": "end_turn"} + assert (await _send_to(server, messages, _prompt(4, session_id, "second question")))[ + "result" + ] == {"stopReason": "end_turn"} + + assert len(factory.specs) == 1 + assert runtime.entered == 1 + assert runtime.requests == [ + RunRequest(prompt=" first question ", history=(), session_id=session_id), + RunRequest( + prompt="second question", + history=( + ConversationTurn(role="user", text=" first question "), + ConversationTurn(role="assistant", text="first answer"), + ), + session_id=session_id, + ), + ] + + projected = factory.specs[0] + assert projected.model.name == "test-model" + assert projected.model.base_url == "http://127.0.0.1:43123/v1" + assert projected.model.api_key_env == acp.ACP_PROVIDER_TOKEN_ENV + assert projected.model.auth is None + assert len(projected.tools) == 1 + tool = projected.tools[0] + assert tool.transport == "streamable-http" + assert tool.url_env is not None + assert factory.environment_snapshots[0][tool.url_env] == _orka_mcp_server()["url"] + assert tool.headers[0].value_env is not None + assert factory.environment_snapshots[0][tool.headers[0].value_env] == "Bearer mcp-session-token" + + generated_names = set(factory.environment_snapshots[0]) + assert all(os.environ.get(name) for name in generated_names) + await server.close() + assert runtime.exited == 1 + assert all(name not in os.environ for name in generated_names) + + asyncio.run(exercise()) + + +def test_child_rejects_a_second_session(monkeypatch): + _set_provider_environment(monkeypatch) + factory = RecordingFactory(lambda: RecordingRuntime([RunResult("unused")])) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + + async def send(message): # noqa: ANN001 + messages.append(dict(message)) + + server = ACPStdioServer(_spec(), factory, send) + await _send_to(server, messages, _initialize()) + assert "sessionId" in (await _send_to(server, messages, _new_session()))["result"] + rejected = await _send_to(server, messages, _new_session(request_id=3)) + + assert rejected["error"]["code"] == -32600 + assert "already owns a session" in rejected["error"]["message"] + assert len(factory.specs) == 1 + await server.close() + + asyncio.run(exercise()) + + +class CancellingRuntime: + def __init__(self, *, swallow_cancellation: bool) -> None: + self.swallow_cancellation = swallow_cancellation + self.started = asyncio.Event() + self.block = asyncio.Event() + self.requests: list[RunRequest] = [] + self.discarded_sessions: list[str] = [] + + async def __aenter__(self) -> RuntimeSession: + return self + + async def __aexit__(self, exc_type, exc, tb): # noqa: ANN001 + return None + + async def run(self, request: RunRequest) -> RunResult: + self.requests.append(request) + if len(self.requests) > 1: + return RunResult("clean answer") + self.started.set() + try: + await self.block.wait() + except asyncio.CancelledError: + if self.swallow_cancellation: + return RunResult("partial answer that must be discarded") + raise + return RunResult("unexpected answer") + + async def discard_session(self, session_id: str) -> None: + self.discarded_sessions.append(session_id) + + +@pytest.mark.parametrize( + ("cancel_method", "swallow_cancellation"), + [("session/cancel", False), ("$/cancel_request", True)], +) +def test_cancellation_returns_cancelled_and_does_not_commit_history( + monkeypatch, + cancel_method: str, + swallow_cancellation: bool, +): + _set_provider_environment(monkeypatch) + runtime = CancellingRuntime(swallow_cancellation=swallow_cancellation) + factory = RecordingFactory(lambda: runtime) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + + async def send(message): # noqa: ANN001 + messages.append(dict(message)) + + server = ACPStdioServer(_spec(), factory, send) + await _send_to(server, messages, _initialize()) + created = await _send_to(server, messages, _new_session()) + session_id = created["result"]["sessionId"] + + await server.accept(_prompt(3, session_id, "cancel me")) + await asyncio.wait_for(runtime.started.wait(), timeout=1) + params = {"sessionId": session_id} if cancel_method == "session/cancel" else {"requestId": 3} + await server.accept( + {"jsonrpc": "2.0", "method": cancel_method, "params": params} + ) + await server.wait_idle() + + assert _response(messages, 3)["result"] == {"stopReason": "cancelled"} + assert not [message for message in messages if message.get("method") == "session/update"] + assert server.sessions[session_id].history == [] + assert runtime.discarded_sessions == [session_id] + + completed = await _send_to(server, messages, _prompt(4, session_id, "try again")) + assert completed["result"] == {"stopReason": "end_turn"} + assert runtime.requests[1].history == () + await server.close() + + asyncio.run(exercise()) + + +def test_runtime_error_is_redacted_and_does_not_commit_history(monkeypatch): + _set_provider_environment(monkeypatch) + runtime = RecordingRuntime( + [ + AgentRunError("provider leaked secret-token", code="ProviderFailure"), + RunResult("recovered"), + ] + ) + factory = RecordingFactory(lambda: runtime) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + + async def send(message): # noqa: ANN001 + messages.append(dict(message)) + + server = ACPStdioServer(_spec(), factory, send) + await _send_to(server, messages, _initialize()) + created = await _send_to(server, messages, _new_session()) + session_id = created["result"]["sessionId"] + + failed = await _send_to(server, messages, _prompt(3, session_id, "fail")) + assert failed["error"] == { + "code": -32000, + "message": "AgentKit runtime prompt failed", + "data": {"code": "ProviderFailure"}, + } + assert "secret-token" not in str(failed) + assert server.sessions[session_id].history == [] + assert runtime.discarded_sessions == [session_id] + + completed = await _send_to(server, messages, _prompt(4, session_id, "recover")) + assert completed["result"] == {"stopReason": "end_turn"} + assert runtime.requests[1].history == () + await server.close() + + asyncio.run(exercise()) + + +def test_prompt_rejects_non_text_content_without_running_model(monkeypatch): + _set_provider_environment(monkeypatch) + runtime = RecordingRuntime([RunResult("must not run")]) + factory = RecordingFactory(lambda: runtime) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + + async def send(message): # noqa: ANN001 + messages.append(dict(message)) + + server = ACPStdioServer(_spec(), factory, send) + await _send_to(server, messages, _initialize()) + created = await _send_to(server, messages, _new_session()) + session_id = created["result"]["sessionId"] + invalid = { + "jsonrpc": "2.0", + "id": 3, + "method": "session/prompt", + "params": { + "sessionId": session_id, + "prompt": [{"type": "image", "data": "ignored", "mimeType": "image/png"}], + }, + } + + failed = await _send_to(server, messages, invalid) + assert failed["error"]["code"] == -32602 + assert "only text" in failed["error"]["message"] + assert runtime.requests == [] + await server.close() + + asyncio.run(exercise()) + + +@pytest.mark.parametrize( + "server", + [ + _orka_mcp_server("https://example.com/mcp"), + {**_orka_mcp_server(), "headers": []}, + {**_orka_mcp_server(), "type": "stdio", "command": "tool"}, + ], +) +def test_session_rejects_unsafe_or_unsupported_mcp_without_building_runtime(monkeypatch, server): + _set_provider_environment(monkeypatch) + factory = RecordingFactory(lambda: RecordingRuntime([RunResult("unused")]), supports_http_mcp=True) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + + async def send(message): # noqa: ANN001 + messages.append(dict(message)) + + dispatcher = ACPStdioServer(_spec(), factory, send) + await _send_to(dispatcher, messages, _initialize()) + failed = await _send_to(dispatcher, messages, _new_session(mcp_servers=[server])) + assert failed["error"]["code"] == -32602 + assert factory.specs == [] + await dispatcher.close() + + asyncio.run(exercise()) + + +def test_failed_runtime_build_restores_projected_environment(monkeypatch): + _set_provider_environment(monkeypatch) + monkeypatch.setattr(acp.secrets, "token_hex", lambda size: "a" * (size * 2)) + prefix = "AGENTKIT_ACP_SESSION_" + "A" * 32 + "_MCP_0" + url_env = prefix + "_URL" + header_env = prefix + "_HEADER_0" + monkeypatch.setenv(url_env, "previous-url") + monkeypatch.setenv(header_env, "previous-header") + + class FailingFactory: + def supports_acp_http_mcp(self) -> bool: + return True + + def build_runtime(self, spec): # noqa: ANN001 + assert os.environ[url_env] == _orka_mcp_server()["url"] + assert os.environ[header_env] == "Bearer mcp-session-token" + raise RuntimeError("build failed") + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + + async def send(message): # noqa: ANN001 + messages.append(dict(message)) + + server = ACPStdioServer(_spec(), FailingFactory(), send) + await _send_to(server, messages, _initialize()) + failed = await _send_to( + server, + messages, + _new_session(mcp_servers=[_orka_mcp_server()]), + ) + assert failed["error"]["code"] == -32603 + assert os.environ[url_env] == "previous-url" + assert os.environ[header_env] == "previous-header" + await server.close() + + asyncio.run(exercise()) + + +def test_runtime_binding_verifies_exact_config_digest_model_and_provider(monkeypatch, tmp_path): + config = tmp_path / "agent.yaml" + config_bytes = b"abiVersion: v0\nmetadata:\n name: exact-bytes\n" + config.write_bytes(config_bytes) + monkeypatch.setenv( + acp.ACP_AGENT_CONFIGURATION_DIGEST_ENV, + "sha256:" + hashlib.sha256(config_bytes).hexdigest(), + ) + monkeypatch.setenv(acp.ACP_MODEL_ENV, "test-model") + _set_provider_environment(monkeypatch) + + acp.validate_acp_runtime_binding(config, _spec()) + + config.write_bytes(config_bytes + b"# changed\n") + with pytest.raises(ACPConfigurationError, match="exact agent config bytes"): + acp.validate_acp_runtime_binding(config, _spec()) + + +@pytest.mark.parametrize( + ("environment_name", "environment_value", "message"), + [ + (acp.ACP_AGENT_CONFIGURATION_DIGEST_ENV, "sha256:short", "lowercase sha256"), + (acp.ACP_MODEL_ENV, "other-model", "does not match model.name"), + (acp.ACP_PROVIDER_BASE_URL_ENV, "https://provider.example.com/v1", "loopback"), + (acp.ACP_PROVIDER_TOKEN_ENV, "", "non-empty"), + ], +) +def test_runtime_binding_rejects_profile_or_provider_mismatch( + monkeypatch, + tmp_path, + environment_name: str, + environment_value: str, + message: str, +): + config = tmp_path / "agent.yaml" + config.write_bytes(b"exact") + monkeypatch.setenv( + acp.ACP_AGENT_CONFIGURATION_DIGEST_ENV, + "sha256:" + hashlib.sha256(b"exact").hexdigest(), + ) + monkeypatch.setenv(acp.ACP_MODEL_ENV, "test-model") + _set_provider_environment(monkeypatch) + monkeypatch.setenv(environment_name, environment_value) + + with pytest.raises(ACPConfigurationError, match=message): + acp.validate_acp_runtime_binding(config, _spec()) + + +@pytest.mark.parametrize( + ("override", "message"), + [ + ({"tools": [{"name": "local", "command": ["local-tool"]}]}, "direct tools"), + ( + { + "brokeredTools": [ + { + "name": "lookup", + "description": "Look up a value.", + "brokeredClass": "read", + "parameters": {"type": "object"}, + } + ] + }, + "brokeredTools", + ), + ( + { + "context": { + "providers": [ + { + "name": "skills", + "type": "skills", + "source": "filesystem", + "path": "/agent/skills", + } + ] + } + }, + "context providers", + ), + ], +) +def test_runtime_binding_rejects_baked_tool_and_context_paths(tmp_path, override, message): + config = tmp_path / "agent.yaml" + config.write_bytes(b"unused") + + with pytest.raises(ACPConfigurationError, match=message): + acp.validate_acp_runtime_binding(config, _spec(**override)) diff --git a/runtimes/common/tests/test_cli_protocol.py b/runtimes/common/tests/test_cli_protocol.py index 793ce4f..3c082b4 100644 --- a/runtimes/common/tests/test_cli_protocol.py +++ b/runtimes/common/tests/test_cli_protocol.py @@ -157,3 +157,36 @@ def test_cli_protocol_flag_sets_agentkit_protocol_for_adapter_runtime(monkeypatc assert captured["port"] == 8080 assert os.environ["AGENTKIT_PROTOCOL"] == "orka" + + +def test_cli_acp_uses_verified_stdio_entrypoint_without_uvicorn(monkeypatch): + captured = {} + spec = _spec() + monkeypatch.setattr(cli, "load", lambda path: spec) + monkeypatch.setattr( + cli, + "validate_acp_runtime_binding", + lambda path, loaded: captured.update({"path": path, "verified": loaded}), + ) + monkeypatch.setattr( + cli, + "run_acp_stdio", + lambda loaded, factory: captured.update({"served": loaded, "factory": factory}), + ) + monkeypatch.setattr( + cli.uvicorn, + "run", + lambda *args, **kwargs: (_ for _ in ()).throw(AssertionError("uvicorn must not run")), + ) + monkeypatch.delenv("AGENTKIT_AUTH_TOKEN", raising=False) + factory = Factory() + + cli.run(factory, ["--config", "/agent/agent.yaml", "--protocol", "acp"]) + + assert captured == { + "path": "/agent/agent.yaml", + "verified": spec, + "served": spec, + "factory": factory, + } + assert os.environ["AGENTKIT_PROTOCOL"] == "acp" diff --git a/runtimes/langgraph/agentkit_serve/agent_factory.py b/runtimes/langgraph/agentkit_serve/agent_factory.py index 52b5b65..019b434 100644 --- a/runtimes/langgraph/agentkit_serve/agent_factory.py +++ b/runtimes/langgraph/agentkit_serve/agent_factory.py @@ -206,6 +206,10 @@ def supports_brokered_coordination() -> bool: return offline_orka_echo_enabled() +def supports_acp_http_mcp() -> bool: + return True + + def build_runtime(spec: AgentSpec) -> RuntimeSession: """Build the runtime session consumed by the shared server.""" if offline_orka_echo_enabled(): diff --git a/runtimes/langgraph/tests/test_guardrails.py b/runtimes/langgraph/tests/test_guardrails.py index 978af2a..32e9e38 100644 --- a/runtimes/langgraph/tests/test_guardrails.py +++ b/runtimes/langgraph/tests/test_guardrails.py @@ -362,3 +362,5 @@ def test_langgraph_orka_offline_echo_bypasses_provider_runtime(monkeypatch): runtime = agent_factory.build_runtime(spec) assert runtime.__class__.__name__ == "OfflineEchoRuntime" + monkeypatch.setenv("AGENTKIT_PROTOCOL", "acp") + assert agent_factory.supports_acp_http_mcp() is True diff --git a/runtimes/microsoft-agent-framework/agentkit_serve/agent_factory.py b/runtimes/microsoft-agent-framework/agentkit_serve/agent_factory.py index 4688bb0..a96bbfc 100644 --- a/runtimes/microsoft-agent-framework/agentkit_serve/agent_factory.py +++ b/runtimes/microsoft-agent-framework/agentkit_serve/agent_factory.py @@ -390,13 +390,33 @@ async def run(self, request: RunRequest) -> RunResult: try: async with lock: self._touch_session(session_id) + session = self.sessions[session_id] include_history = session_id not in self.initialized_sessions - result = await run_agent(self.agent, request, session=session, include_history=include_history) + try: + result = await run_agent( + self.agent, + request, + session=session, + include_history=include_history, + ) + except BaseException: + self._reset_session(session_id) + raise self.initialized_sessions.add(session_id) return result finally: self._release_session_claim(session_id) + async def discard_session(self, session_id: str) -> None: + """Discard framework state for an ACP prompt that did not commit.""" + + lock = self.session_locks.get(session_id) + if lock is None: + return + async with lock: + if self.session_locks.get(session_id) is lock and session_id in self.sessions: + self._reset_session(session_id) + def _session_for(self, session_id: str | None) -> tuple[AgentSession | None, asyncio.Lock | None]: if not session_id: return None, None @@ -421,6 +441,10 @@ def _touch_session(self, session_id: str) -> None: self.most_recent_session_id = session_id self._evict_idle_sessions() + def _reset_session(self, session_id: str) -> None: + self.sessions[session_id] = AgentSession(session_id=session_id) + self.initialized_sessions.discard(session_id) + def _release_session_claim(self, session_id: str) -> None: claims = self.session_claims.get(session_id, 0) if claims <= 1: @@ -605,6 +629,10 @@ def supports_brokered_coordination() -> bool: return offline_orka_echo_enabled() +def supports_acp_http_mcp() -> bool: + return True + + def build_runtime(spec: AgentSpec) -> RuntimeSession: """Build the runtime session consumed by the shared server.""" if offline_orka_echo_enabled(): diff --git a/runtimes/microsoft-agent-framework/tests/test_guardrails.py b/runtimes/microsoft-agent-framework/tests/test_guardrails.py index 6875f9f..022e0fb 100644 --- a/runtimes/microsoft-agent-framework/tests/test_guardrails.py +++ b/runtimes/microsoft-agent-framework/tests/test_guardrails.py @@ -753,7 +753,7 @@ async def fake_run_agent(agent, request, *, session=None, include_history=True): assert include_history_values == [True, False] -def test_failed_first_turn_keeps_explicit_history_on_retry(monkeypatch): +def test_failed_first_turn_replaces_session_and_keeps_explicit_history_on_retry(monkeypatch): from agentkit_serve_common.config import AgentSpec from agentkit_serve_common.conversation import ConversationTurn, RunRequest from agentkit_serve_common.runtime import RunResult @@ -797,7 +797,8 @@ async def exercise(): asyncio.run(exercise()) assert include_history_values == [True, True, False] - assert seen_sessions[0] is seen_sessions[1] is seen_sessions[2] + assert seen_sessions[0] is not seen_sessions[1] + assert seen_sessions[1] is seen_sessions[2] def test_remote_mcp_disables_ping(monkeypatch): @@ -976,3 +977,5 @@ def test_maf_orka_offline_echo_bypasses_provider_runtime(monkeypatch): runtime = agent_factory.build_runtime(spec) assert runtime.__class__.__name__ == "OfflineEchoRuntime" + monkeypatch.setenv("AGENTKIT_PROTOCOL", "acp") + assert agent_factory.supports_acp_http_mcp() is True diff --git a/runtimes/pydantic-ai/agentkit_serve/agent_factory.py b/runtimes/pydantic-ai/agentkit_serve/agent_factory.py index 8ce1e27..15c4e77 100644 --- a/runtimes/pydantic-ai/agentkit_serve/agent_factory.py +++ b/runtimes/pydantic-ai/agentkit_serve/agent_factory.py @@ -191,6 +191,13 @@ def supports_brokered_coordination() -> bool: return offline_orka_echo_enabled() +def supports_acp_http_mcp() -> bool: + return ( + offline_orka_echo_enabled() + or (MCPToolset is not None and StreamableHttpTransport is not None) + ) + + def build_runtime(spec: AgentSpec) -> RuntimeSession: """Build the runtime session consumed by the shared server.""" if offline_orka_echo_enabled(): diff --git a/runtimes/pydantic-ai/tests/test_unsupported_features.py b/runtimes/pydantic-ai/tests/test_unsupported_features.py index 89c640f..2a40fa3 100644 --- a/runtimes/pydantic-ai/tests/test_unsupported_features.py +++ b/runtimes/pydantic-ai/tests/test_unsupported_features.py @@ -43,6 +43,8 @@ def test_pydantic_orka_offline_echo_bypasses_provider_runtime(monkeypatch): runtime = agent_factory.build_runtime(spec) assert runtime.__class__.__name__ == "OfflineEchoRuntime" + monkeypatch.setenv("AGENTKIT_PROTOCOL", "acp") + assert agent_factory.supports_acp_http_mcp() is True def test_pydantic_orka_offline_echo_completes_without_provider(monkeypatch): From d1109651bff1bc8fb5a3572b1677887fb8998da4 Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Thu, 3 Sep 2026 14:25:17 -0700 Subject: [PATCH 02/23] fix(runtime): preserve ACP prompt commit boundary Signed-off-by: Sertac Ozercan --- runtimes/common/agentkit_serve_common/acp.py | 85 +++++++++++-------- runtimes/common/tests/test_acp_protocol.py | 89 ++++++++++++++++++++ 2 files changed, 138 insertions(+), 36 deletions(-) diff --git a/runtimes/common/agentkit_serve_common/acp.py b/runtimes/common/agentkit_serve_common/acp.py index 668d35e..73598c6 100644 --- a/runtimes/common/agentkit_serve_common/acp.py +++ b/runtimes/common/agentkit_serve_common/acp.py @@ -543,7 +543,7 @@ async def _prompt(self, request_key: str, params: Any) -> dict[str, str]: state = self.sessions.get(session_id) if state is None: raise ACPProtocolError(_INVALID_PARAMS, "unknown ACP sessionId") - if state.active_run is not None and not state.active_run.done(): + if state.active_request_key is not None: raise ACPProtocolError(_RUNTIME_ERROR, "ACP session already has an active prompt") prompt = request.get("prompt") if not isinstance(prompt, list) or not prompt: @@ -569,47 +569,60 @@ async def _prompt(self, request_key: str, params: Any) -> dict[str, str]: name="agentkit-acp-runtime-prompt", ) self.prompt_requests[request_key] = state - cancelled = False - runtime_error: BaseException | None = None try: - result = await state.active_run - except asyncio.CancelledError: - cancelled = True - result = None - except BaseException as exc: # noqa: BLE001 - cancellation wins over its fallout. - result = None - runtime_error = exc + cancelled = False + runtime_error: BaseException | None = None + try: + result = await state.active_run + except asyncio.CancelledError: + cancelled = True + result = None + except BaseException as exc: # noqa: BLE001 - cancellation wins over its fallout. + result = None + runtime_error = exc + + if ( + cancelled + or state.cancel_requested + or runtime_error is not None + or result is None + ): + await _discard_runtime_session(state.runtime, session_id) + if cancelled or state.cancel_requested: + return {"stopReason": "cancelled"} + if runtime_error is not None: + raise runtime_error + if result is None: + raise ACPProtocolError(_INTERNAL_ERROR, "runtime returned no prompt result") + try: + if result.text: + await self.send( + { + "jsonrpc": _JSONRPC_VERSION, + "method": _METHOD_SESSION_UPDATE, + "params": { + "sessionId": session_id, + "update": { + "sessionUpdate": "agent_message_chunk", + "content": {"type": "text", "text": result.text}, + }, + }, + } + ) + except BaseException: # noqa: BLE001 - failed output must roll back state. + await _discard_runtime_session(state.runtime, session_id) + raise + if state.cancel_requested: + await _discard_runtime_session(state.runtime, session_id) + return {"stopReason": "cancelled"} + state.history.append(ConversationTurn(role="user", text=run_request.prompt)) + state.history.append(ConversationTurn(role="assistant", text=result.text)) + return {"stopReason": "end_turn"} finally: self.prompt_requests.pop(request_key, None) state.active_request_key = None state.active_run = None - if cancelled or state.cancel_requested or runtime_error is not None or result is None: - await _discard_runtime_session(state.runtime, session_id) - if cancelled or state.cancel_requested: - return {"stopReason": "cancelled"} - if runtime_error is not None: - raise runtime_error - if result is None: - raise ACPProtocolError(_INTERNAL_ERROR, "runtime returned no prompt result") - if result.text: - await self.send( - { - "jsonrpc": _JSONRPC_VERSION, - "method": _METHOD_SESSION_UPDATE, - "params": { - "sessionId": session_id, - "update": { - "sessionUpdate": "agent_message_chunk", - "content": {"type": "text", "text": result.text}, - }, - }, - } - ) - state.history.append(ConversationTurn(role="user", text=run_request.prompt)) - state.history.append(ConversationTurn(role="assistant", text=result.text)) - return {"stopReason": "end_turn"} - def _cancel_state(self, state: _SessionState | None) -> None: if state is None: return diff --git a/runtimes/common/tests/test_acp_protocol.py b/runtimes/common/tests/test_acp_protocol.py index 6f90b0f..06cbc5f 100644 --- a/runtimes/common/tests/test_acp_protocol.py +++ b/runtimes/common/tests/test_acp_protocol.py @@ -453,6 +453,95 @@ async def send(message): # noqa: ANN001 asyncio.run(exercise()) +def test_update_send_failure_discards_runtime_state(monkeypatch): + _set_provider_environment(monkeypatch) + runtime = RecordingRuntime([RunResult("oversized answer"), RunResult("recovered")]) + factory = RecordingFactory(lambda: runtime) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + fail_update = True + + async def send(message): # noqa: ANN001 + nonlocal fail_update + if message.get("method") == "session/update" and fail_update: + fail_update = False + raise RuntimeError("ACP response exceeds the 8 MiB limit") + messages.append(dict(message)) + + server = ACPStdioServer(_spec(), factory, send) + await _send_to(server, messages, _initialize()) + created = await _send_to(server, messages, _new_session()) + session_id = created["result"]["sessionId"] + + failed = await _send_to(server, messages, _prompt(3, session_id, "too large")) + assert failed["error"] == { + "code": -32603, + "message": "AgentKit ACP request failed", + "data": {"code": "RuntimeError"}, + } + assert server.sessions[session_id].history == [] + assert runtime.discarded_sessions == [session_id] + + completed = await _send_to(server, messages, _prompt(4, session_id, "recover")) + assert completed["result"] == {"stopReason": "end_turn"} + assert runtime.requests[1].history == () + await server.close() + + asyncio.run(exercise()) + + +@pytest.mark.parametrize("cancel_method", ["session/cancel", "$/cancel_request"]) +def test_cancellation_while_update_send_waits_does_not_commit_history( + monkeypatch, + cancel_method: str, +): + _set_provider_environment(monkeypatch) + runtime = RecordingRuntime([RunResult("late answer"), RunResult("clean answer")]) + factory = RecordingFactory(lambda: runtime) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + update_started = asyncio.Event() + release_update = asyncio.Event() + + async def send(message): # noqa: ANN001 + if ( + message.get("method") == "session/update" + and not update_started.is_set() + ): + update_started.set() + await release_update.wait() + messages.append(dict(message)) + + server = ACPStdioServer(_spec(), factory, send) + await _send_to(server, messages, _initialize()) + created = await _send_to(server, messages, _new_session()) + session_id = created["result"]["sessionId"] + + await server.accept(_prompt(3, session_id, "cancel during update")) + await asyncio.wait_for(update_started.wait(), timeout=1) + params = ( + {"sessionId": session_id} + if cancel_method == "session/cancel" + else {"requestId": 3} + ) + await server.accept({"jsonrpc": "2.0", "method": cancel_method, "params": params}) + release_update.set() + await server.wait_idle() + + assert _response(messages, 3)["result"] == {"stopReason": "cancelled"} + assert server.sessions[session_id].history == [] + assert runtime.discarded_sessions == [session_id] + + completed = await _send_to(server, messages, _prompt(4, session_id, "try again")) + assert completed["result"] == {"stopReason": "end_turn"} + assert runtime.requests[1].history == () + await server.close() + + asyncio.run(exercise()) + + def test_prompt_rejects_non_text_content_without_running_model(monkeypatch): _set_provider_environment(monkeypatch) runtime = RecordingRuntime([RunResult("must not run")]) From d0bde57bcbfa0f3c5a86e9b7590fdf78c055f969 Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Thu, 3 Sep 2026 14:38:54 -0700 Subject: [PATCH 03/23] fix(acp): drain oversized stdio frames Signed-off-by: Sertac Ozercan --- runtimes/common/agentkit_serve_common/acp.py | 12 ++++++++- runtimes/common/tests/test_acp_protocol.py | 28 ++++++++++++++++++++ 2 files changed, 39 insertions(+), 1 deletion(-) diff --git a/runtimes/common/agentkit_serve_common/acp.py b/runtimes/common/agentkit_serve_common/acp.py index 73598c6..f293187 100644 --- a/runtimes/common/agentkit_serve_common/acp.py +++ b/runtimes/common/agentkit_serve_common/acp.py @@ -687,7 +687,17 @@ async def send(message: Mapping[str, Any]) -> None: line = await asyncio.to_thread(input_stream.readline, _MAX_MESSAGE_BYTES + 2) if not line: break - await server.accept_line(line.rstrip(b"\r\n")) + complete = line.endswith(b"\n") + frame = line.rstrip(b"\r\n") if complete else line + if len(frame) > _MAX_MESSAGE_BYTES: + await server.accept_line(frame) + while not complete: + line = await asyncio.to_thread(input_stream.readline, _MAX_MESSAGE_BYTES + 2) + if not line: + break + complete = line.endswith(b"\n") + continue + await server.accept_line(frame) finally: await server.close() diff --git a/runtimes/common/tests/test_acp_protocol.py b/runtimes/common/tests/test_acp_protocol.py index 06cbc5f..04a7fb0 100644 --- a/runtimes/common/tests/test_acp_protocol.py +++ b/runtimes/common/tests/test_acp_protocol.py @@ -244,6 +244,34 @@ def test_stdio_server_frames_json_rpc_and_closes_on_eof(monkeypatch): assert any(message.get("method") == "session/update" for message in output) +def test_stdio_server_drains_oversized_line_before_parsing_next_frame(monkeypatch): + monkeypatch.setattr(acp, "_MAX_MESSAGE_BYTES", 1024) + injected = json.dumps(_initialize(request_id=99), separators=(",", ":")).encode() + oversized_line = b"x" * (acp._MAX_MESSAGE_BYTES + 2) + injected + b"\n" + valid_line = json.dumps(_initialize(), separators=(",", ":")).encode() + b"\n" + reader = io.BytesIO(oversized_line + valid_line) + writer = io.BytesIO() + + asyncio.run( + acp.serve_acp_stdio( + _spec(), + OfflineEchoRuntimeFactory(), + reader=reader, + writer=writer, + ) + ) + + output = [json.loads(line) for line in writer.getvalue().splitlines()] + oversized_errors = [ + message + for message in output + if message.get("error", {}).get("message") == "ACP message exceeds the 8 MiB limit" + ] + assert len(oversized_errors) == 1 + assert not any(message.get("id") == 99 for message in output) + assert _response(output, 1)["result"]["protocolVersion"] == 1 + + def test_offline_echo_gate_applies_to_acp(monkeypatch): monkeypatch.setenv("AGENTKIT_PROTOCOL", "acp") monkeypatch.setenv("AGENTKIT_ORKA_OFFLINE_ECHO", "true") From d8d41678a3d47a1baa0db94d07c42ab89307f0c7 Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Thu, 3 Sep 2026 15:05:07 -0700 Subject: [PATCH 04/23] fix(acp): chunk assistant output frames Signed-off-by: Sertac Ozercan --- runtimes/common/agentkit_serve_common/acp.py | 24 +++++- runtimes/common/tests/test_acp_protocol.py | 91 ++++++++++++++++++++ 2 files changed, 112 insertions(+), 3 deletions(-) diff --git a/runtimes/common/agentkit_serve_common/acp.py b/runtimes/common/agentkit_serve_common/acp.py index f293187..eda86d5 100644 --- a/runtimes/common/agentkit_serve_common/acp.py +++ b/runtimes/common/agentkit_serve_common/acp.py @@ -19,7 +19,7 @@ from contextlib import redirect_stdout from dataclasses import dataclass, field from pathlib import Path -from typing import Any, Awaitable, BinaryIO, Callable, Mapping +from typing import Any, Awaitable, BinaryIO, Callable, Iterator, Mapping from urllib.parse import urlsplit from .config import AgentSpec @@ -40,6 +40,11 @@ _METHOD_SESSION_UPDATE = "session/update" _METHOD_CANCEL_REQUEST = "$/cancel_request" _MAX_MESSAGE_BYTES = 8 << 20 +# orka.harness.v2 limits assistant-message chunks to 4 KiB of UTF-8 text. +# ACP v1 does not negotiate that limit, so the child must apply it before the +# supervisor maps notifications into harness events. Even worst-case JSON +# escaping remains well below harness v2's 512 KiB default event-line limit. +_MAX_ASSISTANT_MESSAGE_CHUNK_BYTES = 4 << 10 _PARSE_ERROR = -32700 _INVALID_REQUEST = -32600 @@ -117,6 +122,17 @@ async def _discard_runtime_session(runtime: RuntimeSession, session_id: str) -> await result +def _utf8_chunks(value: str, max_bytes: int) -> Iterator[str]: + encoded = value.encode("utf-8") + offset = 0 + while offset < len(encoded): + end = min(offset + max_bytes, len(encoded)) + while end < len(encoded) and encoded[end] & 0xC0 == 0x80: + end -= 1 + yield encoded[offset:end].decode("utf-8") + offset = end + + def _loopback_http_url(value: str, *, name: str) -> str: try: parsed = urlsplit(value) @@ -595,7 +611,9 @@ async def _prompt(self, request_key: str, params: Any) -> dict[str, str]: if result is None: raise ACPProtocolError(_INTERNAL_ERROR, "runtime returned no prompt result") try: - if result.text: + for chunk in _utf8_chunks(result.text, _MAX_ASSISTANT_MESSAGE_CHUNK_BYTES): + if state.cancel_requested: + break await self.send( { "jsonrpc": _JSONRPC_VERSION, @@ -604,7 +622,7 @@ async def _prompt(self, request_key: str, params: Any) -> dict[str, str]: "sessionId": session_id, "update": { "sessionUpdate": "agent_message_chunk", - "content": {"type": "text", "text": result.text}, + "content": {"type": "text", "text": chunk}, }, }, } diff --git a/runtimes/common/tests/test_acp_protocol.py b/runtimes/common/tests/test_acp_protocol.py index 04a7fb0..eda70d9 100644 --- a/runtimes/common/tests/test_acp_protocol.py +++ b/runtimes/common/tests/test_acp_protocol.py @@ -244,6 +244,45 @@ def test_stdio_server_frames_json_rpc_and_closes_on_eof(monkeypatch): assert any(message.get("method") == "session/update" for message in output) +def test_stdio_server_splits_output_larger_than_orka_acp_reader_limit(monkeypatch): + _set_provider_environment(monkeypatch) + monkeypatch.setattr(acp.secrets, "token_hex", lambda size: "c" * (size * 2)) + session_id = "agentkit-" + "c" * 32 + result_text = "x" * ((2 << 20) + 1) + runtime = RecordingRuntime([RunResult(result_text)]) + requests = [_initialize(), _new_session(), _prompt(3, session_id, "large output")] + reader = io.BytesIO( + b"".join( + json.dumps(request, separators=(",", ":")).encode() + b"\n" + for request in requests + ) + ) + writer = io.BytesIO() + + asyncio.run( + acp.serve_acp_stdio( + _spec(), + RecordingFactory(lambda: runtime), + reader=reader, + writer=writer, + ) + ) + + lines = writer.getvalue().splitlines() + output = [json.loads(line) for line in lines] + updates = [message for message in output if message.get("method") == "session/update"] + chunks = [message["params"]["update"]["content"]["text"] for message in updates] + + assert len(updates) > 1 + assert "".join(chunks) == result_text + assert all( + len(chunk.encode()) <= acp._MAX_ASSISTANT_MESSAGE_CHUNK_BYTES # noqa: SLF001 + for chunk in chunks + ) + assert all(len(line) < 512 << 10 for line in lines) + assert _response(output, 3)["result"] == {"stopReason": "end_turn"} + + def test_stdio_server_drains_oversized_line_before_parsing_next_frame(monkeypatch): monkeypatch.setattr(acp, "_MAX_MESSAGE_BYTES", 1024) injected = json.dumps(_initialize(request_id=99), separators=(",", ":")).encode() @@ -519,6 +558,58 @@ async def send(message): # noqa: ANN001 asyncio.run(exercise()) +def test_multibyte_output_is_split_without_partial_commit(monkeypatch): + _set_provider_environment(monkeypatch) + result_text = "a" + "🙂" * (acp._MAX_ASSISTANT_MESSAGE_CHUNK_BYTES // 4 + 1) # noqa: SLF001 + runtime = RecordingRuntime([RunResult(result_text)]) + factory = RecordingFactory(lambda: runtime) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + second_update_started = asyncio.Event() + release_second_update = asyncio.Event() + update_count = 0 + + async def send(message): # noqa: ANN001 + nonlocal update_count + if message.get("method") == "session/update": + update_count += 1 + if update_count == 2: + second_update_started.set() + await release_second_update.wait() + messages.append(dict(message)) + + server = ACPStdioServer(_spec(), factory, send) + await _send_to(server, messages, _initialize()) + created = await _send_to(server, messages, _new_session()) + session_id = created["result"]["sessionId"] + + await server.accept(_prompt(3, session_id, "multibyte output")) + await asyncio.wait_for(second_update_started.wait(), timeout=1) + updates = [message for message in messages if message.get("method") == "session/update"] + assert len(updates) == 1 + assert server.sessions[session_id].history == [] + + release_second_update.set() + await server.wait_idle() + updates = [message for message in messages if message.get("method") == "session/update"] + chunks = [message["params"]["update"]["content"]["text"] for message in updates] + + assert "".join(chunks) == result_text + assert [len(chunk.encode()) for chunk in chunks] == [ + acp._MAX_ASSISTANT_MESSAGE_CHUNK_BYTES - 3, # noqa: SLF001 + 8, + ] + assert server.sessions[session_id].history == [ + ConversationTurn(role="user", text="multibyte output"), + ConversationTurn(role="assistant", text=result_text), + ] + assert _response(messages, 3)["result"] == {"stopReason": "end_turn"} + await server.close() + + asyncio.run(exercise()) + + @pytest.mark.parametrize("cancel_method", ["session/cancel", "$/cancel_request"]) def test_cancellation_while_update_send_waits_does_not_commit_history( monkeypatch, From b0a763d000cc3f2986eb2671bcc97793dd183787 Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Thu, 3 Sep 2026 16:39:30 -0700 Subject: [PATCH 05/23] docs(orka): document adapter digest Signed-off-by: Sertac Ozercan --- docs/orka.md | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/docs/orka.md b/docs/orka.md index 487507e..77faaa0 100644 --- a/docs/orka.md +++ b/docs/orka.md @@ -12,6 +12,7 @@ checkout, layer the v2 supervisor onto that immutable image: ```sh make docker-build-acp-agentkit-runtime \ AGENTKIT_RUNTIME_IMAGE=ghcr.io/acme/fibey@sha256: \ + AGENTKIT_ADAPTER_DIGEST=sha256: \ ACP_AGENTKIT_RUNTIME_IMG=ghcr.io/acme/fibey-orka-v2:dev ``` @@ -42,7 +43,9 @@ Deploy the composed image as an operator-owned v2 supervisor service, configure the standard `ORKA_ACP_*` profile, fence, token-file, and runtime identity settings, then register it with Orka's strict-governed `AgentRuntime` sample. The profile model and `agentConfigurationDigest` must match the baked AgentKit -config, and the registration's `adapterName` must be `agentkit-serve-acp`. +config. The registration's `adapterName` must be `agentkit-serve-acp`, and its +`adapterDigest` must equal the `AGENTKIT_ADAPTER_DIGEST` baked into the composed +image. Orka freezes the AgentRuntime UID, generation, endpoint, profile, authentication Secret versions, and observed instance into each Task binding. It revalidates From 9abf37fa6f5cfd6c69fc6bfb2d003fb667dbf412 Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Thu, 3 Sep 2026 20:26:42 -0700 Subject: [PATCH 06/23] docs(orka): bind adapter identity to source image Signed-off-by: Sertac Ozercan --- docs/orka.md | 11 ++++++----- 1 file changed, 6 insertions(+), 5 deletions(-) diff --git a/docs/orka.md b/docs/orka.md index 77faaa0..f755935 100644 --- a/docs/orka.md +++ b/docs/orka.md @@ -12,7 +12,7 @@ checkout, layer the v2 supervisor onto that immutable image: ```sh make docker-build-acp-agentkit-runtime \ AGENTKIT_RUNTIME_IMAGE=ghcr.io/acme/fibey@sha256: \ - AGENTKIT_ADAPTER_DIGEST=sha256: \ + AGENTKIT_ADAPTER_DIGEST=sha256: \ ACP_AGENTKIT_RUNTIME_IMG=ghcr.io/acme/fibey-orka-v2:dev ``` @@ -32,7 +32,8 @@ The v2 path is strict: - the registered model must equal `model.name` in the baked config; - `agentConfigurationDigest` is `sha256:` plus the SHA-256 of the exact `/agent/agent.yaml` bytes; -- the runtime advertises only the `agentkit-serve-acp` adapter digest; +- the runtime advertises only `agentkit-serve-acp`, with the digest-pinned + AgentKit source image identity as its adapter digest; - Orka sends `AgentConfiguration: null`; the image-bound config is authoritative; - provider calls use the supervisor's loopback proxy, and tools use its one prompt-scoped loopback HTTP MCP server; @@ -43,9 +44,9 @@ Deploy the composed image as an operator-owned v2 supervisor service, configure the standard `ORKA_ACP_*` profile, fence, token-file, and runtime identity settings, then register it with Orka's strict-governed `AgentRuntime` sample. The profile model and `agentConfigurationDigest` must match the baked AgentKit -config. The registration's `adapterName` must be `agentkit-serve-acp`, and its -`adapterDigest` must equal the `AGENTKIT_ADAPTER_DIGEST` baked into the composed -image. +config. The registration's `adapterName` must be `agentkit-serve-acp`. Its +`adapterDigest` and the composition build's `AGENTKIT_ADAPTER_DIGEST` must both +equal the `sha256:` digest from `AGENTKIT_RUNTIME_IMAGE`. Orka freezes the AgentRuntime UID, generation, endpoint, profile, authentication Secret versions, and observed instance into each Task binding. It revalidates From 8e4116d6e961586887b1fbe28411f8ff6f23d31e Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Thu, 3 Sep 2026 20:57:43 -0700 Subject: [PATCH 07/23] fix(acp): verify parsed config bytes atomically Signed-off-by: Sertac Ozercan --- runtimes/common/agentkit_serve_common/acp.py | 18 +++--- runtimes/common/agentkit_serve_common/cli.py | 8 +-- .../common/agentkit_serve_common/config.py | 55 +++++++++++++------ runtimes/common/tests/test_acp_protocol.py | 55 +++++++++++++++++-- runtimes/common/tests/test_cli_protocol.py | 12 ++-- 5 files changed, 107 insertions(+), 41 deletions(-) diff --git a/runtimes/common/agentkit_serve_common/acp.py b/runtimes/common/agentkit_serve_common/acp.py index eda86d5..4388757 100644 --- a/runtimes/common/agentkit_serve_common/acp.py +++ b/runtimes/common/agentkit_serve_common/acp.py @@ -22,7 +22,7 @@ from typing import Any, Awaitable, BinaryIO, Callable, Iterator, Mapping from urllib.parse import urlsplit -from .config import AgentSpec +from .config import AgentSpec, load_with_bytes from .conversation import ConversationTurn, RunRequest from .runtime import AgentRunError, RuntimeFactory, RuntimeSession @@ -177,8 +177,15 @@ def _required_environment(name: str) -> str: raise ACPConfigurationError(exc.message) from exc -def validate_acp_runtime_binding(config_path: str | Path, spec: AgentSpec) -> None: - """Fail closed unless Orka's immutable profile matches the baked config.""" +def load_verified_acp_runtime_binding(config_path: str | Path) -> AgentSpec: + """Read, parse, and verify one immutable ACP agent configuration buffer.""" + spec, config_bytes = load_with_bytes(config_path) + validate_acp_runtime_binding(config_bytes, spec) + return spec + + +def validate_acp_runtime_binding(config_bytes: bytes, spec: AgentSpec) -> None: + """Fail closed unless Orka's immutable profile matches the parsed config bytes.""" if spec.tools: raise ACPConfigurationError("ACP strict mode rejects baked direct tools") @@ -205,11 +212,6 @@ def validate_acp_runtime_binding(config_path: str | Path, spec: AgentSpec) -> No f"{ACP_AGENT_CONFIGURATION_DIGEST_ENV} must be a lowercase sha256 digest" ) from exc - path = Path(config_path) - try: - config_bytes = path.read_bytes() - except OSError as exc: - raise ACPConfigurationError(f"cannot read ACP agent config {path}: {exc}") from exc actual_digest = prefix + hashlib.sha256(config_bytes).hexdigest() if not secrets.compare_digest(expected_digest, actual_digest): raise ACPConfigurationError( diff --git a/runtimes/common/agentkit_serve_common/cli.py b/runtimes/common/agentkit_serve_common/cli.py index c3b9d29..5ec1116 100644 --- a/runtimes/common/agentkit_serve_common/cli.py +++ b/runtimes/common/agentkit_serve_common/cli.py @@ -35,8 +35,8 @@ from .acp import ( ACPConfigurationError, ACPProtocolError, + load_verified_acp_runtime_binding, run_acp_stdio, - validate_acp_runtime_binding, ) from .config import ConfigError, load, load_or_exit from .foundry import create_foundry_app @@ -136,14 +136,14 @@ def run(factory: RuntimeFactory, argv: list[str] | None = None) -> None: # are intentionally adapter-owned and read AGENTKIT_PROTOCOL when a turn later # builds a runtime session. os.environ["AGENTKIT_PROTOCOL"] = protocol - spec = _load_spec_or_exit(args.config, protocol) if protocol == "acp": try: - validate_acp_runtime_binding(args.config, spec) - except (ACPConfigurationError, ACPProtocolError) as exc: + spec = load_verified_acp_runtime_binding(args.config) + except (ConfigError, ACPConfigurationError, ACPProtocolError) as exc: _fail(str(exc)) run_acp_stdio(spec, factory) return + spec = _load_spec_or_exit(args.config, protocol) if spec.brokered_tools and protocol != "foundry": _fail( "brokeredTools require AGENTKIT_PROTOCOL=foundry (or --protocol foundry); " diff --git a/runtimes/common/agentkit_serve_common/config.py b/runtimes/common/agentkit_serve_common/config.py index 6cd7a92..531dcec 100644 --- a/runtimes/common/agentkit_serve_common/config.py +++ b/runtimes/common/agentkit_serve_common/config.py @@ -1245,27 +1245,14 @@ class ConfigError(Exception): """Raised when ``agent.yaml`` is missing, unparseable, or invalid.""" -def load(path: str | Path) -> AgentSpec: - """Read and validate ``agent.yaml`` at ``path``. - - Raises :class:`ConfigError` with a clear, single-line message on any problem - (missing file, non-mapping YAML, schema violation, or unsupported abiVersion). - """ - p = Path(path) - try: - raw = p.read_text(encoding="utf-8") - except FileNotFoundError as exc: - raise ConfigError(f"agent config not found: {p}") from exc - except OSError as exc: - raise ConfigError(f"cannot read agent config {p}: {exc}") from exc - +def _parse_agent_spec(raw: str, source: Path) -> AgentSpec: try: data = safe_load_lossless(raw) except (yaml.YAMLError, ValueError, RecursionError) as exc: - raise ConfigError(f"agent config {p} is not valid YAML: {exc}") from exc + raise ConfigError(f"agent config {source} is not valid YAML: {exc}") from exc if not isinstance(data, dict): - raise ConfigError(f"agent config {p} must be a YAML mapping, got {type(data).__name__}") + raise ConfigError(f"agent config {source} must be a YAML mapping, got {type(data).__name__}") try: spec = AgentSpec.model_validate(data) @@ -1276,18 +1263,50 @@ def load(path: str | Path) -> AgentSpec: for err in exc.errors() ] raise ConfigError( - f"agent config {p} is invalid:\n" + "\n".join(lines) + f"agent config {source} is invalid:\n" + "\n".join(lines) ) from exc if spec.abi_version != ABI_VERSION: raise ConfigError( - f"agent config {p}: unsupported abiVersion {spec.abi_version!r} " + f"agent config {source}: unsupported abiVersion {spec.abi_version!r} " f"(this build of agentkit-serve understands {ABI_VERSION!r})" ) return spec +def load_bytes(raw: bytes, *, source: str | Path) -> AgentSpec: + """Parse and validate one exact ``agent.yaml`` byte buffer.""" + source_path = Path(source) + try: + text = raw.decode("utf-8") + except UnicodeDecodeError as exc: + raise ConfigError(f"agent config {source_path} is not valid UTF-8: {exc}") from exc + return _parse_agent_spec(text, source_path) + + +def load_with_bytes(path: str | Path) -> tuple[AgentSpec, bytes]: + """Read ``agent.yaml`` once and return its validated spec and exact bytes.""" + source = Path(path) + try: + raw = source.read_bytes() + except FileNotFoundError as exc: + raise ConfigError(f"agent config not found: {source}") from exc + except OSError as exc: + raise ConfigError(f"cannot read agent config {source}: {exc}") from exc + return load_bytes(raw, source=source), raw + + +def load(path: str | Path) -> AgentSpec: + """Read and validate ``agent.yaml`` at ``path``. + + Raises :class:`ConfigError` with a clear, single-line message on any problem + (missing file, non-mapping YAML, schema violation, or unsupported abiVersion). + """ + spec, _ = load_with_bytes(path) + return spec + + def validate_required_env(spec: AgentSpec) -> None: """Fail if required runtime env vars declared in the ABI are missing. diff --git a/runtimes/common/tests/test_acp_protocol.py b/runtimes/common/tests/test_acp_protocol.py index eda70d9..eec6104 100644 --- a/runtimes/common/tests/test_acp_protocol.py +++ b/runtimes/common/tests/test_acp_protocol.py @@ -6,6 +6,7 @@ import json import os from collections.abc import Callable +from pathlib import Path from types import TracebackType from typing import Any @@ -773,11 +774,55 @@ def test_runtime_binding_verifies_exact_config_digest_model_and_provider(monkeyp monkeypatch.setenv(acp.ACP_MODEL_ENV, "test-model") _set_provider_environment(monkeypatch) - acp.validate_acp_runtime_binding(config, _spec()) + acp.validate_acp_runtime_binding(config_bytes, _spec()) - config.write_bytes(config_bytes + b"# changed\n") with pytest.raises(ACPConfigurationError, match="exact agent config bytes"): - acp.validate_acp_runtime_binding(config, _spec()) + acp.validate_acp_runtime_binding(config_bytes + b"# changed\n", _spec()) + + +def test_verified_runtime_binding_parses_the_hashed_bytes_after_file_replacement(monkeypatch, tmp_path): + config = tmp_path / "agent.yaml" + verified_bytes = b"""abiVersion: v0 +metadata: + name: verified +model: + provider: openai-compatible + baseURL: https://baked.example.invalid/v1 + name: test-model + apiKeyEnv: TEST +instructions: Verified instructions. +tools: [] +expose: + openai: true + port: 8080 +""" + replaced_bytes = verified_bytes.replace(b"Verified instructions.", b"Replaced instructions.") + config.write_bytes(verified_bytes) + monkeypatch.setenv( + acp.ACP_AGENT_CONFIGURATION_DIGEST_ENV, + "sha256:" + hashlib.sha256(verified_bytes).hexdigest(), + ) + monkeypatch.setenv(acp.ACP_MODEL_ENV, "test-model") + _set_provider_environment(monkeypatch) + + original_read_bytes = Path.read_bytes + reads = 0 + + def read_then_replace(path: Path) -> bytes: + nonlocal reads + raw = original_read_bytes(path) + if path == config: + reads += 1 + config.write_bytes(replaced_bytes) + return raw + + monkeypatch.setattr(Path, "read_bytes", read_then_replace) + + spec = acp.load_verified_acp_runtime_binding(config) + + assert reads == 1 + assert spec.instructions == "Verified instructions." + assert original_read_bytes(config) == replaced_bytes @pytest.mark.parametrize( @@ -807,7 +852,7 @@ def test_runtime_binding_rejects_profile_or_provider_mismatch( monkeypatch.setenv(environment_name, environment_value) with pytest.raises(ACPConfigurationError, match=message): - acp.validate_acp_runtime_binding(config, _spec()) + acp.validate_acp_runtime_binding(b"exact", _spec()) @pytest.mark.parametrize( @@ -849,4 +894,4 @@ def test_runtime_binding_rejects_baked_tool_and_context_paths(tmp_path, override config.write_bytes(b"unused") with pytest.raises(ACPConfigurationError, match=message): - acp.validate_acp_runtime_binding(config, _spec(**override)) + acp.validate_acp_runtime_binding(b"unused", _spec(**override)) diff --git a/runtimes/common/tests/test_cli_protocol.py b/runtimes/common/tests/test_cli_protocol.py index 3c082b4..b560c6e 100644 --- a/runtimes/common/tests/test_cli_protocol.py +++ b/runtimes/common/tests/test_cli_protocol.py @@ -162,12 +162,12 @@ def test_cli_protocol_flag_sets_agentkit_protocol_for_adapter_runtime(monkeypatc def test_cli_acp_uses_verified_stdio_entrypoint_without_uvicorn(monkeypatch): captured = {} spec = _spec() - monkeypatch.setattr(cli, "load", lambda path: spec) - monkeypatch.setattr( - cli, - "validate_acp_runtime_binding", - lambda path, loaded: captured.update({"path": path, "verified": loaded}), - ) + + def load_verified(path): # noqa: ANN001 + captured.update({"path": path, "verified": spec}) + return spec + + monkeypatch.setattr(cli, "load_verified_acp_runtime_binding", load_verified) monkeypatch.setattr( cli, "run_acp_stdio", From 3a4ae0fe4f44b332431aeffdfd3c85ade90215ad Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Thu, 3 Sep 2026 21:59:05 -0700 Subject: [PATCH 08/23] docs(orka): document controller epoch rotation Signed-off-by: Sertac Ozercan --- docs/orka.md | 14 ++++++++++++++ 1 file changed, 14 insertions(+) diff --git a/docs/orka.md b/docs/orka.md index f755935..28a89ef 100644 --- a/docs/orka.md +++ b/docs/orka.md @@ -48,6 +48,20 @@ config. The registration's `adapterName` must be `agentkit-serve-acp`. Its `adapterDigest` and the composition build's `AGENTKIT_ADAPTER_DIGEST` must both equal the `sha256:` digest from `AGENTKIT_RUNTIME_IMAGE`. +Set `ORKA_ACP_CONTROLLER_EPOCH` from Orka's current `ControllerEpoch` record: + +```sh +kubectl -n get cepoch controller-epoch-orka-controller \ + -o jsonpath='{.status.epoch}' +``` + +The supervisor reads the epoch only during startup. The operator that owns this +service must watch the record and restart or replace the supervisor whenever it +changes. Preserve `ORKA_ACP_RUNTIME_INSTANCE_ID` across that restart and issue a +new `ORKA_ACP_SUPERVISOR_BOOT_ID`. Orka keeps a stale-epoch registration not +ready and refuses new Task bindings until authenticated status reports the +current value. AgentKit itself is the ACP child and does not manage this fence. + Orka freezes the AgentRuntime UID, generation, endpoint, profile, authentication Secret versions, and observed instance into each Task binding. It revalidates them before dispatch and recovery mutations. `Task.spec.execution.workspace` From 86563f93712d2ced7b27328dd3557f024c91c825 Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Thu, 3 Sep 2026 22:26:43 -0700 Subject: [PATCH 09/23] fix(runtime): scope failed session rollback to ACP Signed-off-by: Sertac Ozercan --- .../agentkit_serve/agent_factory.py | 16 +++---- .../tests/test_guardrails.py | 46 ++++++++++++++++++- 2 files changed, 51 insertions(+), 11 deletions(-) diff --git a/runtimes/microsoft-agent-framework/agentkit_serve/agent_factory.py b/runtimes/microsoft-agent-framework/agentkit_serve/agent_factory.py index a96bbfc..867bca0 100644 --- a/runtimes/microsoft-agent-framework/agentkit_serve/agent_factory.py +++ b/runtimes/microsoft-agent-framework/agentkit_serve/agent_factory.py @@ -392,16 +392,12 @@ async def run(self, request: RunRequest) -> RunResult: self._touch_session(session_id) session = self.sessions[session_id] include_history = session_id not in self.initialized_sessions - try: - result = await run_agent( - self.agent, - request, - session=session, - include_history=include_history, - ) - except BaseException: - self._reset_session(session_id) - raise + result = await run_agent( + self.agent, + request, + session=session, + include_history=include_history, + ) self.initialized_sessions.add(session_id) return result finally: diff --git a/runtimes/microsoft-agent-framework/tests/test_guardrails.py b/runtimes/microsoft-agent-framework/tests/test_guardrails.py index 022e0fb..3a4084e 100644 --- a/runtimes/microsoft-agent-framework/tests/test_guardrails.py +++ b/runtimes/microsoft-agent-framework/tests/test_guardrails.py @@ -753,7 +753,7 @@ async def fake_run_agent(agent, request, *, session=None, include_history=True): assert include_history_values == [True, False] -def test_failed_first_turn_replaces_session_and_keeps_explicit_history_on_retry(monkeypatch): +def test_discarded_failed_first_turn_replaces_session_and_keeps_explicit_history_on_retry(monkeypatch): from agentkit_serve_common.config import AgentSpec from agentkit_serve_common.conversation import ConversationTurn, RunRequest from agentkit_serve_common.runtime import RunResult @@ -791,6 +791,7 @@ async def fake_run_agent(agent, request, *, session=None, include_history=True): async def exercise(): with pytest.raises(RuntimeError, match="provider unavailable"): await runtime.run(request) + await runtime.discard_session("s1") await runtime.run(request) await runtime.run(request) @@ -801,6 +802,49 @@ async def exercise(): assert seen_sessions[1] is seen_sessions[2] +def test_failed_reused_session_preserves_native_conversation_state(monkeypatch): + from agentkit_serve_common.config import AgentSpec + from agentkit_serve_common.conversation import RunRequest + from agentkit_serve_common.runtime import RunResult + + include_history_values = [] + seen_sessions = [] + + async def fake_run_agent(agent, request, *, session=None, include_history=True): + include_history_values.append(include_history) + seen_sessions.append(session) + if len(seen_sessions) == 2: + raise RuntimeError("provider unavailable") + return RunResult(text="ok") + + monkeypatch.setattr(agent_factory, "run_agent", fake_run_agent) + spec = AgentSpec.model_validate({ + "abiVersion": "v0", + "metadata": {"name": "x"}, + "model": {"provider": "openai-compatible", "baseURL": "https://api.openai.com/v1", "name": "gpt-4o-mini"}, + "instructions": "hi", + "tools": [], + "expose": {"openai": True, "port": 8080}, + }) + runtime = agent_factory.MAFRuntime(spec) + runtime.agent = object() + request = RunRequest(prompt="current", session_id="s1") + + import asyncio + import pytest + + async def exercise(): + await runtime.run(request) + with pytest.raises(RuntimeError, match="provider unavailable"): + await runtime.run(request) + await runtime.run(request) + + asyncio.run(exercise()) + + assert include_history_values == [True, False, False] + assert seen_sessions[0] is seen_sessions[1] is seen_sessions[2] + + def test_remote_mcp_disables_ping(monkeypatch): tool = ToolSpec.model_validate({ "name": "toolbox", From 700e6c80ed7602ab6d554d43cdb73d0e2391a24f Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Fri, 4 Sep 2026 02:19:11 -0700 Subject: [PATCH 10/23] chore(runtime): clarify ACP auth value naming Signed-off-by: Sertac Ozercan --- runtimes/common/agentkit_serve_common/acp.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/runtimes/common/agentkit_serve_common/acp.py b/runtimes/common/agentkit_serve_common/acp.py index 4388757..7f55a1b 100644 --- a/runtimes/common/agentkit_serve_common/acp.py +++ b/runtimes/common/agentkit_serve_common/acp.py @@ -472,9 +472,9 @@ def _project_spec( mcp_servers: list[Any], ) -> tuple[AgentSpec, dict[str, str]]: provider_base_url = os.environ.get(ACP_PROVIDER_BASE_URL_ENV, "") - provider_token = os.environ.get(ACP_PROVIDER_TOKEN_ENV, "") + provider_auth_value = os.environ.get(ACP_PROVIDER_TOKEN_ENV, "") provider_base_url = _loopback_http_url(provider_base_url, name=ACP_PROVIDER_BASE_URL_ENV) - _safe_environment_value(provider_token, name=ACP_PROVIDER_TOKEN_ENV) + _safe_environment_value(provider_auth_value, name=ACP_PROVIDER_TOKEN_ENV) data = self.spec.model_dump(by_alias=True) data["model"]["baseURL"] = provider_base_url From 37d00de525f939b93bc7ece0dda0ae4c1dcdce0c Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Fri, 4 Sep 2026 03:02:37 -0700 Subject: [PATCH 11/23] fix(runtime): complete ACP v1 interoperability Signed-off-by: Sertac Ozercan --- runtimes/common/README.md | 12 +- runtimes/common/agentkit_serve_common/acp.py | 120 ++++-- runtimes/common/tests/test_acp_protocol.py | 403 ++++++++++++++++++- 3 files changed, 495 insertions(+), 40 deletions(-) diff --git a/runtimes/common/README.md b/runtimes/common/README.md index cfa53db..3e0097b 100644 --- a/runtimes/common/README.md +++ b/runtimes/common/README.md @@ -53,11 +53,13 @@ ACP mode requires `AGENTKIT_ACP_AGENT_CONFIGURATION_DIGEST` to equal the equal `model.name`. It replaces the baked model endpoint and credential with `AGENTKIT_ACP_PROVIDER_BASE_URL` and `AGENTKIT_ACP_PROVIDER_TOKEN`. -The child accepts one ACP session, text prompts, cancellation, and at most one -loopback HTTP MCP server carrying a bearer Authorization header. The runtime -keeps successful user and assistant turns for later prompts. It rejects baked -`tools`, `brokeredTools`, and context providers. Orka owns process and workspace -isolation, prompt-scoped MCP authority, provider proxying, and cleanup proof. +The child accepts one ACP session, text and resource-link prompt blocks, +cancellation, and at most one loopback HTTP MCP server carrying a bearer +Authorization header. Resource links are added to the model prompt as labeled +text and are never fetched by the child. The runtime keeps successful user and +assistant turns for later prompts. It rejects baked `tools`, `brokeredTools`, +and context providers. Orka owns process and workspace isolation, prompt-scoped +MCP authority, provider proxying, and cleanup proof. ## Adding a runtime adapter diff --git a/runtimes/common/agentkit_serve_common/acp.py b/runtimes/common/agentkit_serve_common/acp.py index 7f55a1b..a93a8df 100644 --- a/runtimes/common/agentkit_serve_common/acp.py +++ b/runtimes/common/agentkit_serve_common/acp.py @@ -51,7 +51,7 @@ _METHOD_NOT_FOUND = -32601 _INVALID_PARAMS = -32602 _INTERNAL_ERROR = -32603 -_RUNTIME_ERROR = -32000 +_REQUEST_CANCELLED = -32800 _MessageSender = Callable[[Mapping[str, Any]], Awaitable[None]] @@ -246,6 +246,9 @@ def __init__(self, spec: AgentSpec, factory: RuntimeFactory, send: _MessageSende ) self.sessions: dict[str, _SessionState] = {} self.requests: dict[str, asyncio.Task[None]] = {} + self.active_handlers: dict[str, asyncio.Task[None]] = {} + self.started_handlers: set[str] = set() + self.handler_cancellations: set[str] = set() self.prompt_requests: dict[str, _SessionState] = {} self.session_creation_active = False self.closed = False @@ -291,6 +294,8 @@ async def accept(self, message: Any) -> None: name=f"agentkit-acp-{method}", ) self.requests[key] = task + if method != _METHOD_SESSION_PROMPT: + self.active_handlers[key] = task task.add_done_callback(lambda completed, request_key=key: self._request_done(request_key, completed)) # Give prompt dispatch a chance to register its cancellation target before # the reader accepts an immediately following session/cancel notification. @@ -312,6 +317,8 @@ async def close(self) -> None: state.cancel_requested = True if state.active_run is not None and not state.active_run.done(): state.active_run.cancel() + for key in tuple(self.active_handlers): + self._cancel_handler(key) await self.wait_idle() errors: list[BaseException] = [] for state in reversed(tuple(self.sessions.values())): @@ -328,36 +335,56 @@ async def close(self) -> None: def _request_done(self, key: str, task: asyncio.Task[None]) -> None: if self.requests.get(key) is task: self.requests.pop(key, None) + if self.active_handlers.get(key) is task: + self.active_handlers.pop(key, None) + self.started_handlers.discard(key) + self.handler_cancellations.discard(key) # Retrieve exceptions even if stdout failed after the caller disconnected. if not task.cancelled(): task.exception() async def _dispatch_request(self, response_id: Any, key: str, method: str, params: Any) -> None: + handler = asyncio.current_task() + if handler is None: + raise RuntimeError("ACP request dispatch requires an asyncio task") + if method != _METHOD_SESSION_PROMPT: + self.started_handlers.add(key) + result: Any = None + error: tuple[int, str, Mapping[str, Any] | None] | None = None try: - if method == _METHOD_INITIALIZE: - result = self._initialize(params) - elif method == _METHOD_SESSION_NEW: - result = await self._new_session(params) - elif method == _METHOD_SESSION_PROMPT: - result = await self._prompt(key, params) - else: - raise ACPProtocolError(_METHOD_NOT_FOUND, f"unsupported ACP method {method!r}") - except ACPProtocolError as exc: - await self._send_protocol_error(response_id, exc) - return - except AgentRunError as exc: - data = {"code": exc.code or exc.__class__.__name__} - await self._send_error(response_id, _RUNTIME_ERROR, "AgentKit runtime prompt failed", data) - return - except BaseException as exc: # noqa: BLE001 - keep provider details off stdout. - if isinstance(exc, (KeyboardInterrupt, SystemExit)): - raise - await self._send_error( - response_id, - _INTERNAL_ERROR, - "AgentKit ACP request failed", - {"code": exc.__class__.__name__}, - ) + try: + if key in self.handler_cancellations: + raise asyncio.CancelledError + if method == _METHOD_INITIALIZE: + result = self._initialize(params) + elif method == _METHOD_SESSION_NEW: + result = await self._new_session(params) + elif method == _METHOD_SESSION_PROMPT: + result = await self._prompt(key, params) + else: + raise ACPProtocolError(_METHOD_NOT_FOUND, f"unsupported ACP method {method!r}") + except asyncio.CancelledError: + error = (_REQUEST_CANCELLED, "ACP request cancelled", None) + except ACPProtocolError as exc: + error = (exc.code, exc.message, exc.data) + except AgentRunError as exc: + data = {"code": exc.code or exc.__class__.__name__} + error = (_INTERNAL_ERROR, "AgentKit runtime prompt failed", data) + except BaseException as exc: # noqa: BLE001 - keep provider details off stdout. + if isinstance(exc, (KeyboardInterrupt, SystemExit)): + raise + error = ( + _INTERNAL_ERROR, + "AgentKit ACP request failed", + {"code": exc.__class__.__name__}, + ) + finally: + if self.active_handlers.get(key) is handler: + self.active_handlers.pop(key, None) + self.started_handlers.discard(key) + self.handler_cancellations.discard(key) + if error is not None: + await self._send_error(response_id, *error) return await self.send({"jsonrpc": _JSONRPC_VERSION, "id": response_id, "result": result}) @@ -375,7 +402,11 @@ async def _handle_notification(self, method: str, params: Any) -> None: key = _request_key(params["requestId"]) except ACPProtocolError: return - self._cancel_state(self.prompt_requests.get(key)) + state = self.prompt_requests.get(key) + if state is not None: + self._cancel_state(state) + return + self._cancel_handler(key) def _initialize(self, params: Any) -> dict[str, Any]: request = _required_object(params) @@ -562,17 +593,34 @@ async def _prompt(self, request_key: str, params: Any) -> dict[str, str]: if state is None: raise ACPProtocolError(_INVALID_PARAMS, "unknown ACP sessionId") if state.active_request_key is not None: - raise ACPProtocolError(_RUNTIME_ERROR, "ACP session already has an active prompt") + raise ACPProtocolError(_REQUEST_CANCELLED, "ACP session already has an active prompt") prompt = request.get("prompt") if not isinstance(prompt, list) or not prompt: raise ACPProtocolError(_INVALID_PARAMS, "prompt must be a non-empty array") text_blocks: list[str] = [] for index, raw_block in enumerate(prompt): block = _required_object(raw_block, name=f"prompt[{index}]") - if block.get("type") != "text": - raise ACPProtocolError(_INVALID_PARAMS, "ACP mode accepts only text prompt blocks") - text_blocks.append( - _required_prompt_text(block.get("text"), name=f"prompt[{index}].text") + block_type = block.get("type") + if block_type == "text": + text_blocks.append( + _required_prompt_text(block.get("text"), name=f"prompt[{index}].text") + ) + continue + if block_type == "resource_link": + name = _required_string(block.get("name"), name=f"prompt[{index}].name") + uri = _required_string(block.get("uri"), name=f"prompt[{index}].uri") + projected = [f"Resource link: {name}", f"URI: {uri}"] + if block.get("mimeType") is not None: + mime_type = _required_string( + block.get("mimeType"), + name=f"prompt[{index}].mimeType", + ) + projected.append(f"MIME type: {mime_type}") + text_blocks.append("\n".join(projected)) + continue + raise ACPProtocolError( + _INVALID_PARAMS, + "ACP mode accepts only text and resource_link prompt blocks", ) run_request = RunRequest( @@ -650,6 +698,16 @@ def _cancel_state(self, state: _SessionState | None) -> None: if state.active_run is not None and not state.active_run.done(): state.active_run.cancel() + def _cancel_handler(self, key: str) -> None: + handler = self.active_handlers.get(key) + if handler is None or handler.done(): + return + if key in self.handler_cancellations: + return + self.handler_cancellations.add(key) + if key in self.started_handlers: + handler.cancel() + @staticmethod def _install_environment(environment: Mapping[str, str]) -> dict[str, str | None]: previous = {name: os.environ.get(name) for name in environment} diff --git a/runtimes/common/tests/test_acp_protocol.py b/runtimes/common/tests/test_acp_protocol.py index eec6104..fec0c67 100644 --- a/runtimes/common/tests/test_acp_protocol.py +++ b/runtimes/common/tests/test_acp_protocol.py @@ -77,13 +77,25 @@ def _new_session(request_id: int = 2, *, mcp_servers: list[dict[str, Any]] | Non def _prompt(request_id: int, session_id: str, text: str) -> dict[str, Any]: + return _prompt_content( + request_id, + session_id, + [{"type": "text", "text": text}], + ) + + +def _prompt_content( + request_id: int, + session_id: str, + content: list[dict[str, Any]], +) -> dict[str, Any]: return { "jsonrpc": "2.0", "id": request_id, "method": "session/prompt", "params": { "sessionId": session_id, - "prompt": [{"type": "text", "text": text}], + "prompt": content, }, } @@ -113,6 +125,21 @@ async def _send_to( return _response(messages, request["id"]) +async def _wait_for_response( + messages: list[dict[str, Any]], + request_id: int, +) -> dict[str, Any]: + async def wait() -> dict[str, Any]: + while True: + matches = [message for message in messages if message.get("id") == request_id] + if matches: + assert len(matches) == 1 + return matches[0] + await asyncio.sleep(0) + + return await asyncio.wait_for(wait(), timeout=1) + + class RecordingRuntime: def __init__(self, outcomes: list[RunResult | BaseException]) -> None: self.outcomes = list(outcomes) @@ -175,6 +202,40 @@ def build_runtime(self, spec: AgentSpec) -> RuntimeSession: return runtime +class BlockingEnterRuntime: + def __init__(self) -> None: + self.started = asyncio.Event() + self.release = asyncio.Event() + self.exited = 0 + + async def __aenter__(self) -> RuntimeSession: + self.started.set() + await self.release.wait() + return self + + async def __aexit__(self, exc_type, exc, tb): # noqa: ANN001 + self.exited += 1 + return None + + async def run(self, request: RunRequest) -> RunResult: + raise AssertionError(f"unexpected prompt for {request.session_id}") + + +class BlockingExitRuntime(BlockingEnterRuntime): + def __init__(self) -> None: + super().__init__() + self.exit_started = asyncio.Event() + self.release_exit = asyncio.Event() + self.exit_completed = False + + async def __aexit__(self, exc_type, exc, tb): # noqa: ANN001 + self.exited += 1 + self.exit_started.set() + await self.release_exit.wait() + self.exit_completed = True + return None + + def test_offline_echo_round_trip_uses_canonical_acp_shapes(monkeypatch): _set_provider_environment(monkeypatch) @@ -383,6 +444,65 @@ async def send(message): # noqa: ANN001 asyncio.run(exercise()) +def test_resource_links_are_projected_and_preserved_in_history(monkeypatch): + _set_provider_environment(monkeypatch) + runtime = RecordingRuntime([RunResult("first answer"), RunResult("second answer")]) + factory = RecordingFactory(lambda: runtime) + projected_prompt = ( + "Review this resource.\n" + "Resource link: design notes\n" + "URI: https://example.com/design.md\n" + "MIME type: text/markdown\n" + "Summarize it." + ) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + + async def send(message): # noqa: ANN001 + messages.append(dict(message)) + + server = ACPStdioServer(_spec(), factory, send) + await _send_to(server, messages, _initialize()) + created = await _send_to(server, messages, _new_session()) + session_id = created["result"]["sessionId"] + prompt = _prompt_content( + 3, + session_id, + [ + {"type": "text", "text": "Review this resource."}, + { + "type": "resource_link", + "name": "design notes", + "uri": "https://example.com/design.md", + "mimeType": "text/markdown", + }, + {"type": "text", "text": "Summarize it."}, + ], + ) + + assert (await _send_to(server, messages, prompt))["result"] == { + "stopReason": "end_turn" + } + assert ( + await _send_to(server, messages, _prompt(4, session_id, "What changed?")) + )["result"] == {"stopReason": "end_turn"} + assert runtime.requests == [ + RunRequest(prompt=projected_prompt, history=(), session_id=session_id), + RunRequest( + prompt="What changed?", + history=( + ConversationTurn(role="user", text=projected_prompt), + ConversationTurn(role="assistant", text="first answer"), + ), + session_id=session_id, + ), + ] + await server.close() + + asyncio.run(exercise()) + + def test_child_rejects_a_second_session(monkeypatch): _set_provider_environment(monkeypatch) factory = RecordingFactory(lambda: RecordingRuntime([RunResult("unused")])) @@ -406,6 +526,190 @@ async def send(message): # noqa: ANN001 asyncio.run(exercise()) +def test_cancel_request_stops_blocked_session_creation_and_restores_environment(monkeypatch): + _set_provider_environment(monkeypatch) + runtime = BlockingEnterRuntime() + factory = RecordingFactory(lambda: runtime, supports_http_mcp=True) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + + async def send(message): # noqa: ANN001 + messages.append(dict(message)) + + server = ACPStdioServer(_spec(), factory, send) + await _send_to(server, messages, _initialize()) + await server.accept(_new_session(mcp_servers=[_orka_mcp_server()])) + await asyncio.wait_for(runtime.started.wait(), timeout=1) + projected_names = set(factory.environment_snapshots[0]) + assert projected_names + assert all(os.environ.get(name) for name in projected_names) + + await server.accept( + { + "jsonrpc": "2.0", + "method": "$/cancel_request", + "params": {"requestId": 2}, + } + ) + await asyncio.wait_for(server.wait_idle(), timeout=1) + + assert _response(messages, 2)["error"] == { + "code": -32800, + "message": "ACP request cancelled", + } + assert server.sessions == {} + assert runtime.exited == 1 + assert all(name not in os.environ for name in projected_names) + await asyncio.wait_for(server.close(), timeout=1) + + asyncio.run(exercise()) + + +def test_cancel_request_before_session_dispatch_prevents_runtime_start(monkeypatch): + _set_provider_environment(monkeypatch) + runtime = BlockingEnterRuntime() + factory = RecordingFactory(lambda: runtime) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + + async def send(message): # noqa: ANN001 + messages.append(dict(message)) + + server = ACPStdioServer(_spec(), factory, send) + await _send_to(server, messages, _initialize()) + accept_task = asyncio.create_task(server.accept(_new_session())) + await asyncio.sleep(0) + assert not runtime.started.is_set() + + await server.accept( + { + "jsonrpc": "2.0", + "method": "$/cancel_request", + "params": {"requestId": 2}, + } + ) + await accept_task + await asyncio.wait_for(server.wait_idle(), timeout=1) + + assert _response(messages, 2)["error"]["code"] == -32800 + assert factory.runtimes == [] + assert server.sessions == {} + await server.close() + + asyncio.run(exercise()) + + +def test_close_cancels_blocked_session_creation(monkeypatch): + _set_provider_environment(monkeypatch) + runtime = BlockingEnterRuntime() + factory = RecordingFactory(lambda: runtime) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + + async def send(message): # noqa: ANN001 + messages.append(dict(message)) + + server = ACPStdioServer(_spec(), factory, send) + await _send_to(server, messages, _initialize()) + accept_task = asyncio.create_task(server.accept(_new_session())) + await asyncio.sleep(0) + assert not runtime.started.is_set() + + async with asyncio.timeout(1): + await server.close() + await accept_task + + assert _response(messages, 2)["error"]["code"] == -32800 + assert server.sessions == {} + assert factory.runtimes == [] + + asyncio.run(exercise()) + + +def test_close_does_not_interrupt_cancelled_session_cleanup(monkeypatch): + _set_provider_environment(monkeypatch) + runtime = BlockingExitRuntime() + factory = RecordingFactory(lambda: runtime) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + + async def send(message): # noqa: ANN001 + messages.append(dict(message)) + + server = ACPStdioServer(_spec(), factory, send) + await _send_to(server, messages, _initialize()) + await server.accept(_new_session()) + await asyncio.wait_for(runtime.started.wait(), timeout=1) + await server.accept( + { + "jsonrpc": "2.0", + "method": "$/cancel_request", + "params": {"requestId": 2}, + } + ) + await asyncio.wait_for(runtime.exit_started.wait(), timeout=1) + + close_task = asyncio.create_task(server.close()) + await asyncio.sleep(0) + assert not close_task.done() + assert not runtime.exit_completed + + runtime.release_exit.set() + await asyncio.wait_for(close_task, timeout=1) + + assert runtime.exit_completed + assert runtime.exited == 1 + assert _response(messages, 2)["error"]["code"] == -32800 + assert server.sessions == {} + + asyncio.run(exercise()) + + +def test_cancel_request_after_session_commit_does_not_replace_success(monkeypatch): + _set_provider_environment(monkeypatch) + runtime = RecordingRuntime([RunResult("unused")]) + factory = RecordingFactory(lambda: runtime) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + response_started = asyncio.Event() + release_response = asyncio.Event() + + async def send(message): # noqa: ANN001 + if message.get("id") == 2 and "result" in message: + response_started.set() + await release_response.wait() + messages.append(dict(message)) + + server = ACPStdioServer(_spec(), factory, send) + await _send_to(server, messages, _initialize()) + await server.accept(_new_session()) + await asyncio.wait_for(response_started.wait(), timeout=1) + + for request_id in (2, 2, 999): + await server.accept( + { + "jsonrpc": "2.0", + "method": "$/cancel_request", + "params": {"requestId": request_id}, + } + ) + release_response.set() + await asyncio.wait_for(server.wait_idle(), timeout=1) + + created = _response(messages, 2) + assert "error" not in created + assert created["result"]["sessionId"] in server.sessions + assert runtime.entered == 1 + await server.close() + + asyncio.run(exercise()) + + class CancellingRuntime: def __init__(self, *, swallow_cancellation: bool) -> None: self.swallow_cancellation = swallow_cancellation @@ -437,6 +741,56 @@ async def discard_session(self, session_id: str) -> None: self.discarded_sessions.append(session_id) +def test_concurrent_prompt_is_rejected_as_cancelled_without_disturbing_active_prompt( + monkeypatch, +): + _set_provider_environment(monkeypatch) + runtime = CancellingRuntime(swallow_cancellation=False) + factory = RecordingFactory(lambda: runtime) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + + async def send(message): # noqa: ANN001 + messages.append(dict(message)) + + server = ACPStdioServer(_spec(), factory, send) + await _send_to(server, messages, _initialize()) + created = await _send_to(server, messages, _new_session()) + session_id = created["result"]["sessionId"] + + await server.accept(_prompt(3, session_id, "keep running")) + await asyncio.wait_for(runtime.started.wait(), timeout=1) + await server.accept(_prompt(4, session_id, "reject me")) + rejected = await _wait_for_response(messages, 4) + + assert rejected["error"] == { + "code": -32800, + "message": "ACP session already has an active prompt", + } + assert not any(message.get("id") == 3 for message in messages) + assert len(runtime.requests) == 1 + assert server.sessions[session_id].history == [] + + await server.accept( + { + "jsonrpc": "2.0", + "method": "session/cancel", + "params": {"sessionId": session_id}, + } + ) + await asyncio.wait_for(server.wait_idle(), timeout=1) + assert _response(messages, 3)["result"] == {"stopReason": "cancelled"} + assert server.sessions[session_id].history == [] + + completed = await _send_to(server, messages, _prompt(5, session_id, "try again")) + assert completed["result"] == {"stopReason": "end_turn"} + assert runtime.requests[1].history == () + await server.close() + + asyncio.run(exercise()) + + @pytest.mark.parametrize( ("cancel_method", "swallow_cancellation"), [("session/cancel", False), ("$/cancel_request", True)], @@ -505,7 +859,7 @@ async def send(message): # noqa: ANN001 failed = await _send_to(server, messages, _prompt(3, session_id, "fail")) assert failed["error"] == { - "code": -32000, + "code": -32603, "message": "AgentKit runtime prompt failed", "data": {"code": "ProviderFailure"}, } @@ -662,7 +1016,7 @@ async def send(message): # noqa: ANN001 asyncio.run(exercise()) -def test_prompt_rejects_non_text_content_without_running_model(monkeypatch): +def test_prompt_rejects_capability_gated_content_without_running_model(monkeypatch): _set_provider_environment(monkeypatch) runtime = RecordingRuntime([RunResult("must not run")]) factory = RecordingFactory(lambda: runtime) @@ -689,7 +1043,48 @@ async def send(message): # noqa: ANN001 failed = await _send_to(server, messages, invalid) assert failed["error"]["code"] == -32602 - assert "only text" in failed["error"]["message"] + assert "only text and resource_link" in failed["error"]["message"] + assert runtime.requests == [] + await server.close() + + asyncio.run(exercise()) + + +@pytest.mark.parametrize( + "block", + [ + {"type": "resource_link", "uri": "https://example.com/file.txt"}, + {"type": "resource_link", "name": "file"}, + { + "type": "resource_link", + "name": "file", + "uri": "https://example.com/file.txt", + "mimeType": "", + }, + ], +) +def test_prompt_rejects_malformed_resource_links_without_running_model(monkeypatch, block): + _set_provider_environment(monkeypatch) + runtime = RecordingRuntime([RunResult("must not run")]) + factory = RecordingFactory(lambda: runtime) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + + async def send(message): # noqa: ANN001 + messages.append(dict(message)) + + server = ACPStdioServer(_spec(), factory, send) + await _send_to(server, messages, _initialize()) + created = await _send_to(server, messages, _new_session()) + session_id = created["result"]["sessionId"] + + failed = await _send_to( + server, + messages, + _prompt_content(3, session_id, [block]), + ) + assert failed["error"]["code"] == -32602 assert runtime.requests == [] await server.close() From 82671b545a7ae7d205cc80ada1c2bacb78c1d4b9 Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Fri, 4 Sep 2026 03:41:47 -0700 Subject: [PATCH 12/23] fix(acp): keep cancellation responsive under backpressure Signed-off-by: Sertac Ozercan --- runtimes/common/agentkit_serve_common/acp.py | 110 ++++++++- runtimes/common/tests/test_acp_protocol.py | 244 ++++++++++++++++--- 2 files changed, 302 insertions(+), 52 deletions(-) diff --git a/runtimes/common/agentkit_serve_common/acp.py b/runtimes/common/agentkit_serve_common/acp.py index a93a8df..91e49a5 100644 --- a/runtimes/common/agentkit_serve_common/acp.py +++ b/runtimes/common/agentkit_serve_common/acp.py @@ -22,6 +22,7 @@ from typing import Any, Awaitable, BinaryIO, Callable, Iterator, Mapping from urllib.parse import urlsplit +from .adapter_support import _attach_secondary_error, _wait_for_owner_task from .config import AgentSpec, load_with_bytes from .conversation import ConversationTurn, RunRequest from .runtime import AgentRunError, RuntimeFactory, RuntimeSession @@ -70,6 +71,85 @@ class ACPConfigurationError(ValueError): """An ACP startup binding or strict-mode configuration failure.""" +@dataclass +class _PendingStdioWrite: + frame: bytes + completed: asyncio.Future[None] + + +def _retrieve_future_exception(future: asyncio.Future[None]) -> None: + if not future.cancelled(): + future.exception() + + +class _ACPStdioWriter: + """Serialize protocol frames without blocking the asyncio event loop.""" + + def __init__(self, output_stream: BinaryIO) -> None: + self._output_stream = output_stream + self._queue: asyncio.Queue[_PendingStdioWrite | None] = asyncio.Queue() + self._failure: BaseException | None = None + self._closed = False + self._worker = asyncio.create_task(self._run(), name="agentkit-acp-stdio-writer") + + async def send(self, message: Mapping[str, Any]) -> None: + frame = json.dumps(message, ensure_ascii=False, separators=(",", ":")).encode("utf-8") + b"\n" + if len(frame) > _MAX_MESSAGE_BYTES: + raise RuntimeError("ACP response exceeds the 8 MiB limit") + if self._failure is not None: + raise RuntimeError("ACP stdio writer failed") from self._failure + if self._closed: + raise RuntimeError("ACP stdio writer is closed") + + completed = asyncio.get_running_loop().create_future() + self._queue.put_nowait(_PendingStdioWrite(frame=frame, completed=completed)) + try: + await asyncio.shield(completed) + except asyncio.CancelledError: + # The worker owns an enqueued frame until the physical write finishes. + # Keep its completion observable without letting caller cancellation + # release serialization or produce an unhandled future exception. + completed.add_done_callback(_retrieve_future_exception) + raise + + async def close(self, *, preserve: BaseException | None = None) -> None: + if not self._closed: + self._closed = True + self._queue.put_nowait(None) + await _wait_for_owner_task(self._worker, preserve=preserve) + + async def _run(self) -> None: + pending: _PendingStdioWrite | None = None + try: + while True: + pending = await self._queue.get() + if pending is None: + return + await asyncio.to_thread(self._write, pending.frame) + pending.completed.set_result(None) + pending = None + except BaseException as exc: # noqa: BLE001 - wake every sender on terminal worker failure. + self._failure = exc + self._closed = True + if pending is not None and not pending.completed.done(): + pending.completed.set_exception(exc) + self._fail_queued_writes(exc) + raise + + def _write(self, frame: bytes) -> None: + self._output_stream.write(frame) + self._output_stream.flush() + + def _fail_queued_writes(self, error: BaseException) -> None: + while True: + try: + pending = self._queue.get_nowait() + except asyncio.QueueEmpty: + return + if pending is not None and not pending.completed.done(): + pending.completed.set_exception(error) + + @dataclass class _SessionState: session_id: str @@ -749,18 +829,10 @@ async def serve_acp_stdio( input_stream = reader or sys.stdin.buffer output_stream = writer or sys.stdout.buffer - write_lock = asyncio.Lock() - - async def send(message: Mapping[str, Any]) -> None: - encoded = json.dumps(message, ensure_ascii=False, separators=(",", ":")).encode("utf-8") + b"\n" - if len(encoded) > _MAX_MESSAGE_BYTES: - raise RuntimeError("ACP response exceeds the 8 MiB limit") - async with write_lock: - output_stream.write(encoded) - output_stream.flush() - - server = ACPStdioServer(spec, factory, send) + stdio_writer = _ACPStdioWriter(output_stream) + server: ACPStdioServer | None = None try: + server = ACPStdioServer(spec, factory, stdio_writer.send) while True: line = await asyncio.to_thread(input_stream.readline, _MAX_MESSAGE_BYTES + 2) if not line: @@ -777,7 +849,21 @@ async def send(message: Mapping[str, Any]) -> None: continue await server.accept_line(frame) finally: - await server.close() + primary_error = sys.exc_info()[1] + try: + if server is not None: + await server.close() + except BaseException as close_error: + if primary_error is None: + primary_error = close_error + raise + _attach_secondary_error( + primary_error, + close_error, + label="ACP server cleanup also failed", + ) + finally: + await stdio_writer.close(preserve=primary_error) def run_acp_stdio(spec: AgentSpec, factory: RuntimeFactory) -> None: diff --git a/runtimes/common/tests/test_acp_protocol.py b/runtimes/common/tests/test_acp_protocol.py index fec0c67..3e5f0f7 100644 --- a/runtimes/common/tests/test_acp_protocol.py +++ b/runtimes/common/tests/test_acp_protocol.py @@ -5,6 +5,8 @@ import io import json import os +import queue +import threading from collections.abc import Callable from pathlib import Path from types import TracebackType @@ -236,6 +238,63 @@ async def __aexit__(self, exc_type, exc, tb): # noqa: ANN001 return None +class FeedingReader: + def __init__(self) -> None: + self.lines: queue.Queue[bytes] = queue.Queue() + + def readline(self, limit: int) -> bytes: # noqa: ARG002 + return self.lines.get() + + def feed(self, message: dict[str, Any]) -> None: + self.lines.put(json.dumps(message, separators=(",", ":")).encode() + b"\n") + + def close(self) -> None: + self.lines.put(b"") + + +class RecordingWriter: + def __init__(self, *, block_first_update: bool = False) -> None: + self.block_first_update = block_first_update + self.blocked = threading.Event() + self.release = threading.Event() + self.frames: list[bytes] = [] + self.lock = threading.Lock() + + def write(self, frame: bytes) -> int: + message = json.loads(frame) + if ( + self.block_first_update + and message.get("method") == "session/update" + and not self.blocked.is_set() + ): + self.blocked.set() + if not self.release.wait(timeout=5): + raise TimeoutError("test did not release blocked ACP writer") + with self.lock: + self.frames.append(bytes(frame)) + return len(frame) + + def flush(self) -> None: + return None + + def messages(self) -> list[dict[str, Any]]: + with self.lock: + frames = list(self.frames) + return [json.loads(frame) for frame in frames] + + +async def _wait_for_stdio_response(writer: RecordingWriter, request_id: int) -> dict[str, Any]: + async def wait() -> dict[str, Any]: + while True: + matches = [message for message in writer.messages() if message.get("id") == request_id] + if matches: + assert len(matches) == 1 + return matches[0] + await asyncio.sleep(0) + + return await asyncio.wait_for(wait(), timeout=5) + + def test_offline_echo_round_trip_uses_canonical_acp_shapes(monkeypatch): _set_provider_environment(monkeypatch) @@ -278,30 +337,31 @@ async def send(message): # noqa: ANN001 def test_stdio_server_frames_json_rpc_and_closes_on_eof(monkeypatch): _set_provider_environment(monkeypatch) monkeypatch.setattr(acp.secrets, "token_hex", lambda size: "b" * (size * 2)) - session_id = "agentkit-" + "b" * 32 - requests = [ - _initialize(), - _new_session(), - _prompt(3, session_id, "through stdio"), - ] - reader = io.BytesIO( - b"".join( - json.dumps(request, separators=(",", ":")).encode() + b"\n" - for request in requests - ) - ) - writer = io.BytesIO() - asyncio.run( - acp.serve_acp_stdio( - _spec(), - OfflineEchoRuntimeFactory(), - reader=reader, - writer=writer, + async def exercise() -> list[dict[str, Any]]: + reader = FeedingReader() + writer = RecordingWriter() + serve_task = asyncio.create_task( + acp.serve_acp_stdio( + _spec(), + OfflineEchoRuntimeFactory(), + reader=reader, + writer=writer, + ) ) - ) - - output = [json.loads(line) for line in writer.getvalue().splitlines()] + try: + reader.feed(_initialize()) + await _wait_for_stdio_response(writer, 1) + reader.feed(_new_session()) + created = await _wait_for_stdio_response(writer, 2) + reader.feed(_prompt(3, created["result"]["sessionId"], "through stdio")) + await _wait_for_stdio_response(writer, 3) + finally: + reader.close() + await asyncio.wait_for(serve_task, timeout=1) + return writer.messages() + + output = asyncio.run(exercise()) assert _response(output, 3)["result"] == {"stopReason": "end_turn"} assert any(message.get("method") == "session/update" for message in output) @@ -312,26 +372,33 @@ def test_stdio_server_splits_output_larger_than_orka_acp_reader_limit(monkeypatc session_id = "agentkit-" + "c" * 32 result_text = "x" * ((2 << 20) + 1) runtime = RecordingRuntime([RunResult(result_text)]) - requests = [_initialize(), _new_session(), _prompt(3, session_id, "large output")] - reader = io.BytesIO( - b"".join( - json.dumps(request, separators=(",", ":")).encode() + b"\n" - for request in requests - ) - ) - writer = io.BytesIO() - asyncio.run( - acp.serve_acp_stdio( - _spec(), - RecordingFactory(lambda: runtime), - reader=reader, - writer=writer, + async def exercise() -> RecordingWriter: + reader = FeedingReader() + writer = RecordingWriter() + serve_task = asyncio.create_task( + acp.serve_acp_stdio( + _spec(), + RecordingFactory(lambda: runtime), + reader=reader, + writer=writer, + ) ) - ) - - lines = writer.getvalue().splitlines() - output = [json.loads(line) for line in lines] + try: + reader.feed(_initialize()) + await _wait_for_stdio_response(writer, 1) + reader.feed(_new_session()) + await _wait_for_stdio_response(writer, 2) + reader.feed(_prompt(3, session_id, "large output")) + await _wait_for_stdio_response(writer, 3) + finally: + reader.close() + await asyncio.wait_for(serve_task, timeout=1) + return writer + + writer = asyncio.run(exercise()) + lines = list(writer.frames) + output = writer.messages() updates = [message for message in output if message.get("method") == "session/update"] chunks = [message["params"]["update"]["content"]["text"] for message in updates] @@ -373,6 +440,103 @@ def test_stdio_server_drains_oversized_line_before_parsing_next_frame(monkeypatc assert _response(output, 1)["result"]["protocolVersion"] == 1 +def test_stdio_blocked_writer_does_not_block_cancel_dispatch(monkeypatch): + _set_provider_environment(monkeypatch) + runtime = RecordingRuntime([RunResult("late answer"), RunResult("clean answer")]) + factory = RecordingFactory(lambda: runtime) + cancel_dispatched = asyncio.Event() + original_cancel_state = ACPStdioServer._cancel_state # noqa: SLF001 + + def record_cancel(server: ACPStdioServer, state) -> None: # noqa: ANN001 + cancel_dispatched.set() + original_cancel_state(server, state) + + monkeypatch.setattr(ACPStdioServer, "_cancel_state", record_cancel) + + async def exercise() -> None: + reader = FeedingReader() + writer = RecordingWriter(block_first_update=True) + serve_task = asyncio.create_task( + acp.serve_acp_stdio(_spec(), factory, reader=reader, writer=writer) + ) + try: + reader.feed(_initialize()) + await _wait_for_stdio_response(writer, 1) + reader.feed(_new_session()) + created = await _wait_for_stdio_response(writer, 2) + session_id = created["result"]["sessionId"] + + reader.feed(_prompt(3, session_id, "cancel under backpressure")) + assert await asyncio.to_thread(writer.blocked.wait, 1) + reader.feed( + { + "jsonrpc": "2.0", + "method": "session/cancel", + "params": {"sessionId": session_id}, + } + ) + await asyncio.wait_for(cancel_dispatched.wait(), timeout=1) + + writer.release.set() + cancelled = await _wait_for_stdio_response(writer, 3) + assert cancelled["result"] == {"stopReason": "cancelled"} + assert runtime.discarded_sessions == [session_id] + + reader.feed(_prompt(4, session_id, "try again")) + completed = await _wait_for_stdio_response(writer, 4) + assert completed["result"] == {"stopReason": "end_turn"} + assert runtime.requests[1].history == () + finally: + writer.release.set() + reader.close() + await asyncio.wait_for(serve_task, timeout=1) + + asyncio.run(exercise()) + + +def test_stdio_writer_close_waits_for_in_flight_write_after_cancellation(): + async def exercise() -> None: + output = RecordingWriter(block_first_update=True) + writer = acp._ACPStdioWriter(output) # noqa: SLF001 + send_task = asyncio.create_task( + writer.send({"jsonrpc": "2.0", "method": "session/update"}) + ) + assert await asyncio.to_thread(output.blocked.wait, 1) + + close_task = asyncio.create_task(writer.close()) + await asyncio.sleep(0) + close_task.cancel() + await asyncio.sleep(0) + assert not close_task.done() + + output.release.set() + await send_task + with pytest.raises(asyncio.CancelledError): + await close_task + + asyncio.run(exercise()) + + +def test_stdio_server_surfaces_writer_failure(): + class FailingWriter: + def write(self, frame: bytes) -> int: # noqa: ARG002 + raise BrokenPipeError("supervisor closed stdout") + + def flush(self) -> None: + return None + + request = json.dumps(_initialize(), separators=(",", ":")).encode() + b"\n" + with pytest.raises(BrokenPipeError, match="supervisor closed stdout"): + asyncio.run( + acp.serve_acp_stdio( + _spec(), + OfflineEchoRuntimeFactory(), + reader=io.BytesIO(request), + writer=FailingWriter(), + ) + ) + + def test_offline_echo_gate_applies_to_acp(monkeypatch): monkeypatch.setenv("AGENTKIT_PROTOCOL", "acp") monkeypatch.setenv("AGENTKIT_ORKA_OFFLINE_ECHO", "true") From 50a752941ce52e72eb2f9dd73e01ef610106525d Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Fri, 4 Sep 2026 04:07:52 -0700 Subject: [PATCH 13/23] docs(orka): clarify harness v2 runtime policy Signed-off-by: Sertac Ozercan --- docs/orka.md | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/docs/orka.md b/docs/orka.md index 28a89ef..78c215d 100644 --- a/docs/orka.md +++ b/docs/orka.md @@ -46,7 +46,12 @@ settings, then register it with Orka's strict-governed `AgentRuntime` sample. The profile model and `agentConfigurationDigest` must match the baked AgentKit config. The registration's `adapterName` must be `agentkit-serve-acp`. Its `adapterDigest` and the composition build's `AGENTKIT_ADAPTER_DIGEST` must both -equal the `sha256:` digest from `AGENTKIT_RUNTIME_IMAGE`. +equal the `sha256:` digest from `AGENTKIT_RUNTIME_IMAGE`. Set the profile's +`providerKind` to `agentkit` and advertise +`supportsAgentSessionConfiguration: false`. `approvalRequiredTools` must stay +empty because the AgentKit ACP child does not implement permission callbacks. +If the registration allows brokered tools, the Task must submit that exact +`allowedTools` list. Set `ORKA_ACP_CONTROLLER_EPOCH` from Orka's current `ControllerEpoch` record: From 7ada12ef70b305509c1865081d56072771fb3e99 Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Fri, 4 Sep 2026 04:07:52 -0700 Subject: [PATCH 14/23] test(acp): cover canceled stdio write ownership Signed-off-by: Sertac Ozercan --- runtimes/common/tests/test_acp_protocol.py | 66 ++++++++++++++++++++++ 1 file changed, 66 insertions(+) diff --git a/runtimes/common/tests/test_acp_protocol.py b/runtimes/common/tests/test_acp_protocol.py index 3e5f0f7..d55d0f8 100644 --- a/runtimes/common/tests/test_acp_protocol.py +++ b/runtimes/common/tests/test_acp_protocol.py @@ -517,6 +517,72 @@ async def exercise() -> None: asyncio.run(exercise()) +def test_stdio_writer_retains_write_ownership_after_sender_cancellation(): + class BlockingWriter: + def __init__(self) -> None: + self.first_started = threading.Event() + self.second_started = threading.Event() + self.release_first = threading.Event() + self.frames: list[bytes] = [] + self.active_writes = 0 + self.max_active_writes = 0 + self.lock = threading.Lock() + + def write(self, frame: bytes) -> int: + message = json.loads(frame) + with self.lock: + self.active_writes += 1 + self.max_active_writes = max(self.max_active_writes, self.active_writes) + try: + if message["id"] == 1: + self.first_started.set() + if not self.release_first.wait(timeout=5): + raise TimeoutError("test did not release first ACP writer call") + else: + self.second_started.set() + with self.lock: + self.frames.append(bytes(frame)) + return len(frame) + finally: + with self.lock: + self.active_writes -= 1 + + def flush(self) -> None: + return None + + async def exercise() -> None: + output = BlockingWriter() + writer = acp._ACPStdioWriter(output) # noqa: SLF001 + first_send = asyncio.create_task( + writer.send({"jsonrpc": "2.0", "id": 1, "result": {}}) + ) + assert await asyncio.to_thread(output.first_started.wait, 1) + + first_send.cancel() + with pytest.raises(asyncio.CancelledError): + await first_send + + second_send = asyncio.create_task( + writer.send({"jsonrpc": "2.0", "id": 2, "result": {}}) + ) + await asyncio.sleep(0) + assert not second_send.done() + + close_task = asyncio.create_task(writer.close()) + await asyncio.sleep(0) + assert not output.second_started.is_set() + assert not close_task.done() + + output.release_first.set() + await second_send + await close_task + + assert [json.loads(frame)["id"] for frame in output.frames] == [1, 2] + assert output.max_active_writes == 1 + + asyncio.run(exercise()) + + def test_stdio_server_surfaces_writer_failure(): class FailingWriter: def write(self, frame: bytes) -> int: # noqa: ARG002 From 80d10b721791fce7c48eb74a58bb57134783eeeb Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Fri, 4 Sep 2026 04:11:26 -0700 Subject: [PATCH 15/23] fix(runtime): accept Unicode model names Signed-off-by: Sertac Ozercan --- runtimes/common/agentkit_serve_common/acp.py | 2 +- runtimes/common/tests/test_acp_protocol.py | 23 ++++++++++++++++++++ 2 files changed, 24 insertions(+), 1 deletion(-) diff --git a/runtimes/common/agentkit_serve_common/acp.py b/runtimes/common/agentkit_serve_common/acp.py index 91e49a5..4a45b04 100644 --- a/runtimes/common/agentkit_serve_common/acp.py +++ b/runtimes/common/agentkit_serve_common/acp.py @@ -299,7 +299,7 @@ def validate_acp_runtime_binding(config_bytes: bytes, spec: AgentSpec) -> None: ) expected_model = _required_environment(ACP_MODEL_ENV) - if not secrets.compare_digest(expected_model, spec.model.name): + if expected_model != spec.model.name: raise ACPConfigurationError(f"{ACP_MODEL_ENV} does not match model.name in the agent config") provider_base_url = _required_environment(ACP_PROVIDER_BASE_URL_ENV) diff --git a/runtimes/common/tests/test_acp_protocol.py b/runtimes/common/tests/test_acp_protocol.py index d55d0f8..e325cba 100644 --- a/runtimes/common/tests/test_acp_protocol.py +++ b/runtimes/common/tests/test_acp_protocol.py @@ -1405,6 +1405,29 @@ def test_runtime_binding_verifies_exact_config_digest_model_and_provider(monkeyp acp.validate_acp_runtime_binding(config_bytes + b"# changed\n", _spec()) +def test_runtime_binding_accepts_unicode_model_name(monkeypatch): + config_bytes = b"exact" + model_name = "模型" + monkeypatch.setenv( + acp.ACP_AGENT_CONFIGURATION_DIGEST_ENV, + "sha256:" + hashlib.sha256(config_bytes).hexdigest(), + ) + monkeypatch.setenv(acp.ACP_MODEL_ENV, model_name) + _set_provider_environment(monkeypatch) + + acp.validate_acp_runtime_binding( + config_bytes, + _spec( + model={ + "provider": "openai-compatible", + "baseURL": "https://baked.example.invalid/v1", + "name": model_name, + "apiKeyEnv": "BAKED_MODEL_TOKEN", + } + ), + ) + + def test_verified_runtime_binding_parses_the_hashed_bytes_after_file_replacement(monkeypatch, tmp_path): config = tmp_path / "agent.yaml" verified_bytes = b"""abiVersion: v0 From 76e3d5f07dac72433e5a79fcd1ef25c94bd91064 Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Fri, 4 Sep 2026 04:18:21 -0700 Subject: [PATCH 16/23] fix(runtime): pin Foundry responses SDK contract Signed-off-by: Sertac Ozercan --- runtimes/common/pyproject.toml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/runtimes/common/pyproject.toml b/runtimes/common/pyproject.toml index b011082..28ab44e 100644 --- a/runtimes/common/pyproject.toml +++ b/runtimes/common/pyproject.toml @@ -28,11 +28,11 @@ agentkit-foundry-brokered = "agentkit_serve_common.foundry_brokered_cli:main" agentkit-foundry-conformance = "agentkit_serve_common.foundry_conformance:main" [project.optional-dependencies] -foundry-conformance = ["azure-ai-agentserver-responses>=1.0.0b8"] +foundry-conformance = ["azure-ai-agentserver-responses>=1.0.0b8,<2"] # Starlette/FastAPI TestClient currently imports the httpx2 compatibility package. # Keep the Foundry SDK here as well because the common test suite exercises the # optional conformance app. Runtime adapter images install the base package only. -dev = ["pytest>=8.0", "httpx2>=0.1", "azure-ai-agentserver-responses>=1.0.0b8"] +dev = ["pytest>=8.0", "httpx2>=0.1", "azure-ai-agentserver-responses>=1.0.0b8,<2"] [tool.hatch.build.targets.wheel] packages = ["agentkit_serve_common"] From f8bfda4e723895f9cfdb46a115b25a6a56bc6b86 Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Fri, 4 Sep 2026 04:49:19 -0700 Subject: [PATCH 17/23] test(acp): cache recorded stdio messages Signed-off-by: Sertac Ozercan --- runtimes/common/tests/test_acp_protocol.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/runtimes/common/tests/test_acp_protocol.py b/runtimes/common/tests/test_acp_protocol.py index e325cba..aa2235a 100644 --- a/runtimes/common/tests/test_acp_protocol.py +++ b/runtimes/common/tests/test_acp_protocol.py @@ -258,6 +258,7 @@ def __init__(self, *, block_first_update: bool = False) -> None: self.blocked = threading.Event() self.release = threading.Event() self.frames: list[bytes] = [] + self.parsed_messages: list[dict[str, Any]] = [] self.lock = threading.Lock() def write(self, frame: bytes) -> int: @@ -272,6 +273,7 @@ def write(self, frame: bytes) -> int: raise TimeoutError("test did not release blocked ACP writer") with self.lock: self.frames.append(bytes(frame)) + self.parsed_messages.append(message) return len(frame) def flush(self) -> None: @@ -279,8 +281,7 @@ def flush(self) -> None: def messages(self) -> list[dict[str, Any]]: with self.lock: - frames = list(self.frames) - return [json.loads(frame) for frame in frames] + return list(self.parsed_messages) async def _wait_for_stdio_response(writer: RecordingWriter, request_id: int) -> dict[str, Any]: From 56d88e01a23b203fa97abc7a72b3d7d4ceb55d14 Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Sat, 5 Sep 2026 08:11:33 -0700 Subject: [PATCH 18/23] fix(runtime): report tool outcomes and fail closed on MCP errors Emit correlated, redacted tool lifecycle events across the LangGraph, Microsoft Agent Framework, and Pydantic AI adapters. Settle cancellation, normalize MAF provider messages, and propagate fatal MCP errors without replaying tool calls while preserving recoverable admitted errors. Validation: common and framework suites, Go tests, lint, amd64 image checks, and live Orka harness v2 conformance and lifecycle checks. Signed-off-by: Sertac Ozercan --- docs/orka.md | 14 +- runtimes/common/agentkit_serve_common/acp.py | 73 ++++- .../agentkit_serve_common/conversation.py | 19 +- runtimes/common/tests/test_acp_protocol.py | 196 ++++++++++++- .../langgraph/agentkit_serve/agent_factory.py | 48 +++- runtimes/langgraph/tests/test_tool_events.py | 168 +++++++++++ .../agentkit_serve/agent_factory.py | 144 +++++++++- .../tests/test_mcp_failures.py | 263 ++++++++++++++++++ .../tests/test_provider_messages.py | 156 +++++++++++ .../tests/test_tool_events.py | 154 ++++++++++ .../agentkit_serve/agent_factory.py | 36 ++- .../pydantic-ai/tests/test_tool_events.py | 132 +++++++++ 12 files changed, 1385 insertions(+), 18 deletions(-) create mode 100644 runtimes/langgraph/tests/test_tool_events.py create mode 100644 runtimes/microsoft-agent-framework/tests/test_mcp_failures.py create mode 100644 runtimes/microsoft-agent-framework/tests/test_provider_messages.py create mode 100644 runtimes/microsoft-agent-framework/tests/test_tool_events.py create mode 100644 runtimes/pydantic-ai/tests/test_tool_events.py diff --git a/docs/orka.md b/docs/orka.md index 78c215d..e0eeb13 100644 --- a/docs/orka.md +++ b/docs/orka.md @@ -53,11 +53,19 @@ empty because the AgentKit ACP child does not implement permission callbacks. If the registration allows brokered tools, the Task must submit that exact `allowedTools` list. -Set `ORKA_ACP_CONTROLLER_EPOCH` from Orka's current `ControllerEpoch` record: +Set `ORKA_ACP_CONTROLLER_EPOCH` from Orka's current `ControllerEpoch` record. +Select by `spec.name` because the resource name is hashed. This lookup requires +exactly one matching record with a positive integer epoch: ```sh -kubectl -n get cepoch controller-epoch-orka-controller \ - -o jsonpath='{.status.epoch}' +kubectl -n get cepoch -o json | + jq -er ' + [.items[] | select(.spec.name == "orka-controller") | .status.epoch] | + if length == 1 then + .[0] | select(type == "number") | select(. > 0 and . == floor) + else + error("expected exactly one ControllerEpoch for orka-controller") + end' ``` The supervisor reads the epoch only during startup. The operator that owns this diff --git a/runtimes/common/agentkit_serve_common/acp.py b/runtimes/common/agentkit_serve_common/acp.py index 4a45b04..1264ed3 100644 --- a/runtimes/common/agentkit_serve_common/acp.py +++ b/runtimes/common/agentkit_serve_common/acp.py @@ -14,6 +14,7 @@ import ipaddress import json import os +import re import secrets import sys from contextlib import redirect_stdout @@ -24,7 +25,7 @@ from .adapter_support import _attach_secondary_error, _wait_for_owner_task from .config import AgentSpec, load_with_bytes -from .conversation import ConversationTurn, RunRequest +from .conversation import ConversationTurn, RunRequest, ToolCallEvent from .runtime import AgentRunError, RuntimeFactory, RuntimeSession ACP_PROTOCOL_VERSION = 1 @@ -162,6 +163,61 @@ class _SessionState: cancel_requested: bool = False +class _ACPToolObserver: + """Serialize payload-free tool updates within one prompt's lifetime.""" + + def __init__(self, session_id: str, send: _MessageSender) -> None: + self.session_id = session_id + self.send = send + self.pending: dict[str, dict[str, str]] = {} + self.closed = False + self.lock = asyncio.Lock() + + async def observe(self, event: ToolCallEvent) -> None: + async with self.lock: + if self.closed: + return + if not event.tool_call_id or event.status not in {"in_progress", "completed", "failed"}: + raise AgentRunError("runtime emitted an invalid tool lifecycle event") + if event.status == "in_progress": + if event.tool_call_id in self.pending: + raise AgentRunError("runtime repeated an active tool call ID") + # Never forward provider call IDs, arguments, results, or exception + # text. A new ACP ID also avoids collisions across prompt turns. + title = "Tool call" + if re.fullmatch(r"[A-Za-z0-9_.:-]{1,128}", event.tool_name): + title = event.tool_name + tool = {"toolCallId": "tool-" + secrets.token_hex(16), "title": title, "kind": "other"} + self.pending[event.tool_call_id] = tool + update = {"sessionUpdate": "tool_call", **tool, "status": "in_progress"} + else: + tool = self.pending.pop(event.tool_call_id, None) + if tool is None: + raise AgentRunError("runtime finished an unknown tool call") + # Remove before awaiting output, since a cancelled write remains + # owned by the stdio writer and must not get a second terminal. + update = {"sessionUpdate": "tool_call_update", **tool, "status": event.status} + await self._send(update) + + async def close(self) -> bool: + async with self.lock: + self.closed = True + unfinished = bool(self.pending) + while self.pending: + _, tool = self.pending.popitem() + await self._send({"sessionUpdate": "tool_call_update", **tool, "status": "failed"}) + return unfinished + + async def _send(self, update: dict[str, str]) -> None: + await self.send( + { + "jsonrpc": _JSONRPC_VERSION, + "method": _METHOD_SESSION_UPDATE, + "params": {"sessionId": self.session_id, "update": update}, + } + ) + + def _factory_supports_http_mcp(factory: RuntimeFactory) -> bool: capability = getattr(factory, "supports_acp_http_mcp", None) return bool(capability()) if callable(capability) else False @@ -703,10 +759,12 @@ async def _prompt(self, request_key: str, params: Any) -> dict[str, str]: "ACP mode accepts only text and resource_link prompt blocks", ) + tool_observer = _ACPToolObserver(session_id, self.send) run_request = RunRequest( prompt="\n".join(text_blocks), history=tuple(state.history), session_id=session_id, + on_tool_event=tool_observer.observe, ) state.cancel_requested = False state.active_request_key = request_key @@ -727,6 +785,18 @@ async def _prompt(self, request_key: str, params: Any) -> dict[str, str]: result = None runtime_error = exc + try: + unfinished_tools = await tool_observer.close() + if unfinished_tools and runtime_error is None and not cancelled: + runtime_error = AgentRunError("runtime returned before its tool calls settled") + except asyncio.CancelledError: + cancelled = True + except BaseException as exc: # noqa: BLE001 - retain the original runtime failure. + if runtime_error is None: + runtime_error = exc + else: + _attach_secondary_error(runtime_error, exc, label="tool event completion failed") + if ( cancelled or state.cancel_requested @@ -767,6 +837,7 @@ async def _prompt(self, request_key: str, params: Any) -> dict[str, str]: state.history.append(ConversationTurn(role="assistant", text=result.text)) return {"stopReason": "end_turn"} finally: + tool_observer.closed = True self.prompt_requests.pop(request_key, None) state.active_request_key = None state.active_run = None diff --git a/runtimes/common/agentkit_serve_common/conversation.py b/runtimes/common/agentkit_serve_common/conversation.py index f67a803..1c0bd0d 100644 --- a/runtimes/common/agentkit_serve_common/conversation.py +++ b/runtimes/common/agentkit_serve_common/conversation.py @@ -10,7 +10,7 @@ from dataclasses import dataclass, field from datetime import datetime -from typing import Any, Mapping, Protocol, Sequence +from typing import Any, Awaitable, Callable, Literal, Mapping, Protocol, Sequence FORWARDED_ROLES = frozenset({"system", "user", "assistant"}) @@ -27,6 +27,20 @@ class ConversationTurn: text: str +@dataclass(frozen=True) +class ToolCallEvent: + """One tool lifecycle observation, without arguments, results, or errors. + + Adapters await the observer before continuing the run. A call starts with + ``in_progress`` and ends with ``completed`` or ``failed`` under the same ID. + The ID is adapter-local and must not be used as a host authorization token. + """ + + tool_call_id: str + tool_name: str + status: Literal["in_progress", "completed", "failed"] + + @dataclass(frozen=True) class RunRequest: """Framework-neutral request for one agent run. @@ -50,6 +64,9 @@ class RunRequest: turn_id: str | None = None correlation_id: str | None = None metadata: Mapping[str, str] = field(default_factory=dict) + on_tool_event: Callable[[ToolCallEvent], Awaitable[None]] | None = field( + default=None, repr=False, compare=False + ) class OpenAIMessage(Protocol): diff --git a/runtimes/common/tests/test_acp_protocol.py b/runtimes/common/tests/test_acp_protocol.py index aa2235a..50fc270 100644 --- a/runtimes/common/tests/test_acp_protocol.py +++ b/runtimes/common/tests/test_acp_protocol.py @@ -17,7 +17,7 @@ from agentkit_serve_common import acp from agentkit_serve_common.acp import ACPConfigurationError, ACPStdioServer from agentkit_serve_common.config import AgentSpec -from agentkit_serve_common.conversation import ConversationTurn, RunRequest +from agentkit_serve_common.conversation import ConversationTurn, RunRequest, ToolCallEvent from agentkit_serve_common.runtime import ( AgentRunError, OfflineEchoRuntimeFactory, @@ -1144,6 +1144,200 @@ async def send(message): # noqa: ANN001 asyncio.run(exercise()) +def test_tool_events_precede_output_and_keep_ids_private_and_unique(monkeypatch): + _set_provider_environment(monkeypatch) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + + async def send(message): # noqa: ANN001 + messages.append(dict(message)) + + class ToolRuntime(RecordingRuntime): + async def run(self, request: RunRequest) -> RunResult: + self.requests.append(request) + assert request.on_tool_event is not None + for name in ["orka_echo", "orka_recover"]: + await request.on_tool_event(ToolCallEvent(name + "-private-id", name, "in_progress")) + assert messages[-1]["params"]["update"]["title"] == name + # Parallel calls may finish in a different order than they start. + await request.on_tool_event(ToolCallEvent("orka_recover-private-id", "orka_recover", "failed")) + await request.on_tool_event(ToolCallEvent("orka_echo-private-id", "orka_echo", "completed")) + return RunResult("héllo 世界 🌍") + + runtime = ToolRuntime([]) + server = ACPStdioServer(_spec(), RecordingFactory(lambda: runtime), send) + await _send_to(server, messages, _initialize()) + created = await _send_to(server, messages, _new_session()) + session_id = created["result"]["sessionId"] + ids: set[str] = set() + for request_id in [3, 4]: + offset = len(messages) + completed = await _send_to(server, messages, _prompt(request_id, session_id, "tools")) + assert completed["result"] == {"stopReason": "end_turn"} + updates = [message["params"]["update"] for message in messages[offset:-1]] + assert [update["sessionUpdate"] for update in updates] == [ + "tool_call", "tool_call", "tool_call_update", "tool_call_update", "agent_message_chunk" + ] + assert [update["status"] for update in updates[:4]] == [ + "in_progress", "in_progress", "failed", "completed" + ] + assert updates[0]["toolCallId"] == updates[3]["toolCallId"] + assert updates[1]["toolCallId"] == updates[2]["toolCallId"] + current_ids = {updates[0]["toolCallId"], updates[1]["toolCallId"]} + assert len(current_ids) == 2 and ids.isdisjoint(current_ids) + ids.update(current_ids) + for update in updates[:4]: + assert update["kind"] == "other" + assert set(update) == {"sessionUpdate", "toolCallId", "title", "kind", "status"} + assert updates[-1]["content"]["text"] == "héllo 世界 🌍" + assert "private-id" not in json.dumps(messages) + # A retained adapter callback cannot append events after its prompt settles. + offset = len(messages) + await runtime.requests[0].on_tool_event(ToolCallEvent("late", "late", "in_progress")) + assert len(messages) == offset + await server.close() + + asyncio.run(exercise()) + + +@pytest.mark.parametrize("outcome", ["error", "cancel", "incomplete"]) +def test_unsettled_tool_is_failed_before_prompt_ends_and_history_is_discarded(monkeypatch, outcome): + _set_provider_environment(monkeypatch) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + started = asyncio.Event() + + async def send(message): # noqa: ANN001 + messages.append(dict(message)) + + class ToolRuntime(RecordingRuntime): + async def run(self, request: RunRequest) -> RunResult: + self.requests.append(request) + if len(self.requests) > 1: + return RunResult("recovered") + assert request.on_tool_event is not None + await request.on_tool_event(ToolCallEvent("private-call-id", "https://secret.invalid/token", "in_progress")) + started.set() + if outcome == "error": + raise RuntimeError("private tool credentials and result") + if outcome == "cancel": + await asyncio.Future() + return RunResult("must not commit incomplete tools") + + runtime = ToolRuntime([]) + server = ACPStdioServer(_spec(), RecordingFactory(lambda: runtime), send) + await _send_to(server, messages, _initialize()) + created = await _send_to(server, messages, _new_session()) + session_id = created["result"]["sessionId"] + await server.accept(_prompt(3, session_id, "unsettled tool")) + await asyncio.wait_for(started.wait(), timeout=1) + if outcome == "cancel": + await server.accept({"jsonrpc": "2.0", "method": "session/cancel", "params": {"sessionId": session_id}}) + await server.wait_idle() + response = _response(messages, 3) + if outcome == "cancel": + assert response["result"] == {"stopReason": "cancelled"} + else: + assert response["error"]["code"] == -32603 + assert messages[-1] == response + updates = [message["params"]["update"] for message in messages if message.get("method") == "session/update"] + assert [update["status"] for update in updates] == ["in_progress", "failed"] + assert updates[0]["toolCallId"] == updates[1]["toolCallId"] + assert all(update["title"] == "Tool call" for update in updates) + assert "private" not in json.dumps(messages) and "token" not in json.dumps(messages) + assert server.sessions[session_id].history == [] + assert runtime.discarded_sessions == [session_id] + completed = await _send_to(server, messages, _prompt(4, session_id, "recover")) + assert completed["result"] == {"stopReason": "end_turn"} + assert runtime.requests[1].history == () + await server.close() + + asyncio.run(exercise()) + + +@pytest.mark.parametrize("fail_status", ["in_progress", "completed"]) +def test_tool_event_write_failure_discards_prompt_state(monkeypatch, fail_status): + _set_provider_environment(monkeypatch) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + failed_once = False + + async def send(message): # noqa: ANN001 + nonlocal failed_once + update = message.get("params", {}).get("update", {}) + if update.get("status") == fail_status and not failed_once: + failed_once = True + raise RuntimeError("private transport failure") + messages.append(dict(message)) + + class ToolRuntime(RecordingRuntime): + async def run(self, request: RunRequest) -> RunResult: + self.requests.append(request) + assert request.on_tool_event is not None + await request.on_tool_event(ToolCallEvent("call", "echo", "in_progress")) + await request.on_tool_event(ToolCallEvent("call", "echo", "completed")) + return RunResult("done") + + runtime = ToolRuntime([]) + server = ACPStdioServer(_spec(), RecordingFactory(lambda: runtime), send) + await _send_to(server, messages, _initialize()) + created = await _send_to(server, messages, _new_session()) + session_id = created["result"]["sessionId"] + response = await _send_to(server, messages, _prompt(3, session_id, "tools")) + assert failed_once and response["error"]["code"] == -32603 + assert server.sessions[session_id].history == [] + assert runtime.discarded_sessions == [session_id] + updates = [message["params"]["update"] for message in messages if message.get("method") == "session/update"] + assert all(update["sessionUpdate"] != "agent_message_chunk" for update in updates) + assert "private transport failure" not in json.dumps(messages) + await server.close() + + asyncio.run(exercise()) + + +def test_cancellation_during_terminal_tool_write_does_not_emit_conflicting_terminal(monkeypatch): + _set_provider_environment(monkeypatch) + + async def exercise() -> None: + messages: list[dict[str, Any]] = [] + terminal_started = asyncio.Event() + + async def send(message): # noqa: ANN001 + messages.append(dict(message)) + if message.get("params", {}).get("update", {}).get("status") == "completed": + terminal_started.set() + await asyncio.Future() + + class ToolRuntime(RecordingRuntime): + async def run(self, request: RunRequest) -> RunResult: + self.requests.append(request) + assert request.on_tool_event is not None + await request.on_tool_event(ToolCallEvent("call", "echo", "in_progress")) + await request.on_tool_event(ToolCallEvent("call", "echo", "completed")) + return RunResult("must not commit") + + runtime = ToolRuntime([]) + server = ACPStdioServer(_spec(), RecordingFactory(lambda: runtime), send) + await _send_to(server, messages, _initialize()) + created = await _send_to(server, messages, _new_session()) + session_id = created["result"]["sessionId"] + await server.accept(_prompt(3, session_id, "tools")) + await asyncio.wait_for(terminal_started.wait(), timeout=1) + await server.accept({"jsonrpc": "2.0", "method": "session/cancel", "params": {"sessionId": session_id}}) + await asyncio.wait_for(server.wait_idle(), timeout=1) + assert _response(messages, 3)["result"] == {"stopReason": "cancelled"} + updates = [message["params"]["update"] for message in messages if message.get("method") == "session/update"] + assert [update["status"] for update in updates] == ["in_progress", "completed"] + assert server.sessions[session_id].history == [] + assert runtime.discarded_sessions == [session_id] + await server.close() + + asyncio.run(exercise()) + + def test_multibyte_output_is_split_without_partial_commit(monkeypatch): _set_provider_environment(monkeypatch) result_text = "a" + "🙂" * (acp._MAX_ASSISTANT_MESSAGE_CHUNK_BYTES // 4 + 1) # noqa: SLF001 diff --git a/runtimes/langgraph/agentkit_serve/agent_factory.py b/runtimes/langgraph/agentkit_serve/agent_factory.py index 019b434..d7287bf 100644 --- a/runtimes/langgraph/agentkit_serve/agent_factory.py +++ b/runtimes/langgraph/agentkit_serve/agent_factory.py @@ -26,13 +26,16 @@ from __future__ import annotations import asyncio +from collections.abc import Awaitable, Callable from contextlib import AsyncExitStack from datetime import timedelta from types import TracebackType from typing import Any +from uuid import UUID from langchain.agents import create_agent -from langchain_core.messages import AIMessage, BaseMessage, HumanMessage, SystemMessage +from langchain_core.callbacks import AsyncCallbackHandler +from langchain_core.messages import AIMessage, BaseMessage, HumanMessage, SystemMessage, ToolMessage from langchain_mcp_adapters.client import MultiServerMCPClient from langchain_mcp_adapters.tools import load_mcp_tools from langchain_openai import ChatOpenAI @@ -51,7 +54,7 @@ upstream_status_code, ) from agentkit_serve_common.config import AgentSpec, ToolSpec -from agentkit_serve_common.conversation import RunRequest +from agentkit_serve_common.conversation import RunRequest, ToolCallEvent from agentkit_serve_common.runtime import ( AgentRunError, OfflineEchoRuntimeFactory, @@ -339,13 +342,52 @@ def _state_usage(state: Any) -> dict[str, int]: return totals +class _ToolEventCallback(AsyncCallbackHandler): + """Observe real tool execution without forwarding its payload or errors.""" + + raise_error = True + run_inline = True + + def __init__(self, observe: Callable[[ToolCallEvent], Awaitable[None]]) -> None: + self.observe = observe + self.names: dict[UUID, str] = {} + + async def _emit(self, event: ToolCallEvent) -> None: + try: + await self.observe(event) + except Exception: + # LangChain logs callback exceptions even when raise_error is set. + raise AgentRunError("tool lifecycle observer failed") from None + + async def on_tool_start( + self, serialized: dict[str, Any], input_str: str, *, run_id: UUID, **kwargs: Any + ) -> None: + name = serialized.get("name") or "Tool call" + self.names[run_id] = name + await self._emit(ToolCallEvent(str(run_id), name, "in_progress")) + + async def on_tool_end(self, output: Any, *, run_id: UUID, **kwargs: Any) -> None: + name = self.names.pop(run_id) + # Handled MCP failures arrive as ToolMessage(status="error"), without + # invoking on_tool_error, so a completed callback is not always success. + failed = isinstance(output, ToolMessage) and output.status == "error" + await self._emit(ToolCallEvent(str(run_id), name, "failed" if failed else "completed")) + + async def on_tool_error(self, error: BaseException, *, run_id: UUID, **kwargs: Any) -> None: + name = self.names.pop(run_id) + await self._emit(ToolCallEvent(str(run_id), name, "failed")) + + async def run_agent(agent: LangGraphRuntime, request: RunRequest) -> RunResult: """Run the compiled LangGraph once and return a neutral ``RunResult``.""" if agent.graph is None: raise AgentRunError("agent graph is not initialized", status=500, code="AgentNotInitialized") + kwargs = {} + if request.on_tool_event is not None: + kwargs["config"] = {"callbacks": [_ToolEventCallback(request.on_tool_event)]} try: - state = await agent.graph.ainvoke({"messages": _to_messages(request)}) + state = await agent.graph.ainvoke({"messages": _to_messages(request)}, **kwargs) except AgentRunError: raise except Exception as exc: # noqa: BLE001 — normalized for the façade diff --git a/runtimes/langgraph/tests/test_tool_events.py b/runtimes/langgraph/tests/test_tool_events.py new file mode 100644 index 0000000..1a21e8d --- /dev/null +++ b/runtimes/langgraph/tests/test_tool_events.py @@ -0,0 +1,168 @@ +from __future__ import annotations + +import asyncio +from dataclasses import asdict +from types import SimpleNamespace + +import pytest +from langchain.agents import create_agent +from langchain_core.language_models.fake_chat_models import GenericFakeChatModel +from langchain_core.messages import AIMessage +from langchain_core.tools import StructuredTool, ToolException + +from agentkit_serve import agent_factory +from agentkit_serve_common.conversation import RunRequest, ToolCallEvent +from agentkit_serve_common.runtime import AgentRunError + + +class _ToolModel(GenericFakeChatModel): + def bind_tools(self, tools, **kwargs): + return self + + +def _runtime(probe, *, calls=1, handle_tool_error=True): + messages = [ + AIMessage(content="", tool_calls=[{"name": "probe", "args": {}, "id": f"private-provider-id-{i}"}]) + for i in range(calls) + ] + messages.append(AIMessage(content="done")) + tool = StructuredTool.from_function( + coroutine=probe, name="probe", description="Controlled fixture.", handle_tool_error=handle_tool_error + ) + return SimpleNamespace(graph=create_agent(model=_ToolModel(messages=iter(messages)), tools=[tool])) + + +def test_real_tool_execution_awaits_redacted_start_and_result(): + async def exercise(): + events: list[ToolCallEvent] = [] + order = [] + + async def probe(): + order.append("side-effect") + return {"private-payload": ["héllo 世界 🌍", {"nested": True}]} + + async def observe(event): + await asyncio.sleep(0) + events.append(event) + order.append(event.status) + + result = await agent_factory.run_agent(_runtime(probe), RunRequest("probe", on_tool_event=observe)) + assert result.text == "done" + assert order == ["in_progress", "side-effect", "completed"] + assert events[0].tool_call_id == events[1].tool_call_id + assert [event.tool_name for event in events] == ["probe", "probe"] + assert all(set(asdict(event)) == {"tool_call_id", "tool_name", "status"} for event in events) + assert "private" not in str(events) + + asyncio.run(exercise()) + + +def test_handled_tool_error_emits_failure_before_successful_retry(): + async def exercise(): + events = [] + calls = 0 + + async def probe(): + nonlocal calls + calls += 1 + if calls == 1: + raise ToolException("private error details") + return "private successful result" + + async def observe(event): + events.append(event) + + result = await agent_factory.run_agent(_runtime(probe, calls=2), RunRequest("recover", on_tool_event=observe)) + assert result.text == "done" and calls == 2 + assert [event.status for event in events] == ["in_progress", "failed", "in_progress", "completed"] + assert events[0].tool_call_id == events[1].tool_call_id + assert events[2].tool_call_id == events[3].tool_call_id + assert events[0].tool_call_id != events[2].tool_call_id + assert "private" not in str(events) + + asyncio.run(exercise()) + + +def test_unhandled_tool_error_emits_failure_and_propagates(): + async def exercise(): + events = [] + + async def probe(): + raise RuntimeError("private error details") + + async def observe(event): + events.append(event) + + with pytest.raises(AgentRunError): + await agent_factory.run_agent(_runtime(probe), RunRequest("fail", on_tool_event=observe)) + assert [event.status for event in events] == ["in_progress", "failed"] + assert events[0].tool_call_id == events[1].tool_call_id + assert "private" not in str(events) + + asyncio.run(exercise()) + + +def test_cancellation_stops_real_tool_and_leaves_terminal_to_acp(): + async def exercise(): + events = [] + started = asyncio.Event() + stopped = asyncio.Event() + + async def probe(): + started.set() + try: + await asyncio.Future() + finally: + stopped.set() + + async def observe(event): + events.append(event) + + run = asyncio.create_task(agent_factory.run_agent(_runtime(probe), RunRequest("delay", on_tool_event=observe))) + try: + await asyncio.wait_for(started.wait(), 3) + finally: + run.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(run, 3) + assert stopped.is_set() + assert [event.status for event in events] == ["in_progress"] + + asyncio.run(exercise()) + + +@pytest.mark.parametrize("failure_status", ["in_progress", "completed"]) +def test_observer_failure_aborts_run_without_exposing_its_error(failure_status, caplog): + async def exercise(): + calls = [] + + async def probe(): + calls.append("called") + return "private payload" + + async def observe(event): + if event.status == failure_status: + raise RuntimeError("private observer failure") + + with pytest.raises(AgentRunError, match="tool lifecycle observer failed"): + await agent_factory.run_agent(_runtime(probe), RunRequest("probe", on_tool_event=observe)) + assert calls == ([] if failure_status == "in_progress" else ["called"]) + assert "private observer failure" not in caplog.text + + asyncio.run(exercise()) + + +def test_no_observer_preserves_plain_ainvoke_contract(): + async def exercise(): + calls = [] + + class PlainGraph: + async def ainvoke(self, state): + calls.append(state) + return {"messages": [AIMessage(content="plain")]} + + result = await agent_factory.run_agent(SimpleNamespace(graph=PlainGraph()), RunRequest("plain prompt")) + assert result.text == "plain" and len(calls) == 1 + assert calls[0]["messages"][-1].content == "plain prompt" + + asyncio.run(exercise()) diff --git a/runtimes/microsoft-agent-framework/agentkit_serve/agent_factory.py b/runtimes/microsoft-agent-framework/agentkit_serve/agent_factory.py index 867bca0..19fb241 100644 --- a/runtimes/microsoft-agent-framework/agentkit_serve/agent_factory.py +++ b/runtimes/microsoft-agent-framework/agentkit_serve/agent_factory.py @@ -8,27 +8,37 @@ from __future__ import annotations import asyncio +import copy import importlib import inspect import os import time +from collections.abc import Awaitable, Callable from contextlib import AbstractAsyncContextManager, AsyncExitStack +from contextvars import ContextVar from datetime import timedelta from types import TracebackType from urllib.parse import urlsplit +from uuid import uuid4 from agent_framework import ( Agent, AgentSession, + ChatContext, + ChatMiddleware, FileSkillsSource, + FunctionInvocationContext, + FunctionMiddleware, MCPSkillsSource, MCPStdioTool, MCPStreamableHTTPTool, Message, + MiddlewareTermination, SkillsProvider, ) from agent_framework.openai import OpenAIChatCompletionClient from httpx import AsyncClient, URL +from mcp.types import CallToolResult from agentkit_serve_common.adapter_support import ( FORWARDED_ROLES, AsyncExitStackLifecycle, @@ -45,8 +55,9 @@ upstream_status_code, ) from agentkit_serve_common.config import AgentSpec, ContextProviderSpec, ToolSpec -from agentkit_serve_common.conversation import RunRequest +from agentkit_serve_common.conversation import RunRequest, ToolCallEvent from agentkit_serve_common.runtime import ( + AgentRunError, OfflineEchoRuntimeFactory, RunResult, RuntimeSession, @@ -63,6 +74,9 @@ _DEFAULT_FOUNDRY_AUDIENCE = "https://ai.azure.com/.default" _DEFAULT_MCP_REQUEST_TIMEOUT = 120 _DEFAULT_SESSION_CACHE_MAX = 256 +# Invocation tasks inherit the same mutable marker for this run. Separate runs, +# including concurrent Sessions on one Agent, receive independent markers. +_mcp_failures: ContextVar[list[bool] | None] = ContextVar("agentkit_maf_mcp_failures", default=None) def _mcp_request_timeout() -> int: @@ -260,6 +274,49 @@ def _tool_env(tool: ToolSpec) -> dict[str, str]: return declared_tool_env(tool) +class _MCPToolError(Exception): + """A validated MCP result reports an admitted tool execution failure.""" + + +class _MCPProtocolError(Exception): + """An MCP call failed without an admitted tool error result.""" + + +class _MCPCallBoundary: + """Keep MCP protocol failures out of MAF's recoverable tool-error loop.""" + + async def call_tool(self, tool_name, **kwargs): + try: + return await super().call_tool(tool_name, **kwargs) + except _MCPToolError: + raise + except Exception: + # Do not attach upstream exceptions: framework tool logging and + # detailed-error options must never expose transport credentials. + raise _MCPProtocolError("MCP tool protocol failed") from None + + async def _call_tool_with_retries(self, tool_name, filtered_kwargs, meta, parser, span): + # MAF 1.9+ retries tools/call after a lost connection. An attempted call + # may already have executed, so leave retry decisions to the caller. + result = await self.session.call_tool(tool_name, arguments=filtered_kwargs, meta=meta) + return parser(result) + + def _parse_tool_result_from_mcp(self, result): + if not isinstance(result, CallToolResult): + raise _MCPProtocolError("MCP tool protocol failed") + if result.isError: + raise _MCPToolError("MCP tool execution failed") + return super()._parse_tool_result_from_mcp(result) + + +class _MCPStreamableHTTPTool(_MCPCallBoundary, MCPStreamableHTTPTool): + pass + + +class _MCPStdioTool(_MCPCallBoundary, MCPStdioTool): + pass + + def build_tool(tool: ToolSpec, *, stack: AsyncExitStack | None = None): """Create a stdio or Streamable HTTP MCP server for one tool spec.""" timeout = _mcp_request_timeout() @@ -294,7 +351,7 @@ async def inject_headers(request): # noqa: ANN001 "http_client": http_client, } kwargs["request_timeout"] = int(remote_timeout) - mcp_tool = MCPStreamableHTTPTool(**kwargs) + mcp_tool = _MCPStreamableHTTPTool(**kwargs) # Some Streamable HTTP MCP services, including Foundry Toolbox, do not # implement MCP ping. The framework handles request-time connection # errors separately, so skip proactive pings for remote HTTP tools. @@ -310,7 +367,36 @@ async def inject_headers(request): # noqa: ANN001 "tool_name_prefix": tool.name, } kwargs["request_timeout"] = timeout - return MCPStdioTool(**kwargs) + return _MCPStdioTool(**kwargs) + + +class _MCPFailureMiddleware(FunctionMiddleware): + async def process( + self, context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]] + ) -> None: + failures = _mcp_failures.get() + if failures: + raise MiddlewareTermination("MCP tool protocol failed") + try: + await call_next() + except _MCPProtocolError: + if failures is not None: + failures.append(True) + # MiddlewareTermination stops the model loop even on MAF 1.9, + # which predates MiddlewareFailure. run_agent makes it a failure. + raise MiddlewareTermination("MCP tool protocol failed") from None + + +class _ModelMessageMiddleware(ChatMiddleware): + """Keep package and framework author labels out of model speaker fields.""" + + async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None: + # AgentKit's inputs do not define named speakers. MAF adds its agent name + # to history, but chat-to-Responses gateways cannot translate that field. + context.messages = [copy.copy(message) for message in context.messages] + for message in context.messages: + message.author_name = None + await call_next() def build_agent( @@ -327,6 +413,7 @@ def build_agent( name=spec.metadata.name, tools=[build_tool(t, stack=stack) for t in spec.tools], context_providers=context_providers, + middleware=[_ModelMessageMiddleware(), _MCPFailureMiddleware()], ) @@ -687,6 +774,37 @@ def _to_messages(request: RunRequest, *, include_history: bool = True) -> list[M return messages +class _ToolEventMiddleware(FunctionMiddleware): + """Await payload-free observations around each actual tool invocation.""" + + def __init__(self, observe: Callable[[ToolCallEvent], Awaitable[None]]) -> None: + self.observe = observe + self.failed = False + + async def _emit(self, event: ToolCallEvent) -> None: + try: + await self.observe(event) + except Exception: + # MAF absorbs ordinary function exceptions and continues the model + # loop. Stop it, then fail run_agent, even on supported versions + # predating MiddlewareFailure. + self.failed = True + raise MiddlewareTermination("tool lifecycle observer failed") from None + + async def process( + self, context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]] + ) -> None: + call_id = uuid4().hex + name = context.function.name + await self._emit(ToolCallEvent(call_id, name, "in_progress")) + try: + await call_next() + except Exception: + await self._emit(ToolCallEvent(call_id, name, "failed")) + raise + await self._emit(ToolCallEvent(call_id, name, "completed")) + + async def run_agent( agent: Agent, request: RunRequest, @@ -696,8 +814,22 @@ async def run_agent( ) -> RunResult: """Run the MAF agent and return the neutral result shape.""" messages = _to_messages(request, include_history=include_history) + kwargs = {} + observer = None + if request.on_tool_event is not None: + observer = _ToolEventMiddleware(request.on_tool_event) + kwargs["middleware"] = [observer] + failures: list[bool] = [] + token = _mcp_failures.set(failures) try: - result = await agent.run(messages, session=session) - except Exception as exc: # noqa: BLE001 — normalized for the façade - raise normalize_agent_run_error(exc) from exc + try: + result = await agent.run(messages, session=session, **kwargs) + except Exception as exc: # noqa: BLE001 — normalized for the façade + raise normalize_agent_run_error(exc) from exc + finally: + _mcp_failures.reset(token) + if failures: + raise AgentRunError("MCP tool protocol failed") + if observer is not None and observer.failed: + raise AgentRunError("tool lifecycle observer failed") return RunResult(text=_result_text(result), usage=_result_usage(result)) diff --git a/runtimes/microsoft-agent-framework/tests/test_mcp_failures.py b/runtimes/microsoft-agent-framework/tests/test_mcp_failures.py new file mode 100644 index 0000000..4d44a5a --- /dev/null +++ b/runtimes/microsoft-agent-framework/tests/test_mcp_failures.py @@ -0,0 +1,263 @@ +from __future__ import annotations + +import asyncio +import json +import logging +from contextlib import AsyncExitStack +from dataclasses import asdict +from unittest import mock + +import httpx +import pytest +from agent_framework import AgentSession, BaseChatClient, ChatResponse, Content, FunctionInvocationLayer, Message +from anyio import ClosedResourceError +from mcp import types +from mcp.shared.exceptions import McpError + +from agentkit_serve import agent_factory +from agentkit_serve_common.config import AgentSpec, ToolSpec +from agentkit_serve_common.conversation import RunRequest +from agentkit_serve_common.runtime import AgentRunError + + +_PRIVATE = "PRIVATE_MCP_DIAGNOSTIC_91f42" +_PAYLOAD = {"text": "héllo 世界 🌍", "nested": [42, True, None]} + + +class _Session: + def __init__(self, outcomes): + self.outcomes = iter(outcomes) + self.calls = [] + + async def list_tools(self, **kwargs): + return types.ListToolsResult(tools=[types.Tool( + name="probe", + inputSchema={"type": "object", "properties": {"payload": {"type": "object"}}}, + _meta={"fixture.example/authority": "frozen-metadata"}, + )]) + + async def call_tool(self, name, *, arguments, meta): + self.calls.append({"name": name, "arguments": arguments, "meta": meta}) + outcome = next(self.outcomes) + if isinstance(outcome, BaseException): + raise outcome + if callable(outcome): + return await outcome() + return outcome + + +class _Client(FunctionInvocationLayer, BaseChatClient): + def __init__(self, function_name, *, calls=1): + super().__init__() + self.function_name = function_name + self.tool_calls = calls + self.requests = {} + + async def _inner_get_response(self, *, messages, stream, options, **kwargs): + prompt = next(message.text for message in reversed(messages) if message.role == "user") + requests = self.requests.setdefault(prompt, []) + requests.append(messages) + number = len(requests) + if number <= self.tool_calls: + item = Content.from_function_call( + call_id=f"private-provider-id-{prompt}-{number}", name=self.function_name, + arguments={"payload": {**_PAYLOAD, "prompt": prompt}, "undeclared": "must-be-filtered"}, + ) + else: + item = Content.from_text("done") + return ChatResponse(messages=[Message(role="assistant", contents=[item])]) + + +def _success(): + return types.CallToolResult(content=[types.TextContent(type="text", text=json.dumps(_PAYLOAD))]) + + +def _failure(kind): + if kind == "jsonrpc": + return McpError(types.ErrorData(code=-32000, message=_PRIVATE)) + if kind == "closed": + return ClosedResourceError(_PRIVATE) + if kind == "terminated": + return McpError(types.ErrorData(code=-32000, message="session terminated: " + _PRIVATE)) + if kind == "auth": + return httpx.HTTPStatusError(_PRIVATE, request=httpx.Request("POST", "https://example.invalid"), + response=httpx.Response(403)) + if kind == "timeout": + return httpx.ReadTimeout(_PRIVATE) + if kind == "malformed": + try: + types.CallToolResult.model_validate({"content": _PRIVATE}) + except ValueError as exc: + return exc + raise AssertionError("unknown controlled failure") + + +async def _setup(stack, monkeypatch, outcomes, *, transport="streamable-http", calls=1): + monkeypatch.setenv("TEST_MCP_URL", "http://example.invalid/mcp") + tool = ( + ToolSpec(name="fixture", type="mcp", url_env="TEST_MCP_URL", transport=transport) + if transport == "streamable-http" + else ToolSpec(name="fixture", command=["unused-fixture"]) + ) + server = agent_factory.build_tool(tool, stack=stack) + server.session = session = _Session(outcomes) + server._supports_tools = True + server._ping_available = False + server.connect = mock.AsyncMock() + await server.load_tools() + client = _Client(server.functions[0].name, calls=calls) + spec = AgentSpec.model_validate({ + "abiVersion": "v0", "metadata": {"name": "test-package"}, + "model": {"provider": "openai-compatible", "name": "test-model", "baseURL": "http://model.invalid/v1"}, + "instructions": "Follow the user's request.", "tools": [tool.model_dump()], + "expose": {"openai": True, "port": 8080}, + }) + # Keep the real SDK-generated MCP callable, but supply its already-loaded + # function to avoid opening a network transport for the in-memory session. + with mock.patch.object(agent_factory, "build_tool", return_value=server.functions[0]): + agent = await stack.enter_async_context(agent_factory.build_agent(spec, client=client)) + return agent, server, session, client + + +@pytest.mark.parametrize("transport", ["streamable-http", "stdio"]) +@pytest.mark.parametrize("kind", ["jsonrpc", "closed", "terminated", "auth", "timeout", "malformed"]) +@pytest.mark.parametrize("observed", [True, False]) +def test_mcp_protocol_failure_is_fatal_without_retry_or_model_continuation( + monkeypatch, caplog, transport, kind, observed +): + async def exercise(): + events = [] + + async def observe(event): + events.append(event) + + async with AsyncExitStack() as stack: + agent, server, session, client = await _setup(stack, monkeypatch, [_failure(kind), _success()], + transport=transport) + with pytest.raises(AgentRunError, match="^MCP tool protocol failed$") as caught: + await agent_factory.run_agent(agent, RunRequest("fatal", on_tool_event=observe if observed else None)) + assert len(session.calls) == 1 and len(client.requests["fatal"]) == 1 + assert server.connect.await_count == 0 + assert session.calls[0] == { + "name": "probe", "arguments": {"payload": {**_PAYLOAD, "prompt": "fatal"}}, + "meta": {"fixture.example/authority": "frozen-metadata"}, + } + assert caught.value.__cause__ is None + assert [event.status for event in events] == (["in_progress", "failed"] if observed else []) + if observed: + assert events[0].tool_call_id == events[1].tool_call_id + assert all(set(asdict(event)) == {"tool_call_id", "tool_name", "status"} for event in events) + assert "private" not in str(events) + assert _PRIVATE not in caplog.text + + caplog.set_level(logging.DEBUG, logger="agent_framework") + asyncio.run(exercise()) + + +@pytest.mark.parametrize("transport", ["streamable-http", "stdio"]) +def test_admitted_mcp_error_can_recover_with_correlated_events(monkeypatch, caplog, transport): + async def exercise(): + events = [] + + async def observe(event): + events.append(event) + + admitted = types.CallToolResult(content=[types.TextContent(type="text", text=_PRIVATE)], isError=True) + async with AsyncExitStack() as stack: + agent, server, session, client = await _setup(stack, monkeypatch, [admitted, _success()], + transport=transport, calls=2) + result = await agent_factory.run_agent(agent, RunRequest("recover", on_tool_event=observe)) + assert result.text == "done" + assert len(session.calls) == 2 and len(client.requests["recover"]) == 3 + assert server.connect.await_count == 0 + assert [event.status for event in events] == ["in_progress", "failed", "in_progress", "completed"] + assert events[0].tool_call_id == events[1].tool_call_id + assert events[2].tool_call_id == events[3].tool_call_id != events[0].tool_call_id + assert _PRIVATE not in caplog.text + + caplog.set_level(logging.DEBUG, logger="agent_framework") + asyncio.run(exercise()) + + +def test_mcp_cancellation_propagates_without_retry_or_failure_event(monkeypatch): + async def exercise(): + started, stopped = asyncio.Event(), asyncio.Event() + events = [] + + async def pending(): + started.set() + try: + await asyncio.Future() + finally: + stopped.set() + + async def observe(event): + events.append(event) + + async with AsyncExitStack() as stack: + agent, server, session, client = await _setup(stack, monkeypatch, [pending]) + run = asyncio.create_task(agent_factory.run_agent(agent, RunRequest("cancel", on_tool_event=observe))) + try: + await asyncio.wait_for(started.wait(), 3) + finally: + run.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(run, 3) + assert stopped.is_set() and len(session.calls) == 1 and len(client.requests["cancel"]) == 1 + assert [event.status for event in events] == ["in_progress"] and server.connect.await_count == 0 + + asyncio.run(exercise()) + + +@pytest.mark.parametrize("status", ["in_progress", "completed", "failed"]) +def test_mcp_observer_failure_remains_fatal_and_redacted(monkeypatch, caplog, status): + async def exercise(): + async def observe(event): + if event.status == status: + raise RuntimeError(_PRIVATE) + + outcome = (types.CallToolResult(content=[], isError=True) if status == "failed" else _success()) + async with AsyncExitStack() as stack: + agent, server, session, client = await _setup(stack, monkeypatch, [outcome]) + with pytest.raises(AgentRunError, match="tool lifecycle observer failed"): + await agent_factory.run_agent(agent, RunRequest("observer", on_tool_event=observe)) + assert len(session.calls) == (0 if status == "in_progress" else 1) + assert len(client.requests["observer"]) == 1 and server.connect.await_count == 0 + assert _PRIVATE not in caplog.text + + caplog.set_level(logging.DEBUG, logger="agent_framework") + asyncio.run(exercise()) + + +def test_protocol_failure_isolated_across_concurrent_and_subsequent_runs(monkeypatch): + async def exercise(): + held, release = asyncio.Event(), asyncio.Event() + + async def pending_success(): + held.set() + await release.wait() + return _success() + + async with AsyncExitStack() as stack: + agent, _, session, client = await _setup(stack, monkeypatch, [pending_success, _failure("jsonrpc"), _success()]) + good_session = AgentSession(session_id="good-session") + failed_session = AgentSession(session_id="failed-session") + good = asyncio.create_task(agent_factory.run_agent( + agent, RunRequest("concurrent-success"), session=good_session, + )) + try: + await asyncio.wait_for(held.wait(), 3) + with pytest.raises(AgentRunError, match="MCP tool protocol failed"): + await agent_factory.run_agent(agent, RunRequest("concurrent-failure"), session=failed_session) + finally: + release.set() + assert (await asyncio.wait_for(good, 3)).text == "done" + assert (await agent_factory.run_agent( + agent, RunRequest("subsequent-success"), session=good_session, + )).text == "done" + assert len(session.calls) == 3 + assert {name: len(requests) for name, requests in client.requests.items()} == { + "concurrent-success": 2, "concurrent-failure": 1, "subsequent-success": 2, + } + + asyncio.run(exercise()) diff --git a/runtimes/microsoft-agent-framework/tests/test_provider_messages.py b/runtimes/microsoft-agent-framework/tests/test_provider_messages.py new file mode 100644 index 0000000..c6849aa --- /dev/null +++ b/runtimes/microsoft-agent-framework/tests/test_provider_messages.py @@ -0,0 +1,156 @@ +from __future__ import annotations + +import asyncio +import json +from unittest import mock + +import httpx +from agent_framework import Message, tool +from openai import AsyncOpenAI + +from agentkit_serve import agent_factory +from agentkit_serve_common.config import AgentSpec +from agentkit_serve_common.conversation import ConversationTurn, RunRequest + + +def _spec(*, with_tool=False): + return AgentSpec.model_validate({ + "abiVersion": "v0", + "metadata": {"name": "package-agent"}, + "model": { + "provider": "openai-compatible", + "baseURL": "https://provider.example/v1", + "name": "test-model", + }, + "instructions": "Follow the user's instructions.", + "tools": [{"name": "probe", "command": ["unused-fixture"]}] if with_tool else [], + "expose": {"openai": True, "port": 8080}, + }) + + +def _completion(message): + return httpx.Response(200, json={ + "id": "chatcmpl-fixture", + "object": "chat.completion", + "created": 1, + "model": "test-model", + "choices": [{"index": 0, "message": message, + "finish_reason": "tool_calls" if message.get("tool_calls") else "stop"}], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + }) + + +def _reject_speaker_names(body): + # A chat-to-Responses gateway cannot translate named speakers. Package + # metadata must not turn into a speaker identity on subsequent model calls. + assert all("name" not in message for message in body["messages"]) + + +def test_continuation_preserves_history_without_promoting_package_name_to_speaker(): + async def exercise(): + requests = [] + + def respond(request): + body = json.loads(request.content) + _reject_speaker_names(body) + requests.append(body) + return _completion({"role": "assistant", "content": "remembered-é-世界"}) + + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as http: + sdk = AsyncOpenAI(base_url="https://provider.example/v1", api_key="test-key", http_client=http, + max_retries=0) + client = agent_factory.OpenAIChatCompletionClient(model="test-model", async_client=sdk) + with mock.patch.object(agent_factory, "build_client", return_value=client): + async with agent_factory.MAFRuntime(_spec()) as runtime: + first = await runtime.run(RunRequest("Remember é-世界.", session_id="conversation")) + second = await runtime.run(RunRequest( + "Recall it.", session_id="conversation", + history=(ConversationTurn("user", "Remember é-世界."), + ConversationTurn("assistant", first.text)), + )) + assert second.text == first.text == "remembered-é-世界" + assert len(requests) == 2 + assert [(message["role"], message["content"]) for message in requests[1]["messages"]] == [ + ("system", "Follow the user's instructions."), + ("user", "Remember é-世界."), + ("assistant", "remembered-é-世界"), + ("user", "Recall it."), + ] + + asyncio.run(exercise()) + + +def test_tool_loop_preserves_function_name_arguments_and_correlation_without_speaker_name(): + async def exercise(): + requests = [] + calls = [] + events = [] + payload = {"text": "héllo 世界 🌍", "nested": [42, True, None]} + + async def probe(payload: dict): + calls.append(payload) + return json.dumps(payload, ensure_ascii=False) + + async def observe(event): + events.append(event) + + def respond(request): + body = json.loads(request.content) + _reject_speaker_names(body) + requests.append(body) + if len(requests) == 1: + assert body["tools"][0]["function"]["name"] == "probe" + return _completion({"role": "assistant", "content": None, "tool_calls": [{ + "id": "call-fixture", "type": "function", + "function": {"name": "probe", "arguments": json.dumps({"payload": payload})}, + }]}) + assistant, result = body["messages"][-2:] + assert assistant["role"] == "assistant" + assert assistant["tool_calls"][0]["function"]["name"] == "probe" + assert json.loads(assistant["tool_calls"][0]["function"]["arguments"]) == {"payload": payload} + assert result["role"] == "tool" and result["tool_call_id"] == "call-fixture" + assert json.loads(result["content"]) == payload + return _completion({"role": "assistant", "content": "done"}) + + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as http: + sdk = AsyncOpenAI(base_url="https://provider.example/v1", api_key="test-key", http_client=http, + max_retries=0) + client = agent_factory.OpenAIChatCompletionClient(model="test-model", async_client=sdk) + with ( + mock.patch.object(agent_factory, "build_client", return_value=client), + mock.patch.object(agent_factory, "build_tool", return_value=tool(probe, name="probe")), + ): + async with agent_factory.MAFRuntime(_spec(with_tool=True)) as runtime: + result = await runtime.run(RunRequest("Call probe.", session_id="tools", on_tool_event=observe)) + assert result.text == "done" and len(requests) == 2 and calls == [payload] + assert [event.status for event in events] == ["in_progress", "completed"] + assert events[0].tool_call_id == events[1].tool_call_id + + asyncio.run(exercise()) + + +def test_outbound_projection_preserves_original_internal_message_metadata(): + async def exercise(): + original = Message(role="assistant", contents=["remembered-é-世界"], author_name="package-agent") + before = original.to_dict() + requests = [] + + def respond(request): + body = json.loads(request.content) + _reject_speaker_names(body) + requests.append(body) + return _completion({"role": "assistant", "content": "done"}) + + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as http: + sdk = AsyncOpenAI(base_url="https://provider.example/v1", api_key="test-key", http_client=http, + max_retries=0) + client = agent_factory.OpenAIChatCompletionClient(model="test-model", async_client=sdk) + async with agent_factory.build_agent(_spec(), client=client) as agent: + result = await agent.run([original, Message(role="user", contents=["Continue."])]) + assert agent.name == "package-agent" + assert result.text == "done" and len(requests) == 1 + assert original.to_dict() == before + assert any(message["role"] == "assistant" and message["content"] == "remembered-é-世界" + for message in requests[0]["messages"]) + + asyncio.run(exercise()) diff --git a/runtimes/microsoft-agent-framework/tests/test_tool_events.py b/runtimes/microsoft-agent-framework/tests/test_tool_events.py new file mode 100644 index 0000000..bd67a4f --- /dev/null +++ b/runtimes/microsoft-agent-framework/tests/test_tool_events.py @@ -0,0 +1,154 @@ +from __future__ import annotations + +import asyncio +from dataclasses import asdict + +import pytest +from agent_framework import Agent, BaseChatClient, ChatResponse, Content, FunctionInvocationLayer, Message, tool + +from agentkit_serve import agent_factory +from agentkit_serve_common.conversation import RunRequest, ToolCallEvent +from agentkit_serve_common.runtime import AgentRunError + + +class _ToolClient(FunctionInvocationLayer, BaseChatClient): + def __init__(self, calls=1): + super().__init__() + self.remaining = calls + + async def _inner_get_response(self, *, messages, stream, options, **kwargs): + if self.remaining: + self.remaining -= 1 + content = Content.from_function_call( + call_id=f"private-provider-id-{self.remaining}", name="probe", arguments={} + ) + else: + content = Content.from_text("done") + return ChatResponse(messages=[Message(role="assistant", contents=[content])]) + + +def _agent(probe, *, calls=1): + return Agent(client=_ToolClient(calls), tools=[tool(probe, name="probe", description="Controlled fixture.")]) + + +def test_real_tool_execution_awaits_redacted_start_and_result(): + async def exercise(): + events: list[ToolCallEvent] = [] + order = [] + + async def probe(): + order.append("side-effect") + return {"private-payload": ["héllo 世界 🌍", {"nested": True}]} + + async def observe(event): + await asyncio.sleep(0) + events.append(event) + order.append(event.status) + + async with _agent(probe) as agent: + result = await agent_factory.run_agent(agent, RunRequest("probe", on_tool_event=observe)) + assert result.text == "done" + assert order == ["in_progress", "side-effect", "completed"] + assert events[0].tool_call_id == events[1].tool_call_id + assert [event.tool_name for event in events] == ["probe", "probe"] + assert all(set(asdict(event)) == {"tool_call_id", "tool_name", "status"} for event in events) + assert "private" not in str(events) + + asyncio.run(exercise()) + + +def test_failed_tool_emits_failure_before_successful_retry(): + async def exercise(): + events = [] + calls = 0 + + async def probe(): + nonlocal calls + calls += 1 + if calls == 1: + raise RuntimeError("private error details") + return "private successful result" + + async def observe(event): + events.append(event) + + async with _agent(probe, calls=2) as agent: + result = await agent_factory.run_agent(agent, RunRequest("recover", on_tool_event=observe)) + assert result.text == "done" and calls == 2 + assert [event.status for event in events] == ["in_progress", "failed", "in_progress", "completed"] + assert events[0].tool_call_id == events[1].tool_call_id + assert events[2].tool_call_id == events[3].tool_call_id + assert events[0].tool_call_id != events[2].tool_call_id + assert "private" not in str(events) + + asyncio.run(exercise()) + + +def test_cancellation_stops_real_tool_and_leaves_terminal_to_acp(): + async def exercise(): + events = [] + started = asyncio.Event() + stopped = asyncio.Event() + + async def probe(): + started.set() + try: + await asyncio.Future() + finally: + stopped.set() + + async def observe(event): + events.append(event) + + async with _agent(probe) as agent: + run = asyncio.create_task(agent_factory.run_agent(agent, RunRequest("delay", on_tool_event=observe))) + try: + await asyncio.wait_for(started.wait(), 3) + finally: + run.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(run, 3) + assert stopped.is_set() + assert [event.status for event in events] == ["in_progress"] + + asyncio.run(exercise()) + + +@pytest.mark.parametrize("failure_status", ["in_progress", "completed", "failed"]) +def test_observer_failure_aborts_run_without_exposing_its_error(failure_status, caplog): + async def exercise(): + calls = [] + + async def probe(): + calls.append("called") + if failure_status == "failed": + raise RuntimeError("synthetic tool failure") + return "private payload" + + async def observe(event): + if event.status == failure_status: + raise RuntimeError("private observer failure") + + async with _agent(probe) as agent: + with pytest.raises(AgentRunError, match="tool lifecycle observer failed"): + await agent_factory.run_agent(agent, RunRequest("probe", on_tool_event=observe)) + assert calls == ([] if failure_status == "in_progress" else ["called"]) + assert "private observer failure" not in caplog.text + + asyncio.run(exercise()) + + +def test_no_observer_preserves_plain_run_contract(): + async def exercise(): + calls = [] + + class PlainAgent: + async def run(self, messages, *, session): + calls.append((messages, session)) + return type("Result", (), {"text": "plain"})() + + result = await agent_factory.run_agent(PlainAgent(), RunRequest("plain prompt")) + assert result.text == "plain" and len(calls) == 1 + assert calls[0][0][-1].text == "plain prompt" and calls[0][1] is None + + asyncio.run(exercise()) diff --git a/runtimes/pydantic-ai/agentkit_serve/agent_factory.py b/runtimes/pydantic-ai/agentkit_serve/agent_factory.py index 15c4e77..8e7b5cd 100644 --- a/runtimes/pydantic-ai/agentkit_serve/agent_factory.py +++ b/runtimes/pydantic-ai/agentkit_serve/agent_factory.py @@ -24,7 +24,7 @@ from __future__ import annotations from types import TracebackType -from typing import Any +from typing import Any, AsyncIterable from pydantic_ai import Agent @@ -59,7 +59,7 @@ split_tool_command, ) from agentkit_serve_common.config import AgentSpec, ToolSpec -from agentkit_serve_common.conversation import RunRequest +from agentkit_serve_common.conversation import RunRequest, ToolCallEvent from agentkit_serve_common.runtime import ( OfflineEchoRuntimeFactory, RunResult, @@ -263,8 +263,38 @@ def _result_usage(result: object) -> dict[str, int]: async def run_agent(agent: Agent, request: RunRequest) -> RunResult: """Run the pydantic-ai agent and return the neutral result shape.""" message_history = _to_message_history(request) + run_options: dict[str, Any] = {"message_history": message_history} + on_tool_event = request.on_tool_event + if on_tool_event is not None: + from pydantic_ai.messages import ( + AgentStreamEvent, + FunctionToolCallEvent, + FunctionToolResultEvent, + RetryPromptPart, + ) + + async def observe_tools(_: Any, events: AsyncIterable[AgentStreamEvent]) -> None: + async for event in events: + if isinstance(event, FunctionToolCallEvent): + await on_tool_event( + ToolCallEvent(event.part.tool_call_id, event.part.tool_name, "in_progress") + ) + elif isinstance(event, FunctionToolResultEvent): + failed = ( + isinstance(event.part, RetryPromptPart) + or getattr(event.part, "outcome", "success") != "success" + ) + await on_tool_event( + ToolCallEvent( + event.part.tool_call_id, + event.part.tool_name or "", + "failed" if failed else "completed", + ) + ) + + run_options["event_stream_handler"] = observe_tools try: - result = await agent.run(request.prompt, message_history=message_history) + result = await agent.run(request.prompt, **run_options) except Exception as exc: # noqa: BLE001 — normalized for the façade raise normalize_agent_run_error(exc) from exc return RunResult(text=_result_text(result), usage=_result_usage(result)) diff --git a/runtimes/pydantic-ai/tests/test_tool_events.py b/runtimes/pydantic-ai/tests/test_tool_events.py new file mode 100644 index 0000000..1c5b14a --- /dev/null +++ b/runtimes/pydantic-ai/tests/test_tool_events.py @@ -0,0 +1,132 @@ +from __future__ import annotations + +import asyncio +from dataclasses import asdict + +import pytest +from pydantic_ai import Agent, ModelRetry +from pydantic_ai.models.test import TestModel + +from agentkit_serve import agent_factory +from agentkit_serve_common.conversation import RunRequest, ToolCallEvent +from agentkit_serve_common.runtime import AgentRunError + + +def test_real_tool_execution_awaits_start_and_result_observations(): + async def exercise() -> None: + events: list[ToolCallEvent] = [] + order: list[str] = [] + + async def echo() -> dict: + order.append("side-effect") + assert events[-1].status == "in_progress" + return {"nested": ["héllo 世界 🌍", {"private-value": "not an event"}]} + + async def observe(event: ToolCallEvent) -> None: + await asyncio.sleep(0) + events.append(event) + order.append(event.status) + + agent = Agent(TestModel(custom_output_text="done"), tools=[echo]) + async with agent: + result = await agent_factory.run_agent(agent, RunRequest("run echo", on_tool_event=observe)) + assert result.text == "done" + assert order == ["in_progress", "side-effect", "completed"] + assert events[0].tool_call_id == events[1].tool_call_id + assert [event.tool_name for event in events] == ["echo", "echo"] + assert all(set(asdict(event)) == {"tool_call_id", "tool_name", "status"} for event in events) + assert "private-value" not in str(events) + + asyncio.run(exercise()) + + +def test_model_retry_emits_failed_tool_result_then_successful_recovery(): + async def exercise() -> None: + events: list[ToolCallEvent] = [] + calls = 0 + + async def recover() -> str: + nonlocal calls + calls += 1 + if calls == 1: + raise ModelRetry("private error details must not enter tool events") + return "private success payload" + + async def observe(event: ToolCallEvent) -> None: + events.append(event) + + agent = Agent(TestModel(custom_output_text="recovered"), tools=[recover]) + async with agent: + result = await agent_factory.run_agent(agent, RunRequest("recover", on_tool_event=observe)) + assert result.text == "recovered" and calls == 2 + assert [event.status for event in events] == ["in_progress", "failed", "in_progress", "completed"] + assert events[0].tool_call_id == events[1].tool_call_id + assert events[2].tool_call_id == events[3].tool_call_id + assert "private" not in str(events) + + asyncio.run(exercise()) + + +def test_cancellation_interrupts_tool_and_propagates_to_acp_owner(): + async def exercise() -> None: + events: list[ToolCallEvent] = [] + started = asyncio.Event() + stopped = asyncio.Event() + + async def delay() -> str: + started.set() + try: + await asyncio.Future() + finally: + stopped.set() + + async def observe(event: ToolCallEvent) -> None: + events.append(event) + + agent = Agent(TestModel(custom_output_text="unreachable"), tools=[delay]) + async with agent: + run = asyncio.create_task(agent_factory.run_agent(agent, RunRequest("delay", on_tool_event=observe))) + await asyncio.wait_for(started.wait(), timeout=2) + run.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(run, timeout=2) + assert stopped.is_set() + assert [event.status for event in events] == ["in_progress"] + + asyncio.run(exercise()) + + +def test_without_tool_observer_keeps_nonstreaming_run_contract(): + async def exercise() -> None: + calls: list[dict] = [] + + class PlainAgent: + async def run(self, prompt, *, message_history): + calls.append({"prompt": prompt, "history": message_history}) + return type("Result", (), {"output": "plain"})() + + result = await agent_factory.run_agent(PlainAgent(), RunRequest("ordinary HTTP request")) + assert result.text == "plain" + assert calls == [{"prompt": "ordinary HTTP request", "history": []}] + + asyncio.run(exercise()) + + +def test_observer_failure_stops_tool_execution(): + async def exercise() -> None: + calls: list[str] = [] + + async def echo() -> str: + calls.append("unexpected") + return "unreachable" + + async def observe(event: ToolCallEvent) -> None: + raise RuntimeError("private observer failure") + + agent = Agent(TestModel(custom_output_text="unreachable"), tools=[echo]) + async with agent: + with pytest.raises(AgentRunError): + await agent_factory.run_agent(agent, RunRequest("echo", on_tool_event=observe)) + assert calls == [] + + asyncio.run(exercise()) From c9a18070363b9eded0be8ca826d8fc5365afc828 Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Sat, 5 Sep 2026 14:25:01 -0700 Subject: [PATCH 19/23] fix(runtime): harden provider failures and v2 cancellation Signed-off-by: Sertac Ozercan --- pkg/agentkit/config/brokered_unicode_test.go | 67 +++++ pkg/agentkit/config/validate.go | 30 +- .../tests/_brokered_description_cases.py | 18 ++ runtimes/common/tests/test_acp_protocol.py | 8 +- .../agentkit_serve/agent_factory.py | 63 ++-- .../tests/test_mcp_failures.py | 212 +++++++++++++- .../agentkit_serve/agent_factory.py | 59 +++- .../pydantic-ai/tests/test_mcp_failures.py | 275 ++++++++++++++++++ .../tests/test_provider_failures.py | 197 +++++++++++++ 9 files changed, 886 insertions(+), 43 deletions(-) create mode 100644 pkg/agentkit/config/brokered_unicode_test.go create mode 100644 runtimes/pydantic-ai/tests/test_mcp_failures.py create mode 100644 runtimes/pydantic-ai/tests/test_provider_failures.py diff --git a/pkg/agentkit/config/brokered_unicode_test.go b/pkg/agentkit/config/brokered_unicode_test.go new file mode 100644 index 0000000..264f3c0 --- /dev/null +++ b/pkg/agentkit/config/brokered_unicode_test.go @@ -0,0 +1,67 @@ +package config + +import ( + "crypto/sha256" + "fmt" + "strings" + "testing" + "unicode" + "unicode/utf8" +) + +func TestLowerBrokeredTextUnicode15(t *testing.T) { + var scalarValues strings.Builder + for r := rune(0); r <= unicode.MaxRune; r++ { + if utf8.ValidRune(r) { + scalarValues.WriteRune(r) + } + } + // Recorded from Go 1.26.1 strings.ToLower, unicode.Version 15.0.0, + // over every valid Unicode scalar in ascending order. + const want = "137590953b837f1ec8b7c02b3a0425d0789df789a59ad256e7a63416f9fc4c11" + got := fmt.Sprintf("%x", sha256.Sum256([]byte(lowerBrokeredText(scalarValues.String())))) + if got != want { + t.Fatalf("brokered lowercase mapping differs from Unicode 15: got %s, want %s", got, want) + } +} + +func TestValidateBrokeredDescriptionsPreservesUnicode15(t *testing.T) { + for _, test := range []struct { + description string + valid bool + }{ + {description: "Count input tokens\u1c89", valid: true}, + {description: "Count input tokens\ua7cb", valid: true}, + {description: "Count input tokens\ua7ce", valid: true}, + {description: "Count input tokens\U00010d50", valid: true}, + {description: "Count input tokens\U00010d65", valid: true}, + {description: "Count input tokens\U00016ea0", valid: true}, + {description: "Count input tokens\U00016eb8", valid: true}, + {description: "Count input Tokens\U00010d50", valid: false}, + {description: "Count input tokens\u00c9", valid: false}, + {description: "Count input tokens\u0130", valid: false}, + {description: "Count input tokens_\U00010d50", valid: false}, + {description: "Count input tokens\U00010d50 abc123", valid: false}, + {description: "token\U00016ea0=abc123", valid: false}, + {description: "Read {auth\U00016ea0}", valid: false}, + {description: "Bas\u0130c dXNlcjpwYXNz", valid: false}, + } { + t.Run(test.description, func(t *testing.T) { + cfg := validMinimalConfig() + cfg.BrokeredTools = []BrokeredTool{{ + Name: safeLookupToolName, + Description: test.description, + BrokeredClass: BrokeredClassRead, + Parameters: map[string]any{jsonSchemaTypeKey: jsonSchemaTypeObject}, + }} + err := cfg.Validate() + if test.valid { + if err != nil { + t.Fatalf("Unicode 15 description should be accepted: %v", err) + } + } else if err == nil || !strings.Contains(err.Error(), "brokeredTools[0].description") { + t.Fatalf("credential-shaped description should be rejected, got %v", err) + } + }) + } +} diff --git a/pkg/agentkit/config/validate.go b/pkg/agentkit/config/validate.go index 0fbd33e..0fbe992 100644 --- a/pkg/agentkit/config/validate.go +++ b/pkg/agentkit/config/validate.go @@ -16,6 +16,7 @@ import ( "sort" "strconv" "strings" + "unicode" "github.com/sozercan/agentkit/pkg/agentkit/runtimes" "github.com/sozercan/agentkit/pkg/utils" @@ -70,6 +71,21 @@ const ( var brokeredBasicValuePattern = regexp.MustCompile(`(?i)(?:^|[^A-Za-z0-9_])basic[^A-Za-z0-9_]+?([A-Za-z0-9+/]+={0,2})`) +// Brokered descriptions use the Unicode 15 casing contract shared with the +// Python ABI reader. Exclude the mappings introduced in Unicode 16 and 17; +// the full-scalar golden test detects further drift when Go is upgraded. +var brokeredUnicode15Case = unicode.SpecialCase{ + {Lo: 0x1C89, Hi: 0x1C89}, + {Lo: 0xA7CB, Hi: 0xA7CC}, + {Lo: 0xA7CE, Hi: 0xA7CE}, + {Lo: 0xA7D2, Hi: 0xA7D2}, + {Lo: 0xA7D4, Hi: 0xA7D4}, + {Lo: 0xA7DA, Hi: 0xA7DA}, + {Lo: 0xA7DC, Hi: 0xA7DC}, + {Lo: 0x10D50, Hi: 0x10D65}, + {Lo: 0x16EA0, Hi: 0x16EB8}, +} + // Validate reports every problem with the config at once via errors.Join (plan // §16.2 #3 — one report-all validator, not scattered first-error-wins funcs). // @@ -799,19 +815,23 @@ func isSchemaDigest(value string) bool { return err == nil } +func lowerBrokeredText(value string) string { + return strings.ToLowerSpecial(brokeredUnicode15Case, value) +} + func hasUnsafeBrokeredText(value string) bool { - lowered := strings.ToLower(value) + lowered := lowerBrokeredText(value) return hasUnsafeBrokeredDescription(value) || containsBrokeredWord(lowered, "basic") || strings.Contains(lowered, brokeredTokenWord) } func hasUnsafeBrokeredDescription(value string) bool { - lowered := strings.ToLower(value) + lowered := lowerBrokeredText(value) normalized := normalizeKey(lowered) return containsSecretPrefix(value) || strings.Contains(value, "://") || containsBrokeredWord(lowered, "bearer") || containsBrokeredWord(lowered, brokeredSensitiveWord) || containsBrokeredWord(lowered, brokeredSensitivePluralWord) || containsBrokeredBasicAuthReference(value) || containsBrokeredCredentialAssignment(value) || containsBrokeredCredentialReference(value) || strings.Contains(lowered, authorizationKey) || strings.Contains(lowered, "secret") || strings.Contains(lowered, "password") || strings.Contains(lowered, "passphrase") || strings.Contains(lowered, "pwd") || strings.Contains(lowered, "api key") || strings.Contains(lowered, "apikey") || strings.Contains(normalized, "apikey") || strings.Contains(normalized, "xapikey") || strings.Contains(normalized, "subscriptionkey") || strings.Contains(normalized, "xfunctionskey") || strings.Contains(lowered, brokeredUnsafeCookieKey) || strings.Contains(lowered, "set-cookie") || strings.Contains(lowered, "x-api-key") || strings.Contains(lowered, credentialHeaderAPIKey) || strings.Contains(lowered, "subscription-key") || strings.Contains(lowered, "x-functions-key") || strings.Contains(lowered, "ocp-apim-subscription-key") || strings.Contains(lowered, "private key") || strings.Contains(lowered, "privatekey") || strings.Contains(lowered, "key material") || strings.Contains(lowered, ".svc") || strings.Contains(lowered, "cluster.local") } func containsBrokeredBasicAuthReference(value string) bool { - lowered := strings.ToLower(value) + lowered := lowerBrokeredText(value) if !containsBrokeredWord(lowered, "basic") { return false } @@ -831,7 +851,7 @@ func containsBrokeredBasicAuthReference(value string) bool { } } for i, field := range fields { - if !containsBrokeredWord(strings.ToLower(field), "basic") { + if !containsBrokeredWord(lowerBrokeredText(field), "basic") { continue } for j, candidate := range fields[i+1:] { @@ -1022,7 +1042,7 @@ func isBrokeredNumericCount(value string) bool { } func hasBrokeredStructuredKeyShape(value string) bool { - return strings.HasPrefix(strings.TrimLeft(value, "\"'`([{<"), "-") || strings.ContainsAny(value, "_/.[]{}()<>\"'`") || value != strings.ToLower(value) + return strings.HasPrefix(strings.TrimLeft(value, "\"'`([{<"), "-") || strings.ContainsAny(value, "_/.[]{}()<>\"'`") || value != lowerBrokeredText(value) } func containsBrokeredWord(value string, word string) bool { diff --git a/runtimes/common/tests/_brokered_description_cases.py b/runtimes/common/tests/_brokered_description_cases.py index 9312ca3..bc8f15a 100644 --- a/runtimes/common/tests/_brokered_description_cases.py +++ b/runtimes/common/tests/_brokered_description_cases.py @@ -5,7 +5,18 @@ "Count model tokens", "Count model tokens 100", "Count model tokens: 1,000.", + "Count input tokens\u1c89", + "Count input tokens\ua7cb", + "Count input tokens\ua7cc", + "Count input tokens\ua7ce", + "Count input tokens\ua7d2", + "Count input tokens\ua7d4", + "Count input tokens\ua7da", + "Count input tokens\ua7dc", "Count input tokens\U00010d50", + "Count input tokens\U00010d65", + "Count input tokens\U00016ea0", + "Count input tokens\U00016eb8", ) UNSAFE_BROKERED_DESCRIPTIONS = ( @@ -56,6 +67,13 @@ "Finished counting. Model token 123", 'Finished "counting." Model token 123', "Count model tokens 123456789012345678901234567890", + "Count input Tokens\U00010d50", + "Count input tokens\u00c9", + "Count input tokens\u0130", + "Count input tokens_\U00010d50", + "Count input tokens\U00010d50 abc123", + "token\U00016ea0=abc123", + "Read {auth\U00016ea0}", "Read {auth}", "Read (header)", "Read [REDACTED_AUTH_HEADER]", diff --git a/runtimes/common/tests/test_acp_protocol.py b/runtimes/common/tests/test_acp_protocol.py index 50fc270..a24421d 100644 --- a/runtimes/common/tests/test_acp_protocol.py +++ b/runtimes/common/tests/test_acp_protocol.py @@ -260,6 +260,8 @@ def __init__(self, *, block_first_update: bool = False) -> None: self.frames: list[bytes] = [] self.parsed_messages: list[dict[str, Any]] = [] self.lock = threading.Lock() + self.loop = asyncio.get_running_loop() + self.updated = asyncio.Event() def write(self, frame: bytes) -> int: message = json.loads(frame) @@ -274,6 +276,7 @@ def write(self, frame: bytes) -> int: with self.lock: self.frames.append(bytes(frame)) self.parsed_messages.append(message) + self.loop.call_soon_threadsafe(self.updated.set) return len(frame) def flush(self) -> None: @@ -287,11 +290,14 @@ def messages(self) -> list[dict[str, Any]]: async def _wait_for_stdio_response(writer: RecordingWriter, request_id: int) -> dict[str, Any]: async def wait() -> dict[str, Any]: while True: + # Sleeping until a write avoids starving the writer thread with + # repeated list scans. Clear before reading to retain racing writes. + writer.updated.clear() matches = [message for message in writer.messages() if message.get("id") == request_id] if matches: assert len(matches) == 1 return matches[0] - await asyncio.sleep(0) + await writer.updated.wait() return await asyncio.wait_for(wait(), timeout=5) diff --git a/runtimes/microsoft-agent-framework/agentkit_serve/agent_factory.py b/runtimes/microsoft-agent-framework/agentkit_serve/agent_factory.py index 19fb241..5518786 100644 --- a/runtimes/microsoft-agent-framework/agentkit_serve/agent_factory.py +++ b/runtimes/microsoft-agent-framework/agentkit_serve/agent_factory.py @@ -74,9 +74,15 @@ _DEFAULT_FOUNDRY_AUDIENCE = "https://ai.azure.com/.default" _DEFAULT_MCP_REQUEST_TIMEOUT = 120 _DEFAULT_SESSION_CACHE_MAX = 256 -# Invocation tasks inherit the same mutable marker for this run. Separate runs, -# including concurrent Sessions on one Agent, receive independent markers. -_mcp_failures: ContextVar[list[bool] | None] = ContextVar("agentkit_maf_mcp_failures", default=None) +# Invocation tasks share a fatal-error signal with their run owner. Concurrent +# Sessions on the same Agent have independent signals. +_run_failure: ContextVar[asyncio.Future[str] | None] = ContextVar("agentkit_maf_run_failure", default=None) + + +def _fail_run(message: str) -> None: + failure = _run_failure.get() + if failure is not None and not failure.done(): + failure.set_result(message) def _mcp_request_timeout() -> int: @@ -374,14 +380,13 @@ class _MCPFailureMiddleware(FunctionMiddleware): async def process( self, context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]] ) -> None: - failures = _mcp_failures.get() - if failures: - raise MiddlewareTermination("MCP tool protocol failed") + failure = _run_failure.get() + if failure is not None and failure.done(): + raise MiddlewareTermination(failure.result()) try: await call_next() except _MCPProtocolError: - if failures is not None: - failures.append(True) + _fail_run("MCP tool protocol failed") # MiddlewareTermination stops the model loop even on MAF 1.9, # which predates MiddlewareFailure. run_agent makes it a failure. raise MiddlewareTermination("MCP tool protocol failed") from None @@ -779,7 +784,6 @@ class _ToolEventMiddleware(FunctionMiddleware): def __init__(self, observe: Callable[[ToolCallEvent], Awaitable[None]]) -> None: self.observe = observe - self.failed = False async def _emit(self, event: ToolCallEvent) -> None: try: @@ -788,7 +792,7 @@ async def _emit(self, event: ToolCallEvent) -> None: # MAF absorbs ordinary function exceptions and continues the model # loop. Stop it, then fail run_agent, even on supported versions # predating MiddlewareFailure. - self.failed = True + _fail_run("tool lifecycle observer failed") raise MiddlewareTermination("tool lifecycle observer failed") from None async def process( @@ -815,21 +819,32 @@ async def run_agent( """Run the MAF agent and return the neutral result shape.""" messages = _to_messages(request, include_history=include_history) kwargs = {} - observer = None if request.on_tool_event is not None: - observer = _ToolEventMiddleware(request.on_tool_event) - kwargs["middleware"] = [observer] - failures: list[bool] = [] - token = _mcp_failures.set(failures) + kwargs["middleware"] = [_ToolEventMiddleware(request.on_tool_event)] + failure: asyncio.Future[str] = asyncio.get_running_loop().create_future() + token = _run_failure.set(failure) + + async def execute(): + return await agent.run(messages, session=session, **kwargs) + + running = asyncio.create_task(execute()) try: try: - result = await agent.run(messages, session=session, **kwargs) - except Exception as exc: # noqa: BLE001 — normalized for the façade - raise normalize_agent_run_error(exc) from exc + # MiddlewareTermination stops the next model step, but supported + # MAF versions first join all calls in the current batch. Cancel + # and join the SDK run so a fatal call also stops pending siblings. + await asyncio.wait((running, failure), return_when=asyncio.FIRST_COMPLETED) + if not failure.done(): + try: + result = await running + except Exception as exc: # noqa: BLE001 — normalized for the façade + raise normalize_agent_run_error(exc) from exc + finally: + running.cancel() + await asyncio.gather(running, return_exceptions=True) + if failure.done(): + raise AgentRunError(failure.result()) + return RunResult(text=_result_text(result), usage=_result_usage(result)) finally: - _mcp_failures.reset(token) - if failures: - raise AgentRunError("MCP tool protocol failed") - if observer is not None and observer.failed: - raise AgentRunError("tool lifecycle observer failed") - return RunResult(text=_result_text(result), usage=_result_usage(result)) + failure.cancel() + _run_failure.reset(token) diff --git a/runtimes/microsoft-agent-framework/tests/test_mcp_failures.py b/runtimes/microsoft-agent-framework/tests/test_mcp_failures.py index 4d44a5a..c0f8819 100644 --- a/runtimes/microsoft-agent-framework/tests/test_mcp_failures.py +++ b/runtimes/microsoft-agent-framework/tests/test_mcp_failures.py @@ -47,10 +47,11 @@ async def call_tool(self, name, *, arguments, meta): class _Client(FunctionInvocationLayer, BaseChatClient): - def __init__(self, function_name, *, calls=1): + def __init__(self, function_name, *, calls=1, parallel_calls=1): super().__init__() self.function_name = function_name self.tool_calls = calls + self.parallel_calls = parallel_calls self.requests = {} async def _inner_get_response(self, *, messages, stream, options, **kwargs): @@ -59,13 +60,13 @@ async def _inner_get_response(self, *, messages, stream, options, **kwargs): requests.append(messages) number = len(requests) if number <= self.tool_calls: - item = Content.from_function_call( - call_id=f"private-provider-id-{prompt}-{number}", name=self.function_name, + items = [Content.from_function_call( + call_id=f"private-provider-id-{prompt}-{number}-{slot}", name=self.function_name, arguments={"payload": {**_PAYLOAD, "prompt": prompt}, "undeclared": "must-be-filtered"}, - ) + ) for slot in range(self.parallel_calls)] else: - item = Content.from_text("done") - return ChatResponse(messages=[Message(role="assistant", contents=[item])]) + items = [Content.from_text("done")] + return ChatResponse(messages=[Message(role="assistant", contents=items)]) def _success(): @@ -92,7 +93,7 @@ def _failure(kind): raise AssertionError("unknown controlled failure") -async def _setup(stack, monkeypatch, outcomes, *, transport="streamable-http", calls=1): +async def _setup(stack, monkeypatch, outcomes, *, transport="streamable-http", calls=1, parallel_calls=1): monkeypatch.setenv("TEST_MCP_URL", "http://example.invalid/mcp") tool = ( ToolSpec(name="fixture", type="mcp", url_env="TEST_MCP_URL", transport=transport) @@ -105,7 +106,7 @@ async def _setup(stack, monkeypatch, outcomes, *, transport="streamable-http", c server._ping_available = False server.connect = mock.AsyncMock() await server.load_tools() - client = _Client(server.functions[0].name, calls=calls) + client = _Client(server.functions[0].name, calls=calls, parallel_calls=parallel_calls) spec = AgentSpec.model_validate({ "abiVersion": "v0", "metadata": {"name": "test-package"}, "model": {"provider": "openai-compatible", "name": "test-model", "baseURL": "http://model.invalid/v1"}, @@ -154,6 +155,201 @@ async def observe(event): asyncio.run(exercise()) +@pytest.mark.parametrize("transport", ["streamable-http", "stdio"]) +@pytest.mark.parametrize("kind", ["jsonrpc", "closed", "auth", "timeout"]) +@pytest.mark.parametrize("observed", [True, False]) +def test_fatal_parallel_call_cancels_and_joins_pending_sibling(monkeypatch, caplog, transport, kind, observed): + async def exercise(): + started, cancelled, release_cleanup, stopped = (asyncio.Event() for _ in range(4)) + events = [] + + async def fatal(): + await started.wait() + raise _failure(kind) + + async def pending(): + started.set() + try: + await asyncio.Future() + except asyncio.CancelledError: + cancelled.set() + await release_cleanup.wait() + stopped.set() + raise + + async def observe(event): + events.append(event) + + async with AsyncExitStack() as stack: + agent, server, session, client = await _setup( + stack, monkeypatch, [fatal, pending], transport=transport, parallel_calls=2, + ) + run = asyncio.create_task(agent_factory.run_agent( + agent, RunRequest("fatal-batch", on_tool_event=observe if observed else None), + )) + try: + await asyncio.wait_for(cancelled.wait(), 3) + assert not run.done(), "fatal run returned before sibling cleanup finished" + release_cleanup.set() + with pytest.raises(AgentRunError, match="^MCP tool protocol failed$"): + await asyncio.wait_for(asyncio.shield(run), 3) + assert stopped.is_set(), "fatal run returned before its sibling was cancelled" + finally: + release_cleanup.set() + run.cancel() + await asyncio.gather(run, return_exceptions=True) + assert len(session.calls) == 2 and len(client.requests["fatal-batch"]) == 1 + assert server.connect.await_count == 0 + assert all(call == { + "name": "probe", "arguments": {"payload": {**_PAYLOAD, "prompt": "fatal-batch"}}, + "meta": {"fixture.example/authority": "frozen-metadata"}, + } for call in session.calls) + if observed: + assert [event.status for event in events] == ["in_progress", "in_progress", "failed"] + assert events[0].tool_call_id == events[2].tool_call_id != events[1].tool_call_id + assert "private" not in str(events) + else: + assert events == [] + assert _PRIVATE not in caplog.text + + caplog.set_level(logging.DEBUG, logger="agent_framework") + asyncio.run(exercise()) + + +@pytest.mark.parametrize("transport", ["streamable-http", "stdio"]) +def test_admitted_parallel_error_does_not_cancel_successful_sibling(monkeypatch, caplog, transport): + async def exercise(): + started, failed_event, release, cancelled = (asyncio.Event() for _ in range(4)) + events = [] + + async def admitted(): + await started.wait() + return types.CallToolResult(content=[types.TextContent(type="text", text=_PRIVATE)], isError=True) + + async def pending_success(): + started.set() + try: + await release.wait() + except asyncio.CancelledError: + cancelled.set() + raise + return _success() + + async def observe(event): + events.append(event) + if event.status == "failed": + failed_event.set() + + async with AsyncExitStack() as stack: + agent, server, session, client = await _setup( + stack, monkeypatch, [admitted, pending_success], transport=transport, parallel_calls=2, + ) + run = asyncio.create_task(agent_factory.run_agent( + agent, RunRequest("recover-batch", on_tool_event=observe), + )) + try: + await asyncio.wait_for(failed_event.wait(), 3) + assert not run.done() and not cancelled.is_set() + release.set() + assert (await asyncio.wait_for(asyncio.shield(run), 3)).text == "done" + finally: + release.set() + run.cancel() + await asyncio.gather(run, return_exceptions=True) + assert not cancelled.is_set() + assert len(session.calls) == 2 and len(client.requests["recover-batch"]) == 2 + assert server.connect.await_count == 0 + assert [event.status for event in events] == ["in_progress", "in_progress", "failed", "completed"] + assert events[0].tool_call_id == events[2].tool_call_id + assert events[1].tool_call_id == events[3].tool_call_id != events[0].tool_call_id + assert _PRIVATE not in caplog.text and "private" not in str(events) + + caplog.set_level(logging.DEBUG, logger="agent_framework") + asyncio.run(exercise()) + + +def test_caller_cancellation_stops_all_parallel_calls(monkeypatch): + async def exercise(): + started = asyncio.Event() + calls, cancelled, events = [], [], [] + + async def pending(): + calls.append(len(calls)) + if len(calls) == 2: + started.set() + try: + await asyncio.Future() + except asyncio.CancelledError: + cancelled.append(True) + raise + + async def observe(event): + events.append(event) + + async with AsyncExitStack() as stack: + agent, server, session, client = await _setup(stack, monkeypatch, [pending, pending], parallel_calls=2) + run = asyncio.create_task(agent_factory.run_agent( + agent, RunRequest("cancel-batch", on_tool_event=observe), + )) + try: + await asyncio.wait_for(started.wait(), 3) + finally: + run.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(run, 3) + assert len(cancelled) == len(session.calls) == 2 + assert len(client.requests["cancel-batch"]) == 1 and server.connect.await_count == 0 + assert [event.status for event in events] == ["in_progress", "in_progress"] + + asyncio.run(exercise()) + + +@pytest.mark.parametrize("status", ["in_progress", "completed", "failed"]) +def test_parallel_observer_failure_cancels_pending_sibling(monkeypatch, caplog, status): + async def exercise(): + started, stopped = asyncio.Event(), asyncio.Event() + starts = [] + + async def completed(): + await started.wait() + return types.CallToolResult(content=[], isError=True) if status == "failed" else _success() + + async def pending(): + started.set() + try: + await asyncio.Future() + except asyncio.CancelledError: + stopped.set() + raise + + async def observe(event): + if event.status == "in_progress": + starts.append(event.tool_call_id) + if event.status == status and event.tool_call_id == starts[0]: + await started.wait() + raise RuntimeError(_PRIVATE) + + outcomes = [pending] if status == "in_progress" else [completed, pending] + async with AsyncExitStack() as stack: + agent, server, session, client = await _setup(stack, monkeypatch, outcomes, parallel_calls=2) + run = asyncio.create_task(agent_factory.run_agent( + agent, RunRequest("observer-batch", on_tool_event=observe), + )) + try: + with pytest.raises(AgentRunError, match="^tool lifecycle observer failed$"): + await asyncio.wait_for(asyncio.shield(run), 3) + assert stopped.is_set() + finally: + run.cancel() + await asyncio.gather(run, return_exceptions=True) + assert len(session.calls) == (1 if status == "in_progress" else 2) + assert len(client.requests["observer-batch"]) == 1 and server.connect.await_count == 0 + assert _PRIVATE not in caplog.text + + caplog.set_level(logging.DEBUG, logger="agent_framework") + asyncio.run(exercise()) + + @pytest.mark.parametrize("transport", ["streamable-http", "stdio"]) def test_admitted_mcp_error_can_recover_with_correlated_events(monkeypatch, caplog, transport): async def exercise(): diff --git a/runtimes/pydantic-ai/agentkit_serve/agent_factory.py b/runtimes/pydantic-ai/agentkit_serve/agent_factory.py index 8e7b5cd..6f161df 100644 --- a/runtimes/pydantic-ai/agentkit_serve/agent_factory.py +++ b/runtimes/pydantic-ai/agentkit_serve/agent_factory.py @@ -23,10 +23,11 @@ from __future__ import annotations +from contextlib import asynccontextmanager from types import TracebackType -from typing import Any, AsyncIterable +from typing import Any, AsyncIterable, AsyncIterator -from pydantic_ai import Agent +from pydantic_ai import Agent, ModelRetry try: # pydantic-ai 1.x from pydantic_ai.mcp import MCPServerStdio @@ -41,9 +42,11 @@ try: from fastmcp.client.transports.stdio import StdioTransport from fastmcp.client.transports.http import StreamableHttpTransport + from fastmcp.exceptions import ToolError except ImportError: # pragma: no cover - older dependency set without FastMCP transports StdioTransport = None # type: ignore[assignment] StreamableHttpTransport = None # type: ignore[assignment] +from pydantic_ai.models import StreamedResponse from pydantic_ai.models.openai import OpenAIChatModel from pydantic_ai.providers.openai import OpenAIProvider @@ -61,6 +64,7 @@ from agentkit_serve_common.config import AgentSpec, ToolSpec from agentkit_serve_common.conversation import RunRequest, ToolCallEvent from agentkit_serve_common.runtime import ( + AgentRunError, OfflineEchoRuntimeFactory, RunResult, RuntimeSession, @@ -86,13 +90,46 @@ def validate_supported_spec(spec: AgentSpec) -> None: raise AgentBuildError("pydantic-ai runtime does not support context providers") +class _OpenAIChatModel(OpenAIChatModel): + @asynccontextmanager + async def request_stream(self, *args: Any, **kwargs: Any) -> AsyncIterator[StreamedResponse]: + async with super().request_stream(*args, **kwargs) as response: + yield response + # EOF alone must not commit partial text or execute partial tool + # calls. Check each model response before the agent advances. + if not response.cancelled and not ( + response.finish_reason or (response.provider_details or {}).get("finish_reason") + ): + raise AgentRunError("Model stream ended before completion") + + def build_model(spec: AgentSpec) -> OpenAIChatModel: """Construct the OpenAI-compatible chat model pointed at ``model.baseURL``.""" provider = OpenAIProvider( base_url=spec.model.base_url, api_key=resolve_api_key(spec), ) - return OpenAIChatModel(spec.model.name, provider=provider) + return _OpenAIChatModel(spec.model.name, provider=provider) + + +async def _process_mcp_tool_call(ctx: Any, call_tool: Any, name: str, args: dict[str, Any]) -> Any: + try: + return await call_tool(name, args) + except ToolError: + # FastMCP raises ToolError only for an admitted isError result. Preserve + # model recovery for that case, without forwarding upstream diagnostics. + raise ModelRetry("MCP tool execution failed") from None + except ExceptionGroup as exc: + # FastMCP can group completed tool errors during session teardown. + # Mixed protocol/transport failures must still stop the run. + _, remaining = exc.split(ToolError) + if remaining is None: + raise ModelRetry("MCP tool execution failed") from None + raise AgentRunError("MCP tool protocol failed") from None + except Exception: + # Recent Pydantic AI versions also retry JSON-RPC errors by default. + # Authorization, protocol and transport failures must end this run. + raise AgentRunError("MCP tool protocol failed") from None def build_tool_server(tool: ToolSpec) -> Any: @@ -110,7 +147,13 @@ def build_tool_server(tool: ToolSpec) -> Any: url, httpx_client_factory=same_origin_mcp_httpx_client_factory(tool, url, timeout=timeout), ) - return MCPToolset(transport, init_timeout=timeout, read_timeout=timeout).prefixed(tool.name) + return MCPToolset( + transport, + init_timeout=timeout, + read_timeout=timeout, + tool_error_behavior="error", + process_tool_call=_process_mcp_tool_call, + ).prefixed(tool.name) command, args = split_tool_command(tool, example='["npx", "-y", "..."]') @@ -143,7 +186,13 @@ def build_tool_server(tool: ToolSpec) -> Any: # the stdio subprocess should be torn down instead of kept alive. keep_alive=False, ) - return MCPToolset(transport, init_timeout=timeout, read_timeout=timeout).prefixed(tool.name) + return MCPToolset( + transport, + init_timeout=timeout, + read_timeout=timeout, + tool_error_behavior="error", + process_tool_call=_process_mcp_tool_call, + ).prefixed(tool.name) def build_agent(spec: AgentSpec) -> Agent: diff --git a/runtimes/pydantic-ai/tests/test_mcp_failures.py b/runtimes/pydantic-ai/tests/test_mcp_failures.py new file mode 100644 index 0000000..3e4e58f --- /dev/null +++ b/runtimes/pydantic-ai/tests/test_mcp_failures.py @@ -0,0 +1,275 @@ +from __future__ import annotations + +import asyncio +from contextlib import asynccontextmanager +from types import SimpleNamespace + +import httpx +import pytest +from fastmcp import Client +from fastmcp.exceptions import ToolError +from mcp import types +from mcp.shared.exceptions import McpError +from pydantic_ai import Agent +from pydantic_ai.models.test import TestModel + +from agentkit_serve import agent_factory +from agentkit_serve_common.config import ToolSpec +from agentkit_serve_common.conversation import RunRequest +from agentkit_serve_common.runtime import AgentRunError + + +PRIVATE = "PRIVATE_MCP_DIAGNOSTIC_0ca0cb" + + +class CountingModel(TestModel): + requests = 0 + + async def request(self, *args, **kwargs): + self.requests += 1 + return await super().request(*args, **kwargs) + + @asynccontextmanager + async def request_stream(self, *args, **kwargs): + self.requests += 1 + async with super().request_stream(*args, **kwargs) as response: + yield response + + +def _failure(kind): + if kind == "jsonrpc": + return McpError(types.ErrorData(code=-32000, message=PRIVATE)) + if kind == "auth": + return httpx.HTTPStatusError( + PRIVATE, + request=httpx.Request("POST", "https://example.invalid"), + response=httpx.Response(403), + ) + if kind == "timeout": + return httpx.ReadTimeout(PRIVATE) + if kind == "malformed": + try: + types.CallToolResult.model_validate({"content": PRIVATE}) + except ValueError as exc: + return exc + if kind == "grouped": + return ExceptionGroup("private transport group", [_failure("jsonrpc")]) + if kind == "mixed": + return ExceptionGroup( + "private mixed group", [ToolError(PRIVATE), _failure("jsonrpc")] + ) + raise AssertionError("unknown controlled failure") + + +def _setup(monkeypatch, outcomes, *, transport="streamable-http", tools=1): + monkeypatch.setenv("TEST_MCP_URL", "http://example.invalid/mcp") + spec = ( + ToolSpec( + name="fixture", type="mcp", url_env="TEST_MCP_URL", transport=transport + ) + if transport == "streamable-http" + else ToolSpec(name="fixture", command=["unused-fixture"]) + ) + toolset = agent_factory.build_tool_server(spec) + client = toolset.wrapped.client + initialized = types.InitializeResult( + protocolVersion="2025-11-25", + capabilities=types.ServerCapabilities(tools=types.ToolsCapability()), + serverInfo=types.Implementation(name="agentkit-mcp-fixture", version="1"), + ) + + async def enter(self): + return self + + async def exit(self, *args): + return None + + async def list_tools(): + return [ + types.Tool( + name=f"probe_{index}", inputSchema={"type": "object", "properties": {}} + ) + for index in range(tools) + ] + + # Keep real FastMCP result parsing, Pydantic toolsets, and agent execution. + # Substitute connection establishment and the remote MCP responses only. + monkeypatch.setattr(Client, "__aenter__", enter) + monkeypatch.setattr(Client, "__aexit__", exit) + monkeypatch.setattr(Client, "initialize_result", property(lambda self: initialized)) + monkeypatch.setattr( + Client, + "session", + property( + lambda self: SimpleNamespace( + _tool_output_schemas={}, + list_tools=list_tools, + ) + ), + ) + monkeypatch.setattr(client, "list_tools", list_tools) + pending = iter(outcomes) + calls = [] + + async def call_tool_mcp(name, arguments, **kwargs): + calls.append({"name": name, "arguments": arguments, "meta": kwargs.get("meta")}) + result = next(pending) + if isinstance(result, BaseException): + raise result + if callable(result): + return await result() + return result + + monkeypatch.setattr(client, "call_tool_mcp", call_tool_mcp) + model = CountingModel(custom_output_text="done") + return Agent(model, toolsets=[toolset]), model, calls + + +def _success(): + return types.CallToolResult( + content=[types.TextContent(type="text", text="héllo 世界 🌍")] + ) + + +@pytest.mark.parametrize("transport", ["streamable-http", "stdio"]) +@pytest.mark.parametrize( + "kind", ["jsonrpc", "auth", "timeout", "malformed", "grouped", "mixed"] +) +@pytest.mark.parametrize("observed", [True, False]) +def test_protocol_failures_end_the_run_without_retry_or_model_continuation( + monkeypatch, transport, kind, observed +): + async def exercise(): + events = [] + + async def observe(event): + events.append(event) + + agent, model, calls = _setup( + monkeypatch, [_failure(kind), _success()], transport=transport + ) + async with agent: + with pytest.raises( + AgentRunError, match="MCP tool protocol failed" + ) as caught: + await agent_factory.run_agent( + agent, + RunRequest("fatal", on_tool_event=observe if observed else None), + ) + assert len(calls) == 1 and model.requests == 1 + assert PRIVATE not in str(caught.value) + assert [event.status for event in events] == ( + ["in_progress"] if observed else [] + ) + + asyncio.run(exercise()) + + +@pytest.mark.parametrize("transport", ["streamable-http", "stdio"]) +@pytest.mark.parametrize("observed", [True, False]) +@pytest.mark.parametrize("grouped", [True, False]) +def test_admitted_iserror_still_allows_model_recovery( + monkeypatch, transport, observed, grouped +): + async def exercise(): + events = [] + + async def observe(event): + events.append(event) + + admitted = types.CallToolResult( + content=[types.TextContent(type="text", text=PRIVATE)], isError=True + ) + if grouped: + admitted = ExceptionGroup( + "private group", [ExceptionGroup("nested group", [ToolError(PRIVATE)])] + ) + agent, model, calls = _setup( + monkeypatch, [admitted, _success()], transport=transport + ) + async with agent: + result = await agent_factory.run_agent( + agent, + RunRequest("recover", on_tool_event=observe if observed else None), + ) + assert result.text == "done" and len(calls) == 2 and model.requests == 3 + assert [event.status for event in events] == ( + ["in_progress", "failed", "in_progress", "completed"] if observed else [] + ) + assert PRIVATE not in str(events) + + asyncio.run(exercise()) + + +def test_grouped_cancellation_is_not_converted_to_model_retry(monkeypatch): + async def exercise(): + cancellation = BaseExceptionGroup( + "cancelled transport", [ToolError(PRIVATE), asyncio.CancelledError()] + ) + agent, model, calls = _setup(monkeypatch, [cancellation]) + async with agent: + with pytest.raises(BaseExceptionGroup) as caught: + await agent_factory.run_agent(agent, RunRequest("cancel")) + assert caught.value.subgroup(asyncio.CancelledError) is not None + assert len(calls) == 1 and model.requests == 1 + + asyncio.run(exercise()) + + +def test_parallel_protocol_failure_cancels_other_call_without_model_continuation( + monkeypatch, +): + async def exercise(): + started, stopped = asyncio.Event(), asyncio.Event() + + async def pending(): + started.set() + try: + await asyncio.Future() + finally: + stopped.set() + + async def fail(): + await started.wait() + raise _failure("jsonrpc") + + async def observe(event): + pass + + agent, model, calls = _setup(monkeypatch, [pending, fail], tools=2) + async with agent: + with pytest.raises(AgentRunError, match="MCP tool protocol failed"): + await asyncio.wait_for( + agent_factory.run_agent( + agent, RunRequest("parallel", on_tool_event=observe) + ), + 3, + ) + assert len(calls) == 2 and model.requests == 1 and stopped.is_set() + + asyncio.run(exercise()) + + +def test_cancellation_is_not_converted_to_model_retry(monkeypatch): + async def exercise(): + started, stopped = asyncio.Event(), asyncio.Event() + + async def pending(): + started.set() + try: + await asyncio.Future() + finally: + stopped.set() + + agent, model, calls = _setup(monkeypatch, [pending]) + async with agent: + run = asyncio.create_task( + agent_factory.run_agent(agent, RunRequest("cancel")) + ) + await asyncio.wait_for(started.wait(), 3) + run.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(run, 3) + assert len(calls) == 1 and model.requests == 1 and stopped.is_set() + + asyncio.run(exercise()) diff --git a/runtimes/pydantic-ai/tests/test_provider_failures.py b/runtimes/pydantic-ai/tests/test_provider_failures.py new file mode 100644 index 0000000..e3cf950 --- /dev/null +++ b/runtimes/pydantic-ai/tests/test_provider_failures.py @@ -0,0 +1,197 @@ +from __future__ import annotations + +import asyncio +import json + +import httpx +import pytest +from openai import AsyncOpenAI +from pydantic_ai import Agent +from pydantic_ai.providers.openai import OpenAIProvider + +from agentkit_serve import agent_factory +from agentkit_serve_common.config import AgentSpec +from agentkit_serve_common.conversation import RunRequest +from agentkit_serve_common.runtime import AgentRunError + + +PRIVATE = "private-incomplete-model-output" + + +def _frame(delta, finish=None): + return ( + "data: " + + json.dumps( + { + "id": "completion-test", + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "delta": delta, "finish_reason": finish}], + } + ) + + "\n\n" + ).encode() + + +def _response(*, tool=False, complete=True, done=True): + delta = ( + { + "tool_calls": [ + { + "index": 0, + "id": "call-probe", + "type": "function", + "function": {"name": "probe", "arguments": "{}"}, + } + ] + } + if tool + else {"content": PRIVATE} + ) + content = _frame({"role": "assistant"}) + _frame(delta) + if complete: + content += _frame({}, "tool_calls" if tool else "stop") + if done: + content += b"data: [DONE]\n\n" + return httpx.Response( + 200, headers={"content-type": "text/event-stream"}, content=content + ) + + +def _model(monkeypatch, client): + provider = OpenAIProvider( + openai_client=AsyncOpenAI( + api_key="synthetic-test-key", + base_url="https://fixture.invalid/v1", + http_client=client, + max_retries=0, + ) + ) + monkeypatch.setattr(agent_factory, "OpenAIProvider", lambda **kwargs: provider) + return agent_factory.build_model( + AgentSpec.model_validate( + { + "abiVersion": "v0", + "metadata": {"name": "provider-test"}, + "model": { + "provider": "openai-compatible", + "baseURL": "https://fixture.invalid/v1", + "name": "gpt-4o-mini", + "apiKeyEnv": "OPENAI_API_KEY", + }, + "instructions": "Use the probe.", + "tools": [], + "expose": {"openai": True, "port": 8080}, + } + ) + ) + + +@pytest.mark.parametrize("after_tool", [True, False]) +@pytest.mark.parametrize("incomplete_part", ["text", "tool"]) +def test_incomplete_stream_fails_before_committing_text_or_executing_its_tools( + monkeypatch, after_tool, incomplete_part +): + async def exercise(): + requests, invocations, events = [], [], [] + + async def probe(): + invocations.append("probe") + return "success" + + async def observe(event): + events.append(event) + + def handle(request): + requests.append(json.loads(request.content)) + if after_tool and len(requests) == 1: + return _response(tool=True) + if len(requests) > (2 if after_tool else 1): + return _response() + return _response(tool=incomplete_part == "tool", complete=False, done=False) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client: + agent = Agent(_model(monkeypatch, client), tools=[probe]) + async with agent: + with pytest.raises(AgentRunError) as caught: + await agent_factory.run_agent( + agent, RunRequest("test", on_tool_event=observe) + ) + assert PRIVATE not in str(caught.value) + assert len(requests) == (2 if after_tool else 1) + assert invocations == (["probe"] if after_tool else []) + assert [event.status for event in events] == ( + ["in_progress", "completed"] if after_tool else [] + ) + + asyncio.run(exercise()) + + +@pytest.mark.parametrize("done", [True, False]) +def test_explicit_completion_reason_preserves_tool_loop_and_output(monkeypatch, done): + async def exercise(): + requests, invocations, events = [], [], [] + + async def probe(): + invocations.append("probe") + return "success" + + async def observe(event): + events.append(event) + + def handle(request): + requests.append(json.loads(request.content)) + return _response(tool=len(requests) == 1, done=done) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client: + agent = Agent(_model(monkeypatch, client), tools=[probe]) + async with agent: + result = await agent_factory.run_agent( + agent, RunRequest("test", on_tool_event=observe) + ) + assert result.text == PRIVATE + assert len(requests) == 2 and invocations == ["probe"] + assert [event.status for event in events] == ["in_progress", "completed"] + + asyncio.run(exercise()) + + +def test_cancelling_incomplete_stream_preserves_cancellation(monkeypatch): + async def exercise(): + started, closed = asyncio.Event(), asyncio.Event() + + class PendingStream(httpx.AsyncByteStream): + async def __aiter__(self): + yield _frame({"content": PRIVATE}) + started.set() + await asyncio.Future() + + async def aclose(self): + closed.set() + + async def observe(event): + pass + + def handle(request): + return httpx.Response( + 200, + headers={"content-type": "text/event-stream"}, + stream=PendingStream(), + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client: + agent = Agent(_model(monkeypatch, client)) + async with agent: + run = asyncio.create_task( + agent_factory.run_agent( + agent, RunRequest("test", on_tool_event=observe) + ) + ) + await asyncio.wait_for(started.wait(), 3) + run.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(run, 3) + assert closed.is_set() + + asyncio.run(exercise()) From b8685f42e7248667d844f74774a18bb9b48257ee Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Thu, 10 Sep 2026 15:19:44 -0700 Subject: [PATCH 20/23] ci: use latest Vekil image for live E2E Signed-off-by: Sertac Ozercan --- .github/workflows/ci.yml | 2 +- docs/development.md | 2 +- scripts/live-copilot-agent-e2e.sh | 2 +- 3 files changed, 3 insertions(+), 3 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 8364e65..33fc3ee 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -408,7 +408,7 @@ jobs: timeout-minutes: 75 env: COPILOT_GITHUB_TOKEN: ${{ secrets.COPILOT_GITHUB_TOKEN }} - VEKIL_IMAGE: ghcr.io/sozercan/vekil@sha256:d13edeedf7bec319da8eb3ea4949a4d0802e244c14765a347e62e1b8b7be8e3d + VEKIL_IMAGE: ghcr.io/sozercan/vekil:latest steps: - name: Decide whether live Vekil/Copilot E2E can run id: gate diff --git a/docs/development.md b/docs/development.md index f7f9702..f701962 100644 --- a/docs/development.md +++ b/docs/development.md @@ -122,7 +122,7 @@ live model provider. The optional live job runs `scripts/live-copilot-agent-e2e.sh` when `COPILOT_GITHUB_TOKEN` is available and the run is allowed to access repository -secrets. It uses the pinned `ghcr.io/sozercan/vekil` image to provide an +secrets. It uses `ghcr.io/sozercan/vekil:latest` to provide an OpenAI-compatible endpoint and validates a real built AgentKit container through `/v1/chat/completions`. diff --git a/scripts/live-copilot-agent-e2e.sh b/scripts/live-copilot-agent-e2e.sh index 6c41255..953eb34 100755 --- a/scripts/live-copilot-agent-e2e.sh +++ b/scripts/live-copilot-agent-e2e.sh @@ -20,7 +20,7 @@ work_dir="$(mktemp -d "${RUNNER_TEMP:-${TMPDIR:-/tmp}}/agentkit-live-copilot.XXX copilot_token="${COPILOT_GITHUB_TOKEN:-}" vekil_cache_dir="${VEKIL_CACHE_DIR:-${HOME:-}/.config/vekil}" -vekil_image="${VEKIL_IMAGE:-ghcr.io/sozercan/vekil@sha256:d13edeedf7bec319da8eb3ea4949a4d0802e244c14765a347e62e1b8b7be8e3d}" +vekil_image="${VEKIL_IMAGE:-ghcr.io/sozercan/vekil:latest}" vekil_container_name="${VEKIL_CONTAINER_NAME:-agentkit-vekil}" vekil_host_port="${VEKIL_HOST_PORT:-1337}" vekil_container_port="${VEKIL_CONTAINER_PORT:-1337}" From ffb456f8b4b92c6fda6170b27ad4fde36fcd1b96 Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Thu, 10 Sep 2026 15:29:09 -0700 Subject: [PATCH 21/23] fix(ci): report Vekil readiness failures Signed-off-by: Sertac Ozercan --- scripts/live-copilot-agent-e2e.sh | 1 + 1 file changed, 1 insertion(+) diff --git a/scripts/live-copilot-agent-e2e.sh b/scripts/live-copilot-agent-e2e.sh index 953eb34..1f822dc 100755 --- a/scripts/live-copilot-agent-e2e.sh +++ b/scripts/live-copilot-agent-e2e.sh @@ -126,6 +126,7 @@ wait_for_vekil_ready() { sleep 2 done + log "Vekil readiness response: $(curl -sS --max-time 10 "${url}" 2>&1 | redact || true)" die "Vekil /readyz never became available at ${url}" } From 860892441be630da9982ffaee754351e40332a8b Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Thu, 10 Sep 2026 16:27:52 -0700 Subject: [PATCH 22/23] ci: test live Copilot E2E with Vekil v0.14.1 Signed-off-by: Sertac Ozercan --- .github/workflows/ci.yml | 2 +- docs/development.md | 2 +- scripts/live-copilot-agent-e2e.sh | 2 +- 3 files changed, 3 insertions(+), 3 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 33fc3ee..5033bff 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -408,7 +408,7 @@ jobs: timeout-minutes: 75 env: COPILOT_GITHUB_TOKEN: ${{ secrets.COPILOT_GITHUB_TOKEN }} - VEKIL_IMAGE: ghcr.io/sozercan/vekil:latest + VEKIL_IMAGE: ghcr.io/sozercan/vekil:v0.14.1@sha256:2fa0558f6304cc6ed1fb5b0135f62f12f28f1cdd0a8c057c4283414bceac1362 steps: - name: Decide whether live Vekil/Copilot E2E can run id: gate diff --git a/docs/development.md b/docs/development.md index f701962..515dbff 100644 --- a/docs/development.md +++ b/docs/development.md @@ -122,7 +122,7 @@ live model provider. The optional live job runs `scripts/live-copilot-agent-e2e.sh` when `COPILOT_GITHUB_TOKEN` is available and the run is allowed to access repository -secrets. It uses `ghcr.io/sozercan/vekil:latest` to provide an +secrets. It uses `ghcr.io/sozercan/vekil:v0.14.1`, pinned by digest, to provide an OpenAI-compatible endpoint and validates a real built AgentKit container through `/v1/chat/completions`. diff --git a/scripts/live-copilot-agent-e2e.sh b/scripts/live-copilot-agent-e2e.sh index 1f822dc..7a2b04e 100755 --- a/scripts/live-copilot-agent-e2e.sh +++ b/scripts/live-copilot-agent-e2e.sh @@ -20,7 +20,7 @@ work_dir="$(mktemp -d "${RUNNER_TEMP:-${TMPDIR:-/tmp}}/agentkit-live-copilot.XXX copilot_token="${COPILOT_GITHUB_TOKEN:-}" vekil_cache_dir="${VEKIL_CACHE_DIR:-${HOME:-}/.config/vekil}" -vekil_image="${VEKIL_IMAGE:-ghcr.io/sozercan/vekil:latest}" +vekil_image="${VEKIL_IMAGE:-ghcr.io/sozercan/vekil:v0.14.1@sha256:2fa0558f6304cc6ed1fb5b0135f62f12f28f1cdd0a8c057c4283414bceac1362}" vekil_container_name="${VEKIL_CONTAINER_NAME:-agentkit-vekil}" vekil_host_port="${VEKIL_HOST_PORT:-1337}" vekil_container_port="${VEKIL_CONTAINER_PORT:-1337}" From e2433f71b7d4552154d25e84b3e6e9304b400605 Mon Sep 17 00:00:00 2001 From: Sertac Ozercan Date: Thu, 10 Sep 2026 23:14:57 -0700 Subject: [PATCH 23/23] ci: update Vekil to v0.14.3 Signed-off-by: Sertac Ozercan --- .github/workflows/ci.yml | 2 +- docs/development.md | 2 +- scripts/live-copilot-agent-e2e.sh | 2 +- 3 files changed, 3 insertions(+), 3 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 5033bff..f95d786 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -408,7 +408,7 @@ jobs: timeout-minutes: 75 env: COPILOT_GITHUB_TOKEN: ${{ secrets.COPILOT_GITHUB_TOKEN }} - VEKIL_IMAGE: ghcr.io/sozercan/vekil:v0.14.1@sha256:2fa0558f6304cc6ed1fb5b0135f62f12f28f1cdd0a8c057c4283414bceac1362 + VEKIL_IMAGE: ghcr.io/sozercan/vekil:v0.14.3@sha256:996b628fbe8c7a35d33e9d6bb855f2613228fc5c9b09498dae6ea6b208a0071b steps: - name: Decide whether live Vekil/Copilot E2E can run id: gate diff --git a/docs/development.md b/docs/development.md index 515dbff..47caf65 100644 --- a/docs/development.md +++ b/docs/development.md @@ -122,7 +122,7 @@ live model provider. The optional live job runs `scripts/live-copilot-agent-e2e.sh` when `COPILOT_GITHUB_TOKEN` is available and the run is allowed to access repository -secrets. It uses `ghcr.io/sozercan/vekil:v0.14.1`, pinned by digest, to provide an +secrets. It uses `ghcr.io/sozercan/vekil:v0.14.3`, pinned by digest, to provide an OpenAI-compatible endpoint and validates a real built AgentKit container through `/v1/chat/completions`. diff --git a/scripts/live-copilot-agent-e2e.sh b/scripts/live-copilot-agent-e2e.sh index 7a2b04e..c3dc3f9 100755 --- a/scripts/live-copilot-agent-e2e.sh +++ b/scripts/live-copilot-agent-e2e.sh @@ -20,7 +20,7 @@ work_dir="$(mktemp -d "${RUNNER_TEMP:-${TMPDIR:-/tmp}}/agentkit-live-copilot.XXX copilot_token="${COPILOT_GITHUB_TOKEN:-}" vekil_cache_dir="${VEKIL_CACHE_DIR:-${HOME:-}/.config/vekil}" -vekil_image="${VEKIL_IMAGE:-ghcr.io/sozercan/vekil:v0.14.1@sha256:2fa0558f6304cc6ed1fb5b0135f62f12f28f1cdd0a8c057c4283414bceac1362}" +vekil_image="${VEKIL_IMAGE:-ghcr.io/sozercan/vekil:v0.14.3@sha256:996b628fbe8c7a35d33e9d6bb855f2613228fc5c9b09498dae6ea6b208a0071b}" vekil_container_name="${VEKIL_CONTAINER_NAME:-agentkit-vekil}" vekil_host_port="${VEKIL_HOST_PORT:-1337}" vekil_container_port="${VEKIL_CONTAINER_PORT:-1337}"