From 1aa8578f7403ee6247f13040aee24a3e49cad597 Mon Sep 17 00:00:00 2001 From: Shanmukha Pasumarthy Date: Tue, 15 Sep 2026 21:38:11 +0530 Subject: [PATCH] fix(agentserver): partition in-memory response state by user Scope response, item, history, and legacy execution state by the trusted platform user key. Preserve unkeyed local use in a separate anonymous partition and cover caller isolation, collisions, and cleanup. Authored-by: GitHub Copilot CLI 1.0.81-0 Model: GPT-6 Astra (gpt-6-astra) Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 9e8471c6-fef7-46a5-9b42-d9e87f236721 --- .../CHANGELOG.md | 2 + .../azure-ai-agentserver-responses/README.md | 15 ++ .../azure-ai-agentserver-responses/api.md | 36 ++- .../ai/agentserver/responses/store/_base.py | 5 +- .../ai/agentserver/responses/store/_memory.py | 149 +++++++---- .../test_user_isolation_enforcement.py | 60 +++++ .../unit/test_in_memory_provider_crud.py | 246 +++++++++++++++++- 7 files changed, 454 insertions(+), 59 deletions(-) diff --git a/sdk/agentserver/azure-ai-agentserver-responses/CHANGELOG.md b/sdk/agentserver/azure-ai-agentserver-responses/CHANGELOG.md index 8f9837aa7d44..9d6f20e51803 100644 --- a/sdk/agentserver/azure-ai-agentserver-responses/CHANGELOG.md +++ b/sdk/agentserver/azure-ai-agentserver-responses/CHANGELOG.md @@ -4,6 +4,8 @@ ### Bugs Fixed +- Partition `InMemoryResponseProvider` responses, items, history, and legacy replay + state by the platform user key, keeping anonymous local state separate. - Scoped durable multi-turn task IDs with `FOUNDRY_AGENT_SESSION_GUID` when available, preventing recreated same-name sessions from colliding with task tombstones. Existing pre-rollout active chains remain resumable through a diff --git a/sdk/agentserver/azure-ai-agentserver-responses/README.md b/sdk/agentserver/azure-ai-agentserver-responses/README.md index ffd2b670cf3a..4d8eca22861b 100644 --- a/sdk/agentserver/azure-ai-agentserver-responses/README.md +++ b/sdk/agentserver/azure-ai-agentserver-responses/README.md @@ -255,6 +255,21 @@ app = ResponsesAgentServerHost(options=options) ## Troubleshooting +### In-memory storage identity + +`InMemoryResponseProvider` partitions response envelopes, input/output items, and +history by the exact `PlatformContext.user_id_key` supplied by a trusted host. +Omitting context or setting the user key to `None` selects a separate anonymous +partition for local development, not unrestricted access to named users' data. +Empty and whitespace keys remain distinct; `call_id` does not affect the partition. +Pass the same user key on every related provider operation, including the legacy +execution/replay helpers when used directly. + +The provider does not authenticate callers. Hosts must establish trustworthy +platform context before accessing it. This storage boundary is not complete +end-to-end authorization: runtime routing and the process-wide SSE stream registry +are separate from the provider's envelope, item, and history storage. + ### Common errors - **400 Bad Request**: The request body failed validation. Check that optional fields such as `model` (when provided) are valid and that `input` items are well-formed. diff --git a/sdk/agentserver/azure-ai-agentserver-responses/api.md b/sdk/agentserver/azure-ai-agentserver-responses/api.md index 17bde6a3b87c..ebb639373fea 100644 --- a/sdk/agentserver/azure-ai-agentserver-responses/api.md +++ b/sdk/agentserver/azure-ai-agentserver-responses/api.md @@ -275,14 +275,16 @@ namespace azure.ai.agentserver.responses response_id: str, event: StreamEventRecord, *, - ttl_seconds: int | None = ... + ttl_seconds: int | None = ..., + context: PlatformContext | None = ... ) -> bool: ... async def create_execution( self, execution: ResponseExecution, *, - ttl_seconds: int | None = ... + ttl_seconds: int | None = ..., + context: PlatformContext | None = ... ) -> None: ... async def create_response( @@ -294,7 +296,12 @@ namespace azure.ai.agentserver.responses context: PlatformContext | None = ... ) -> None: ... - async def delete(self, response_id: str) -> bool: ... + async def delete( + self, + response_id: str, + *, + context: PlatformContext | None = ... + ) -> bool: ... async def delete_response( self, @@ -303,7 +310,12 @@ namespace azure.ai.agentserver.responses context: PlatformContext | None = ... ) -> None: ... - async def get_execution(self, response_id: str) -> ResponseExecution | None: ... + async def get_execution( + self, + response_id: str, + *, + context: PlatformContext | None = ... + ) -> ResponseExecution | None: ... async def get_history_item_ids( self, @@ -332,7 +344,12 @@ namespace azure.ai.agentserver.responses context: PlatformContext | None = ... ) -> list[OutputItem | None]: ... - async def get_replay_events(self, response_id: str) -> list[StreamEventRecord] | None: ... + async def get_replay_events( + self, + response_id: str, + *, + context: PlatformContext | None = ... + ) -> list[StreamEventRecord] | None: ... async def get_response( self, @@ -351,7 +368,8 @@ namespace azure.ai.agentserver.responses self, response_id: str, *, - ttl_seconds: int | None = ... + ttl_seconds: int | None = ..., + context: PlatformContext | None = ... ) -> bool: ... async def set_response_snapshot( @@ -359,7 +377,8 @@ namespace azure.ai.agentserver.responses response_id: str, response: ResponseObject, *, - ttl_seconds: int | None = ... + ttl_seconds: int | None = ..., + context: PlatformContext | None = ... ) -> bool: ... async def transition_execution_status( @@ -367,7 +386,8 @@ namespace azure.ai.agentserver.responses response_id: str, next_status: ResponseStatus, *, - ttl_seconds: int | None = ... + ttl_seconds: int | None = ..., + context: PlatformContext | None = ... ) -> bool: ... async def update_response( diff --git a/sdk/agentserver/azure-ai-agentserver-responses/azure/ai/agentserver/responses/store/_base.py b/sdk/agentserver/azure-ai-agentserver-responses/azure/ai/agentserver/responses/store/_base.py index d4bba871cc8f..e393d4b6525c 100644 --- a/sdk/agentserver/azure-ai-agentserver-responses/azure/ai/agentserver/responses/store/_base.py +++ b/sdk/agentserver/azure-ai-agentserver-responses/azure/ai/agentserver/responses/store/_base.py @@ -51,8 +51,9 @@ class ResponseProviderProtocol(Protocol): Every operation accepts an optional ``context`` parameter (S-018). Implementations MUST use it to partition data in multi-tenant - deployments. When ``None``, the provider operates without tenant - scoping (suitable for local development). + deployments. Providers supporting anonymous local development must keep + unkeyed state separate from named users' state; ``None`` must not bypass + partitioning. """ async def create_response( diff --git a/sdk/agentserver/azure-ai-agentserver-responses/azure/ai/agentserver/responses/store/_memory.py b/sdk/agentserver/azure-ai-agentserver-responses/azure/ai/agentserver/responses/store/_memory.py index a28b3820bb9a..89f6c6a1e879 100644 --- a/sdk/agentserver/azure-ai-agentserver-responses/azure/ai/agentserver/responses/store/_memory.py +++ b/sdk/agentserver/azure-ai-agentserver-responses/azure/ai/agentserver/responses/store/_memory.py @@ -20,6 +20,13 @@ _DEFAULT_REPLAY_EVENT_TTL_SECONDS: int = 600 """Minimum per-event replay TTL (10 minutes) per spec B35.""" +_StoreKey = tuple[str | None, str] + + +def _store_key(identifier: str, context: PlatformContext | None) -> _StoreKey: + """Keep absent identity anonymous; preserve present keys exactly, including empty strings.""" + return (context.user_id_key if context is not None else None, identifier) + class _StoreEntry: """Container for one response execution and its replay state.""" @@ -55,15 +62,21 @@ class InMemoryResponseProvider(ResponseProviderProtocol): process-wide ``azure.ai.agentserver.core.streaming.streams`` registry, configured at host startup; this provider stores only response envelopes, input items, and history pointers. + + State is partitioned by the trusted ``PlatformContext.user_id_key``. + Missing context or a ``None`` user key selects a separate anonymous + partition for local use, never an unrestricted lookup. Empty and whitespace + keys remain distinct, matching ``PlatformContext`` semantics. The provider + does not authenticate callers or interpret ``call_id``. """ def __init__(self) -> None: """Initialize in-memory state and an async mutation lock.""" - self._entries: Dict[str, _StoreEntry] = {} + self._entries: Dict[_StoreKey, _StoreEntry] = {} self._lock = asyncio.Lock() - self._item_store: Dict[str, OutputItem] = {} - self._conversation_responses: defaultdict[str, list[str]] = defaultdict(list) - self._stream_events: Dict[str, list[ResponseStreamEvent]] = {} + self._item_store: Dict[_StoreKey, OutputItem] = {} + self._conversation_responses: defaultdict[_StoreKey, list[str]] = defaultdict(list) + self._stream_events: Dict[_StoreKey, list[ResponseStreamEvent]] = {} @contextlib.asynccontextmanager async def _locked(self) -> AsyncIterator[None]: @@ -102,7 +115,7 @@ async def create_response( """ response_id = str(response.get("id")) async with self._locked(): - entry = self._entries.get(response_id) + entry = self._entries.get(_store_key(response_id, context)) if entry is not None and not entry.deleted: raise ResponseAlreadyExistsError(response_id) @@ -112,12 +125,12 @@ async def create_response( item_id = self._extract_item_id(item) if item_id is None: continue - self._item_store[item_id] = deepcopy(item) + self._item_store[_store_key(item_id, context)] = deepcopy(item) input_ids.append(item_id) history_ids = list(history_item_ids) if history_item_ids is not None else [] - output_ids = self._store_output_items_unlocked(response) - self._entries[response_id] = _StoreEntry( + output_ids = self._store_output_items_unlocked(response, context=context) + self._entries[_store_key(response_id, context)] = _StoreEntry( execution=ResponseExecution( response_id=response_id, mode_flags=self._resolve_mode_flags_from_response(response), @@ -132,7 +145,7 @@ async def create_response( conversation_id = get_conversation_id(response) if conversation_id is not None: - self._conversation_responses[conversation_id].append(response_id) + self._conversation_responses[_store_key(conversation_id, context)].append(response_id) async def get_response(self, response_id: str, *, context: PlatformContext | None = None) -> ResponseObject: """Retrieve one response envelope by identifier. @@ -146,7 +159,7 @@ async def get_response(self, response_id: str, *, context: PlatformContext | Non :raises KeyError: If the response does not exist or has been deleted. """ async with self._locked(): - entry = self._entries.get(response_id) + entry = self._entries.get(_store_key(response_id, context)) if entry is None or entry.deleted or entry.response is None: raise KeyError(f"response '{response_id}' not found") return deepcopy(entry.response) @@ -166,13 +179,13 @@ async def update_response(self, response: ResponseObject, *, context: PlatformCo """ response_id = str(response.get("id")) async with self._locked(): - entry = self._entries.get(response_id) + entry = self._entries.get(_store_key(response_id, context)) if entry is None or entry.deleted: raise KeyError(f"response '{response_id}' not found") entry.response = deepcopy(response) entry.execution.set_response_snapshot(deepcopy(response)) - entry.output_item_ids = self._store_output_items_unlocked(response) + entry.output_item_ids = self._store_output_items_unlocked(response, context=context) async def delete_response(self, response_id: str, *, context: PlatformContext | None = None) -> None: """Delete a stored response envelope by identifier. @@ -187,7 +200,7 @@ async def delete_response(self, response_id: str, *, context: PlatformContext | :raises KeyError: If the response does not exist or has already been deleted. """ async with self._locked(): - entry = self._entries.get(response_id) + entry = self._entries.get(_store_key(response_id, context)) if entry is None or entry.deleted: raise KeyError(f"response '{response_id}' not found") entry.deleted = True @@ -226,7 +239,7 @@ async def get_input_items( :raises ValueError: If the response has been deleted. """ async with self._locked(): - entry = self._entries.get(response_id) + entry = self._entries.get(_store_key(response_id, context)) if entry is None: raise KeyError(f"response '{response_id}' not found") if entry.deleted: @@ -251,9 +264,9 @@ async def get_input_items( safe_limit = max(1, min(100, int(limit))) return [ - deepcopy(self._item_store[item_id]) + deepcopy(self._item_store[key]) for item_id in ordered_ids[:safe_limit] - if item_id in self._item_store + if (key := _store_key(item_id, context)) in self._item_store ] async def get_items( @@ -275,7 +288,8 @@ async def get_items( """ async with self._locked(): return [ - deepcopy(self._item_store[item_id]) if item_id in self._item_store else None for item_id in item_ids + deepcopy(self._item_store[key]) if (key := _store_key(item_id, context)) in self._item_store else None + for item_id in item_ids ] async def get_history_item_ids( @@ -308,7 +322,7 @@ async def get_history_item_ids( resolved: list[str] = [] if previous_response_id is not None: - entry = self._entries.get(previous_response_id) + entry = self._entries.get(_store_key(previous_response_id, context)) if entry is not None and not entry.deleted: # Resolve history chain for the previous response: # return historyItemIds + inputItemIds + outputItemIds of the previous response @@ -317,8 +331,8 @@ async def get_history_item_ids( resolved.extend(entry.output_item_ids or []) if conversation_id is not None: - for response_id in self._conversation_responses.get(conversation_id, []): - entry = self._entries.get(response_id) + for response_id in self._conversation_responses.get(_store_key(conversation_id, context), []): + entry = self._entries.get(_store_key(response_id, context)) if entry is None or entry.deleted: continue resolved.extend(entry.history_item_ids or []) @@ -331,35 +345,48 @@ async def get_history_item_ids( # preserving chronological order in the returned slice. return resolved[-limit:] - async def create_execution(self, execution: ResponseExecution, *, ttl_seconds: int | None = None) -> None: + async def create_execution( + self, + execution: ResponseExecution, + *, + ttl_seconds: int | None = None, + context: PlatformContext | None = None, + ) -> None: """Create a new execution and replay container for ``execution.response_id``. :param execution: The execution state to store. :type execution: ~azure.ai.agentserver.responses.models.runtime.ResponseExecution :keyword int or None ttl_seconds: Optional time-to-live in seconds for automatic expiration. + :keyword context: Platform context for partitioning; omitted context selects anonymous state. + :paramtype context: ~azure.ai.agentserver.responses.PlatformContext | None :rtype: None :raises ValueError: If an entry with the same response ID already exists. """ async with self._locked(): - if execution.response_id in self._entries: + key = _store_key(execution.response_id, context) + if key in self._entries: raise ValueError(f"response '{execution.response_id}' already exists") - self._entries[execution.response_id] = _StoreEntry( + self._entries[key] = _StoreEntry( execution=deepcopy(execution), replay=_StreamReplayState(response_id=execution.response_id), expires_at=self._compute_expiry(ttl_seconds), ) - async def get_execution(self, response_id: str) -> ResponseExecution | None: + async def get_execution( + self, response_id: str, *, context: PlatformContext | None = None + ) -> ResponseExecution | None: """Get a defensive copy of execution state for ``response_id`` if present. :param response_id: The unique identifier of the response execution to retrieve. :type response_id: str + :keyword context: Platform context for partitioning; omitted context selects anonymous state. + :paramtype context: ~azure.ai.agentserver.responses.PlatformContext | None :returns: A deep copy of the execution state, or ``None`` if not found. :rtype: ~azure.ai.agentserver.responses.models.runtime.ResponseExecution | None """ async with self._locked(): - entry = self._entries.get(response_id) + entry = self._entries.get(_store_key(response_id, context)) if entry is None: return None return deepcopy(entry.execution) @@ -370,6 +397,7 @@ async def set_response_snapshot( response: ResponseObject, *, ttl_seconds: int | None = None, + context: PlatformContext | None = None, ) -> bool: """Set the latest response snapshot for an existing response execution. @@ -378,11 +406,13 @@ async def set_response_snapshot( :param response: The response snapshot to associate with the execution. :type response: ~azure.ai.agentserver.responses.models._generated.Response :keyword int or None ttl_seconds: Optional time-to-live in seconds to refresh expiration. + :keyword context: Platform context for partitioning; omitted context selects anonymous state. + :paramtype context: ~azure.ai.agentserver.responses.PlatformContext | None :returns: ``True`` if the entry was found and updated, ``False`` otherwise. :rtype: bool """ async with self._locked(): - entry = self._entries.get(response_id) + entry = self._entries.get(_store_key(response_id, context)) if entry is None: return False @@ -396,6 +426,7 @@ async def transition_execution_status( next_status: ResponseStatus, *, ttl_seconds: int | None = None, + context: PlatformContext | None = None, ) -> bool: """Transition execution state while preserving lifecycle invariants. @@ -404,11 +435,13 @@ async def transition_execution_status( :param next_status: The target status to transition to. :type next_status: ~azure.ai.agentserver.responses.models.runtime.ResponseStatus :keyword int or None ttl_seconds: Optional time-to-live in seconds to refresh expiration. + :keyword context: Platform context for partitioning; omitted context selects anonymous state. + :paramtype context: ~azure.ai.agentserver.responses.PlatformContext | None :returns: ``True`` if the entry was found and transitioned, ``False`` otherwise. :rtype: bool """ async with self._locked(): - entry = self._entries.get(response_id) + entry = self._entries.get(_store_key(response_id, context)) if entry is None: return False @@ -416,18 +449,26 @@ async def transition_execution_status( self._apply_ttl_unlocked(entry, ttl_seconds) return True - async def set_cancel_requested(self, response_id: str, *, ttl_seconds: int | None = None) -> bool: + async def set_cancel_requested( + self, + response_id: str, + *, + ttl_seconds: int | None = None, + context: PlatformContext | None = None, + ) -> bool: """Mark cancellation requested and enforce lifecycle-safe cancel transitions. :param response_id: The unique identifier of the response to cancel. :type response_id: str :keyword int or None ttl_seconds: Optional time-to-live in seconds to refresh expiration. + :keyword context: Platform context for partitioning; omitted context selects anonymous state. + :paramtype context: ~azure.ai.agentserver.responses.PlatformContext | None :returns: ``True`` if the entry was found and cancel was applied, ``False`` otherwise. :rtype: bool :raises ValueError: If the execution is already terminal in a non-cancelled state. """ async with self._locked(): - entry = self._entries.get(response_id) + entry = self._entries.get(_store_key(response_id, context)) if entry is None: return False @@ -469,6 +510,7 @@ async def append_stream_event( event: StreamEventRecord, *, ttl_seconds: int | None = None, + context: PlatformContext | None = None, ) -> bool: """Append one stream event to replay state for an existing execution. @@ -477,11 +519,13 @@ async def append_stream_event( :param event: The stream event record to append. :type event: ~azure.ai.agentserver.responses.models.runtime.StreamEventRecord :keyword int or None ttl_seconds: Optional time-to-live in seconds to refresh expiration. + :keyword context: Platform context for partitioning; omitted context selects anonymous state. + :paramtype context: ~azure.ai.agentserver.responses.PlatformContext | None :returns: ``True`` if the entry was found and the event was appended, ``False`` otherwise. :rtype: bool """ async with self._locked(): - entry = self._entries.get(response_id) + entry = self._entries.get(_store_key(response_id, context)) if entry is None: return False @@ -489,7 +533,9 @@ async def append_stream_event( self._apply_ttl_unlocked(entry, ttl_seconds) return True - async def get_replay_events(self, response_id: str) -> list[StreamEventRecord] | None: + async def get_replay_events( + self, response_id: str, *, context: PlatformContext | None = None + ) -> list[StreamEventRecord] | None: """Get defensive copies of replay events for ``response_id``, filtering out expired events. Events older than the entry's ``replay_event_ttl_seconds`` (default 600s / 10 minutes, @@ -497,18 +543,20 @@ async def get_replay_events(self, response_id: str) -> list[StreamEventRecord] | :param response_id: The unique identifier of the response whose events to retrieve. :type response_id: str + :keyword context: Platform context for partitioning; omitted context selects anonymous state. + :paramtype context: ~azure.ai.agentserver.responses.PlatformContext | None :returns: A list of deep-copied stream event records, or ``None`` if not found. :rtype: list[~azure.ai.agentserver.responses.models.runtime.StreamEventRecord] | None """ async with self._locked(): - entry = self._entries.get(response_id) + entry = self._entries.get(_store_key(response_id, context)) if entry is None: return None cutoff = datetime.now(timezone.utc) - timedelta(seconds=entry.replay_event_ttl_seconds) live = [e for e in entry.replay.events if e.emitted_at >= cutoff] return deepcopy(live) - async def delete(self, response_id: str) -> bool: + async def delete(self, response_id: str, *, context: PlatformContext | None = None) -> bool: """Delete all state for a response ID if present. Removes the entry entirely from the store (unlike ``delete_response`` @@ -516,15 +564,18 @@ async def delete(self, response_id: str) -> bool: :param response_id: The unique identifier of the response to remove. :type response_id: str + :keyword context: Platform context for partitioning; omitted context selects anonymous state. + :paramtype context: ~azure.ai.agentserver.responses.PlatformContext | None :returns: ``True`` if an entry was found and removed, ``False`` otherwise. :rtype: bool """ async with self._locked(): - self._stream_events.pop(response_id, None) - return self._entries.pop(response_id, None) is not None + key = _store_key(response_id, context) + self._stream_events.pop(key, None) + return self._entries.pop(key, None) is not None async def purge_expired(self, *, now: datetime | None = None) -> int: - """Remove expired entries and return count. + """Remove expired entries across all partitions and return count. :keyword ~datetime.datetime or None now: Optional override for the current time (useful for testing). :returns: The number of expired entries that were removed. @@ -569,33 +620,35 @@ def _purge_expired_unlocked(self, *, now: datetime | None = None) -> int: :rtype: int """ current_time = now or datetime.now(timezone.utc) - expired_ids = [ - response_id - for response_id, entry in self._entries.items() + expired_keys = [ + key + for key, entry in self._entries.items() if entry.expires_at is not None and entry.expires_at <= current_time ] - for response_id in expired_ids: - del self._entries[response_id] - self._stream_events.pop(response_id, None) + for key in expired_keys: + del self._entries[key] + self._stream_events.pop(key, None) # Prune orphaned stream events that have no corresponding entry. # Legacy bookkeeping — kept structurally so the in-memory provider # still tracks its expiration loop unchanged. Stream events are # now persisted by the SDK ``streams`` registry, not here. - orphaned_ids = [rid for rid in self._stream_events if rid not in self._entries] - for rid in orphaned_ids: - del self._stream_events[rid] + orphaned_keys = [key for key in self._stream_events if key not in self._entries] + for key in orphaned_keys: + del self._stream_events[key] - return len(expired_ids) + return len(expired_keys) - def _store_output_items_unlocked(self, response: ResponseObject) -> list[str]: + def _store_output_items_unlocked(self, response: ResponseObject, *, context: PlatformContext | None) -> list[str]: """Extract output items from a response, store them in the item store, and return their IDs. Must be called while holding ``self._lock``. :param response: The response envelope whose output items should be stored. :type response: ~azure.ai.agentserver.responses.models._generated.Response + :keyword context: Platform context for partitioning. + :paramtype context: ~azure.ai.agentserver.responses.PlatformContext | None :returns: Ordered list of output item IDs. :rtype: list[str] """ @@ -606,7 +659,7 @@ def _store_output_items_unlocked(self, response: ResponseObject) -> list[str]: for item in output: item_id = self._extract_item_id(item) if item_id is not None: - self._item_store[item_id] = deepcopy(item) + self._item_store[_store_key(item_id, context)] = deepcopy(item) output_ids.append(item_id) return output_ids diff --git a/sdk/agentserver/azure-ai-agentserver-responses/tests/contract/test_user_isolation_enforcement.py b/sdk/agentserver/azure-ai-agentserver-responses/tests/contract/test_user_isolation_enforcement.py index 7bbdaac1b29f..c89295c1fa12 100644 --- a/sdk/agentserver/azure-ai-agentserver-responses/tests/contract/test_user_isolation_enforcement.py +++ b/sdk/agentserver/azure-ai-agentserver-responses/tests/contract/test_user_isolation_enforcement.py @@ -15,12 +15,15 @@ import asyncio import json as _json from typing import Any +from unittest.mock import AsyncMock import pytest from starlette.testclient import TestClient from azure.ai.agentserver.responses import ResponsesAgentServerHost from azure.ai.agentserver.responses._id_generator import IdGenerator +from azure.ai.agentserver.responses._response_context import PlatformContext +from azure.ai.agentserver.responses.store._memory import InMemoryResponseProvider from azure.ai.agentserver.responses.streaming._event_stream import ResponseEventStream from tests._helpers import poll_until @@ -209,6 +212,63 @@ def _build_async_client(handler: Any) -> _AsyncAsgiClient: return _AsyncAsgiClient(app) +@pytest.mark.asyncio +@pytest.mark.parametrize( + "owner_key,other_key,evict", + [ + ("key_A", "key_B", False), + ("key_A", None, False), + ("key_A", "key_B", True), + ("key_A", None, True), + (None, "key_A", True), + ], +) +async def test_memory_provider_isolation_before_and_after_runtime_eviction( + owner_key: str | None, other_key: str | None, evict: bool, monkeypatch: pytest.MonkeyPatch +) -> None: + """Envelope/input provider fallbacks retain the caller's partition after eviction.""" + provider = InMemoryResponseProvider() + host = ResponsesAgentServerHost(store=provider) + host.response_handler(_noop_handler) + client = _AsyncAsgiClient(host) + runtime = host._orchestrator._runtime_state + try_evict = runtime.try_evict + # Hold eager terminal eviction so both sides of the provider fallback are exercised. + monkeypatch.setattr(runtime, "try_evict", AsyncMock(return_value=False)) + owner_headers = {"x-agent-user-id": owner_key} if owner_key is not None else {} + other_headers = {"x-agent-user-id": other_key} if other_key is not None else {} + created = await client.post( + "/responses", + json_body={"model": "m", "input": [{"role": "user", "content": "private input"}], "store": True}, + headers={**owner_headers, "x-agent-foundry-call-id": "creation-call"}, + ) + assert created.status_code == 200 + response_id = created.json()["id"] + monkeypatch.setattr(runtime, "try_evict", try_evict) + assert await runtime.get(response_id) is not None + if evict: + assert await runtime.try_evict(response_id) + assert await runtime.get(response_id) is None + + path = f"/responses/{response_id}" + for method, endpoint in [("GET", path), ("GET", f"{path}/input_items"), ("DELETE", path)]: + denied = await client.request(method, endpoint, headers=other_headers) + assert denied.status_code == 404, denied.body + with pytest.raises(KeyError): + await provider.update_response(created.json(), context=PlatformContext(user_id_key=other_key)) + + later_headers = {**owner_headers, "x-agent-foundry-call-id": "later-call"} + fetched = await client.get(path, headers=later_headers) + assert fetched.status_code == 200 + inputs = await client.get(f"{path}/input_items", headers=later_headers) + assert inputs.status_code == 200 + assert inputs.json()["data"][0]["content"][0]["text"] == "private input" + deleted = await client.request("DELETE", path, headers=later_headers) + assert deleted.status_code == 200 + with pytest.raises(KeyError): + await provider.get_response(response_id, context=PlatformContext(user_id_key=owner_key)) + + # ── GET with isolation ──────────────────────────────────── diff --git a/sdk/agentserver/azure-ai-agentserver-responses/tests/unit/test_in_memory_provider_crud.py b/sdk/agentserver/azure-ai-agentserver-responses/tests/unit/test_in_memory_provider_crud.py index 9b36e9b0f77a..fde30be6eed3 100644 --- a/sdk/agentserver/azure-ai-agentserver-responses/tests/unit/test_in_memory_provider_crud.py +++ b/sdk/agentserver/azure-ai-agentserver-responses/tests/unit/test_in_memory_provider_crud.py @@ -15,8 +15,10 @@ import pytest +from azure.ai.agentserver.responses._response_context import PlatformContext, ResponseContext from azure.ai.agentserver.responses.models import ResponseObject -from azure.ai.agentserver.responses.models.runtime import StreamEventRecord +from azure.ai.agentserver.responses.models.runtime import ResponseExecution, ResponseModeFlags, StreamEventRecord +from azure.ai.agentserver.responses.store import ResponseAlreadyExistsError from azure.ai.agentserver.responses.store._memory import InMemoryResponseProvider # --------------------------------------------------------------------------- @@ -62,6 +64,248 @@ def _output_message(item_id: str, text: str) -> dict[str, Any]: } +_USER_KEYS = ["user_A", "user_B", None, "", " ", " user_A "] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("owner_key,other_key", [(a, b) for a in _USER_KEYS for b in _USER_KEYS if a != b]) +async def test_partitions__foreign_response_and_items_are_missing(owner_key: str | None, other_key: str | None) -> None: + provider = InMemoryResponseProvider() + owner = PlatformContext(user_id_key=owner_key, call_id="shared-call") + other = PlatformContext(user_id_key=other_key, call_id="shared-call") + response = _response("owned", output=[_output_message("output", "private")], conversation_id="conversation") + await provider.create_response(response, [_input_item("input", "private")], None, context=owner) + + for operation in ( + provider.get_response("owned", context=other), + provider.update_response(_response("owned", status="failed"), context=other), + provider.delete_response("owned", context=other), + provider.get_input_items("owned", after="input", before="output", context=other), + ): + with pytest.raises(KeyError, match="not found"): + await operation + assert await provider.get_items(["input", "output"], context=other) == [None, None] + assert await provider.get_history_item_ids("owned", None, 100, context=other) == [] + assert await provider.get_history_item_ids(None, "conversation", 100, context=other) == [] + assert await provider.get_response("owned", context=owner) == response + + +@pytest.mark.asyncio +async def test_partitions__colliding_response_item_and_conversation_ids() -> None: + provider = InMemoryResponseProvider() + contexts = [PlatformContext(user_id_key=key, call_id="create") for key in _USER_KEYS] + for index, context in enumerate(contexts): + await provider.create_response( + _response("same", output=[_output_message("output", str(index))], conversation_id="same"), + [_input_item("input", str(index))], + None, + context=context, + ) + with pytest.raises(ResponseAlreadyExistsError): + await provider.create_response(_response("same"), None, None, context=context) + await provider.create_response( + _response(f"next_{index}", conversation_id="same"), + [_input_item(f"next_input_{index}", str(index))], + await provider.get_history_item_ids("same", None, 100, context=context), + context=context, + ) + + for index, context in enumerate(contexts): + later = PlatformContext(user_id_key=context.user_id_key, call_id="different-call") + items = await provider.get_items(["input", "output", "missing"], context=later) + assert items == [_input_item("input", str(index)), _output_message("output", str(index)), None] + history = await provider.get_history_item_ids(None, "same", 100, context=later) + assert history == ["input", "output", "input", "output", f"next_input_{index}"] + assert await provider.get_history_item_ids(f"next_{index}", None, 100, context=later) == [ + "input", + "output", + f"next_input_{index}", + ] + assert await provider.get_input_items( + f"next_{index}", ascending=True, limit=1, after="input", context=later + ) == [_output_message("output", str(index))] + assert await provider.get_input_items(f"next_{index}", before="input", context=later) == [ + _input_item(f"next_input_{index}", str(index)), + _output_message("output", str(index)), + ] + # An ID from another partition must behave like any other unknown cursor. + foreign_cursor = f"next_input_{(index + 1) % len(contexts)}" + assert await provider.get_input_items("same", after=foreign_cursor, context=later) == [ + _input_item("input", str(index)) + ] + items[0]["content"][0]["text"] = "mutated" + assert (await provider.get_items(["input"], context=later))[0] == _input_item("input", str(index)) + await provider.update_response( + _response("same", output=[_output_message("output", f"updated_{index}")]), context=later + ) + + for index, context in enumerate(contexts): + assert (await provider.get_items(["output"], context=context))[0] == _output_message( + "output", f"updated_{index}" + ) + + await provider.delete_response("same", context=contexts[0]) + with pytest.raises(ValueError, match="deleted"): + await provider.get_input_items("same", context=contexts[0]) + for context in contexts[1:]: + assert (await provider.get_response("same", context=context))["id"] == "same" + assert await provider.get_history_item_ids("same", None, 100, context=context) == ["input", "output"] + + +@pytest.mark.asyncio +async def test_partitions__missing_context_and_unkeyed_context_share_anonymous_crud() -> None: + provider = InMemoryResponseProvider() + unkeyed = PlatformContext(call_id="opaque") + await provider.create_response(_response("anonymous"), [_input_item("input", "local")], None) + assert (await provider.get_response("anonymous", context=unkeyed))["id"] == "anonymous" + await provider.update_response(_response("anonymous", status="failed"), context=unkeyed) + assert (await provider.get_response("anonymous"))["status"] == "failed" + assert await provider.get_items(["input"], context=unkeyed) == [_input_item("input", "local")] + assert await provider.get_input_items("anonymous", context=unkeyed) == [_input_item("input", "local")] + assert await provider.get_history_item_ids("anonymous", None, 100, context=unkeyed) == ["input"] + await provider.delete_response("anonymous", context=unkeyed) + with pytest.raises(KeyError): + await provider.get_response("anonymous") + await provider.create_response(_response("anonymous"), None, None, context=unkeyed) + assert (await provider.get_response("anonymous"))["id"] == "anonymous" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("reader_key", ["user_A", "user_B", None]) +@pytest.mark.parametrize("prefetched", [False, True]) +async def test_partitions__response_context_resolves_only_owned_references( + reader_key: str | None, prefetched: bool +) -> None: + provider = InMemoryResponseProvider() + owner = PlatformContext(user_id_key="user_A") + item = _input_item("owned_item", "private") + await provider.create_response(_response("owned", conversation_id="conversation"), [item], None, context=owner) + reader = PlatformContext(user_id_key=reader_key, call_id="next-request") + ctx = ResponseContext( + response_id="next", + mode_flags=ResponseModeFlags(stream=False, store=True, background=False), + provider=provider, + input_items=[{"type": "item_reference", "id": "owned_item"}], + previous_response_id="owned", + conversation_id="conversation", + platform_context=reader, + prefetched_history_ids=["owned_item"] if prefetched else None, + ) + inputs = await ctx.get_input_items() + history = await ctx.get_history() + if reader_key == "user_A": + assert len(inputs) == 1 + assert inputs[0]["content"] == item["content"] + assert history and all(entry == item for entry in history) + else: + assert inputs == () + assert history == () + # Foreign history pointers supplied by a caller never resolve foreign payloads. + await provider.create_response(_response("next"), None, ["owned_item"], context=reader) + assert await provider.get_input_items("next", context=reader) == ([item] if reader_key == "user_A" else []) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("other_key", ["user_B", None]) +async def test_partitions__legacy_helpers_do_not_bypass_identity(other_key: str | None) -> None: + provider = InMemoryResponseProvider() + owner = PlatformContext(user_id_key="user_A") + other = PlatformContext(user_id_key=other_key) + execution = ResponseExecution( + response_id="execution", + mode_flags=ResponseModeFlags(stream=True, store=True, background=True), + ) + await provider.create_execution(execution, context=owner) + event = StreamEventRecord(sequence_number=0, event_type="response.created", payload={"owner": "A"}) + assert await provider.get_execution("execution", context=other) is None + assert await provider.get_replay_events("execution", context=other) is None + assert not await provider.set_response_snapshot("execution", _response("execution"), context=other) + assert not await provider.transition_execution_status("execution", "in_progress", context=other) + assert not await provider.set_cancel_requested("execution", context=other) + assert not await provider.append_stream_event("execution", event, context=other) + assert not await provider.delete("execution", context=other) + assert await provider.append_stream_event("execution", event, context=owner) + assert (await provider.get_replay_events("execution", context=owner))[0].payload == {"owner": "A"} + assert (await provider.get_execution("execution", context=owner)).status == execution.status + + +@pytest.mark.asyncio +async def test_partitions__legacy_execution_replay_expiry_and_cleanup_collisions() -> None: + provider = InMemoryResponseProvider() + contexts = [PlatformContext(user_id_key="user_A"), PlatformContext(user_id_key="user_B"), None] + execution = ResponseExecution( + response_id="same", + mode_flags=ResponseModeFlags(stream=True, store=True, background=True), + ) + for index, context in enumerate(contexts): + await provider.create_execution(execution, context=context) + with pytest.raises(ValueError, match="already exists"): + await provider.create_execution(execution, context=context) + assert await provider.set_response_snapshot("same", _response("same", status="in_progress"), context=context) + assert await provider.transition_execution_status("same", "in_progress", context=context) + old_event = StreamEventRecord( + sequence_number=0, + event_type="response.created", + payload={"owner": index}, + emitted_at=datetime.now(timezone.utc) - timedelta(seconds=601), + ) + live_event = StreamEventRecord( + sequence_number=1, + event_type="response.in_progress", + payload={"owner": index}, + ) + assert await provider.append_stream_event("same", old_event, context=context) + assert await provider.append_stream_event("same", live_event, context=context) + assert (await provider.get_execution("same", context=context)).status == "in_progress" + replay = await provider.get_replay_events("same", context=context) + assert len(replay) == 1 + assert replay[0].payload == {"owner": index} + replay[0].payload["owner"] = "mutated" + assert (await provider.get_replay_events("same", context=context))[0].payload == {"owner": index} + + # Legacy event bookkeeping shares the same composite keys as response entries. + provider._stream_events[("user_A", "same")] = [] + provider._stream_events[("user_B", "same")] = [] + provider._stream_events[(None, "same")] = [] + provider._stream_events[("user_A", "orphan")] = [] + assert await provider.set_cancel_requested("same", ttl_seconds=10, context=contexts[0]) + assert (await provider.get_execution("same", context=contexts[0])).cancel_requested + assert not (await provider.get_execution("same", context=contexts[1])).cancel_requested + assert await provider.purge_expired(now=datetime.now(timezone.utc) + timedelta(seconds=11)) == 1 + assert await provider.get_execution("same", context=contexts[0]) is None + assert ("user_A", "same") not in provider._stream_events + assert ("user_A", "orphan") not in provider._stream_events + for context in contexts[1:]: + assert (await provider.get_execution("same", context=context)).status == "in_progress" + assert len(await provider.get_replay_events("same", context=context)) == 1 + assert await provider.delete("same", context=contexts[1]) + assert ("user_B", "same") not in provider._stream_events + assert (None, "same") in provider._stream_events + assert await provider.get_execution("same") is not None + # Exercise automatic purge on normal lookups, not only the explicit maintenance method. + provider._entries[(None, "same")].expires_at = datetime.now(timezone.utc) - timedelta(seconds=1) + assert await provider.get_execution("same") is None + assert provider._stream_events == {} + + +@pytest.mark.asyncio +async def test_partitions__expired_response_leaves_other_users_history_intact() -> None: + provider = InMemoryResponseProvider() + contexts = [PlatformContext(user_id_key="user_A"), PlatformContext(user_id_key="user_B"), None] + for context in contexts: + await provider.create_response( + _response("same", conversation_id="conversation"), [_input_item("input", "text")], None, context=context + ) + assert await provider.set_response_snapshot("same", _response("same"), ttl_seconds=10, context=contexts[0]) + assert await provider.purge_expired(now=datetime.now(timezone.utc) + timedelta(seconds=11)) == 1 + with pytest.raises(KeyError): + await provider.get_response("same", context=contexts[0]) + assert await provider.get_history_item_ids(None, "conversation", 100, context=contexts[0]) == [] + for context in contexts[1:]: + assert await provider.get_history_item_ids(None, "conversation", 100, context=context) == ["input"] + assert await provider.get_input_items("same", context=context) == [_input_item("input", "text")] + + # =========================================================================== # Create # ===========================================================================