From b9a7c420877e4a1e9f00180a05d85da7b5da12d0 Mon Sep 17 00:00:00 2001 From: eavanvalkenburg Date: Wed, 30 Sep 2026 11:15:27 +0200 Subject: [PATCH 1/2] [BREAKING] Python: add parsed durable Foundry Invocations agent runs --- python/packages/foundry_hosting/README.md | 30 + .../__init__.py | 4 +- .../_invocations.py | 266 +++++- .../_request.py | 39 +- .../foundry_hosting/tests/test_invocations.py | 778 ++++++++++++++++-- .../tests/test_invocations_int.py | 34 +- .../foundry_hosting/tests/test_request.py | 28 +- .../foundry-hosted-agents/README.md | 2 +- .../invocations/basic/README.md | 142 ++-- .../invocations/basic/main.py | 50 +- 10 files changed, 1177 insertions(+), 196 deletions(-) diff --git a/python/packages/foundry_hosting/README.md b/python/packages/foundry_hosting/README.md index ab8fb4b5cc1..f4bf70bc297 100644 --- a/python/packages/foundry_hosting/README.md +++ b/python/packages/foundry_hosting/README.md @@ -268,6 +268,36 @@ and user ID. Consumers must use it as a whole and must not parse it or depend on representation. Repeated requests for the same identifier pair restore the saved session. Locally, the platform session ID is used unchanged. +### Invocations agent requests and wire compatibility + +By default, `POST /invocations` accepts `{"message": "Hi", "options": {}, "stream": false}`. Applications can supply +a sync or async `parse_request(request)` returning a typed `InvocationRun(messages, options, stream)` to accept their +own JSON shape and MAF `Message` inputs. A sync or async `prepare_options(request, options)` hook can filter or replace +a **copy** of this turn's caller options without changing the agent's `default_options`. It must return a mapping +with string keys. The host rejects reserved platform/session fields, `store`, `extra_body`, and private continuation +fields after the hook; callers cannot select another sandbox or enable downstream service continuation through +runtime options. When an agent cannot accept runtime options, `unsupported_options="warn"` (default) logs and ignores +them; `"ignore"` silently drops them and `"error"` rejects them. For request-scoped factories, unsupported options +discovered after streaming starts are reported as an SSE `error` event with `status: 400`. + +**New default:** non-streaming success is JSON `{"response": "..."}`. Streaming success is real SSE +`event: delta` with `{"text": "..."}`, followed by `event: done` with the platform sandbox `session_id`. +The `done` event is sent only after the final MAF `ResponseStream` is finalized and the session is persisted. +Client validation errors return JSON HTTP 400 before streaming where possible. Provider errors are logged and +sanitized as JSON HTTP 500 or an SSE `error`; a cross-process ETag conflict is JSON HTTP 409 or an SSE `error` +with `code: "session_conflict"` and `status: 409`. A stream may emit deltas before an error. The host serializes +same-session requests in one process, but a CAS conflict can still follow external tool effects in separate +processes; it does not guarantee exactly-once execution. See the +[Invocations agent/parser example](../../samples/04-hosting/foundry-hosted-agents/invocations/basic/). + +**Deprecated opt-in:** set `InvocationsHostServer(agent, legacy_wire_format=True)` only for existing clients +that must keep the previous plain-text non-streaming response and raw text-chunk streaming format. The host logs +and emits a deprecation warning once on construction. Errors are never returned as successful text: non-stream +failures still use JSON error statuses; a post-start legacy stream failure terminates the stream rather than +injecting unexpected SSE framing. Migrate all opted-in clients to JSON/SSE; remove the compatibility mode only +after those callers have migrated and a separate, deliberate breaking-change decision, never by silently +switching an opted-in deployment. + Both hosts accept `agent_session_store_provider` to select a `StoreProvider[SessionStore]`. Session state must support `AgentSession` serialization. Use `register_state_type()` codecs for custom types; unsupported live objects fail during persistence. Restored sessions preserve diff --git a/python/packages/foundry_hosting/agent_framework_foundry_hosting/__init__.py b/python/packages/foundry_hosting/agent_framework_foundry_hosting/__init__.py index e5b59394e8a..952a8c0a8d0 100644 --- a/python/packages/foundry_hosting/agent_framework_foundry_hosting/__init__.py +++ b/python/packages/foundry_hosting/agent_framework_foundry_hosting/__init__.py @@ -5,7 +5,7 @@ if TYPE_CHECKING: from ._invocations import InvocationsHostServer - from ._request import HostedResponseRequest + from ._request import HostedResponseRequest, InvocationRun from ._responses import ResponsesHostServer from ._scope import FoundryRequestScope from ._state_store import ( @@ -36,6 +36,7 @@ "FoundryRequestScope": "._scope", "FoundryToolbox": "._toolbox", "HostedResponseRequest": "._request", + "InvocationRun": "._request", "FunctionApprovalStore": "._state_store", "FunctionApprovalStoreProvider": "._state_store", "InvocationsHostServer": "._invocations", @@ -55,6 +56,7 @@ "FunctionApprovalStore", "FunctionApprovalStoreProvider", "HostedResponseRequest", + "InvocationRun", "InvocationsHostServer", "ResponsesHostServer", "StoreProvider", diff --git a/python/packages/foundry_hosting/agent_framework_foundry_hosting/_invocations.py b/python/packages/foundry_hosting/agent_framework_foundry_hosting/_invocations.py index 417a175eb3b..85b46b11d01 100644 --- a/python/packages/foundry_hosting/agent_framework_foundry_hosting/_invocations.py +++ b/python/packages/foundry_hosting/agent_framework_foundry_hosting/_invocations.py @@ -2,28 +2,62 @@ from __future__ import annotations +import asyncio +import inspect import json import logging import sys -from collections.abc import Awaitable, Callable +import warnings +import weakref +from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping from contextlib import AbstractAsyncContextManager, AsyncExitStack, asynccontextmanager +from contextvars import Token +from copy import deepcopy +from typing import cast from agent_framework import AgentSession, ResponseStream, SessionStore, SupportsAgentRun from agent_framework._telemetry import mark_feature_used -from azure.ai.agentserver.core import FoundryAgentRequestContext, get_request_context +from azure.ai.agentserver.core import ( + FoundryAgentRequestContext, + get_request_context, + reset_request_context, + set_request_context, +) +from azure.ai.agentserver.core.storage import FoundryStorageConflictError, FoundryStoragePreconditionError from azure.ai.agentserver.invocations import InvocationAgentServerHost from starlette.requests import Request -from starlette.responses import Response, StreamingResponse +from starlette.responses import JSONResponse, Response, StreamingResponse from starlette.types import Receive, Scope, Send -from typing_extensions import Any, AsyncGenerator +from typing_extensions import Any -from ._agent_source import is_agent, resolve_agent, validate_agent_source +from ._agent_source import AgentSource, is_agent, resolve_agent, validate_agent_source from ._feature_usage import FeatureIndex +from ._request import InvocationRun, UnsupportedOptions, validate_request_options, validate_unsupported_options from ._scope import FoundryRequestScope from ._state_store import AgentSessionStoreProvider, StoreProvider logger = logging.getLogger(__name__) +InvocationParser = Callable[[Request], InvocationRun | Awaitable[InvocationRun]] +InvocationOptionsHook = Callable[[Request, dict[str, Any]], Mapping[str, Any] | Awaitable[Mapping[str, Any]]] + + +class _UnsupportedAgentOptions(TypeError): + """The agent cannot accept the caller's run options under the selected policy.""" + + +def _sse(event: str, data: Mapping[str, Any]) -> str: + return f"event: {event}\ndata: {json.dumps(data)}\n\n" + + +def _invocation_failure(exc: Exception) -> tuple[str, int, str]: + cause: BaseException | None = exc + while cause is not None: + if isinstance(cause, (FoundryStorageConflictError, FoundryStoragePreconditionError)): + return "Another request advanced this agent session; reload before retrying.", 409, "session_conflict" + cause = cause.__cause__ + return "Agent invocation failed.", 500, "invocation_failed" + class _InvocationStreamingResponse(StreamingResponse): def __init__(self, content: AsyncGenerator[str]) -> None: @@ -49,14 +83,18 @@ async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: class InvocationsHostServer(InvocationAgentServerHost): - """An invocations server host for an agent.""" + """Host an agent with durable sessions and application-defined Invocations input.""" def __init__( self, - agent: SupportsAgentRun | Callable[[], SupportsAgentRun | Awaitable[SupportsAgentRun]], + agent: AgentSource, *, openapi_spec: dict[str, Any] | None = None, agent_session_store_provider: StoreProvider[SessionStore] | None = None, + parse_request: InvocationParser | None = None, + prepare_options: InvocationOptionsHook | None = None, + unsupported_options: UnsupportedOptions = "warn", + legacy_wire_format: bool = False, **kwargs: Any, ) -> None: """Initialize an InvocationsHostServer. @@ -68,21 +106,44 @@ def __init__( agent_session_store_provider: Provider for conversation session storage. Defaults to Foundry storage when hosted and the SDK's file-backed storage locally. New default stores expire sessions 30 days after their last write. Custom providers control their own retention. + parse_request: Optional sync or async parser returning an `InvocationRun` from application JSON. + Without one, accepts a JSON object with `message`, optional `options`, and optional `stream`. + prepare_options: Optional sync or async hook to filter or replace a copy of caller run options. + unsupported_options: `"warn"` (default), `"ignore"`, or `"error"` for agents without runtime options. + legacy_wire_format: Opt into the deprecated plain-text response and raw streaming chunks instead of + the default JSON response and framed `delta`/`done`/`error` server-sent events. **kwargs: Additional keyword arguments. - - This host will expect the request to be a JSON body with a "message" field. - The response contains the agent's text, or streamed text when "stream" is true. """ validate_agent_source(agent) + if parse_request is not None and not callable(parse_request): + raise TypeError("parse_request must be a callable.") + if prepare_options is not None and not callable(prepare_options): + raise TypeError("prepare_options must be a callable.") + if not isinstance(legacy_wire_format, bool): + raise TypeError("legacy_wire_format must be a boolean.") super().__init__(openapi_spec=openapi_spec, **kwargs) self._agent = agent self._owns_request_agent = not is_agent(agent) + self._parse_request = parse_request + self._prepare_options = prepare_options + self._unsupported_options = validate_unsupported_options(unsupported_options) + self._legacy_wire_format = legacy_wire_format self._session_store_provider = ( AgentSessionStoreProvider(store_name="invocation_sessions") if agent_session_store_provider is None else agent_session_store_provider ) + self._session_locks: weakref.WeakValueDictionary[str | tuple[str, str], asyncio.Lock] = ( + weakref.WeakValueDictionary() + ) + if legacy_wire_format: + message = ( + "legacy_wire_format=True is deprecated; migrate Invocations clients to JSON responses " + "and framed SSE events before removing this compatibility mode." + ) + warnings.warn(message, DeprecationWarning, stacklevel=2) + logger.warning("DEPRECATION: %s", message) self.invoke_handler(self._handle_invoke) mark_feature_used(FeatureIndex.FOUNDRY_HOSTING) @@ -199,50 +260,167 @@ async def _request_session( f"session persistence also failed: {str(exc) or type(exc).__name__}" ) from exc + async def _parse(self, request: Request) -> InvocationRun: + if self._parse_request is not None: + result = self._parse_request(request) + parsed = await result if inspect.isawaitable(result) else result + if not isinstance(parsed, InvocationRun): + raise TypeError("parse_request must return InvocationRun.") + return parsed + + payload = await request.json() + if not isinstance(payload, dict): + raise ValueError("The invocation must be a JSON object.") + body = cast(Mapping[str, Any], payload) + message = body.get("message") + stream = body.get("stream", False) + options = body.get("options", {}) + if not isinstance(message, str): + raise ValueError("message must be a string.") + if not isinstance(stream, bool): + raise ValueError("stream must be a boolean.") + if not isinstance(options, dict): + raise ValueError("options must be an object.") + return InvocationRun( + messages=message if stream else [message], options=cast(Mapping[str, Any], options), stream=stream + ) + + async def _options(self, request: Request, parsed: InvocationRun) -> dict[str, Any]: + options = deepcopy(dict(parsed.options)) + if self._prepare_options is not None: + result = self._prepare_options(request, options) + if inspect.isawaitable(result): + result = await result + if not isinstance(result, Mapping) or any(not isinstance(key, str) for key in result): + raise TypeError("prepare_options must return a mapping of MAF run options with string keys.") + options = deepcopy(dict(result)) + validate_request_options(options) + return options + + def _agent_kwargs(self, agent: SupportsAgentRun, options: dict[str, Any]) -> dict[str, Any]: + if not options: + return {} + try: + inspect.signature(agent.run).bind_partial(options=options) + except (TypeError, ValueError): + if self._unsupported_options == "error": + raise _UnsupportedAgentOptions("The hosted agent does not accept caller runtime options.") from None + if self._unsupported_options == "warn": + logger.warning("Agent doesn't support runtime options. They will be ignored.") + return {} + return {"options": options} + + @staticmethod + async def _close_interrupted_stream(stream: object) -> None: + if isinstance(stream, ResponseStream): + close = getattr(cast(object, stream), "close", None) + if close is None: + logger.warning("The installed core cannot close an interrupted agent stream.") + return + else: + close = getattr(stream, "aclose", None) + if close is not None: + await close() + async def _handle_invoke(self, request: Request) -> Response: """Invoke the agent with the given request.""" context = get_request_context() try: hosted_scope = self._hosted_scope(request, context) if self.config.is_hosted else None partition_key = self._partition_key(context=context, scope=hosted_scope) - except Exception as e: - return Response(content=str(e), status_code=500) + except RuntimeError as exc: + logger.error("Failed to resolve Invocations session: %s", exc) + return JSONResponse({"error": str(exc)}, status_code=500) - data = await request.json() - - stream = data.get("stream", False) - user_message = data.get("message", None) - if user_message is None: - error = "Missing 'message' in request" - if stream: - return StreamingResponse(content=error, status_code=400) - return Response(content=error, status_code=400) + try: + parsed = await self._parse(request) + options = await self._options(request, parsed) + except (TypeError, ValueError) as exc: + return JSONResponse({"error": str(exc)}, status_code=400) + except Exception: + logger.exception("Failed to prepare Invocations request") + return JSONResponse({"error": "Failed to prepare invocation request."}, status_code=500) - if stream: + if parsed.stream: + if options and self._unsupported_options == "error" and is_agent(self._agent): + try: + self._agent_kwargs(self._agent, options) + except _UnsupportedAgentOptions as exc: + return JSONResponse({"error": str(exc)}, status_code=400) async def stream_response() -> AsyncGenerator[str]: - async with ( - self._request_session(partition_key, context, hosted_scope=hosted_scope) as session, - self._request_agent() as agent, - ): - stream = agent.run(user_message, session=session, stream=True) - try: - async for update in stream: - if update.text: - yield update.text - finally: - if isinstance(stream, ResponseStream): - await stream.close() - else: - close = getattr(stream, "aclose", None) - if close is not None: - await close() + token: Token[FoundryAgentRequestContext] | None = set_request_context(context) + try: + lock = self._session_locks.setdefault(partition_key, asyncio.Lock()) + async with lock, self._request_agent() as agent: + run_kwargs = self._agent_kwargs(agent, options) + async with self._request_session(partition_key, context, hosted_scope=hosted_scope) as session: + stream = agent.run(parsed.messages, session=session, stream=True, **run_kwargs) + completed = False + try: + async for update in stream: + if update.text: + frame = ( + update.text + if self._legacy_wire_format + else _sse("delta", {"text": update.text}) + ) + # The SDK may close a suspended generator from another task. + reset_request_context(token) + token = None + yield frame + token = set_request_context(context) + if isinstance(stream, ResponseStream): + await stream.get_final_response() + completed = True + finally: + if not completed: + if token is None: + token = set_request_context(context) + try: + await self._close_interrupted_stream(stream) + except Exception: + logger.exception("Failed to close interrupted Invocations agent stream") + if not self._legacy_wire_format: + session_id = partition_key[0] if isinstance(partition_key, tuple) else partition_key + if token is not None: + reset_request_context(token) + token = None + yield _sse("done", {"session_id": session_id}) + except _UnsupportedAgentOptions as exc: + if self._legacy_wire_format: + raise + if token is not None: + reset_request_context(token) + token = None + yield _sse("error", {"message": str(exc), "code": "unsupported_options", "status": 400}) + except Exception as exc: + logger.exception("Invocations agent stream failed") + if self._legacy_wire_format: + raise + message, status, code = _invocation_failure(exc) + if token is not None: + reset_request_context(token) + token = None + yield _sse("error", {"message": message, "code": code, "status": status}) + finally: + if token is not None: + reset_request_context(token) return _InvocationStreamingResponse(stream_response()) - async with ( - self._request_session(partition_key, context, hosted_scope=hosted_scope) as session, - self._request_agent() as agent, - ): - response = await agent.run([user_message], session=session) - return Response(content=response.text) + try: + lock = self._session_locks.setdefault(partition_key, asyncio.Lock()) + async with lock, self._request_agent() as agent: + run_kwargs = self._agent_kwargs(agent, options) + async with self._request_session(partition_key, context, hosted_scope=hosted_scope) as session: + response = await agent.run(parsed.messages, session=session, **run_kwargs) + if self._legacy_wire_format: + return Response(content=response.text) + return JSONResponse({"response": response.text}) + except _UnsupportedAgentOptions as exc: + return JSONResponse({"error": str(exc)}, status_code=400) + except Exception as exc: + logger.exception("Invocations agent request failed") + message, status, _ = _invocation_failure(exc) + return JSONResponse({"error": message}, status_code=status) diff --git a/python/packages/foundry_hosting/agent_framework_foundry_hosting/_request.py b/python/packages/foundry_hosting/agent_framework_foundry_hosting/_request.py index 1d645031436..c7e1fc89389 100644 --- a/python/packages/foundry_hosting/agent_framework_foundry_hosting/_request.py +++ b/python/packages/foundry_hosting/agent_framework_foundry_hosting/_request.py @@ -1,15 +1,17 @@ # Copyright (c) Microsoft. All rights reserved. -"""Request-scoped Responses options and a view for developer hooks.""" +"""Foundry request models and per-turn options for protocol hosts.""" from __future__ import annotations import inspect -from collections.abc import Awaitable, Callable, Mapping +from collections.abc import Awaitable, Callable, Mapping, Sequence from copy import deepcopy +from dataclasses import dataclass, field from types import MappingProxyType from typing import Any, Literal, TypeAlias, cast +from agent_framework import AgentRunInputs, Content, Message from azure.ai.agentserver.responses import ResponseContext from azure.ai.agentserver.responses.models import CreateResponse, Item @@ -17,6 +19,39 @@ UnsupportedOptions: TypeAlias = Literal["ignore", "warn", "error"] + +@dataclass(frozen=True) +class InvocationRun: + """Messages, per-turn options, and streaming intent parsed from an Invocations request. + + Args: + messages: MAF message input for this turn, including typed `Message` or `Content` values. + options: Caller generation options for this turn, separate from agent defaults. + stream: Whether to stream framed SSE events rather than return a JSON response. + + Raises: + TypeError: If the messages, options, or streaming intent have invalid types. + """ + + messages: AgentRunInputs + options: Mapping[str, Any] = field(default_factory=lambda: dict[str, Any]()) + stream: bool = False + + def __post_init__(self) -> None: + if not ( + isinstance(self.messages, (str, Content, Message)) + or ( + isinstance(self.messages, Sequence) + and all(isinstance(message, (str, Content, Message)) for message in self.messages) + ) + ): + raise TypeError("InvocationRun.messages must be a string, Content, Message, or a sequence of them.") + if not isinstance(self.options, Mapping) or any(not isinstance(key, str) for key in self.options): + raise TypeError("InvocationRun.options must be a mapping with string keys.") + if not isinstance(self.stream, bool): + raise TypeError("InvocationRun.stream must be a boolean.") + + _HOST_CONTROLLED_FIELDS = frozenset({ "agent", "agent_reference", diff --git a/python/packages/foundry_hosting/tests/test_invocations.py b/python/packages/foundry_hosting/tests/test_invocations.py index 6e7dd818f68..15028c37059 100644 --- a/python/packages/foundry_hosting/tests/test_invocations.py +++ b/python/packages/foundry_hosting/tests/test_invocations.py @@ -13,11 +13,13 @@ import asyncio import json +import logging +import uuid import weakref from collections.abc import AsyncIterator, Awaitable, Callable, Iterator, Mapping, Sequence from contextlib import contextmanager from itertools import product -from typing import cast +from typing import Literal, cast from unittest.mock import AsyncMock, MagicMock, patch import httpx @@ -53,7 +55,7 @@ from starlette.responses import Response, StreamingResponse from typing_extensions import Any -from agent_framework_foundry_hosting import InvocationsHostServer, StoreProvider +from agent_framework_foundry_hosting import InvocationRun, InvocationsHostServer, StoreProvider from agent_framework_foundry_hosting._state_store import FoundryAgentSessionStore pytestmark = pytest.mark.filterwarnings("ignore:.*SessionStore is experimental.*") @@ -98,6 +100,7 @@ def __init__( self._response = response self._stream_updates = stream_updates or [] self.calls: list[dict[str, Any]] = [] + self.default_options: dict[str, Any] = {} self._update_session = update_session def run( @@ -108,7 +111,12 @@ def run( session: AgentSession | None = None, **kwargs: Any, ) -> Any: - self.calls.append({"messages": messages, "stream": stream, "session": session}) + self.calls.append({ + "messages": messages, + "stream": stream, + "session": session, + "options": kwargs.get("options"), + }) if session is not None and self._update_session is not None: self._update_session(session) if stream: @@ -194,6 +202,39 @@ def _make_agent( return _FakeAgent(response=response, stream_updates=stream_updates, update_session=update_session) +class _NoOptionsAgent: + """Implement SupportsAgentRun without accepting the optional run-options argument.""" + + def __init__(self) -> None: + self._agent = _make_agent(response_text="ok", stream_texts=["ok"]) + self.id = self._agent.id + self.name = self._agent.name + self.description = self._agent.description + self.calls = self._agent.calls + + def run( + self, + messages: Any = None, + *, + stream: bool = False, + session: AgentSession | None = None, + function_invocation_kwargs: Mapping[str, Any] | None = None, + client_kwargs: Mapping[str, Any] | None = None, + ) -> Any: + return self._agent.run(messages, stream=stream, session=session) + + def create_session(self, *, session_id: str | None = None) -> AgentSession: + return self._agent.create_session(session_id=session_id) + + def get_session( + self, + service_session_id: str | ServiceSessionId, + *, + session_id: str | None = None, + ) -> AgentSession: + return self._agent.get_session(service_session_id, session_id=session_id) + + def _make_request(payload: dict[str, Any], *, agent_session_id: str | None = None) -> Request: """Build a mock Starlette request whose ``json()`` returns ``payload``.""" request = MagicMock(spec=Request) @@ -227,6 +268,26 @@ async def _collect_stream(response: StreamingResponse) -> str: return "".join(chunks) +def _parse_sse_events(body: str) -> list[tuple[str, dict[str, Any]]]: + events: list[tuple[str, dict[str, Any]]] = [] + for frame in body.split("\n\n"): + if frame.startswith("event: "): + event, data = frame.split("\n", 1) + events.append((event.removeprefix("event: "), json.loads(data.removeprefix("data: ")))) + return events + + +async def _success_text(response: Response) -> str: + assert response.status_code == 200 + if isinstance(response, StreamingResponse): + events = _parse_sse_events(await _collect_stream(response)) + assert [event for event, _ in events][-1] == "done" + assert all(event in ("delta", "done") for event, _ in events) + return "".join(data["text"] for event, data in events if event == "delta") + assert response.media_type == "application/json" + return json.loads(bytes(response.body))["response"] + + # endregion @@ -293,10 +354,7 @@ def update_session(session: AgentSession) -> None: response = await server._handle_invoke( # pyright: ignore[reportPrivateUsage] _make_request({"message": "next", "stream": stream}, agent_session_id=sandbox_id) ) - if isinstance(response, StreamingResponse): - assert await _collect_stream(response) == "ok" - else: - assert bytes(response.body).decode() == "ok" + assert await _success_text(response) == "ok" assert agent.calls[-1]["session"].state["turn"] == round_number assert len(store_names) == 12 @@ -353,7 +411,7 @@ async def receive() -> Any: async def send(message: Any) -> None: expected = b": keep-alive\n\n" if keep_alive else b"first" - if message["type"] == "http.response.body" and message.get("body") == expected: + if message["type"] == "http.response.body" and expected in message.get("body", b""): assert entered.is_set() if spec_version == "2.4": raise OSError("disconnected while idle") @@ -433,10 +491,7 @@ async def complete() -> ChatResponse: response = await server._handle_invoke( # pyright: ignore[reportPrivateUsage] _make_request({"message": message, "stream": stream}) ) - if isinstance(response, StreamingResponse): - assert await _collect_stream(response) == "ok" - else: - assert bytes(response.body).decode() == "ok" + assert await _success_text(response) == "ok" assert executed == ["called"] assert len(client.calls) == 3 @@ -493,17 +548,21 @@ def update_session(session: AgentSession) -> None: agent = _make_agent(response_text="ok", stream_texts=["ok"], update_session=update_session) server = InvocationsHostServer(agent, agent_session_store_provider=_SessionStoreProvider(store)) - expected = ( - "Invocation failed: run failed; session persistence also failed: save failed" - if failure == "run_and_save" - else (f"{failure} failed") - ) - with _request_context(session_id="session"), pytest.raises((ValueError, RuntimeError), match=expected): + with _request_context(session_id="session"): result = await server._handle_invoke( # pyright: ignore[reportPrivateUsage] _make_request({"message": "hello", "stream": stream}) ) if isinstance(result, StreamingResponse): - await _collect_stream(result) + events = _parse_sse_events(await _collect_stream(result)) + assert events[-1] == ( + "error", + {"message": "Agent invocation failed.", "code": "invocation_failed", "status": 500}, + ) + assert not any(event == "done" for event, _ in events) + assert [event for event, _ in events[:-1]] == (["delta"] if failure == "save" else []) + else: + assert result.status_code == 500 + assert json.loads(bytes(result.body)) == {"error": "Agent invocation failed."} if failure == "load": assert agent.calls == [] @@ -575,6 +634,606 @@ def create_agent(name: str) -> _FakeAgent: with pytest.raises(TypeError, match="agent callable must accept no arguments"): InvocationsHostServer(cast(Any, create_agent)) + @pytest.mark.parametrize( + ("parameter", "value", "error"), + [ + ("parse_request", 42, "parse_request must be a callable"), + ("prepare_options", 42, "prepare_options must be a callable"), + ("unsupported_options", "discard", "unsupported_options"), + ("legacy_wire_format", "false", "legacy_wire_format must be a boolean"), + ], + ) + def test_rejects_invalid_configuration(self, parameter: str, value: Any, error: str) -> None: + with pytest.raises((TypeError, ValueError), match=error): + InvocationsHostServer(_make_agent(response_text="ok"), **{parameter: value}) + + +# endregion + + +class TestParsedRequests: + @pytest.mark.parametrize("stream", [False, True]) + async def test_parser_and_hook_use_typed_messages_without_mutating_defaults(self, stream: bool) -> None: + source_options: dict[str, Any] = {"temperature": 0.8, "nested": {"tag": "original"}, "store": True} + agent = _make_agent(response_text="ok", stream_texts=["ok"]) + agent.default_options = {"store": False, "temperature": 0.2} + request = _make_request({"prompt": "hello"}) + + async def parse(incoming: Request) -> InvocationRun: + assert incoming is request + body = await incoming.json() + return InvocationRun( + messages=[Message(role="user", contents=[Content.from_text(body["prompt"])])], + options=source_options, + stream=stream, + ) + + async def prepare(incoming: Request, options: dict[str, Any]) -> dict[str, Any]: + assert incoming is request + options["nested"]["tag"] = "changed" + options.pop("store") + return {"temperature": options["temperature"]} + + server = InvocationsHostServer(agent, parse_request=parse, prepare_options=prepare) + with _request_context(session_id="parsed"): + response = await server._handle_invoke(request) # pyright: ignore[reportPrivateUsage] + assert await _success_text(response) == "ok" + + assert agent.calls[0]["messages"][0].text == "hello" + assert agent.calls[0]["options"] == {"temperature": 0.8} + assert source_options == {"temperature": 0.8, "nested": {"tag": "original"}, "store": True} + assert agent.default_options == {"store": False, "temperature": 0.2} + + @pytest.mark.parametrize( + ("payload", "error"), + [ + ([], "JSON object"), + ({}, "message"), + ({"message": None}, "message"), + ({"message": "hello", "stream": "true"}, "stream"), + ({"message": "hello", "options": []}, "options"), + ], + ) + async def test_default_parser_rejects_invalid_payload_without_using_storage(self, payload: Any, error: str) -> None: + agent = _make_agent(response_text="ok") + provider = _SessionStoreProvider(_mock_session_store()) + server = InvocationsHostServer(agent, agent_session_store_provider=provider) + request = _make_request(payload) + with _request_context(session_id="invalid"): + response = await server._handle_invoke(request) # pyright: ignore[reportPrivateUsage] + assert response.status_code == 400 + assert error in json.loads(bytes(response.body))["error"] + assert agent.calls == [] + assert provider.contexts == [] + + @pytest.mark.parametrize("result", [None, "hello", {"messages": "hello"}]) + async def test_custom_parser_must_return_invocation_run(self, result: Any) -> None: + agent = _make_agent(response_text="ok") + server = InvocationsHostServer(agent, parse_request=lambda _request: result) + with _request_context(session_id="invalid"): + response = await server._handle_invoke(_make_request({"message": "hi"})) # pyright: ignore[reportPrivateUsage] + assert response.status_code == 400 + assert json.loads(bytes(response.body)) == {"error": "parse_request must return InvocationRun."} + assert agent.calls == [] + + async def test_parser_failure_is_logged_but_not_exposed(self, caplog: pytest.LogCaptureFixture) -> None: + def broken_parser(_request: Request) -> InvocationRun: + raise RuntimeError("internal parser secret") + + server = InvocationsHostServer(_make_agent(response_text="ok"), parse_request=broken_parser) + with _request_context(session_id="invalid"): + response = await server._handle_invoke(_make_request({"message": "hi"})) # pyright: ignore[reportPrivateUsage] + assert response.status_code == 500 + assert json.loads(bytes(response.body)) == {"error": "Failed to prepare invocation request."} + assert "internal parser secret" in caplog.text + + @pytest.mark.parametrize( + "options", + [ + {"store": True}, + {"session_id": "forged"}, + {"agent_session_id": "forged"}, + {"previous_response_id": "private"}, + {"continuation_token": {"secret": "private"}}, + {"extra_body": {"store": True}}, + {"user": "forged"}, + ], + ) + async def test_reserved_options_are_rejected_before_agent_run(self, options: dict[str, Any]) -> None: + agent = _make_agent(response_text="ok") + server = InvocationsHostServer(agent) + with _request_context(session_id="options"): + response = await server._handle_invoke( # pyright: ignore[reportPrivateUsage] + _make_request({"message": "hi", "options": options}) + ) + assert response.status_code == 400 + assert "host-controlled fields" in json.loads(bytes(response.body))["error"] + assert agent.calls == [] + + @pytest.mark.parametrize("result", [None, [], {1: "invalid"}, {"session_id": "forged"}]) + async def test_options_hook_must_return_safe_string_keyed_mapping(self, result: Any) -> None: + agent = _make_agent(response_text="ok") + server = InvocationsHostServer(agent, prepare_options=lambda _request, _options: result) + with _request_context(session_id="invalid-options"): + response = await server._handle_invoke(_make_request({"message": "hi"})) # pyright: ignore[reportPrivateUsage] + assert response.status_code == 400 + assert agent.calls == [] + + @pytest.mark.parametrize("stream", [False, True]) + @pytest.mark.parametrize("policy", ["ignore", "warn", "error"]) + async def test_unsupported_options_policy( + self, policy: Literal["ignore", "warn", "error"], stream: bool, caplog: pytest.LogCaptureFixture + ) -> None: + agent = _NoOptionsAgent() + server = InvocationsHostServer(agent, unsupported_options=policy) + with _request_context(session_id="unsupported"): + response = await server._handle_invoke( # pyright: ignore[reportPrivateUsage] + _make_request({"message": "hello", "options": {"temperature": 0.8}, "stream": stream}) + ) + if policy == "error": + assert response.status_code == 400 + assert not isinstance(response, StreamingResponse) + assert "does not accept" in json.loads(bytes(response.body))["error"] + assert agent.calls == [] + else: + assert await _success_text(response) == "ok" + assert agent.calls[-1]["options"] is None + + assert ("Agent doesn't support runtime options" in caplog.text) is (policy == "warn") + + async def test_factory_options_failure_is_framed_before_first_delta(self) -> None: + provider = _SessionStoreProvider(_mock_session_store()) + server = InvocationsHostServer( + lambda: _NoOptionsAgent(), + agent_session_store_provider=provider, + unsupported_options="error", + ) + with _request_context(session_id="factory-options"): + response = await server._handle_invoke( # pyright: ignore[reportPrivateUsage] + _make_request({"message": "hi", "stream": True, "options": {"temperature": 0.8}}) + ) + assert isinstance(response, StreamingResponse) + assert _parse_sse_events(await _collect_stream(response)) == [ + ( + "error", + { + "message": "The hosted agent does not accept caller runtime options.", + "code": "unsupported_options", + "status": 400, + }, + ) + ] + assert provider.contexts == [] + + +class TestSerializedSessions: + @pytest.mark.parametrize("stream", [False, True]) + async def test_same_host_serializes_overlapping_turns(self, stream: bool) -> None: + first_started = asyncio.Event() + release_first = asyncio.Event() + started: list[int] = [] + store = SessionStore() + + class CountingAgent(_FakeAgent): + def run( + self, + messages: Any = None, + *, + stream: bool = False, + session: AgentSession | None = None, + **kwargs: Any, + ) -> Any: + assert session is not None + + async def run_once() -> AgentResponse: + turn = session.state.get("turn", 0) + 1 + started.append(turn) + if turn == 1: + first_started.set() + await release_first.wait() + session.state["turn"] = turn + return AgentResponse(messages=[Message("assistant", str(turn))]) + + async def updates() -> AsyncIterator[AgentResponseUpdate]: + turn = session.state.get("turn", 0) + 1 + started.append(turn) + if turn == 1: + first_started.set() + await release_first.wait() + session.state["turn"] = turn + yield AgentResponseUpdate(contents=[Content.from_text(str(turn))]) + + return updates() if stream else run_once() + + server = InvocationsHostServer(CountingAgent(), agent_session_store_provider=_SessionStoreProvider(store)) + + async def invoke() -> str: + with _request_context(session_id="serialized"): + response = await server._handle_invoke( # pyright: ignore[reportPrivateUsage] + _make_request({"message": "next", "stream": stream}) + ) + return await _success_text(response) + + tasks = [asyncio.create_task(invoke())] + try: + await asyncio.wait_for(first_started.wait(), timeout=5) + tasks.append(asyncio.create_task(invoke())) + await asyncio.sleep(0) + assert started == [1] + release_first.set() + assert await asyncio.wait_for(asyncio.gather(*tasks), timeout=5) == ["1", "2"] + assert started == [1, 2] + saved = await store.get("serialized") + assert saved is not None + assert saved.state == {"turn": 2} + finally: + release_first.set() + for task in tasks: + if not task.done(): + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) + + async def test_disconnect_holds_session_lock_until_partial_save_finishes(self) -> None: + save_started = asyncio.Event() + release_save = asyncio.Event() + + class BlockingStore(SessionStore): + def __init__(self) -> None: + super().__init__() + self.saves = 0 + + async def set(self, session_id: str, session: AgentSession) -> None: + self.saves += 1 + if self.saves == 1: + save_started.set() + await release_save.wait() + await super().set(session_id, session) + + def increment(session: AgentSession) -> None: + session.state["turn"] = session.state.get("turn", 0) + 1 + + store = BlockingStore() + agent = _make_agent(response_text="ok", stream_texts=["partial"], update_session=increment) + server = InvocationsHostServer(agent, agent_session_store_provider=_SessionStoreProvider(store)) + + async def invoke_next() -> str: + with _request_context(session_id="disconnect-serialized"): + response = await server._handle_invoke( # pyright: ignore[reportPrivateUsage] + _make_request({"message": "next"}) + ) + return await _success_text(response) + + with _request_context(session_id="disconnect-serialized"): + response = await server._handle_invoke( # pyright: ignore[reportPrivateUsage] + _make_request({"message": "first", "stream": True}) + ) + assert isinstance(response, StreamingResponse) + iterator = cast(Any, response.body_iterator) + assert _parse_sse_events(await anext(iterator)) == [("delta", {"text": "partial"})] + closing = asyncio.create_task(iterator.aclose()) + next_turn: asyncio.Task[str] | None = None + try: + await asyncio.wait_for(save_started.wait(), timeout=5) + next_turn = asyncio.create_task(invoke_next()) + await asyncio.sleep(0) + assert len(agent.calls) == 1 + release_save.set() + await asyncio.wait_for(closing, timeout=5) + assert await asyncio.wait_for(next_turn, timeout=5) == "ok" + assert agent.calls[-1]["session"].state == {"turn": 2} + finally: + release_save.set() + if not closing.done(): + closing.cancel() + if next_turn is not None and not next_turn.done(): + next_turn.cancel() + await asyncio.gather(closing, *(task for task in (next_turn,) if task is not None), return_exceptions=True) + + async def test_http_route_restores_session_across_hosts_and_factory_agents(self) -> None: + store = SessionStore() + agents: list[_FakeAgent] = [] + + def increment(session: AgentSession) -> None: + session.state["turn"] = session.state.get("turn", 0) + 1 + + def create_agent() -> _FakeAgent: + agent = _make_agent(response_text="ok", stream_texts=["ok"], update_session=increment) + agents.append(agent) + return agent + + session_id = f"asgi-{uuid.uuid4().hex}" + hosts = [ + InvocationsHostServer(create_agent, agent_session_store_provider=_SessionStoreProvider(store)), + InvocationsHostServer(create_agent, agent_session_store_provider=_SessionStoreProvider(store)), + ] + for index, host in enumerate(hosts): + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=host), base_url="http://test") as client: + response = await client.post( + "/invocations", + params={"agent_session_id": session_id}, + json={"message": f"turn-{index + 1}"}, + ) + assert response.status_code == 200 + assert response.json() == {"response": "ok"} + assert response.headers["x-agent-session-id"] == session_id + assert agents[-1].calls[0]["session"].state == {"turn": index + 1} + + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=hosts[1]), base_url="http://test") as client: + streamed = await client.post( + "/invocations", + params={"agent_session_id": session_id}, + json={"message": "turn-3", "stream": True}, + ) + assert streamed.status_code == 200 + assert _parse_sse_events(streamed.text) == [ + ("delta", {"text": "ok"}), + ("done", {"session_id": session_id}), + ] + assert agents[-1].calls[0]["session"].state == {"turn": 3} + assert len(agents) == 3 + saved = await store.get(session_id) + assert saved is not None + assert saved.state == {"turn": 3} + + +class TestStreamingFailuresAndCompatibility: + async def test_completed_response_stream_finalizes_on_released_core_without_close( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + store = SessionStore() + + class FinalizingAgent(_FakeAgent): + def run( + self, messages: Any = None, *, stream: bool = False, session: AgentSession | None = None, **kwargs: Any + ) -> Any: + assert session is not None + session.state["started"] = True + + async def updates() -> AsyncIterator[AgentResponseUpdate]: + yield AgentResponseUpdate(contents=[Content.from_text("ready")], role="assistant") + + def record_finalization(response: AgentResponse[Any]) -> None: + session.state["finalized"] = True + + return ResponseStream( + updates(), + finalizer=AgentResponse.from_updates, + result_hooks=[record_finalization], + ) + + monkeypatch.delattr(ResponseStream, "close") + server = InvocationsHostServer(FinalizingAgent(), agent_session_store_provider=_SessionStoreProvider(store)) + with _request_context(session_id="finalize"): + response = await server._handle_invoke( # pyright: ignore[reportPrivateUsage] + _make_request({"message": "hello", "stream": True}) + ) + assert await _success_text(response) == "ready" + + persisted = await store.get("finalize") + assert persisted is not None + assert persisted.state == {"started": True, "finalized": True} + + @pytest.mark.parametrize("emit_delta", [False, True]) + async def test_provider_stream_failure_is_framed_and_persists_safe_state( + self, emit_delta: bool, caplog: pytest.LogCaptureFixture + ) -> None: + store = SessionStore() + + class FailingAgent(_FakeAgent): + def run( + self, messages: Any = None, *, stream: bool = False, session: AgentSession | None = None, **kwargs: Any + ) -> Any: + async def updates() -> AsyncIterator[AgentResponseUpdate]: + assert session is not None + session.state["partial"] = True + if emit_delta: + yield AgentResponseUpdate(contents=[Content.from_text("partial")]) + raise ValueError("private provider details") + + return updates() + + server = InvocationsHostServer(FailingAgent(), agent_session_store_provider=_SessionStoreProvider(store)) + with _request_context(session_id="partial"): + response = await server._handle_invoke( # pyright: ignore[reportPrivateUsage] + _make_request({"message": "hello", "stream": True}) + ) + assert isinstance(response, StreamingResponse) + body = await _collect_stream(response) + + events = _parse_sse_events(body) + assert events == ([("delta", {"text": "partial"})] if emit_delta else []) + [ + ("error", {"message": "Agent invocation failed.", "code": "invocation_failed", "status": 500}) + ] + assert "private provider details" not in body + assert "private provider details" in caplog.text + persisted = await store.get("partial") + assert persisted is not None + assert persisted.state == {"partial": True} + + async def test_factory_failure_before_first_sse_event_is_sanitized(self, caplog: pytest.LogCaptureFixture) -> None: + async def unavailable_agent() -> _FakeAgent: + raise RuntimeError("secret agent initialization failure") + + provider = _SessionStoreProvider(_mock_session_store()) + server = InvocationsHostServer(unavailable_agent, agent_session_store_provider=provider) + with _request_context(session_id="factory-failure"): + response = await server._handle_invoke( # pyright: ignore[reportPrivateUsage] + _make_request({"message": "hello", "stream": True}) + ) + assert isinstance(response, StreamingResponse) + assert _parse_sse_events(await _collect_stream(response)) == [ + ("error", {"message": "Agent invocation failed.", "code": "invocation_failed", "status": 500}) + ] + assert "secret agent initialization failure" in caplog.text + assert provider.contexts == [] + + async def test_http_stream_failure_before_first_delta_is_an_error_event(self) -> None: + async def unavailable_agent() -> _FakeAgent: + raise RuntimeError("private factory details") + + server = InvocationsHostServer(unavailable_agent) + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=server), base_url="http://test") as client: + response = await client.post( + "/invocations", + params={"agent_session_id": f"failure-{uuid.uuid4().hex}"}, + json={"message": "hello", "stream": True}, + ) + assert response.status_code == 200 + assert response.headers["content-type"].startswith("text/event-stream") + assert _parse_sse_events(response.text) == [ + ("error", {"message": "Agent invocation failed.", "code": "invocation_failed", "status": 500}) + ] + assert "private factory details" not in response.text + + async def test_cross_task_stream_close_restores_context_and_saves_session(self) -> None: + store = SessionStore() + + class ContextAgent(_FakeAgent): + def run( + self, messages: Any = None, *, stream: bool = False, session: AgentSession | None = None, **kwargs: Any + ) -> Any: + async def updates() -> AsyncIterator[AgentResponseUpdate]: + assert session is not None + assert get_request_context().session_id == "cross-task" + session.state["started"] = True + try: + yield AgentResponseUpdate(contents=[Content.from_text("first")]) + finally: + assert get_request_context().session_id == "cross-task" + session.state["closed"] = True + + return updates() + + server = InvocationsHostServer(ContextAgent(), agent_session_store_provider=_SessionStoreProvider(store)) + with _request_context(session_id="cross-task"): + response = await server._handle_invoke( # pyright: ignore[reportPrivateUsage] + _make_request({"message": "hello", "stream": True}) + ) + assert isinstance(response, StreamingResponse) + previous_id = get_request_context().session_id + iterator = cast(Any, response.body_iterator) + assert _parse_sse_events(await anext(iterator)) == [("delta", {"text": "first"})] + assert get_request_context().session_id == previous_id + await asyncio.create_task(iterator.aclose()) + assert get_request_context().session_id == previous_id + persisted = await store.get("cross-task") + assert persisted is not None + assert persisted.state == {"started": True, "closed": True} + + async def test_finalizer_failure_after_delta_never_emits_done(self) -> None: + store = SessionStore() + + class FailingFinalizerAgent(_FakeAgent): + def run( + self, messages: Any = None, *, stream: bool = False, session: AgentSession | None = None, **kwargs: Any + ) -> Any: + async def updates() -> AsyncIterator[AgentResponseUpdate]: + assert session is not None + session.state["partial"] = True + yield AgentResponseUpdate(contents=[Content.from_text("one")], role="assistant") + + def fail_finalization(_updates: Sequence[AgentResponseUpdate]) -> AgentResponse: + raise RuntimeError("private finalizer details") + + return ResponseStream(updates(), finalizer=fail_finalization) + + server = InvocationsHostServer( + FailingFinalizerAgent(), agent_session_store_provider=_SessionStoreProvider(store) + ) + with _request_context(session_id="finalizer-failure"): + response = await server._handle_invoke( # pyright: ignore[reportPrivateUsage] + _make_request({"message": "hello", "stream": True}) + ) + assert isinstance(response, StreamingResponse) + events = _parse_sse_events(await _collect_stream(response)) + assert events == [ + ("delta", {"text": "one"}), + ("error", {"message": "Agent invocation failed.", "code": "invocation_failed", "status": 500}), + ] + persisted = await store.get("finalizer-failure") + assert persisted is not None + assert persisted.state == {"partial": True} + + async def test_save_failure_after_delta_does_not_report_success(self, caplog: pytest.LogCaptureFixture) -> None: + store = _mock_session_store() + store.set.side_effect = RuntimeError("private write failure") + server = InvocationsHostServer( + _make_agent(stream_texts=["one"]), agent_session_store_provider=_SessionStoreProvider(store) + ) + with _request_context(session_id="save-failure"): + response = await server._handle_invoke( # pyright: ignore[reportPrivateUsage] + _make_request({"message": "hello", "stream": True}) + ) + assert isinstance(response, StreamingResponse) + events = _parse_sse_events(await _collect_stream(response)) + assert events == [ + ("delta", {"text": "one"}), + ("error", {"message": "Agent invocation failed.", "code": "invocation_failed", "status": 500}), + ] + assert "private write failure" in caplog.text + + async def test_invalid_json_uses_client_status_on_actual_http_route(self) -> None: + agent = _make_agent(response_text="ok") + server = InvocationsHostServer(agent) + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=server), base_url="http://test") as client: + response = await client.post( + "/invocations", + params={"agent_session_id": f"invalid-{uuid.uuid4().hex}"}, + content=b'{"message":', + headers={"content-type": "application/json"}, + ) + assert response.status_code == 400 + assert response.headers["content-type"] == "application/json" + assert "error" in response.json() + assert agent.calls == [] + + async def test_legacy_wire_format_is_explicit_once_warned_and_unchanged( + self, caplog: pytest.LogCaptureFixture + ) -> None: + with pytest.warns(DeprecationWarning, match="legacy_wire_format"): + server = InvocationsHostServer( + _make_agent(response_text="Hello!", stream_texts=["Hel", "lo", "!"]), legacy_wire_format=True + ) + with _request_context(session_id="legacy"): + for _ in range(2): + plain = await server._handle_invoke(_make_request({"message": "hi"})) # pyright: ignore[reportPrivateUsage] + assert bytes(plain.body).decode() == "Hello!" + streamed = await server._handle_invoke( # pyright: ignore[reportPrivateUsage] + _make_request({"message": "hi", "stream": True}) + ) + assert isinstance(streamed, StreamingResponse) + assert await _collect_stream(streamed) == "Hello!" + + deprecations = [ + record + for record in caplog.records + if record.levelno == logging.WARNING and "DEPRECATION:" in record.message + ] + assert len(deprecations) == 1 + + async def test_legacy_stream_failure_terminates_instead_of_sending_new_frame(self) -> None: + class FailingAgent(_FakeAgent): + def run( + self, messages: Any = None, *, stream: bool = False, session: AgentSession | None = None, **kwargs: Any + ) -> Any: + async def updates() -> AsyncIterator[AgentResponseUpdate]: + yield AgentResponseUpdate(contents=[Content.from_text("first")]) + raise RuntimeError("provider failed") + + return updates() + + with pytest.warns(DeprecationWarning): + server = InvocationsHostServer(FailingAgent(), legacy_wire_format=True) + with _request_context(session_id="legacy-failure"): + response = await server._handle_invoke( # pyright: ignore[reportPrivateUsage] + _make_request({"message": "hi", "stream": True}) + ) + assert isinstance(response, StreamingResponse) + iterator = cast(Any, response.body_iterator) + assert await anext(iterator) == "first" + with pytest.raises(RuntimeError, match="provider failed"): + await anext(iterator) + # endregion @@ -725,10 +1384,7 @@ async def test_hosted_routed_query_works_without_configured_session(self, stream with _request_context(call_id="call-1", user_id="user-1", session_id="sandbox-1"): response = await server._handle_invoke(request) # pyright: ignore[reportPrivateUsage] - if isinstance(response, StreamingResponse): - assert await _collect_stream(response) == "ok" - else: - assert bytes(response.body).decode() == "ok" + assert await _success_text(response) == "ok" assert agent.calls[0]["session"].session_id == '["sandbox-1","user-1"]' @@ -748,9 +1404,7 @@ async def test_hosted_routed_stream_captures_request_context(self, stream: bool) response = await server._handle_invoke(request) # pyright: ignore[reportPrivateUsage] if isinstance(response, StreamingResponse): assert provider.contexts == [] - assert await _collect_stream(response) == "ok" - else: - assert bytes(response.body).decode() == "ok" + assert await _success_text(response) == "ok" assert provider.contexts == [context] store.set.assert_awaited_once() @@ -855,23 +1509,48 @@ async def updates() -> AsyncIterator[AgentResponseUpdate]: host.config.is_hosted = True host.config.session_id = "" - async def invoke(host: InvocationsHostServer, call_id: str) -> str: + async def invoke(host: InvocationsHostServer, call_id: str) -> tuple[int, str]: with _request_context(call_id=call_id, user_id=user_id, session_id=sandbox_id): response = await host._handle_invoke( # pyright: ignore[reportPrivateUsage] _make_request({"message": "next", "stream": stream}, agent_session_id=sandbox_id) ) if isinstance(response, StreamingResponse): - return await _collect_stream(response) - return bytes(response.body).decode() + return response.status_code, await _collect_stream(response) + return response.status_code, bytes(response.body).decode() tasks = [asyncio.create_task(invoke(host, f"call-{index + 1}")) for index, host in enumerate(hosts)] try: await asyncio.wait_for(asyncio.gather(*(event.wait() for event in started)), timeout=5) released[0].set() - assert await asyncio.wait_for(tasks[0], timeout=5) == "ok" + first_status, first_body = await asyncio.wait_for(tasks[0], timeout=5) + assert first_status == 200 + if stream: + assert _parse_sse_events(first_body)[-1][0] == "done" + else: + assert json.loads(first_body) == {"response": "ok"} released[1].set() - with pytest.raises(RuntimeError, match="Another request advanced this agent session"): - await asyncio.wait_for(tasks[1], timeout=5) + second_status, second_body = await asyncio.wait_for(tasks[1], timeout=5) + if stream: + assert second_status == 200 + assert _parse_sse_events(second_body) == [ + ( + "delta", + {"text": "ok"}, + ), + ( + "error", + { + "message": "Another request advanced this agent session; reload before retrying.", + "code": "session_conflict", + "status": 409, + }, + ), + ] + else: + assert second_status == 409 + assert json.loads(second_body) == { + "error": "Another request advanced this agent session; reload before retrying." + } finally: for event in released: event.set() @@ -940,7 +1619,7 @@ async def test_factory_agent_context_lifetime_non_streaming(self) -> None: with _request_context(session_id="sess-1"): result = await server._handle_invoke(_make_request({"message": "one"})) # pyright: ignore[reportPrivateUsage] - assert bytes(result.body).decode() == "ok" + assert await _success_text(result) == "ok" assert events == ["enter", "run", "exit"] async def test_factory_agent_context_lifetime_until_stream_closes(self) -> None: @@ -958,7 +1637,7 @@ async def test_factory_agent_context_lifetime_until_stream_closes(self) -> None: assert isinstance(response, StreamingResponse) iterator = cast(Any, response.body_iterator) - assert await anext(iterator) == "one" + assert _parse_sse_events(await anext(iterator)) == [("delta", {"text": "one"})] await iterator.aclose() assert events == ["enter", "run", "stream_close", "exit"] @@ -981,8 +1660,8 @@ def create_agent() -> _FakeAgent: _make_request({"message": "two"}) ) - assert bytes(first.body).decode() == "agent-1" - assert bytes(second.body).decode() == "agent-2" + assert await _success_text(first) == "agent-1" + assert await _success_text(second) == "agent-2" assert len(agents) == 2 assert agents[0] is not agents[1] assert agents[0].calls[0]["session"] is not agents[1].calls[0]["session"] @@ -1005,10 +1684,7 @@ async def test_reusing_session_restores_identifier(self, hosted: bool, stream: b ): for _ in range(2): response = await server._handle_invoke(request) # pyright: ignore[reportPrivateUsage] - if isinstance(response, StreamingResponse): - assert await _collect_stream(response) == "ok" - else: - assert bytes(response.body).decode() == "ok" + assert await _success_text(response) == "ok" assert agent.calls[0]["session"] is not agent.calls[1]["session"] assert agent.calls[0]["session"].session_id == agent.calls[1]["session"].session_id == expected_id @@ -1027,7 +1703,8 @@ async def test_missing_message_streaming_returns_400(self) -> None: request = _make_request({"stream": True}) with _request_context(session_id="sess-1"): response = await server._handle_invoke(request) # pyright: ignore[reportPrivateUsage] - assert isinstance(response, StreamingResponse) + assert isinstance(response, Response) + assert response.media_type == "application/json" assert response.status_code == 400 async def test_partition_key_failure_returns_500(self) -> None: @@ -1068,7 +1745,7 @@ async def test_non_streaming_returns_agent_text(self) -> None: assert isinstance(response, Response) assert response.status_code == 200 - assert bytes(response.body).decode() == "Hello!" + assert await _success_text(response) == "Hello!" # Agent is called with the message wrapped in a list and the restored session. assert agent.calls[0]["messages"] == ["Hi"] assert agent.calls[0]["stream"] is False @@ -1083,7 +1760,12 @@ async def test_streaming_yields_update_text(self) -> None: assert isinstance(response, StreamingResponse) assert response.media_type == "text/event-stream" - assert await _collect_stream(response) == "Hello!" + assert _parse_sse_events(await _collect_stream(response)) == [ + ("delta", {"text": "Hel"}), + ("delta", {"text": "lo"}), + ("delta", {"text": "!"}), + ("done", {"session_id": "sess-1"}), + ] assert agent.calls[0]["messages"] == "Hi" assert agent.calls[0]["stream"] is True @@ -1138,10 +1820,7 @@ def update_session(session: AgentSession) -> None: response = await server._handle_invoke( # pyright: ignore[reportPrivateUsage] _make_request({"message": "Hi", "stream": stream}) ) - if isinstance(response, StreamingResponse): - assert await _collect_stream(response) == "ok" - else: - assert bytes(response.body).decode() == "ok" + assert await _success_text(response) == "ok" assert response.status_code == 200 session = agent.calls[-1]["session"] @@ -1158,10 +1837,7 @@ def update_session(session: AgentSession) -> None: response = await server._handle_invoke( # pyright: ignore[reportPrivateUsage] _make_request({"message": "Continue", "stream": stream}) ) - if isinstance(response, StreamingResponse): - assert await _collect_stream(response) == "ok" - else: - assert bytes(response.body).decode() == "ok" + assert await _success_text(response) == "ok" assert response.status_code == 200 assert agent.calls[-1]["session"] is not session diff --git a/python/packages/foundry_hosting/tests/test_invocations_int.py b/python/packages/foundry_hosting/tests/test_invocations_int.py index 136760da74f..ba6f6126a44 100644 --- a/python/packages/foundry_hosting/tests/test_invocations_int.py +++ b/python/packages/foundry_hosting/tests/test_invocations_int.py @@ -6,11 +6,10 @@ ASGITransport — no real server process is started. The agent talks to a real Foundry project endpoint so every test requires valid credentials. -The invocations protocol is intentionally simple: a request is a JSON body with -a ``message`` field (and an optional ``stream`` flag). Non-streaming responses -return the agent's answer as plain text; streaming responses return the answer -as a ``text/event-stream`` of text chunks. Session continuity is keyed off the -``agent_session_id`` query parameter. +The default request is a JSON object with a ``message`` field and optional +``stream`` flag. Non-streaming responses return ``{"response": "..."}``; +streaming responses contain framed ``delta`` and ``done`` SSE events. +Session continuity is keyed off the ``agent_session_id`` query parameter. Required environment variables: FOUNDRY_PROJECT_ENDPOINT - The Azure AI Foundry project endpoint URL. @@ -20,11 +19,12 @@ from __future__ import annotations import os +import uuid from typing import Annotated, Any import httpx import pytest -from agent_framework import Agent, tool +from agent_framework import Agent, InMemoryHistoryProvider, tool from agent_framework.foundry import FoundryChatClient from azure.identity import AzureCliCredential @@ -54,6 +54,7 @@ def server() -> InvocationsHostServer: agent = Agent( client=client, # ty: ignore[invalid-argument-type] instructions="You are a concise assistant. Keep answers very short (one or two sentences).", + context_providers=[InMemoryHistoryProvider()], default_options={"store": False}, # pyrefly: ignore[bad-argument-type] ) @@ -121,18 +122,21 @@ async def test_simple_text_non_streaming(self, server: InvocationsHostServer) -> resp = await _post_invocation(server, message="Say hello in exactly three words.", stream=False) assert resp.status_code == 200 - assert len(resp.text) > 0 + assert resp.headers["content-type"] == "application/json" + assert resp.json()["response"] @pytest.mark.flaky @pytest.mark.integration @skip_if_foundry_hosting_integration_tests_disabled async def test_simple_text_streaming(self, server: InvocationsHostServer) -> None: - """Streaming: send a message and receive text chunks as an event stream.""" + """Streaming: send a message and receive framed SSE events.""" resp = await _post_invocation(server, message="Say hello in exactly three words.", stream=True) assert resp.status_code == 200 assert "text/event-stream" in resp.headers["content-type"] - assert len(resp.text) > 0 + assert "event: delta" in resp.text + assert "event: done" in resp.text + assert "event: error" not in resp.text @pytest.mark.flaky @pytest.mark.integration @@ -144,6 +148,7 @@ async def test_missing_message_returns_400(self, server: InvocationsHostServer) resp = await client.post("/invocations", json={"stream": False}, timeout=120) assert resp.status_code == 400 + assert "message" in resp.json()["error"] # --------------------------------------------------------------------------- @@ -159,7 +164,7 @@ class TestMultiTurn: @skip_if_foundry_hosting_integration_tests_disabled async def test_two_turn_conversation(self, server: InvocationsHostServer) -> None: """Turn 1 establishes context; turn 2 recalls it via the same session.""" - session_id = "int-test-session-two-turn" + session_id = f"int-test-session-two-turn-{uuid.uuid4().hex}" resp1 = await _post_invocation( server, @@ -176,14 +181,14 @@ async def test_two_turn_conversation(self, server: InvocationsHostServer) -> Non session_id=session_id, ) assert resp2.status_code == 200 - assert "blue" in resp2.text.lower() + assert "blue" in resp2.json()["response"].lower() @pytest.mark.flaky @pytest.mark.integration @skip_if_foundry_hosting_integration_tests_disabled async def test_multi_turn_streaming(self, server: InvocationsHostServer) -> None: """Multi-turn conversation with streaming on the second turn.""" - session_id = "int-test-session-stream" + session_id = f"int-test-session-stream-{uuid.uuid4().hex}" resp1 = await _post_invocation( server, @@ -201,6 +206,8 @@ async def test_multi_turn_streaming(self, server: InvocationsHostServer) -> None ) assert resp2.status_code == 200 assert "text/event-stream" in resp2.headers["content-type"] + assert "event: delta" in resp2.text + assert "event: done" in resp2.text assert "42" in resp2.text @@ -224,7 +231,7 @@ async def test_tool_call_non_streaming(self, server_with_tools: InvocationsHostS ) assert resp.status_code == 200 - assert "72" in resp.text + assert "72" in resp.json()["response"] @pytest.mark.flaky @pytest.mark.integration @@ -239,4 +246,5 @@ async def test_tool_call_streaming(self, server_with_tools: InvocationsHostServe assert resp.status_code == 200 assert "text/event-stream" in resp.headers["content-type"] + assert "event: done" in resp.text assert "72" in resp.text diff --git a/python/packages/foundry_hosting/tests/test_request.py b/python/packages/foundry_hosting/tests/test_request.py index 289b927c865..87310a858b1 100644 --- a/python/packages/foundry_hosting/tests/test_request.py +++ b/python/packages/foundry_hosting/tests/test_request.py @@ -1,6 +1,6 @@ # Copyright (c) Microsoft. All rights reserved. -"""Only the Responses request view and option policy are part of this slice.""" +"""Responses and Invocations request models and option policies.""" from __future__ import annotations @@ -11,7 +11,7 @@ from azure.ai.agentserver.responses import ResponseContext from azure.ai.agentserver.responses.models import CreateResponse -from agent_framework_foundry_hosting import HostedResponseRequest +from agent_framework_foundry_hosting import HostedResponseRequest, InvocationRun from agent_framework_foundry_hosting._request import ( prepare_response_options, response_run_options, @@ -133,3 +133,27 @@ def test_supported_unsupported_option_modes(mode: str) -> None: def test_unknown_unsupported_option_mode_fails_at_construction() -> None: with pytest.raises(ValueError, match="unsupported_options"): validate_unsupported_options("silent") + + +def test_invocation_run_accepts_typed_messages_with_independent_defaults() -> None: + first = InvocationRun(messages="hello") + second = InvocationRun(messages="world", stream=True) + assert first.messages == "hello" + assert first.options == second.options == {} + assert first.options is not second.options + assert second.stream is True + + +@pytest.mark.parametrize( + ("kwargs", "error"), + [ + ({"messages": None}, "messages"), + ({"messages": [object()]}, "messages"), + ({"messages": "hello", "options": "not an object"}, "options"), + ({"messages": "hello", "options": {1: "wrong key"}}, "options"), + ({"messages": "hello", "stream": "true"}, "stream"), + ], +) +def test_invocation_run_rejects_invalid_parser_values(kwargs: dict[str, Any], error: str) -> None: + with pytest.raises(TypeError, match=error): + InvocationRun(**cast(Any, kwargs)) diff --git a/python/samples/04-hosting/foundry-hosted-agents/README.md b/python/samples/04-hosting/foundry-hosted-agents/README.md index f9c9ef13e44..500d7c5e9a1 100644 --- a/python/samples/04-hosting/foundry-hosted-agents/README.md +++ b/python/samples/04-hosting/foundry-hosted-agents/README.md @@ -57,7 +57,7 @@ See [Using deployed agent](responses/using_deployed_agent.py) for service-create | # | Sample | Description | |---|--------|-------------| -| 1 | [Basic](invocations/basic/) | A minimal agent demonstrating basic request/response using the invocations protocol. | +| 1 | [Basic](invocations/basic/) | An Invocations agent with a custom JSON parser, durable MAF history, and JSON/SSE responses. | | 2 | [Break Glass](invocations/break_glass/) | An agent demonstrating a "break glass" scenario where customizations of the API behaviors are needed, allowing for more direct control over how requests and responses are handled by the hosting layer. | | 3 | [Telegram](invocations/telegram/) | A Telegram bot routed through API Management to a direct-code hosted agent, with streaming responses and durable Cosmos DB history. | diff --git a/python/samples/04-hosting/foundry-hosted-agents/invocations/basic/README.md b/python/samples/04-hosting/foundry-hosted-agents/invocations/basic/README.md index b0297ec194e..baed33d8836 100644 --- a/python/samples/04-hosting/foundry-hosted-agents/invocations/basic/README.md +++ b/python/samples/04-hosting/foundry-hosted-agents/invocations/basic/README.md @@ -1,88 +1,82 @@ -# What this sample demonstrates - -An [Agent Framework](https://github.com/microsoft/agent-framework) agent -hosted using the **Invocations protocol** with session management. Unlike -Responses, Invocations does **not** provide built-in conversation history. -The host persists its MAF `AgentSession` using the default file-based store -locally and Foundry storage when hosted, under the separate -`invocation_sessions` logical store. This basic agent does not configure a -history provider; add one if the model should remember previous messages. - -## How It Works - -### Model Integration - -The agent uses `FoundryChatClient` to create a Responses client from the -project endpoint and model deployment. When a request arrives, the host -restores (or creates) a MAF session, runs the agent with the user message -and session context, and persists its state. The agent supports streaming -and non-streaming response modes. - -See [main.py](main.py) for the full implementation. - -### Agent Hosting - -The agent is hosted using the [Agent Framework](https://github.com/microsoft/agent-framework) with the `InvocationsHostServer`, which provisions a REST API endpoint compatible with the Azure AI Invocations protocol. - -## Running the Agent Host - -Follow the instructions in the [Running the Agent Host Locally](../../README.md#running-the-agent-host-locally) section of the README in the parent directory to run the agent host. - -## Interacting with the agent - -> Depending on how you run the agent host, you can invoke the agent using `curl` (`Invoke-WebRequest` in PowerShell) or `azd`. Please refer to the [parent README](../../README.md) for more details. Use this README for sample queries you can send to the agent. - -Send a POST request to the server with a JSON body containing a "message" field to interact with the agent. For example: - -```bash -curl -X POST http://localhost:8088/invocations -i -H "Content-Type: application/json" -d '{"message": "Hi"}' -``` - -Or with streaming: +# Invocations agent with a custom request parser + +[`main.py`](main.py) hosts an Agent Framework agent using the Foundry +Invocations protocol. The sample accepts application JSON with a `prompt` +field rather than the host's default `message` field. Its `parse_request` +callback returns `InvocationRun(messages, options, stream)`, and +`prepare_options` allows only `temperature` and `max_tokens` from the caller. +The host also rejects platform IDs, `store`, and private continuation options +even if a hook tries to return them. The agent keeps `store=False` as a +developer default; caller options cannot enable service-managed history. + +**Invocations does not store conversation history.** The sample's +`InMemoryHistoryProvider` keeps model messages in the MAF `AgentSession`, which +the host saves in its separate `invocation_sessions` store. That store uses +files locally and Foundry state storage when hosted. A later turn, even in a +replacement host process, restores the history from the same trusted user and +sandbox. Sandbox files remain in the Foundry sandbox; they are not part of +the MAF session. + +## Run and invoke + +Follow the [parent guide](../../README.md#running-the-agent-host-locally) +to run the sample locally. Send a non-streaming request: ```bash -curl -X POST http://localhost:8088/invocations -i -H "Content-Type: application/json" -d '{"message": "Hi", "stream": true}' +curl -i -X POST http://localhost:8088/invocations \ + -H "Content-Type: application/json" \ + -d '{"prompt": "Hi", "options": {"temperature": 0.3}}' ``` -The server responds with text. The `-i` flag in the `curl` command includes -HTTP response headers, including the session ID that can be reused on later -requests. Here is an example: +The default wire format is JSON, for example: -``` -HTTP/1.1 200 -content-length: 34 +```http +HTTP/1.1 200 OK content-type: application/json -x-agent-invocation-id: ec04d020-a0e7-441e-ae83-db75635a9f83 x-agent-session-id: 9370b9d4-cd13-4436-a57f-03b843ac0e17 -x-platform-server: azure-ai-agentserver-core/2.0.0a20260410006 (python/3.12) -date: Fri, 17 Apr 2026 23:46:44 GMT -server: hypercorn-h11 -Hi! How can I help? +{"response":"Hi! How can I help?"} ``` -### Multi-turn conversation - -To reuse the same sandbox and MAF session (not model message history in this -basic example), take the session ID from the previous response header and -include it in the URL query for the next request: +To continue in that sandbox, put the returned **platform** +`x-agent-session-id` in the Invocations **query parameter**: ```bash -curl -X POST http://localhost:8088/invocations?agent_session_id=9370b9d4-cd13-4436-a57f-03b843ac0e17 -i -H "Content-Type: application/json" -d '{"message": "How are you?"}' +curl -i -N -X POST \ + "http://localhost:8088/invocations?agent_session_id=9370b9d4-cd13-4436-a57f-03b843ac0e17" \ + -H "Content-Type: application/json" \ + -d '{"prompt": "What did I say earlier?", "stream": true}' +``` + +Streaming uses framed server-sent events (`text/event-stream`): + +```text +event: delta +data: {"text": "You said hi."} + +event: done +data: {"session_id": "9370b9d4-cd13-4436-a57f-03b843ac0e17"} ``` -On Foundry, the `agent_session_id` **query parameter** routes an invocation to -the corresponding sandbox; an ID in the JSON body does not route it. The host -also requires platform user and call IDs. If the hosted container does not -receive `FOUNDRY_AGENT_SESSION_ID`, supply an explicit routed query ID even on -the first invocation (for example, [create a hosted session first](https://learn.microsoft.com/azure/foundry/agents/how-to/manage-hosted-sessions)). -Without either source, the host rejects the request rather than persisting -state under a generated ID. When the environment variable is present, a -different query ID is rejected. Locally, the SDK still provides a single-user -fallback for requests without an ID. -See the [state store guide](../../../../../packages/foundry_hosting/README.md#state-store) -for namespace and upgrade details. - -## Deploying the Agent to Foundry - -To host the agent on Foundry, follow the instructions in the [Deploying the Agent to Foundry](../../README.md#deploying-the-agent-to-foundry) section of the README in the parent directory. +The `done` ID is the **sandbox ID**, not the MAF `AgentSession.session_id`. +On Foundry, the query ID routes the request to the sandbox; a body field does +not. The host requires trusted platform user and call IDs. If +`FOUNDRY_AGENT_SESSION_ID` is absent, even the first hosted request requires +an explicit query ID matching the routed request context (for example, +[create a hosted session first](https://learn.microsoft.com/azure/foundry/agents/how-to/manage-hosted-sessions)). +If the environment ID is present, a different query ID is rejected. Local +requests retain the SDK's single-user generated-ID fallback. See the +[state-store guide](../../../../../packages/foundry_hosting/README.md#state-store) +for scope and retention details. + +Malformed requests return JSON client errors. Streaming failures produce +`event: error` instead of `done`; a competing host's session write is reported +as a conflict, not silently overwritten. Do not blindly retry calls with +non-idempotent tools. Existing clients that still require the old plain-text +response and raw text-chunk stream can set `legacy_wire_format=True` on the +host temporarily; the host warns once and the mode is deprecated. Migrate +clients to JSON and framed SSE before removing that opt-in in a deliberate +breaking change. + +Follow the [parent deployment guide](../../README.md#deploying-the-agent-to-foundry) +when deploying this example to Foundry. diff --git a/python/samples/04-hosting/foundry-hosted-agents/invocations/basic/main.py b/python/samples/04-hosting/foundry-hosted-agents/invocations/basic/main.py index 730600d4f98..e6b98f2f159 100644 --- a/python/samples/04-hosting/foundry-hosted-agents/invocations/basic/main.py +++ b/python/samples/04-hosting/foundry-hosted-agents/invocations/basic/main.py @@ -1,17 +1,45 @@ # Copyright (c) Microsoft. All rights reserved. +from __future__ import annotations + import os +from typing import Any -from agent_framework import Agent -from agent_framework.foundry import FoundryChatClient, InvocationsHostServer +from agent_framework import Agent, InMemoryHistoryProvider +from agent_framework.foundry import FoundryChatClient +from agent_framework_foundry_hosting import InvocationRun, InvocationsHostServer from azure.identity import DefaultAzureCredential from dotenv import load_dotenv +from starlette.requests import Request + +"""Host an Invocations agent with an application JSON parser and persisted MAF history.""" -# Load environment variables from .env file load_dotenv() -def main(): +# 1. Map the application's JSON payload to typed MAF inputs. +async def parse_request(request: Request) -> InvocationRun: + """Map application-specific JSON to a validated agent turn.""" + payload: Any = await request.json() + if not isinstance(payload, dict) or not isinstance(payload.get("prompt"), str): + raise ValueError("prompt must be a string.") + options = payload.get("options", {}) + if not isinstance(options, dict): + raise ValueError("options must be an object.") + stream = payload.get("stream", False) + if not isinstance(stream, bool): + raise ValueError("stream must be a boolean.") + return InvocationRun(messages=payload["prompt"], options=options, stream=stream) + + +# 2. Allow only caller-controlled generation settings. +def prepare_options(_request: Request, options: dict[str, Any]) -> dict[str, Any]: + """Allow only the caller's generation controls; keep storage developer-owned.""" + return {name: value for name, value in options.items() if name in {"temperature", "max_tokens"}} + + +# 3. Persist the agent's own conversation history without provider-managed storage. +def main() -> None: client = FoundryChatClient( project_endpoint=os.environ["FOUNDRY_PROJECT_ENDPOINT"], model=os.environ["AZURE_AI_MODEL_DEPLOYMENT_NAME"], @@ -21,15 +49,21 @@ def main(): agent = Agent( client=client, instructions="You are a friendly assistant. Keep your answers brief.", - # History will be managed by the hosting infrastructure, thus there - # is no need to store history by the service. Learn more at: - # https://developers.openai.com/api/reference/resources/responses/methods/create + context_providers=[InMemoryHistoryProvider()], default_options={"store": False}, ) - server = InvocationsHostServer(agent) + server = InvocationsHostServer( + agent, + parse_request=parse_request, + prepare_options=prepare_options, + unsupported_options="error", + ) server.run() if __name__ == "__main__": main() + +# Expected non-streaming response: {"response": ""} +# Streaming emits event: delta frames, then event: done after the MAF session is saved. From 1eaea68c82437d4f49118f8d563b8ba03585f272 Mon Sep 17 00:00:00 2001 From: eavanvalkenburg Date: Wed, 30 Sep 2026 12:01:07 +0200 Subject: [PATCH 2/2] Python: reject unsafe Invocations agent control options --- python/packages/foundry_hosting/README.md | 7 +- .../_invocations.py | 18 ++- .../foundry_hosting/tests/test_invocations.py | 125 +++++++++++++++++- 3 files changed, 142 insertions(+), 8 deletions(-) diff --git a/python/packages/foundry_hosting/README.md b/python/packages/foundry_hosting/README.md index f4bf70bc297..79434081b98 100644 --- a/python/packages/foundry_hosting/README.md +++ b/python/packages/foundry_hosting/README.md @@ -276,7 +276,12 @@ own JSON shape and MAF `Message` inputs. A sync or async `prepare_options(reques a **copy** of this turn's caller options without changing the agent's `default_options`. It must return a mapping with string keys. The host rejects reserved platform/session fields, `store`, `extra_body`, and private continuation fields after the hook; callers cannot select another sandbox or enable downstream service continuation through -runtime options. When an agent cannot accept runtime options, `unsupported_options="warn"` (default) logs and ignores +runtime options. It also rejects agent execution controls: `additional_function_arguments`, `function_invocation_kwargs`, +`client_kwargs`, `middleware`, `session`, `tools`, `instructions`, `compaction_strategy`, and `tokenizer`. Trusted tool +arguments belong in developer-configured agent defaults or middleware/factories, not in request options or hook output. +A hook may strip denied caller fields before validation; allowed generation and provider-specific options remain +available. These restrictions apply to both wire formats and every `unsupported_options` policy. +When an agent cannot accept runtime options, `unsupported_options="warn"` (default) logs and ignores them; `"ignore"` silently drops them and `"error"` rejects them. For request-scoped factories, unsupported options discovered after streaming starts are reported as an SSE `error` event with `status: 400`. diff --git a/python/packages/foundry_hosting/agent_framework_foundry_hosting/_invocations.py b/python/packages/foundry_hosting/agent_framework_foundry_hosting/_invocations.py index 85b46b11d01..a6f81fee53b 100644 --- a/python/packages/foundry_hosting/agent_framework_foundry_hosting/_invocations.py +++ b/python/packages/foundry_hosting/agent_framework_foundry_hosting/_invocations.py @@ -41,6 +41,18 @@ InvocationParser = Callable[[Request], InvocationRun | Awaitable[InvocationRun]] InvocationOptionsHook = Callable[[Request, dict[str, Any]], Mapping[str, Any] | Awaitable[Mapping[str, Any]]] +_AGENT_CONTROLLED_FIELDS = frozenset({ + "additional_function_arguments", + "client_kwargs", + "compaction_strategy", + "function_invocation_kwargs", + "instructions", + "middleware", + "session", + "tokenizer", + "tools", +}) + class _UnsupportedAgentOptions(TypeError): """The agent cannot accept the caller's run options under the selected policy.""" @@ -108,7 +120,8 @@ def __init__( after their last write. Custom providers control their own retention. parse_request: Optional sync or async parser returning an `InvocationRun` from application JSON. Without one, accepts a JSON object with `message`, optional `options`, and optional `stream`. - prepare_options: Optional sync or async hook to filter or replace a copy of caller run options. + prepare_options: Optional sync or async hook to filter or replace a copy of caller generation options. + Tool context and agent execution controls must remain in developer-owned agent configuration. unsupported_options: `"warn"` (default), `"ignore"`, or `"error"` for agents without runtime options. legacy_wire_format: Opt into the deprecated plain-text response and raw streaming chunks instead of the default JSON response and framed `delta`/`done`/`error` server-sent events. @@ -295,6 +308,9 @@ async def _options(self, request: Request, parsed: InvocationRun) -> dict[str, A raise TypeError("prepare_options must return a mapping of MAF run options with string keys.") options = deepcopy(dict(result)) validate_request_options(options) + reserved = _AGENT_CONTROLLED_FIELDS.intersection(options) + if reserved: + raise ValueError(f"Invocations options cannot set agent-controlled fields: {', '.join(sorted(reserved))}.") return options def _agent_kwargs(self, agent: SupportsAgentRun, options: dict[str, Any]) -> dict[str, Any]: diff --git a/python/packages/foundry_hosting/tests/test_invocations.py b/python/packages/foundry_hosting/tests/test_invocations.py index 15028c37059..8375d41daf6 100644 --- a/python/packages/foundry_hosting/tests/test_invocations.py +++ b/python/packages/foundry_hosting/tests/test_invocations.py @@ -654,9 +654,19 @@ def test_rejects_invalid_configuration(self, parameter: str, value: Any, error: class TestParsedRequests: @pytest.mark.parametrize("stream", [False, True]) async def test_parser_and_hook_use_typed_messages_without_mutating_defaults(self, stream: bool) -> None: - source_options: dict[str, Any] = {"temperature": 0.8, "nested": {"tag": "original"}, "store": True} + source_options: dict[str, Any] = { + "temperature": 0.8, + "reasoning_effort": "low", + "nested": {"tag": "original"}, + "store": True, + "additional_function_arguments": {"user_id": "forged"}, + } agent = _make_agent(response_text="ok", stream_texts=["ok"]) - agent.default_options = {"store": False, "temperature": 0.2} + agent.default_options = { + "store": False, + "temperature": 0.2, + "additional_function_arguments": {"user_id": "trusted"}, + } request = _make_request({"prompt": "hello"}) async def parse(incoming: Request) -> InvocationRun: @@ -672,7 +682,7 @@ async def prepare(incoming: Request, options: dict[str, Any]) -> dict[str, Any]: assert incoming is request options["nested"]["tag"] = "changed" options.pop("store") - return {"temperature": options["temperature"]} + return {"temperature": options["temperature"], "reasoning_effort": options["reasoning_effort"]} server = InvocationsHostServer(agent, parse_request=parse, prepare_options=prepare) with _request_context(session_id="parsed"): @@ -680,9 +690,19 @@ async def prepare(incoming: Request, options: dict[str, Any]) -> dict[str, Any]: assert await _success_text(response) == "ok" assert agent.calls[0]["messages"][0].text == "hello" - assert agent.calls[0]["options"] == {"temperature": 0.8} - assert source_options == {"temperature": 0.8, "nested": {"tag": "original"}, "store": True} - assert agent.default_options == {"store": False, "temperature": 0.2} + assert agent.calls[0]["options"] == {"temperature": 0.8, "reasoning_effort": "low"} + assert source_options == { + "temperature": 0.8, + "reasoning_effort": "low", + "nested": {"tag": "original"}, + "store": True, + "additional_function_arguments": {"user_id": "forged"}, + } + assert agent.default_options == { + "store": False, + "temperature": 0.2, + "additional_function_arguments": {"user_id": "trusted"}, + } @pytest.mark.parametrize( ("payload", "error"), @@ -750,6 +770,99 @@ async def test_reserved_options_are_rejected_before_agent_run(self, options: dic assert "host-controlled fields" in json.loads(bytes(response.body))["error"] assert agent.calls == [] + @pytest.mark.parametrize("stream", [False, True]) + @pytest.mark.parametrize("source", ["request", "parser", "hook"]) + @pytest.mark.parametrize( + ("field", "value"), + [ + ("additional_function_arguments", {"user_id": "forged", "tenant_id": "forged"}), + ("client_kwargs", {"middleware": []}), + ("compaction_strategy", {}), + ("function_invocation_kwargs", {"user_id": "forged"}), + ("instructions", "Caller-controlled instructions"), + ("middleware", []), + ("session", {"session_id": "forged"}), + ("tokenizer", {}), + ("tools", []), + ], + ) + async def test_agent_control_options_are_rejected_before_storage_or_agent_creation( + self, stream: bool, source: str, field: str, value: Any + ) -> None: + agent = _make_agent(response_text="ok", stream_texts=["ok"]) + factory_calls: list[str] = [] + provider = _SessionStoreProvider(_mock_session_store()) + options = {field: value} + + def create_agent() -> _FakeAgent: + factory_calls.append("created") + return agent + + server = InvocationsHostServer( + create_agent, + agent_session_store_provider=provider, + parse_request=(lambda _request: InvocationRun(messages="hi", options=options, stream=stream)) + if source == "parser" + else None, + prepare_options=(lambda _request, _options: options) if source == "hook" else None, + ) + with _request_context(session_id="unsafe-options"): + response = await server._handle_invoke( # pyright: ignore[reportPrivateUsage] + _make_request({"message": "hi", "stream": stream, "options": options if source == "request" else {}}) + ) + assert response.status_code == 400 + assert response.media_type == "application/json" + assert json.loads(bytes(response.body)) == { + "error": f"Invocations options cannot set agent-controlled fields: {field}." + } + assert factory_calls == [] + assert agent.calls == [] + assert provider.contexts == [] + + @pytest.mark.parametrize("stream", [False, True]) + @pytest.mark.parametrize("legacy_wire_format", [False, True]) + @pytest.mark.parametrize("policy", ["ignore", "warn", "error"]) + async def test_http_route_rejects_tool_context_options_under_every_wire_and_options_policy( + self, stream: bool, legacy_wire_format: bool, policy: Literal["ignore", "warn", "error"] + ) -> None: + agent = _make_agent(response_text="ok", stream_texts=["ok"]) + provider = _SessionStoreProvider(_mock_session_store()) + agent.default_options = {"additional_function_arguments": {"user_id": "trusted", "tenant_id": "trusted"}} + + def create_host() -> InvocationsHostServer: + return InvocationsHostServer( + agent, + agent_session_store_provider=provider, + unsupported_options=policy, + legacy_wire_format=legacy_wire_format, + ) + + if legacy_wire_format: + with pytest.warns(DeprecationWarning, match="legacy_wire_format"): + server = create_host() + else: + server = create_host() + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=server), base_url="http://test") as client: + response = await client.post( + "/invocations", + params={"agent_session_id": f"unsafe-options-{uuid.uuid4().hex}"}, + json={ + "message": "hi", + "stream": stream, + "options": {"additional_function_arguments": {"user_id": "forged", "tenant_id": "forged"}}, + }, + ) + assert response.status_code == 400 + assert response.headers["content-type"] == "application/json" + assert response.json() == { + "error": "Invocations options cannot set agent-controlled fields: additional_function_arguments." + } + assert agent.default_options == { + "additional_function_arguments": {"user_id": "trusted", "tenant_id": "trusted"} + } + assert agent.calls == [] + assert provider.contexts == [] + @pytest.mark.parametrize("result", [None, [], {1: "invalid"}, {"session_id": "forged"}]) async def test_options_hook_must_return_safe_string_keyed_mapping(self, result: Any) -> None: agent = _make_agent(response_text="ok")