diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 8364e65..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@sha256:d13edeedf7bec319da8eb3ea4949a4d0802e244c14765a347e62e1b8b7be8e3d + 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/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/development.md b/docs/development.md index f7f9702..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 the pinned `ghcr.io/sozercan/vekil` image 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/docs/orka.md b/docs/orka.md index 2af06f7..e0eeb13 100644 --- a/docs/orka.md +++ b/docs/orka.md @@ -1,4 +1,91 @@ -# 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: \ + AGENTKIT_ADAPTER_DIGEST=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 `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; +- 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. 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 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. +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 -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 +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` +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 +327,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/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/README.md b/runtimes/common/README.md index 3c9816c..3e0097b 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,27 @@ 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 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 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..1264ed3 --- /dev/null +++ b/runtimes/common/agentkit_serve_common/acp.py @@ -0,0 +1,945 @@ +"""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 re +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, 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, ToolCallEvent +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 +# 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 +_METHOD_NOT_FOUND = -32601 +_INVALID_PARAMS = -32602 +_INTERNAL_ERROR = -32603 +_REQUEST_CANCELLED = -32800 + +_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 _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 + 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 + + +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 + + +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 _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) + 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 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") + 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 + + 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 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.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 + + 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 + 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. + 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() + 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())): + 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) + 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: + 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}) + + 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 + 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) + 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_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_auth_value, 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_request_key is not None: + 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}]") + 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", + ) + + 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 + state.active_run = asyncio.create_task( + state.runtime.run(run_request), + name="agentkit-acp-runtime-prompt", + ) + self.prompt_requests[request_key] = state + try: + 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 + + 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 + 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: + for chunk in _utf8_chunks(result.text, _MAX_ASSISTANT_MESSAGE_CHUNK_BYTES): + if state.cancel_requested: + break + await self.send( + { + "jsonrpc": _JSONRPC_VERSION, + "method": _METHOD_SESSION_UPDATE, + "params": { + "sessionId": session_id, + "update": { + "sessionUpdate": "agent_message_chunk", + "content": {"type": "text", "text": chunk}, + }, + }, + } + ) + 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: + tool_observer.closed = True + self.prompt_requests.pop(request_key, None) + state.active_request_key = None + state.active_run = None + + 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() + + 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} + 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 + 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: + break + 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: + 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: + """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..5ec1116 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, + load_verified_acp_runtime_binding, + run_acp_stdio, +) 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) @@ -130,6 +136,13 @@ 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 + if protocol == "acp": + try: + 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( 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/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/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/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"] 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 new file mode 100644 index 0000000..a24421d --- /dev/null +++ b/runtimes/common/tests/test_acp_protocol.py @@ -0,0 +1,1746 @@ +from __future__ import annotations + +import asyncio +import hashlib +import io +import json +import os +import queue +import threading +from collections.abc import Callable +from pathlib import Path +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, ToolCallEvent +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 _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": content, + }, + } + + +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"]) + + +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) + 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 + + +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 + + +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.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) + 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)) + self.parsed_messages.append(message) + self.loop.call_soon_threadsafe(self.updated.set) + return len(frame) + + def flush(self) -> None: + return None + + def messages(self) -> list[dict[str, Any]]: + with self.lock: + return list(self.parsed_messages) + + +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 writer.updated.wait() + + return await asyncio.wait_for(wait(), timeout=5) + + +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)) + + 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, + ) + ) + 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) + + +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)]) + + 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, + ) + ) + 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] + + 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() + 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_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_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 + 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") + + 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_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")])) + + 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()) + + +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 + 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) + + +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)], +) +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": -32603, + "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_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()) + + +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 + 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, + 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_capability_gated_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 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() + + 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_bytes, _spec()) + + with pytest.raises(ACPConfigurationError, match="exact agent config bytes"): + 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 +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( + ("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(b"exact", _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(b"unused", _spec(**override)) diff --git a/runtimes/common/tests/test_cli_protocol.py b/runtimes/common/tests/test_cli_protocol.py index 793ce4f..b560c6e 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() + + 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", + 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..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, @@ -206,6 +209,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(): @@ -335,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_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/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 4688bb0..5518786 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,15 @@ _DEFAULT_FOUNDRY_AUDIENCE = "https://ai.azure.com/.default" _DEFAULT_MCP_REQUEST_TIMEOUT = 120 _DEFAULT_SESSION_CACHE_MAX = 256 +# 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: @@ -260,6 +280,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 +357,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 +373,35 @@ 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: + failure = _run_failure.get() + if failure is not None and failure.done(): + raise MiddlewareTermination(failure.result()) + try: + await call_next() + except _MCPProtocolError: + _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 + + +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 +418,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()], ) @@ -390,13 +482,29 @@ 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) + result = await run_agent( + self.agent, + request, + session=session, + include_history=include_history, + ) 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 +529,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 +717,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(): @@ -663,6 +779,36 @@ 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 + + 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. + _fail_run("tool lifecycle observer failed") + 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, @@ -672,8 +818,33 @@ async def run_agent( ) -> RunResult: """Run the MAF agent and return the neutral result shape.""" messages = _to_messages(request, include_history=include_history) + kwargs = {} + if request.on_tool_event is not None: + 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: - 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 - return RunResult(text=_result_text(result), usage=_result_usage(result)) + try: + # 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: + failure.cancel() + _run_failure.reset(token) diff --git a/runtimes/microsoft-agent-framework/tests/test_guardrails.py b/runtimes/microsoft-agent-framework/tests/test_guardrails.py index 6875f9f..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_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,12 +791,57 @@ 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) asyncio.run(exercise()) assert include_history_values == [True, True, False] + assert seen_sessions[0] is not seen_sessions[1] + 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] @@ -976,3 +1021,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/microsoft-agent-framework/tests/test_mcp_failures.py b/runtimes/microsoft-agent-framework/tests/test_mcp_failures.py new file mode 100644 index 0000000..c0f8819 --- /dev/null +++ b/runtimes/microsoft-agent-framework/tests/test_mcp_failures.py @@ -0,0 +1,459 @@ +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, 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): + 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: + 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: + items = [Content.from_text("done")] + return ChatResponse(messages=[Message(role="assistant", contents=items)]) + + +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, 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) + 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, 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"}, + "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"]) +@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(): + 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 8ce1e27..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 +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 @@ -59,8 +62,9 @@ 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 ( + 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: @@ -191,6 +240,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(): @@ -256,8 +312,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_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()) 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()) 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): diff --git a/scripts/live-copilot-agent-e2e.sh b/scripts/live-copilot-agent-e2e.sh index 6c41255..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@sha256:d13edeedf7bec319da8eb3ea4949a4d0802e244c14765a347e62e1b8b7be8e3d}" +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}" @@ -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}" }