diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/__init__.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/__init__.py index 813f800eb..026bd32b2 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/__init__.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/__init__.py @@ -104,8 +104,22 @@ # Storage from .storage.store_item import StoreItem -from .storage import Storage -from .storage.memory_storage import MemoryStorage +from .storage import ( + Storage, + StorageDeleteOptions, + StorageDeleteResult, + StorageDeleteResults, + StorageOperationStatus, + StorageProvider, + StorageReadResult, + StorageReadResults, + StorageV2, + StorageWriteMode, + StorageWriteOptions, + StorageWriteResult, + StorageWriteResults, +) +from .storage.memory_storage import MemoryStorage, MemoryStorageV2 # Error Resources from .errors import error_resources, ErrorMessage, ErrorResources @@ -186,7 +200,20 @@ "UserState", "StoreItem", "Storage", + "StorageV2", + "StorageProvider", + "StorageOperationStatus", + "StorageWriteMode", + "StorageWriteOptions", + "StorageDeleteOptions", + "StorageReadResult", + "StorageReadResults", + "StorageWriteResult", + "StorageWriteResults", + "StorageDeleteResult", + "StorageDeleteResults", "MemoryStorage", + "MemoryStorageV2", "AgenticUserAuthorization", "Authorization", "MiddlewareSet", diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/_oauth/_flow_storage_client.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/_oauth/_flow_storage_client.py index 9011c1633..ce1dc2c47 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/_oauth/_flow_storage_client.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/_oauth/_flow_storage_client.py @@ -1,7 +1,13 @@ # Copyright (c) Microsoft Corporation. All rights reserved. # Licensed under the MIT License. -from ..storage import Storage +from ..storage import Storage, StorageProvider +from ..storage.storage_compatibility import ( + as_storage_v2, + assert_storage_delete_succeeded, + assert_storage_write_succeeded, + get_storage_read_value, +) from ._flow_state import _FlowState @@ -31,8 +37,8 @@ def __init__( self, channel_id: str, user_id: str, - storage: Storage, - cache_class: type[Storage] | None = None, + storage: StorageProvider, + cache_class: type[StorageProvider] | None = None, ): """ Args: @@ -66,26 +72,40 @@ def key(self, auth_handler_id: str) -> str: async def read(self, auth_handler_id: str) -> _FlowState | None: """Reads the flow state for a specific authentication handler.""" key: str = self.key(auth_handler_id) - data = await self._cache.read([key], target_cls=_FlowState) - if key not in data: - data = await self._storage.read([key], target_cls=_FlowState) - if key not in data: + cached = await as_storage_v2(self._cache).read([key], target_cls=_FlowState) + data = get_storage_read_value(cached, key) + if data is None: + results = await as_storage_v2(self._storage).read( + [key], target_cls=_FlowState + ) + data = get_storage_read_value(results, key) + if data is None: return None - await self._cache.write({key: data[key]}) - return data.get(key) + cached_results = await as_storage_v2(self._cache).write({key: data}) + assert_storage_write_succeeded(cached_results, [key]) + return data async def write(self, value: _FlowState) -> None: """Saves the flow state for a specific authentication handler.""" key: str = self.key(value.auth_handler_id) - cached_state = await self._cache.read([key], target_cls=_FlowState) - if not cached_state or cached_state.get(key, None) != value: - await self._cache.write({key: value}) - await self._storage.write({key: value}) + cached_results = await as_storage_v2(self._cache).read( + [key], target_cls=_FlowState + ) + cached_state = get_storage_read_value(cached_results, key) + if cached_state != value: + cache_write = await as_storage_v2(self._cache).write({key: value}) + assert_storage_write_succeeded(cache_write, [key]) + storage_write = await as_storage_v2(self._storage).write({key: value}) + assert_storage_write_succeeded(storage_write, [key]) async def delete(self, auth_handler_id: str) -> None: """Deletes the flow state for a specific authentication handler.""" key: str = self.key(auth_handler_id) - cached_state = await self._cache.read([key], target_cls=_FlowState) - if cached_state: - await self._cache.delete([key]) - await self._storage.delete([key]) + cached_state = await as_storage_v2(self._cache).read( + [key], target_cls=_FlowState + ) + if get_storage_read_value(cached_state, key) is not None: + cache_delete = await as_storage_v2(self._cache).delete([key]) + assert_storage_delete_succeeded(cache_delete, [key]) + storage_delete = await as_storage_v2(self._storage).delete([key]) + assert_storage_delete_succeeded(storage_delete, [key]) diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/app_options.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/app_options.py index 283168f26..8c8254d41 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/app_options.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/app_options.py @@ -10,7 +10,7 @@ from microsoft_agents.hosting.core.app.oauth import AuthHandler from microsoft_agents.hosting.core.authorization import Connections -from microsoft_agents.hosting.core.storage import Storage +from microsoft_agents.hosting.core.storage import StorageProvider # from .auth import AuthOptions from .typing_indicator import TypingOptions @@ -33,7 +33,7 @@ class ApplicationOptions: Optional. `AgentApplication` ID of the bot. """ - storage: Optional[Storage] = None + storage: Optional[StorageProvider] = None """ Optional. `Storage` provider to use for the application. """ diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/oauth/_handlers/_authorization_handler.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/oauth/_handlers/_authorization_handler.py index 99b9760ad..bebff2a4c 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/oauth/_handlers/_authorization_handler.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/oauth/_handlers/_authorization_handler.py @@ -10,7 +10,7 @@ from microsoft_agents.activity import TokenResponse from ....turn_context import TurnContext -from ....storage import Storage +from ....storage import StorageProvider from ....authorization import Connections from ..auth_handler import AuthHandler from .._sign_in_response import _SignInResponse @@ -21,13 +21,13 @@ class _AuthorizationHandler(ABC): """Base class for different authorization strategies.""" - _storage: Storage + _storage: StorageProvider _connection_manager: Connections _handler: AuthHandler def __init__( self, - storage: Storage, + storage: StorageProvider, connection_manager: Connections, auth_handler: Optional[AuthHandler] = None, *, diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/oauth/_handlers/agentic_user_authorization.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/oauth/_handlers/agentic_user_authorization.py index 5b1dafd05..c64de855f 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/oauth/_handlers/agentic_user_authorization.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/oauth/_handlers/agentic_user_authorization.py @@ -13,7 +13,7 @@ from ...._oauth import _FlowStateTag from .._sign_in_response import _SignInResponse from ._authorization_handler import _AuthorizationHandler -from ....storage import Storage +from ....storage import StorageProvider from ....authorization import Connections from ..auth_handler import AuthHandler from ..telemetry import spans @@ -26,7 +26,7 @@ class AgenticUserAuthorization(_AuthorizationHandler): def __init__( self, - storage: Storage, + storage: StorageProvider, connection_manager: Connections, auth_handler: Optional[AuthHandler] = None, *, diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/oauth/_handlers/connector_user_authorization.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/oauth/_handlers/connector_user_authorization.py index 0e57433a0..b2de01b93 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/oauth/_handlers/connector_user_authorization.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/oauth/_handlers/connector_user_authorization.py @@ -13,7 +13,7 @@ from ...._oauth._flow_state import _FlowStateTag from ....turn_context import TurnContext -from ....storage import Storage +from ....storage import StorageProvider from ....authorization import Connections from ..auth_handler import AuthHandler from ._authorization_handler import _AuthorizationHandler @@ -30,7 +30,7 @@ class ConnectorUserAuthorization(_AuthorizationHandler): def __init__( self, - storage: Storage, + storage: StorageProvider, connection_manager: Connections, auth_handler: Optional[AuthHandler] = None, *, diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/oauth/authorization.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/oauth/authorization.py index 7d0300025..7190fe02c 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/oauth/authorization.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/oauth/authorization.py @@ -23,7 +23,13 @@ ) from ...turn_context import TurnContext -from ...storage import Storage +from ...storage import StorageProvider +from ...storage.storage_compatibility import ( + as_storage_v2, + assert_storage_delete_succeeded, + assert_storage_write_succeeded, + get_storage_read_value, +) from ...authorization import Connections from ..._oauth import _FlowStateTag from ..state import TurnState @@ -58,13 +64,13 @@ class _AuthInterceptResult: class Authorization: """Class responsible for managing authorization flows.""" - _storage: Storage + _storage: StorageProvider _connection_manager: Connections _handlers: dict[str, _AuthorizationHandler] def __init__( self, - storage: Storage, + storage: StorageProvider, connection_manager: Connections, auth_handlers: dict[str, AuthHandler] | None = None, auto_sign_in: bool = False, @@ -180,7 +186,10 @@ async def _load_sign_in_state(self, context: TurnContext) -> _SignInState | None :rtype: :class:`microsoft_agents.hosting.core.app.oauth._sign_in_state._SignInState` | None """ key = self._sign_in_state_key(context) - return (await self._storage.read([key], target_cls=_SignInState)).get(key) + results = await as_storage_v2(self._storage).read( + [key], target_cls=_SignInState + ) + return get_storage_read_value(results, key) async def _save_sign_in_state( self, context: TurnContext, state: _SignInState @@ -193,7 +202,8 @@ async def _save_sign_in_state( :type state: :class:`microsoft_agents.hosting.core.app.oauth._sign_in_state._SignInState` """ key = self._sign_in_state_key(context) - await self._storage.write({key: state}) + results = await as_storage_v2(self._storage).write({key: state}) + assert_storage_write_succeeded(results, [key]) async def _delete_sign_in_state(self, context: TurnContext) -> None: """Delete the sign-in state from storage for the given context. @@ -202,7 +212,8 @@ async def _delete_sign_in_state(self, context: TurnContext) -> None: :type context: :class:`microsoft_agents.hosting.core.turn_context.TurnContext` """ key = self._sign_in_state_key(context) - await self._storage.delete([key]) + results = await as_storage_v2(self._storage).delete([key]) + assert_storage_delete_succeeded(results, [key]) @staticmethod def _cache_key(context: TurnContext, handler_id: str) -> str: diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/proactive/proactive.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/proactive/proactive.py index f3557830a..5fb175ce2 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/proactive/proactive.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/proactive/proactive.py @@ -9,7 +9,13 @@ from microsoft_agents.activity import Activity, ResourceResponse from microsoft_agents.hosting.core.app.state.turn_state import TurnState -from microsoft_agents.hosting.core.storage import Storage +from microsoft_agents.hosting.core.storage import StorageProvider +from microsoft_agents.hosting.core.storage.storage_compatibility import ( + as_storage_v2, + assert_storage_delete_succeeded, + assert_storage_write_succeeded, + get_storage_read_value, +) from .conversation import Conversation from .create_conversation_options import CreateConversationOptions @@ -76,7 +82,7 @@ def _storage_key(conversation_id: str) -> str: return f"{_STORAGE_KEY_PREFIX}{conversation_id}" @property - def _storage(self) -> Storage: + def _storage(self) -> StorageProvider: storage = self._options.storage or self._app.options.storage if not storage: raise RuntimeError( @@ -126,7 +132,8 @@ async def store_conversation( conversation.validate() key = self._storage_key(conversation.conversation_reference.conversation.id) logger.debug("Storing conversation with key: %s", key) - await self._storage.write({key: conversation}) + results = await as_storage_v2(self._storage).write({key: conversation}) + assert_storage_write_succeeded(results, [key]) async def get_conversation(self, conversation_id: str) -> Conversation | None: """ @@ -141,8 +148,10 @@ async def get_conversation(self, conversation_id: str) -> Conversation | None: """ with spans.ProactiveGetConversation(conversation_id) as span: key = self._storage_key(conversation_id) - results = await self._storage.read([key], target_cls=Conversation) - conversation = results.get(key) + results = await as_storage_v2(self._storage).read( + [key], target_cls=Conversation + ) + conversation = get_storage_read_value(results, key) span.share(found=conversation is not None) return conversation @@ -156,7 +165,8 @@ async def delete_conversation(self, conversation_id: str) -> None: with spans.ProactiveDeleteConversation(conversation_id): key = self._storage_key(conversation_id) logger.debug("Deleting conversation with key: %s", key) - await self._storage.delete([key]) + results = await as_storage_v2(self._storage).delete([key]) + assert_storage_delete_succeeded(results, [key]) # ------------------------------------------------------------------ # Send a single activity diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/proactive/proactive_options.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/proactive/proactive_options.py index 9780b5da3..42c3b38ca 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/proactive/proactive_options.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/proactive/proactive_options.py @@ -7,7 +7,7 @@ from dataclasses import dataclass -from microsoft_agents.hosting.core.storage import Storage +from microsoft_agents.hosting.core.storage import StorageProvider @dataclass @@ -24,7 +24,7 @@ class ProactiveOptions: :type fail_on_unsigned_in_connections: bool """ - storage: Storage | None = None + storage: StorageProvider | None = None """Storage used to persist Conversation objects.""" fail_on_unsigned_in_connections: bool = True diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/state/conversation_state.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/state/conversation_state.py index 4287862e9..9a0653f9f 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/state/conversation_state.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/state/conversation_state.py @@ -8,7 +8,7 @@ from typing import Type -from microsoft_agents.hosting.core.storage import Storage, StoreItem +from microsoft_agents.hosting.core.storage import StorageProvider, StoreItem from microsoft_agents.hosting.core.turn_context import TurnContext from microsoft_agents.hosting.core.state import AgentState @@ -23,7 +23,7 @@ class ConversationState(AgentState): CONTEXT_SERVICE_KEY = "ConversationState" - def __init__(self, storage: Storage) -> None: + def __init__(self, storage: StorageProvider) -> None: """ Initialize ConversationState with a key and optional properties. diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/state/turn_state.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/state/turn_state.py index 0f17366b0..24a7cc89f 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/state/turn_state.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/app/state/turn_state.py @@ -9,7 +9,7 @@ from typing import Any, Optional, Type, TypeVar, Callable import asyncio -from microsoft_agents.hosting.core.storage import Storage +from microsoft_agents.hosting.core.storage import StorageProvider from microsoft_agents.hosting.core.turn_context import TurnContext from microsoft_agents.hosting.core.app.state.conversation_state import ConversationState @@ -52,7 +52,9 @@ def __init__(self, *agent_states: AgentState) -> None: self._scopes[TempState.SCOPE_NAME] = TempState() @classmethod - def with_storage(cls, storage: Storage, *agent_states: AgentState) -> "TurnState": + def with_storage( + cls, storage: StorageProvider, *agent_states: AgentState + ) -> "TurnState": """ Creates TurnState with default ConversationState and UserState. @@ -277,7 +279,7 @@ async def save(self, turn_context: TurnContext, force: bool = False) -> None: ] await asyncio.gather(*tasks) - async def load(self, context: TurnContext, storage: Storage) -> "TurnState": + async def load(self, context: TurnContext, storage: StorageProvider) -> "TurnState": """ Loads a TurnState instance with the default states. diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/client/conversation_id_factory.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/client/conversation_id_factory.py index 455e6cecd..fabc29e74 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/client/conversation_id_factory.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/client/conversation_id_factory.py @@ -2,29 +2,43 @@ # Licensed under the MIT License. from uuid import uuid4 -from functools import partial - +from typing import Any from microsoft_agents.activity import AgentsModel -from microsoft_agents.hosting.core.storage import Storage, StoreItem +from microsoft_agents.hosting.core.storage import StorageProvider, StoreItem +from microsoft_agents.hosting.core.storage.storage_compatibility import ( + as_storage_v2, + assert_storage_delete_succeeded, + assert_storage_write_succeeded, + get_storage_read_value, +) from .agent_conversation_reference import AgentConversationReference from .conversation_id_factory_protocol import ConversationIdFactoryProtocol -def _implement_store_item_for_agents_model_cls(model_instance: AgentsModel): +def _implement_store_item_for_agents_model_cls(model_instance: AgentsModel) -> None: instance_cls = type(model_instance) if not isinstance(model_instance, StoreItem): - instance_cls = type(model_instance) + + def store_item_to_json(instance: AgentsModel) -> dict[str, Any]: + return instance.model_dump(mode="json", exclude_none=True) + + @classmethod + def from_json_to_store_item( + cls: type[AgentsModel], data: dict[str, Any] + ) -> AgentsModel: + return cls.model_validate(data) + setattr( instance_cls, "store_item_to_json", - partial(model_instance.model_dump, mode="json", exclude_none=True), + store_item_to_json, ) - instance_cls.from_json_to_store_item = classmethod(instance_cls.model_validate) + instance_cls.from_json_to_store_item = from_json_to_store_item class ConversationIdFactory(ConversationIdFactoryProtocol): - def __init__(self, storage: Storage) -> None: + def __init__(self, storage: StorageProvider) -> None: if not storage: raise ValueError("ConversationIdFactory.__init__(): storage cannot be None") self._storage = storage @@ -46,7 +60,8 @@ async def create_conversation_id(self, options) -> str: _implement_store_item_for_agents_model_cls(agent_conversation_reference) conversation_info = {agent_conversation_id: agent_conversation_reference} - await self._storage.write(conversation_info) + results = await as_storage_v2(self._storage).write(conversation_info) + assert_storage_write_succeeded(results, [agent_conversation_id]) return agent_conversation_id @@ -58,11 +73,14 @@ async def get_agent_conversation_reference( "ConversationIdFactory.get_agent_conversation_reference(): agent_conversation_id cannot be None" ) - storage_record = await self._storage.read( + storage_record = await as_storage_v2(self._storage).read( [agent_conversation_id], target_cls=AgentConversationReference ) - - return storage_record[agent_conversation_id] + result = get_storage_read_value(storage_record, agent_conversation_id) + if result is None: + raise KeyError(agent_conversation_id) + return result async def delete_conversation_reference(self, agent_conversation_id): - await self._storage.delete([agent_conversation_id]) + results = await as_storage_v2(self._storage).delete([agent_conversation_id]) + assert_storage_delete_succeeded(results, [agent_conversation_id]) diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/state/agent_state.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/state/agent_state.py index 4fc82ef16..62b2706ad 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/state/agent_state.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/state/agent_state.py @@ -8,7 +8,17 @@ import logging from typing import Callable, Type -from microsoft_agents.hosting.core.storage import Storage, StoreItem +from microsoft_agents.hosting.core.storage import ( + StorageProvider, + StorageWriteOptions, + StoreItem, +) +from microsoft_agents.hosting.core.storage.storage_compatibility import ( + as_storage_v2, + assert_storage_delete_succeeded, + assert_storage_write_succeeded, + get_storage_read_value, +) from .state_property_accessor import StatePropertyAccessor from ..turn_context import TurnContext @@ -21,13 +31,18 @@ class CachedAgentState(StoreItem): Internal cached Agent state. """ - def __init__(self, state: dict[str, StoreItem | dict] | None = None): + def __init__( + self, + state: dict[str, StoreItem | dict] | None = None, + version: str | None = None, + ): if state: self.state = state self.hash = self.compute_hash() else: self.state = {} self.hash = hash(str({})) + self.version = version @property def has_state(self) -> bool: @@ -73,7 +88,7 @@ class AgentState: You can define additional scopes for your agent. """ - def __init__(self, storage: Storage, context_service_key: str): + def __init__(self, storage: StorageProvider, context_service_key: str): """ Initializes a new instance of the :class:`microsoft_agents.hosting.core.state.agent_state.AgentState` class. @@ -90,7 +105,7 @@ def __init__(self, storage: Storage, context_service_key: str): :raises: It raises an argument null exception. """ self.state_key = "state" - self._storage = storage + self._storage = as_storage_v2(storage) self._context_service_key = context_service_key self._cached_state: CachedAgentState | None = None @@ -155,7 +170,9 @@ async def load(self, turn_context: TurnContext, force: bool = False) -> None: if self._should_load(turn_context, force): items = await self._storage.read([storage_key], target_cls=CachedAgentState) - val = items.get(storage_key, CachedAgentState()) + val = get_storage_read_value(items, storage_key) or CachedAgentState() + result = items.get(storage_key) + val.version = result.version if result is not None else None self._cached_state = val turn_context.turn_state[self._context_service_key] = val @@ -191,7 +208,19 @@ async def save(self, turn_context: TurnContext, force: bool = False) -> None: if force or (cached_state is not None and cached_state.is_changed): storage_key = self.get_storage_key(turn_context) changes: dict[str, StoreItem] = {storage_key: cached_state} - await self._storage.write(changes) + results = await self._storage.write( + changes, StorageWriteOptions(expected_version=cached_state.version) + ) + write_result = results.get(storage_key) + if write_result is not None and write_result.status.value != "succeeded": + status = write_result.status.name.title().replace("_", "") + raise RuntimeError( + f"AgentState '{self._context_service_key}' could not save key " + f"'{storage_key}' because another turn updated the state first " + f"(status: {status}). This turn's state changes were not saved." + ) + assert_storage_write_succeeded(results, list(changes)) + cached_state.version = results[storage_key].version cached_state.hash = cached_state.compute_hash() def clear(self, turn_context: TurnContext | None = None) -> None: @@ -229,7 +258,8 @@ async def delete(self, turn_context: TurnContext) -> None: turn_context.turn_state.pop(self._context_service_key) storage_key = self.get_storage_key(turn_context) - await self._storage.delete({storage_key}) + results = await self._storage.delete([storage_key]) + assert_storage_delete_succeeded(results, [storage_key]) @abstractmethod def get_storage_key( diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/state/user_state.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/state/user_state.py index 758a1aec3..0dac6c48a 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/state/user_state.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/state/user_state.py @@ -1,7 +1,7 @@ # Copyright (c) Microsoft Corporation. All rights reserved. # Licensed under the MIT License. -from microsoft_agents.hosting.core.storage import Storage +from microsoft_agents.hosting.core.storage import StorageProvider from ..turn_context import TurnContext from .agent_state import AgentState @@ -16,7 +16,7 @@ class UserState(AgentState): "UserState: channel_id and/or conversation missing from context.activity." ) - def __init__(self, storage: Storage, namespace=""): + def __init__(self, storage: StorageProvider, namespace=""): """ Creates a new UserState instance. :param storage: diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/storage/__init__.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/storage/__init__.py index b31e36dcf..49b867262 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/storage/__init__.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/storage/__init__.py @@ -2,8 +2,25 @@ # Licensed under the MIT License. from .store_item import StoreItem -from .storage import Storage, AsyncStorageBase -from .memory_storage import MemoryStorage +from .storage import ( + AsyncStorageBase, + AsyncStorageBaseV2, + Storage, + StorageDeleteOptions, + StorageDeleteResult, + StorageDeleteResults, + StorageOperationStatus, + StorageProvider, + StorageReadResult, + StorageReadResults, + StorageV2, + StorageWriteMode, + StorageWriteOptions, + StorageWriteResult, + StorageWriteResults, + is_store_item, +) +from .memory_storage import MemoryStorage, MemoryStorageV2 from .transcript import ( TranscriptInfo, @@ -19,8 +36,23 @@ __all__ = [ "StoreItem", "Storage", + "StorageV2", + "StorageProvider", + "StorageOperationStatus", + "StorageWriteMode", + "StorageWriteOptions", + "StorageDeleteOptions", + "StorageReadResult", + "StorageReadResults", + "StorageWriteResult", + "StorageWriteResults", + "StorageDeleteResult", + "StorageDeleteResults", + "is_store_item", "AsyncStorageBase", + "AsyncStorageBaseV2", "MemoryStorage", + "MemoryStorageV2", "TranscriptInfo", "TranscriptLogger", "ConsoleTranscriptLogger", diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/storage/memory_storage.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/storage/memory_storage.py index 5e962ecda..f34181b4d 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/storage/memory_storage.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/storage/memory_storage.py @@ -2,79 +2,240 @@ # Licensed under the MIT License. from asyncio import Lock -from typing import TypeVar +from copy import deepcopy +from typing import TypeVar, cast from ._type_aliases import JSON -from .storage import Storage +from .storage import ( + Storage, + StorageDeleteOptions, + StorageDeleteResult, + StorageDeleteResults, + StorageOperationStatus, + StorageReadResult, + StorageReadResults, + StorageV2, + StorageWriteMode, + StorageWriteOptions, + StorageWriteResult, + StorageWriteResults, + is_store_item, +) +from .storage_compatibility import ( + validate_expected_version, + validate_storage_v2_changes, + validate_storage_v2_keys, +) from .store_item import StoreItem +from .telemetry import spans StoreItemT = TypeVar("StoreItemT", bound=StoreItem) -class MemoryStorage(Storage): - """In-memory storage implementation for testing and development purposes.""" - - def __init__(self, state: dict[str, JSON] | None = None): - """Initializes the MemoryStorage with an optional initial state. +class _MemoryStore: + """Shared synchronized persistence mechanics for the two public adapters.""" - :param state: An optional dictionary representing the initial state of the storage. - :raises ValueError: If state is not a dictionary or None. - """ + def __init__(self, state: dict[str, JSON] | None = None) -> None: self._memory: dict[str, JSON] = state or {} + self._versions: dict[str, str] = {} + self._next_version = 1 self._lock = Lock() async def read( - self, keys: list[str], *, target_cls: type[StoreItemT], **kwargs - ) -> dict[str, StoreItemT]: - """Reads items from the in-memory storage. - - :param keys: A list of keys to read from the storage. - :param target_cls: The class type of the items to be read. Must be a subclass of StoreItem. - :return: A dictionary mapping keys to their corresponding StoreItem instances. - :raises ValueError: If keys are empty. - """ - - if not keys: - raise ValueError("Storage.read(): Keys are required when reading.") - - result: dict[str, StoreItemT] = {} + self, + keys: list[str], + *, + target_cls: type[StoreItemT], + copy_data: bool, + ) -> StorageReadResults[StoreItemT]: + results: StorageReadResults[StoreItemT] = StorageReadResults() + async with self._lock: + for key in keys: + if key not in self._memory: + results[key] = cast( + StorageReadResult[StoreItemT], + StorageReadResult( + key=key, status=StorageOperationStatus.NOT_FOUND + ), + ) + continue + data = self._memory[key] + if copy_data: + data = deepcopy(data) + results[key] = cast( + StorageReadResult[StoreItemT], + StorageReadResult( + key=key, + status=StorageOperationStatus.SUCCEEDED, + value=cast( + StoreItemT, target_cls.from_json_to_store_item(data) + ), + version=self._versions.get(key), + ), + ) + return results + + async def write( + self, + changes: dict[str, StoreItem], + options: StorageWriteOptions, + *, + copy_data: bool, + ) -> StorageWriteResults: + results = StorageWriteResults() + async with self._lock: + for key, value in changes.items(): + exists = key in self._memory + current_version = self._versions.get(key) + if options.mode == StorageWriteMode.CREATE_ONLY and exists: + results[key] = StorageWriteResult( + key=key, + status=StorageOperationStatus.CONFLICT, + version=current_version, + ) + continue + if ( + options.expected_version is not None + and options.expected_version != current_version + ): + results[key] = StorageWriteResult( + key=key, + status=StorageOperationStatus.CONDITION_NOT_MET, + version=current_version, + ) + continue + if options.mode == StorageWriteMode.REPLACE and not exists: + results[key] = StorageWriteResult( + key=key, status=StorageOperationStatus.NOT_FOUND + ) + continue + + data = value.store_item_to_json() + self._memory[key] = deepcopy(data) if copy_data else data + version = self._new_version() + self._versions[key] = version + results[key] = StorageWriteResult( + key=key, + status=StorageOperationStatus.SUCCEEDED, + version=version, + ) + return results + + async def delete( + self, keys: list[str], options: StorageDeleteOptions + ) -> StorageDeleteResults: + results = StorageDeleteResults() async with self._lock: for key in keys: - if key == "": - raise ValueError("MemoryStorage.read(): key cannot be empty") - if key in self._memory: - result[key] = target_cls.from_json_to_store_item(self._memory[key]) + if key not in self._memory: + status = ( + StorageOperationStatus.CONDITION_NOT_MET + if options.expected_version is not None + else StorageOperationStatus.NOT_FOUND + ) + results[key] = StorageDeleteResult(key=key, status=status) + continue + current_version = self._versions.get(key) + if ( + options.expected_version is not None + and options.expected_version != current_version + ): + results[key] = StorageDeleteResult( + key=key, + status=StorageOperationStatus.CONDITION_NOT_MET, + version=current_version, + ) + continue + self._memory.pop(key) + self._versions.pop(key, None) + results[key] = StorageDeleteResult( + key=key, + status=StorageOperationStatus.SUCCEEDED, + version=current_version, + ) + return results + + def _new_version(self) -> str: + version = str(self._next_version) + self._next_version += 1 + return version - return result - async def write(self, changes: dict[str, StoreItem]): - """Writes items to the in-memory storage. +class MemoryStorage(Storage): + """Legacy in-memory storage adapter for testing and development.""" - :param changes: A dictionary mapping keys to StoreItem instances to be written to the storage. - :raises ValueError: If changes is None or any key is empty. - """ - if not changes: - raise ValueError("MemoryStorage.write(): changes cannot be None") + def __init__(self, state: dict[str, JSON] | None = None) -> None: + self._store = _MemoryStore(state) - async with self._lock: - for key in changes: - if key == "": - raise ValueError("MemoryStorage.write(): key cannot be empty") - self._memory[key] = changes[key].store_item_to_json() + async def read( + self, keys: list[str], *, target_cls: type[StoreItemT], **kwargs: object + ) -> dict[str, StoreItemT]: + if not keys: + raise ValueError("Storage.read(): Keys are required when reading.") + if any(not key for key in keys): + raise ValueError("MemoryStorage.read(): key cannot be empty") + with spans.StorageRead(len(keys)): + results = await self._store.read( + keys, target_cls=target_cls, copy_data=False + ) + return { + key: value for key in keys if (value := results.get_value(key)) is not None + } + + async def write(self, changes: dict[str, StoreItem]) -> None: + if not changes: + raise ValueError("MemoryStorage.write(): changes cannot be empty") + if any(not key for key in changes): + raise ValueError("MemoryStorage.write(): key cannot be empty") + with spans.StorageWrite(len(changes)): + results = await self._store.write( + changes, StorageWriteOptions(), copy_data=False + ) + results.assert_succeeded(changes) + + async def delete(self, keys: list[str]) -> None: + if not keys: + raise ValueError("Storage.delete(): Keys are required when deleting.") + if any(not key for key in keys): + raise ValueError("MemoryStorage.delete(): key cannot be empty") + with spans.StorageDelete(len(keys)): + results = await self._store.delete(keys, StorageDeleteOptions()) + results.assert_succeeded(keys, allow_not_found=True) - async def delete(self, keys: list[str]): - """Deletes items from the in-memory storage. - :param keys: A list of keys to delete from the storage. - :raises ValueError: If keys is empty or any key is empty. - """ +class MemoryStorageV2(StorageV2): + """In-memory Storage V2 adapter with per-key results and versions.""" - if not keys: - raise ValueError("Storage.delete(): Keys are required when deleting.") + def __init__(self, state: dict[str, JSON] | None = None) -> None: + self._store = _MemoryStore(state) - async with self._lock: - for key in keys: - if key == "": - raise ValueError("MemoryStorage.delete(): key cannot be empty") - if key in self._memory: - del self._memory[key] + async def read( + self, keys: list[str], *, target_cls: type[StoreItemT], **kwargs: object + ) -> StorageReadResults[StoreItemT]: + validate_storage_v2_keys(keys) + with spans.StorageRead(len(keys)): + return await self._store.read(keys, target_cls=target_cls, copy_data=True) + + async def write( + self, + changes: dict[str, StoreItem], + options: StorageWriteOptions | None = None, + ) -> StorageWriteResults: + validate_storage_v2_changes(changes) + write_options = options or StorageWriteOptions() + validate_expected_version(write_options.expected_version) + if any(not is_store_item(value) for value in changes.values()): + raise ValueError("Storage V2 values must implement store_item_to_json().") + with spans.StorageWrite(len(changes)): + return await self._store.write(changes, write_options, copy_data=True) + + async def delete( + self, + keys: list[str], + options: StorageDeleteOptions | None = None, + ) -> StorageDeleteResults: + validate_storage_v2_keys(keys) + delete_options = options or StorageDeleteOptions() + validate_expected_version(delete_options.expected_version) + with spans.StorageDelete(len(keys)): + return await self._store.delete(keys, delete_options) diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/storage/storage.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/storage/storage.py index 5b4161b07..da14dce9f 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/storage/storage.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/storage/storage.py @@ -1,7 +1,10 @@ # Copyright (c) Microsoft Corporation. All rights reserved. # Licensed under the MIT License. -from typing import TypeVar +from dataclasses import dataclass +from enum import Enum +from collections.abc import Iterable +from typing import Generic, NoReturn, TypeAlias, TypeVar from abc import ABC, abstractmethod from asyncio import gather @@ -11,6 +14,131 @@ StoreItemT = TypeVar("StoreItemT", bound=StoreItem) +class StorageOperationStatus(str, Enum): + """Outcome of one version 2 storage operation.""" + + SUCCEEDED = "succeeded" + NOT_FOUND = "notFound" + CONFLICT = "conflict" + CONDITION_NOT_MET = "conditionNotMet" + + +class StorageWriteMode(str, Enum): + """Write mode for a version 2 storage operation.""" + + UPSERT = "upsert" + CREATE_ONLY = "createOnly" + REPLACE = "replace" + + +def is_store_item(value: object) -> bool: + """Return whether a value can be serialized by a storage provider.""" + return callable(getattr(value, "store_item_to_json", None)) + + +@dataclass(frozen=True, slots=True) +class StorageWriteOptions: + """Options applied to every item in a version 2 write operation.""" + + mode: StorageWriteMode = StorageWriteMode.UPSERT + expected_version: str | None = None + + def __post_init__(self) -> None: + if not isinstance(self.mode, StorageWriteMode): + raise ValueError(f'Storage V2 write mode "{self.mode}" is not supported.') + + +@dataclass(frozen=True, slots=True) +class StorageDeleteOptions: + """Options applied to every item in a version 2 delete operation.""" + + expected_version: str | None = None + + +@dataclass(frozen=True, slots=True) +class StorageReadResult(Generic[StoreItemT]): + """Result for one version 2 read operation.""" + + key: str + status: StorageOperationStatus + value: StoreItemT | None = None + version: str | None = None + + +@dataclass(frozen=True, slots=True) +class StorageWriteResult: + """Result for one version 2 write operation.""" + + key: str + status: StorageOperationStatus + version: str | None = None + + +@dataclass(frozen=True, slots=True) +class StorageDeleteResult: + """Result for one version 2 delete operation.""" + + key: str + status: StorageOperationStatus + version: str | None = None + + +def _raise_result_error( + operation: str, key: str, status: StorageOperationStatus | None +) -> NoReturn: + value = status.value if status is not None else "missing" + raise RuntimeError( + f'Storage V2 {operation} failed for key "{key}" with status "{value}".' + ) + + +class StorageReadResults(dict[str, StorageReadResult[StoreItemT]], Generic[StoreItemT]): + """Per-key results returned by a Storage V2 read operation.""" + + def get_value(self, key: str) -> StoreItemT | None: + """Return a successful value, map not-found to ``None``, or raise.""" + result = self.get(key) + if result is not None and result.status == StorageOperationStatus.NOT_FOUND: + return None + if result is not None and result.status == StorageOperationStatus.SUCCEEDED: + return result.value + _raise_result_error("read", key, result.status if result else None) + + +class StorageWriteResults(dict[str, StorageWriteResult]): + """Per-key results returned by a Storage V2 write operation.""" + + def assert_succeeded(self, keys: Iterable[str] | None = None) -> None: + """Raise unless every requested write succeeded.""" + for key in keys if keys is not None else self: + result = self.get(key) + if result is None or result.status != StorageOperationStatus.SUCCEEDED: + _raise_result_error( + "write", key, result.status if result is not None else None + ) + + +class StorageDeleteResults(dict[str, StorageDeleteResult]): + """Per-key results returned by a Storage V2 delete operation.""" + + def assert_succeeded( + self, + keys: Iterable[str] | None = None, + *, + allow_not_found: bool = False, + ) -> None: + """Raise unless every requested delete has an accepted outcome.""" + accepted = {StorageOperationStatus.SUCCEEDED} + if allow_not_found: + accepted.add(StorageOperationStatus.NOT_FOUND) + for key in keys if keys is not None else self: + result = self.get(key) + if result is None or result.status not in accepted: + _raise_result_error( + "delete", key, result.status if result is not None else None + ) + + class Storage(ABC): """Abstract base class for storage implementations.""" @@ -45,6 +173,43 @@ async def delete(self, keys: list[str]) -> None: pass +class StorageV2(ABC): + """Version 2 storage interface. + + Each operation returns a result for every requested key. Values remain + :class:`StoreItem` instances because the Python SDK requires an explicit + deserialization type for reads. + """ + + @abstractmethod + async def read( + self, keys: list[str], *, target_cls: type[StoreItemT], **kwargs + ) -> StorageReadResults[StoreItemT]: + """Reads items and returns one result per requested key.""" + pass + + @abstractmethod + async def write( + self, + changes: dict[str, StoreItem], + options: StorageWriteOptions | None = None, + ) -> StorageWriteResults: + """Writes items and returns one result per requested key.""" + pass + + @abstractmethod + async def delete( + self, + keys: list[str], + options: StorageDeleteOptions | None = None, + ) -> StorageDeleteResults: + """Deletes items and returns one result per requested key.""" + pass + + +StorageProvider: TypeAlias = Storage | StorageV2 + + class AsyncStorageBase(Storage): """Base class for asynchronous storage implementations with operations that work on single items. The bulk operations are implemented in terms @@ -135,3 +300,89 @@ async def delete(self, keys: list[str]) -> None: await self.initialize() await gather(*[self._delete_item(key) for key in keys]) + + +class AsyncStorageBaseV2(StorageV2): + """Build V2 bulk operations from provider-specific single-item operations.""" + + async def initialize(self) -> None: + """Initialize the backing storage when required by the provider.""" + pass + + @abstractmethod + async def _read_item( + self, key: str, *, target_cls: type[StoreItemT], **kwargs + ) -> StorageReadResult[StoreItemT]: + """Read one item and return its result.""" + pass + + async def read( + self, keys: list[str], *, target_cls: type[StoreItemT], **kwargs + ) -> StorageReadResults[StoreItemT]: + if any(not key.strip() for key in keys): + raise ValueError("Storage V2 keys must be non-empty strings.") + if not keys: + return StorageReadResults() + with spans.StorageRead(len(keys)): + await self.initialize() + results = await gather( + *(self._read_item(key, target_cls=target_cls, **kwargs) for key in keys) + ) + return StorageReadResults((result.key, result) for result in results) + + @abstractmethod + async def _write_item( + self, key: str, value: StoreItem, options: StorageWriteOptions + ) -> StorageWriteResult: + """Write one item and return its result.""" + pass + + async def write( + self, + changes: dict[str, StoreItem], + options: StorageWriteOptions | None = None, + ) -> StorageWriteResults: + if any(not key.strip() for key in changes): + raise ValueError("Storage V2 keys must be non-empty strings.") + if not changes: + return StorageWriteResults() + if any(not is_store_item(value) for value in changes.values()): + raise ValueError("Storage V2 values must implement store_item_to_json().") + write_options = options or StorageWriteOptions() + if write_options.expected_version == "": + raise ValueError("Storage V2 expected_version cannot be empty.") + with spans.StorageWrite(len(changes)): + await self.initialize() + results = await gather( + *( + self._write_item(key, value, write_options) + for key, value in changes.items() + ) + ) + return StorageWriteResults((result.key, result) for result in results) + + @abstractmethod + async def _delete_item( + self, key: str, options: StorageDeleteOptions + ) -> StorageDeleteResult: + """Delete one item and return its result.""" + pass + + async def delete( + self, + keys: list[str], + options: StorageDeleteOptions | None = None, + ) -> StorageDeleteResults: + if any(not key.strip() for key in keys): + raise ValueError("Storage V2 keys must be non-empty strings.") + if not keys: + return StorageDeleteResults() + delete_options = options or StorageDeleteOptions() + if delete_options.expected_version == "": + raise ValueError("Storage V2 expected_version cannot be empty.") + with spans.StorageDelete(len(keys)): + await self.initialize() + results = await gather( + *(self._delete_item(key, delete_options) for key in keys) + ) + return StorageDeleteResults((result.key, result) for result in results) diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/storage/storage_compatibility.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/storage/storage_compatibility.py new file mode 100644 index 000000000..15276b769 --- /dev/null +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/storage/storage_compatibility.py @@ -0,0 +1,179 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +"""Compatibility helpers for Storage V1 and Storage V2.""" + +from __future__ import annotations + +from typing import Any, NoReturn + +from .storage import ( + Storage, + StorageDeleteOptions, + StorageDeleteResults, + StorageDeleteResult, + StorageOperationStatus, + StorageProvider, + StoreItemT, + StorageReadResult, + StorageReadResults, + StorageV2, + StorageWriteMode, + StorageWriteOptions, + StorageWriteResults, + StorageWriteResult, +) +from .store_item import StoreItem + + +def is_storage_v2(storage: StorageProvider) -> bool: + """Return ``True`` only for a provider that implements the V2 interface.""" + return isinstance(storage, StorageV2) + + +def as_storage_v2(storage: StorageProvider) -> StorageV2: + """Convert a supported provider to the V2 interface.""" + if is_storage_v2(storage): + return storage + return _StorageToStorageV2Adapter(storage) + + +def get_storage_read_value( + results: StorageReadResults[StoreItemT] | None, key: str +) -> StoreItemT | None: + """Return a successful V2 value, map not-found to ``None``, or raise.""" + if results is None: + _raise_result_error("read", key, None) + return results.get_value(key) + + +def assert_storage_write_succeeded( + results: StorageWriteResults | None, keys: list[str] +) -> None: + """Raise unless every V2 write result succeeded.""" + if results is None: + _raise_result_error("write", keys[0] if keys else "", None) + results.assert_succeeded(keys) + + +def assert_storage_delete_succeeded( + results: StorageDeleteResults | None, keys: list[str] +) -> None: + """Raise unless every V2 delete kept V1 idempotent semantics.""" + if results is None: + _raise_result_error("delete", keys[0] if keys else "", None) + results.assert_succeeded(keys, allow_not_found=True) + + +def validate_storage_v2_keys(keys: list[str]) -> None: + """Validate V2 key input.""" + if any(not key.strip() for key in keys): + raise ValueError("Storage V2 keys must be non-empty strings.") + + +def validate_storage_v2_changes(changes: dict[str, object]) -> None: + """Validate V2 change keys.""" + if any(not key.strip() for key in changes): + raise ValueError("Storage V2 keys must be non-empty strings.") + + +def validate_expected_version(expected_version: str | None) -> None: + """Validate an optional V2 version token.""" + if expected_version == "": + raise ValueError("Storage V2 expected_version cannot be empty.") + + +class _StorageToStorageV2Adapter(StorageV2): + """Adapt a legacy provider where V2 behavior is safely available.""" + + def __init__(self, storage: Storage): + self._storage = storage + + async def read( + self, + keys: list[str], + *, + target_cls: type[StoreItemT], + **kwargs: Any, + ) -> StorageReadResults[StoreItemT]: + validate_storage_v2_keys(keys) + if not keys: + return StorageReadResults() + items = await self._storage.read(keys, target_cls=target_cls, **kwargs) + return StorageReadResults( + { + key: StorageReadResult( + key=key, + status=( + StorageOperationStatus.SUCCEEDED + if key in items + else StorageOperationStatus.NOT_FOUND + ), + value=items.get(key), + ) + for key in keys + } + ) + + async def write( + self, + changes: dict[str, StoreItem], + options: StorageWriteOptions | None = None, + ) -> StorageWriteResults: + validate_storage_v2_changes(changes) + if not changes: + return StorageWriteResults() + options = options or StorageWriteOptions() + validate_expected_version(options.expected_version) + if options.mode != StorageWriteMode.UPSERT: + raise NotImplementedError( + 'Legacy storage does not support the V2 storage option "mode".' + ) + if options.expected_version is not None: + raise NotImplementedError( + "Legacy storage does not support the V2 storage option " + '"expected_version".' + ) + await self._storage.write(changes) + return StorageWriteResults( + { + key: StorageWriteResult( + key=key, status=StorageOperationStatus.SUCCEEDED + ) + for key in changes + } + ) + + async def delete( + self, + keys: list[str], + options: StorageDeleteOptions | None = None, + ) -> StorageDeleteResults: + validate_storage_v2_keys(keys) + if not keys: + return StorageDeleteResults() + options = options or StorageDeleteOptions() + validate_expected_version(options.expected_version) + if options.expected_version is not None: + raise NotImplementedError( + "Legacy storage does not support the V2 storage option " + '"expected_version".' + ) + await self._storage.delete(keys) + return StorageDeleteResults( + { + key: StorageDeleteResult( + key=key, status=StorageOperationStatus.SUCCEEDED + ) + for key in keys + } + ) + + +def _raise_result_error( + operation: str, key: str, status: StorageOperationStatus | None +) -> NoReturn: + value = status.value if status is not None else "missing" + raise RuntimeError( + f'Storage V2 {operation} failed for key "{key}" with status "{value}".' + ) diff --git a/libraries/microsoft-agents-hosting-core/readme.md b/libraries/microsoft-agents-hosting-core/readme.md index e0e114afa..8a5a9c757 100644 --- a/libraries/microsoft-agents-hosting-core/readme.md +++ b/libraries/microsoft-agents-hosting-core/readme.md @@ -246,6 +246,36 @@ async def on_error(context: TurnContext, error: Exception): ## Key Classes Reference +### Storage V2 + +Storage providers use V1 by default. Select V2 when your application needs a +result for each key and optimistic concurrency: + +```python +from microsoft_agents.hosting.core.storage import ( + MemoryStorageV2, + StorageWriteOptions, + StorageWriteMode, +) + +storage = MemoryStorageV2() +results = await storage.write( + {"profile": profile}, + StorageWriteOptions(mode=StorageWriteMode.CREATE_ONLY), +) + +if results["profile"].status.value == "succeeded": + version = results["profile"].version +``` + +V2 operations return `succeeded`, `notFound`, `conflict`, or +`conditionNotMet` for each requested key. The result `version` is a storage +concurrency token. It is separate from model data. + +Storage provider authors can inherit `AsyncStorageBaseV2` and implement the +single-item `_read_item`, `_write_item`, and `_delete_item` hooks. The base +class supplies the asynchronous bulk `read`, `write`, and `delete` methods. + ### Core Classes - **`AgentApplication`** - Main application class with fluent API - **`ActivityHandler`** - Base class for inheritance-based agents @@ -284,4 +314,4 @@ async def on_error(context: TurnContext, error: Exception): |Semantic Kernel Integration|A weather agent built with Semantic Kernel|[semantic-kernel-multiturn](https://github.com/microsoft/Agents/blob/main/samples/python/semantic-kernel-multiturn/README.md)| |Streaming Agent|Streams OpenAI responses|[azure-ai-streaming](https://github.com/microsoft/Agents/blob/main/samples/python/azureai-streaming/README.md)| |Copilot Studio Client|Console app to consume a Copilot Studio Agent|[copilotstudio-client](https://github.com/microsoft/Agents/blob/main/samples/python/copilotstudio-client/README.md)| -|Cards Agent|Agent that uses rich cards to enhance conversation design |[cards](https://github.com/microsoft/Agents/blob/main/samples/python/cards/README.md)| \ No newline at end of file +|Cards Agent|Agent that uses rich cards to enhance conversation design |[cards](https://github.com/microsoft/Agents/blob/main/samples/python/cards/README.md)| diff --git a/libraries/microsoft-agents-hosting-core/setup.py b/libraries/microsoft-agents-hosting-core/setup.py index 56079d3ee..6dd610014 100644 --- a/libraries/microsoft-agents-hosting-core/setup.py +++ b/libraries/microsoft-agents-hosting-core/setup.py @@ -20,6 +20,7 @@ "opentelemetry-sdk>=1.27.0", "aiohttp>=3.11.11", "yarl>=1.17.0,<2.0", + "typing-extensions>=4.12.0", "azure-core", ], ) diff --git a/libraries/microsoft-agents-storage-blob/microsoft_agents/storage/blob/__init__.py b/libraries/microsoft-agents-storage-blob/microsoft_agents/storage/blob/__init__.py index 1c6767a2a..0da326751 100644 --- a/libraries/microsoft-agents-storage-blob/microsoft_agents/storage/blob/__init__.py +++ b/libraries/microsoft-agents-storage-blob/microsoft_agents/storage/blob/__init__.py @@ -1,4 +1,4 @@ -from .blob_storage import BlobStorage +from .blob_storage import BlobStorage, BlobStorageV2 from .blob_storage_config import BlobStorageConfig -__all__ = ["BlobStorage", "BlobStorageConfig"] +__all__ = ["BlobStorage", "BlobStorageV2", "BlobStorageConfig"] diff --git a/libraries/microsoft-agents-storage-blob/microsoft_agents/storage/blob/blob_storage.py b/libraries/microsoft-agents-storage-blob/microsoft_agents/storage/blob/blob_storage.py index 8dfcf8376..2ac9da825 100644 --- a/libraries/microsoft-agents-storage-blob/microsoft_agents/storage/blob/blob_storage.py +++ b/libraries/microsoft-agents-storage-blob/microsoft_agents/storage/blob/blob_storage.py @@ -1,14 +1,27 @@ import json -from typing import TypeVar +from dataclasses import dataclass +from typing import Any, TypeVar, cast from io import BytesIO +from azure.core import MatchConditions from azure.storage.blob.aio import ( - ContainerClient, BlobServiceClient, ) -from microsoft_agents.hosting.core.storage import StoreItem -from microsoft_agents.hosting.core.storage.storage import AsyncStorageBase +from microsoft_agents.hosting.core.storage import ( + StoreItem, + StorageDeleteOptions, + StorageDeleteResult, + StorageOperationStatus, + StorageReadResult, + StorageWriteMode, + StorageWriteOptions, + StorageWriteResult, +) +from microsoft_agents.hosting.core.storage.storage import ( + AsyncStorageBase, + AsyncStorageBaseV2, +) from microsoft_agents.hosting.core.storage._type_aliases import JSON from microsoft_agents.hosting.core.storage.error_handling import ( ignore_error, @@ -21,57 +34,134 @@ StoreItemT = TypeVar("StoreItemT", bound=StoreItem) -class BlobStorage(AsyncStorageBase): - """A Blob Storage provider for storing StoreItem objects in Azure Blob Storage.""" +@dataclass(frozen=True, slots=True) +class _BlobDocument: + content: bytes + version: str | None - def __init__(self, config: BlobStorageConfig): - """Initialize the BlobStorage with the given configuration. - :param config: BlobStorageConfig object containing the configuration for the blob storage. - :raises ValueError: If the container name is not provided in the configuration. - """ +class _BlobStorageBackend: + """Own Azure Blob lifecycle and version-neutral blob operations.""" + def __init__(self, config: BlobStorageConfig) -> None: if not config.container_name: raise ValueError(str(blob_storage_errors.BlobContainerNameRequired)) - - self.config = config - - self._blob_service_client: BlobServiceClient = self._create_client() - self._container_client: ContainerClient = ( - self._blob_service_client.get_container_client(config.container_name) + self._config = config + self._blob_service_client = self._create_client() + self._container_client = self._blob_service_client.get_container_client( + config.container_name ) - self._initialized: bool = False + self._initialized = False def _create_client(self) -> BlobServiceClient: - """Creates a BlobServiceClient based on the provided configuration. - :return: An instance of BlobServiceClient. - :raises ValueError: If the configuration is invalid. - """ - if self.config.url: # connect with URL and credentials - if not self.config.credential: + if self._config.url: + if not self._config.credential: raise ValueError( blob_storage_errors.InvalidConfiguration.format( "Credential is required when using a custom service URL" ) ) return BlobServiceClient( - account_url=self.config.url, credential=self.config.credential - ) - - else: # connect with connection string - return BlobServiceClient.from_connection_string( - self.config.connection_string + account_url=self._config.url, credential=self._config.credential ) + return BlobServiceClient.from_connection_string(self._config.connection_string) async def initialize(self) -> None: - """Initializes the storage container""" if not self._initialized: - # This should only happen once - assuming this is a singleton. await ignore_error( self._container_client.create_container(), is_status_code_error(409) ) self._initialized = True + async def read(self, key: str, **kwargs: Any) -> _BlobDocument | None: + blob_client = self._container_client.get_blob_client(key) + try: + downloader = await blob_client.download_blob(**kwargs) + return _BlobDocument( + content=await downloader.readall(), + version=_etag_from(downloader.properties), + ) + except Exception as error: # noqa: BLE001 + if _status_code(error) == 404: + return None + raise + + async def write( + self, + key: str, + content: bytes, + *, + overwrite: bool, + etag: str | None = None, + match_condition: MatchConditions | None = None, + ) -> str | None: + upload_options: dict[str, Any] = {"overwrite": overwrite} + if etag is not None: + upload_options["etag"] = etag + if match_condition is not None: + upload_options["match_condition"] = match_condition + response = await self._container_client.get_blob_client(key).upload_blob( + BytesIO(content), length=len(content), **upload_options + ) + return _etag_from(response) + + async def delete( + self, + key: str, + *, + etag: str | None = None, + match_condition: MatchConditions | None = None, + ) -> None: + delete_options: dict[str, Any] = {} + if etag is not None: + delete_options["etag"] = etag + if match_condition is not None: + delete_options["match_condition"] = match_condition + await self._container_client.get_blob_client(key).delete_blob(**delete_options) + + async def get_version(self, key: str) -> str | None: + try: + properties = await self._container_client.get_blob_client( + key + ).get_blob_properties() + return _etag_from(properties) + except Exception as error: # noqa: BLE001 + if _status_code(error) == 404: + return None + raise + + async def close(self) -> None: + await self._container_client.close() + await self._blob_service_client.close() + + +def _etag_from(properties: Any) -> str | None: + if isinstance(properties, dict): + return properties.get("etag") + return getattr(properties, "etag", None) + + +def _status_code(error: Exception) -> int | None: + return getattr(error, "status_code", None) + + +class BlobStorage(AsyncStorageBase): + """Legacy Azure Blob storage adapter.""" + + def __init__(self, config: BlobStorageConfig): + """Initialize the BlobStorage with the given configuration. + + :param config: BlobStorageConfig object containing the configuration for the blob storage. + :raises ValueError: If the container name is not provided in the configuration. + """ + + self.config = config + self._backend = _BlobStorageBackend(config) + + async def initialize(self) -> None: + """Initializes the storage container""" + await self._backend.initialize() + async def _read_item( self, key: str, *, target_cls: type[StoreItemT], **kwargs ) -> tuple[str | None, StoreItemT | None]: @@ -81,17 +171,14 @@ async def _read_item( :param target_cls: The type of the StoreItem to deserialize into. :return: A tuple containing the key and the deserialized StoreItem, or (None, None) if not found. """ - item = await ignore_error( - self._container_client.download_blob(blob=key, timeout=5), - is_status_code_error(404), - ) - if not item: + read_options = {"timeout": 5, **kwargs} + document = await self._backend.read(key, **read_options) + if document is None: return None, None - item_rep: bytes = await item.readall() - item_JSON: JSON = json.loads(item_rep) + item_JSON: JSON = json.loads(document.content) try: - return key, target_cls.from_json_to_store_item(item_JSON) + return key, cast(StoreItemT, target_cls.from_json_to_store_item(item_JSON)) except AttributeError as error: raise TypeError( f"BlobStorage.read_item(): could not deserialize blob item into {target_cls} class. Error: {error}" @@ -111,13 +198,7 @@ async def _write_item(self, key: str, item: StoreItem) -> None: ) item_rep_bytes = json.dumps(item_JSON).encode("utf-8") - # getting the length is important for performance with large blobs - await self._container_client.upload_blob( - name=key, - data=BytesIO(item_rep_bytes), - overwrite=True, - length=len(item_rep_bytes), - ) + await self._backend.write(key, item_rep_bytes, overwrite=True) async def _delete_item(self, key: str) -> None: """Deletes an item from blob storage. @@ -125,11 +206,166 @@ async def _delete_item(self, key: str) -> None: :param key: The key of the item to delete. :raises ValueError: If the deletion fails for reasons other than the item not existing. """ - await ignore_error( - self._container_client.delete_blob(blob=key), is_status_code_error(404) + try: + await self._backend.delete(key) + except Exception as error: # noqa: BLE001 + if _status_code(error) != 404: + raise + + async def close(self) -> None: + """Close the Azure clients owned by this adapter.""" + await self._backend.close() + + +class BlobStorageV2(AsyncStorageBaseV2): + """Azure Blob Storage V2 adapter with per-key results.""" + + def __init__(self, config: BlobStorageConfig): + self.config = config + self._backend = _BlobStorageBackend(config) + + async def initialize(self) -> None: + await self._backend.initialize() + + async def _read_item( + self, key: str, *, target_cls: type[StoreItemT], **kwargs: Any + ) -> StorageReadResult[StoreItemT]: + read_options = {"timeout": 5, **kwargs} + document = await self._backend.read(key, **read_options) + if document is None: + return cast( + StorageReadResult[StoreItemT], + StorageReadResult(key=key, status=StorageOperationStatus.NOT_FOUND), + ) + value = cast( + StoreItemT, + target_cls.from_json_to_store_item(json.loads(document.content)), + ) + return cast( + StorageReadResult[StoreItemT], + StorageReadResult( + key=key, + status=StorageOperationStatus.SUCCEEDED, + value=value, + version=document.version, + ), ) - async def _close(self) -> None: - """Cleans up the storage resources.""" - await self._container_client.close() - await self._blob_service_client.close() + async def _write_item( + self, + key: str, + value: StoreItem, + options: StorageWriteOptions, + ) -> StorageWriteResult: + payload = json.dumps(value.store_item_to_json()).encode("utf-8") + try: + condition_version = ( + options.expected_version + if options.expected_version is not None + else "*" if options.mode == StorageWriteMode.REPLACE else None + ) + version = await self._backend.write( + key, + payload, + overwrite=options.mode != StorageWriteMode.CREATE_ONLY, + etag=condition_version, + match_condition=( + MatchConditions.IfNotModified + if condition_version is not None + else None + ), + ) + return StorageWriteResult( + key=key, + status=StorageOperationStatus.SUCCEEDED, + version=version, + ) + except Exception as error: # noqa: BLE001 + status_code = _status_code(error) + if ( + options.mode == StorageWriteMode.CREATE_ONLY + and status_code in (409, 412) + and options.expected_version is not None + ): + current_version = await self._backend.get_version(key) + return StorageWriteResult( + key=key, + status=( + StorageOperationStatus.CONFLICT + if current_version is not None + else StorageOperationStatus.CONDITION_NOT_MET + ), + version=current_version, + ) + if options.mode == StorageWriteMode.CREATE_ONLY and ( + status_code == 409 + or (status_code == 412 and options.expected_version is None) + ): + return StorageWriteResult( + key=key, + status=StorageOperationStatus.CONFLICT, + version=await self._backend.get_version(key), + ) + if status_code == 404: + return StorageWriteResult( + key=key, + status=( + StorageOperationStatus.CONDITION_NOT_MET + if options.expected_version is not None + else StorageOperationStatus.NOT_FOUND + ), + ) + if status_code == 412: + return StorageWriteResult( + key=key, + status=( + StorageOperationStatus.NOT_FOUND + if options.mode == StorageWriteMode.REPLACE + and options.expected_version is None + else StorageOperationStatus.CONDITION_NOT_MET + ), + version=await self._backend.get_version(key), + ) + raise + + async def _delete_item( + self, + key: str, + options: StorageDeleteOptions, + ) -> StorageDeleteResult: + try: + await self._backend.delete( + key, + etag=options.expected_version, + match_condition=( + MatchConditions.IfNotModified + if options.expected_version is not None + else None + ), + ) + return StorageDeleteResult( + key=key, + status=StorageOperationStatus.SUCCEEDED, + version=options.expected_version, + ) + except Exception as error: # noqa: BLE001 + if _status_code(error) == 404: + return StorageDeleteResult( + key=key, + status=( + StorageOperationStatus.CONDITION_NOT_MET + if options.expected_version is not None + else StorageOperationStatus.NOT_FOUND + ), + ) + if _status_code(error) == 412: + return StorageDeleteResult( + key=key, + status=StorageOperationStatus.CONDITION_NOT_MET, + version=await self._backend.get_version(key), + ) + raise + + async def close(self) -> None: + """Close the Azure clients owned by this adapter.""" + await self._backend.close() diff --git a/libraries/microsoft-agents-storage-blob/readme.md b/libraries/microsoft-agents-storage-blob/readme.md index bfa41490f..256e32758 100644 --- a/libraries/microsoft-agents-storage-blob/readme.md +++ b/libraries/microsoft-agents-storage-blob/readme.md @@ -230,10 +230,15 @@ config = BlobStorageConfig( ## Key Classes Reference -- **`BlobStorage`** - Main storage implementation using Azure Blob Storage +- **`BlobStorage`** - Legacy storage implementation using Azure Blob Storage +- **`BlobStorageV2`** - Storage V2 implementation with per-key results and optimistic concurrency - **`BlobStorageConfig`** - Configuration settings for connection and authentication - **`StoreItem`** - Base class for data models (inherit to create custom types) +`BlobStorage` and `BlobStorageV2` are separate public implementations. They +share only private Azure client lifecycle and raw blob operations; each class +owns its contract-specific read, write, and delete behavior. + # Quick Links - 📦 [All SDK Packages on PyPI](https://pypi.org/search/?q=microsoft-agents) @@ -251,4 +256,4 @@ config = BlobStorageConfig( |Semantic Kernel Integration|A weather agent built with Semantic Kernel|[semantic-kernel-multiturn](https://github.com/microsoft/Agents/blob/main/samples/python/semantic-kernel-multiturn/README.md)| |Streaming Agent|Streams OpenAI responses|[azure-ai-streaming](https://github.com/microsoft/Agents/blob/main/samples/python/azureai-streaming/README.md)| |Copilot Studio Client|Console app to consume a Copilot Studio Agent|[copilotstudio-client](https://github.com/microsoft/Agents/blob/main/samples/python/copilotstudio-client/README.md)| -|Cards Agent|Agent that uses rich cards to enhance conversation design |[cards](https://github.com/microsoft/Agents/blob/main/samples/python/cards/README.md)| \ No newline at end of file +|Cards Agent|Agent that uses rich cards to enhance conversation design |[cards](https://github.com/microsoft/Agents/blob/main/samples/python/cards/README.md)| diff --git a/libraries/microsoft-agents-storage-cosmos/microsoft_agents/storage/cosmos/__init__.py b/libraries/microsoft-agents-storage-cosmos/microsoft_agents/storage/cosmos/__init__.py index 00d54854d..7aa7d6423 100644 --- a/libraries/microsoft-agents-storage-cosmos/microsoft_agents/storage/cosmos/__init__.py +++ b/libraries/microsoft-agents-storage-cosmos/microsoft_agents/storage/cosmos/__init__.py @@ -1,7 +1,8 @@ -from .cosmos_db_storage import CosmosDBStorage +from .cosmos_db_storage import CosmosDBStorage, CosmosDBStorageV2 from .cosmos_db_storage_config import CosmosDBStorageConfig __all__ = [ "CosmosDBStorage", + "CosmosDBStorageV2", "CosmosDBStorageConfig", ] diff --git a/libraries/microsoft-agents-storage-cosmos/microsoft_agents/storage/cosmos/cosmos_db_storage.py b/libraries/microsoft-agents-storage-cosmos/microsoft_agents/storage/cosmos/cosmos_db_storage.py index 68bd7d563..207551888 100644 --- a/libraries/microsoft-agents-storage-cosmos/microsoft_agents/storage/cosmos/cosmos_db_storage.py +++ b/libraries/microsoft-agents-storage-cosmos/microsoft_agents/storage/cosmos/cosmos_db_storage.py @@ -1,13 +1,14 @@ # Copyright (c) Microsoft Corporation. All rights reserved. # Licensed under the MIT License. -from typing import TypeVar +from typing import Any, TypeVar, cast import asyncio from azure.cosmos import ( documents, CosmosDict, ) +from azure.core import MatchConditions from azure.cosmos.aio import ( ContainerProxy, CosmosClient, @@ -16,7 +17,18 @@ import azure.cosmos.exceptions as cosmos_exceptions from azure.cosmos.partition_key import NonePartitionKeyValue -from microsoft_agents.hosting.core.storage import AsyncStorageBase, StoreItem +from microsoft_agents.hosting.core.storage import ( + AsyncStorageBaseV2, + AsyncStorageBase, + StoreItem, + StorageDeleteOptions, + StorageDeleteResult, + StorageOperationStatus, + StorageReadResult, + StorageWriteMode, + StorageWriteOptions, + StorageWriteResult, +) from microsoft_agents.hosting.core.storage._type_aliases import JSON from microsoft_agents.hosting.core.storage.error_handling import ignore_error from microsoft_agents.storage.cosmos.errors import storage_errors @@ -31,32 +43,19 @@ ) -class CosmosDBStorage(AsyncStorageBase): - """A CosmosDB based storage provider using partitioning""" - - def __init__(self, config: CosmosDBStorageConfig): - """Create the storage object. - - :param config: - """ - super().__init__() +class _CosmosStorageBackend: + """Own Cosmos DB lifecycle and version-neutral document operations.""" + def __init__(self, config: CosmosDBStorageConfig) -> None: CosmosDBStorageConfig.validate_cosmos_db_config(config) - - self._config: CosmosDBStorageConfig = config - self._client: CosmosClient = self._create_client() + self._config = config + self._client = self._create_client() self._database: DatabaseProxy | None = None self._container: ContainerProxy | None = None - self._compatability_mode_partition_key: bool = False - # Lock used for synchronizing container creation - self._lock: asyncio.Lock = asyncio.Lock() + self._compatability_mode_partition_key = False + self._lock = asyncio.Lock() def _create_client(self) -> CosmosClient: - """Create a CosmosClient based on the configuration. - - :return: A CosmosClient instance. - :raises ValueError: If the configuration is invalid. - """ if self._config.url: if not self._config.credential: raise ValueError( @@ -67,13 +66,9 @@ def _create_client(self) -> CosmosClient: return CosmosClient( url=self._config.url, credential=self._config.credential ) - connection_policy = self._config.cosmos_client_options.get( "connection_policy", documents.ConnectionPolicy() ) - - # kwargs 'connection_verify' is to handle CosmosClient overwriting the - # ConnectionPolicy.DisableSSLVerification value. return CosmosClient( self._config.cosmos_db_endpoint, self._config.auth_key, @@ -87,83 +82,86 @@ def _create_client(self) -> CosmosClient: ) def _sanitize(self, key: str) -> str: - """Sanitize the key for use in CosmosDB.""" return sanitize_key( key, self._config.key_suffix, self._config.compatibility_mode ) - async def _read_item( - self, key: str, *, target_cls: type[StoreItemT], **kwargs - ) -> tuple[str | None, StoreItemT | None]: - """Read an item from the storage. - - :param key: The key of the item to read. - :param target_cls: The type of the item to read. - :return: A tuple containing the real key and the item, or (None, None) if not found. - :raises ValueError: If the key is empty. - """ + def _get_partition_key(self, key: str): + return NonePartitionKeyValue if self._compatability_mode_partition_key else key + def _document(self, key: str, content: JSON) -> CosmosDict: if key == "": raise ValueError(str(storage_errors.CosmosDbKeyCannotBeEmpty)) + return { + "id": self._sanitize(key), + "realId": key, + "document": content, + } - escaped_key: str = self._sanitize(key) - read_item_response: CosmosDict | None = await ignore_error( - self._container.read_item( - escaped_key, self._get_partition_key(escaped_key) - ), - cosmos_resource_not_found, - ) - if read_item_response is None: - return None, None - - doc: JSON | None = read_item_response.get("document") - if doc is None: - return read_item_response["realId"], None - return read_item_response["realId"], target_cls.from_json_to_store_item(doc) - - async def _write_item(self, key: str, item: StoreItem) -> None: - """Write an item to the storage. - - :param key: The key of the item to write. - :param item: The item to write. - :raises ValueError: If the key is empty. - """ + async def read(self, key: str, **kwargs: Any) -> CosmosDict: if key == "": raise ValueError(str(storage_errors.CosmosDbKeyCannotBeEmpty)) + escaped_key = self._sanitize(key) + return await self._container.read_item( + escaped_key, self._get_partition_key(escaped_key), **kwargs + ) - escaped_key: str = self._sanitize(key) - - doc = { - "id": escaped_key, - "realId": key, # to retrieve the raw key later - "document": item.store_item_to_json(), + async def try_read(self, key: str) -> CosmosDict | None: + try: + return await self.read(key) + except Exception as error: # noqa: BLE001 + if _status_code(error) == 404: + return None + raise + + async def create(self, key: str, content: JSON) -> CosmosDict: + return await self._container.create_item(body=self._document(key, content)) + + async def upsert(self, key: str, content: JSON) -> CosmosDict: + return await self._container.upsert_item(body=self._document(key, content)) + + async def replace( + self, + key: str, + content: JSON, + *, + etag: str | None = None, + match_condition: MatchConditions | None = None, + ) -> CosmosDict: + escaped_key = self._sanitize(key) + replace_options: dict[str, Any] = { + "item": escaped_key, + "body": self._document(key, content), } - await self._container.upsert_item(body=doc) - - async def _delete_item(self, key: str) -> None: - """Delete an item from the storage. - - :param key: The key of the item to delete. - :raises ValueError: If the key is empty. - """ + if etag is not None: + replace_options["etag"] = etag + if match_condition is not None: + replace_options["match_condition"] = match_condition + return await self._container.replace_item(**replace_options) + + async def delete( + self, + key: str, + *, + etag: str | None = None, + match_condition: MatchConditions | None = None, + ) -> None: if key == "": raise ValueError(str(storage_errors.CosmosDbKeyCannotBeEmpty)) - - escaped_key: str = self._sanitize(key) - - await ignore_error( - self._container.delete_item( - escaped_key, self._get_partition_key(escaped_key) - ), - cosmos_resource_not_found, + escaped_key = self._sanitize(key) + delete_options: dict[str, Any] = {} + if etag is not None: + delete_options["etag"] = etag + if match_condition is not None: + delete_options["match_condition"] = match_condition + await self._container.delete_item( + escaped_key, + self._get_partition_key(escaped_key), + **delete_options, ) async def _create_container(self) -> None: - """Create the container if it does not exist.""" - partition_key = { - "paths": ["/id"], - "kind": documents.PartitionKind.Hash, - } + partition_key = {"paths": ["/id"], "kind": documents.PartitionKind.Hash} try: kwargs = {} if self._config.container_throughput: @@ -171,48 +169,249 @@ async def _create_container(self) -> None: self._container = await self._database.create_container( self._config.container_id, partition_key, **kwargs ) - except Exception as err: + except Exception: self._container = self._database.get_container_client( self._config.container_id ) properties = await self._container.read() - # if "partitionKey" not in properties: - # self._compatability_mode_partition_key = True - # else: - # containers created had no partition key, so the default was "/_partitionKey" paths = properties["partitionKey"]["paths"] if "/_partitionKey" in paths: self._compatability_mode_partition_key = True elif "/id" not in paths: raise Exception( storage_errors.InvalidConfiguration.format( - f"Custom Partition Key Paths are not supported. {self._config.container_id} has a custom Partition Key Path of {paths[0]}." + "Custom Partition Key Paths are not supported. " + f"{self._config.container_id} has a custom Partition " + f"Key Path of {paths[0]}." ) ) async def initialize(self) -> None: - """Initialize the storage provider.""" if not self._container: async with self._lock: - # in case another async task attempted to initialize just before acquiring the lock if self._container: return - if not self._database: self._database = await self._client.create_database_if_not_exists( self._config.database_id ) - await self._create_container() - def _get_partition_key(self, key: str): - """Get the partition key for the given key, considering compatibility mode. + async def close(self) -> None: + await self._client.close() + + +def _status_code(error: Exception) -> int | None: + return getattr(error, "status_code", None) + + +class CosmosDBStorage(AsyncStorageBase): + """Legacy Cosmos DB storage adapter.""" + + def __init__(self, config: CosmosDBStorageConfig): + """Create the storage object.""" + self._config = config + self._backend = _CosmosStorageBackend(config) - :param key: The key for which to get the partition key. - :return: The partition key value. + async def _read_item( + self, key: str, *, target_cls: type[StoreItemT], **kwargs + ) -> tuple[str | None, StoreItemT | None]: + """Read an item from the storage. + + :param key: The key of the item to read. + :param target_cls: The type of the item to read. + :return: A tuple containing the real key and the item, or (None, None) if not found. + :raises ValueError: If the key is empty. """ - return NonePartitionKeyValue if self._compatability_mode_partition_key else key - async def _close(self) -> None: - """Close the storage provider.""" - await self._client.close() + read_item_response: CosmosDict | None = await ignore_error( + self._backend.read(key, **kwargs), + cosmos_resource_not_found, + ) + if read_item_response is None: + return None, None + + doc: JSON | None = read_item_response.get("document") + if doc is None: + return read_item_response["realId"], None + return read_item_response["realId"], cast( + StoreItemT, target_cls.from_json_to_store_item(doc) + ) + + async def _write_item(self, key: str, item: StoreItem) -> None: + """Write an item to the storage. + + :param key: The key of the item to write. + :param item: The item to write. + :raises ValueError: If the key is empty. + """ + await self._backend.upsert(key, item.store_item_to_json()) + + async def _delete_item(self, key: str) -> None: + """Delete an item from the storage. + + :param key: The key of the item to delete. + :raises ValueError: If the key is empty. + """ + await ignore_error( + self._backend.delete(key), + cosmos_resource_not_found, + ) + + async def initialize(self) -> None: + await self._backend.initialize() + + async def close(self) -> None: + """Close the Azure client owned by this adapter.""" + await self._backend.close() + + +class CosmosDBStorageV2(AsyncStorageBaseV2): + """Cosmos DB Storage V2 adapter with per-key results.""" + + def __init__(self, config: CosmosDBStorageConfig): + self._config = config + self._backend = _CosmosStorageBackend(config) + + async def _read_item( + self, key: str, *, target_cls: type[StoreItemT], **kwargs: Any + ) -> StorageReadResult[StoreItemT]: + try: + document = await self._backend.read(key, **kwargs) + return cast( + StorageReadResult[StoreItemT], + StorageReadResult( + key=key, + status=StorageOperationStatus.SUCCEEDED, + value=cast( + StoreItemT, + target_cls.from_json_to_store_item(document["document"]), + ), + version=document.get("_etag"), + ), + ) + except Exception as error: # noqa: BLE001 + if _status_code(error) == 404: + return cast( + StorageReadResult[StoreItemT], + StorageReadResult(key=key, status=StorageOperationStatus.NOT_FOUND), + ) + raise + + async def _write_item( + self, + key: str, + value: StoreItem, + options: StorageWriteOptions, + ) -> StorageWriteResult: + if ( + options.mode == StorageWriteMode.CREATE_ONLY + and options.expected_version is not None + ): + current = await self._backend.try_read(key) + current_version = current.get("_etag") if current else None + return StorageWriteResult( + key=key, + status=( + StorageOperationStatus.CONFLICT + if current is not None + else StorageOperationStatus.CONDITION_NOT_MET + ), + version=current_version, + ) + content = value.store_item_to_json() + try: + if options.mode == StorageWriteMode.CREATE_ONLY: + response = await self._backend.create(key, content) + elif ( + options.mode == StorageWriteMode.REPLACE + or options.expected_version is not None + ): + response = await self._backend.replace( + key, + content, + etag=options.expected_version, + match_condition=( + MatchConditions.IfNotModified + if options.expected_version is not None + else None + ), + ) + else: + response = await self._backend.upsert(key, content) + return StorageWriteResult( + key=key, + status=StorageOperationStatus.SUCCEEDED, + version=response.get("_etag"), + ) + except Exception as error: # noqa: BLE001 + status_code = _status_code(error) + if options.mode == StorageWriteMode.CREATE_ONLY and status_code == 409: + return StorageWriteResult( + key=key, + status=StorageOperationStatus.CONFLICT, + version=(await self._backend.try_read(key) or {}).get("_etag"), + ) + if status_code == 404: + return StorageWriteResult( + key=key, + status=( + StorageOperationStatus.CONDITION_NOT_MET + if options.expected_version is not None + else StorageOperationStatus.NOT_FOUND + ), + ) + if status_code == 412: + return StorageWriteResult( + key=key, + status=StorageOperationStatus.CONDITION_NOT_MET, + version=(await self._backend.try_read(key) or {}).get("_etag"), + ) + raise + + async def _delete_item( + self, + key: str, + options: StorageDeleteOptions, + ) -> StorageDeleteResult: + try: + await self._backend.delete( + key, + etag=options.expected_version, + match_condition=( + MatchConditions.IfNotModified + if options.expected_version is not None + else None + ), + ) + return StorageDeleteResult( + key=key, + status=StorageOperationStatus.SUCCEEDED, + version=options.expected_version, + ) + except Exception as error: # noqa: BLE001 + status_code = _status_code(error) + if status_code == 404: + return StorageDeleteResult( + key=key, + status=( + StorageOperationStatus.CONDITION_NOT_MET + if options.expected_version is not None + else StorageOperationStatus.NOT_FOUND + ), + ) + if status_code == 412: + return StorageDeleteResult( + key=key, + status=StorageOperationStatus.CONDITION_NOT_MET, + version=(await self._backend.try_read(key) or {}).get("_etag"), + ) + raise + + async def initialize(self) -> None: + """Initialize the storage provider.""" + await self._backend.initialize() + + async def close(self) -> None: + """Close the Azure client owned by this adapter.""" + await self._backend.close() diff --git a/libraries/microsoft-agents-storage-cosmos/microsoft_agents/storage/cosmos/cosmos_db_storage_config.py b/libraries/microsoft-agents-storage-cosmos/microsoft_agents/storage/cosmos/cosmos_db_storage_config.py index 596a10360..4123628e7 100644 --- a/libraries/microsoft-agents-storage-cosmos/microsoft_agents/storage/cosmos/cosmos_db_storage_config.py +++ b/libraries/microsoft-agents-storage-cosmos/microsoft_agents/storage/cosmos/cosmos_db_storage_config.py @@ -1,5 +1,4 @@ import json - from azure.core.credentials_async import AsyncTokenCredential from microsoft_agents.storage.cosmos.errors import storage_errors @@ -64,7 +63,9 @@ def __init__( self.credential: AsyncTokenCredential | None = credential @staticmethod - def validate_cosmos_db_config(config: "CosmosDBStorageConfig") -> None: + def validate_cosmos_db_config( + config: "CosmosDBStorageConfig", + ) -> None: """Validate the CosmosDBConfig object. This is used prior to the creation of the CosmosDBStorage object.""" diff --git a/libraries/microsoft-agents-storage-cosmos/readme.md b/libraries/microsoft-agents-storage-cosmos/readme.md index 5a3b33810..14d319c87 100644 --- a/libraries/microsoft-agents-storage-cosmos/readme.md +++ b/libraries/microsoft-agents-storage-cosmos/readme.md @@ -181,6 +181,12 @@ Additionally we provide a Copilot Studio Client, to interact with Agents created pip install microsoft-agents-storage-cosmos ``` +## Storage version + +Use `CosmosDBStorage` for the legacy storage contract or `CosmosDBStorageV2` +for per-key operation results and optimistic concurrency tokens. Both classes +accept the same `CosmosDBStorageConfig`. + ## Environment Setup @@ -206,10 +212,15 @@ Install and run the Azure Cosmos DB Emulator for local testing: ## Key Classes Reference -- **`CosmosDBStorage`** - Main storage implementation using Azure Cosmos DB +- **`CosmosDBStorage`** - Legacy storage implementation using Azure Cosmos DB +- **`CosmosDBStorageV2`** - Storage V2 implementation with per-key results and optimistic concurrency - **`CosmosDBStorageConfig`** - Configuration settings for connection and behavior - **`StoreItem`** - Base class for data models (inherit to create custom types) +`CosmosDBStorage` and `CosmosDBStorageV2` are separate public implementations. +They share only private Azure client lifecycle and raw document operations; +each class owns its contract-specific read, write, and delete behavior. + # Quick Links - 📦 [All SDK Packages on PyPI](https://pypi.org/search/?q=microsoft-agents) @@ -227,4 +238,4 @@ Install and run the Azure Cosmos DB Emulator for local testing: |Semantic Kernel Integration|A weather agent built with Semantic Kernel|[semantic-kernel-multiturn](https://github.com/microsoft/Agents/blob/main/samples/python/semantic-kernel-multiturn/README.md)| |Streaming Agent|Streams OpenAI responses|[azure-ai-streaming](https://github.com/microsoft/Agents/blob/main/samples/python/azureai-streaming/README.md)| |Copilot Studio Client|Console app to consume a Copilot Studio Agent|[copilotstudio-client](https://github.com/microsoft/Agents/blob/main/samples/python/copilotstudio-client/README.md)| -|Cards Agent|Agent that uses rich cards to enhance conversation design |[cards](https://github.com/microsoft/Agents/blob/main/samples/python/cards/README.md)| \ No newline at end of file +|Cards Agent|Agent that uses rich cards to enhance conversation design |[cards](https://github.com/microsoft/Agents/blob/main/samples/python/cards/README.md)| diff --git a/test_samples/app_style/README.md b/test_samples/app_style/README.md index 4c79ffb03..eb409837d 100644 --- a/test_samples/app_style/README.md +++ b/test_samples/app_style/README.md @@ -42,3 +42,4 @@ Invoke-RestMethod -Method POST -Uri "http://localhost:5199/api/sendmessage" -Con ``` When `TOKENVALIDATION__ENABLED` is `true`, add an `Authorization: Bearer ` header to each call. The proactive endpoints will respond with JSON payloads describing success or validation errors. + diff --git a/tests/hosting_core/state/test_agent_state.py b/tests/hosting_core/state/test_agent_state.py index 195e309c3..60925cb8e 100644 --- a/tests/hosting_core/state/test_agent_state.py +++ b/tests/hosting_core/state/test_agent_state.py @@ -16,7 +16,16 @@ from microsoft_agents.hosting.core.state.user_state import UserState from microsoft_agents.hosting.core.app.state.conversation_state import ConversationState from microsoft_agents.hosting.core.turn_context import TurnContext -from microsoft_agents.hosting.core.storage import Storage, StoreItem, MemoryStorage +from microsoft_agents.hosting.core.storage import ( + Storage, + StorageV2, + StorageOperationStatus, + StorageWriteResult, + StorageWriteResults, + StoreItem, + MemoryStorage, + MemoryStorageV2, +) from microsoft_agents.activity import ( Activity, ActivityTypes, @@ -174,7 +183,7 @@ async def test_state_save_no_changes(self): await self.user_state.load(self.context) # Save without making changes - should not call storage - storage_mock = MagicMock(spec=Storage) + storage_mock = MagicMock(spec=StorageV2) storage_mock.write = AsyncMock() self.user_state._storage = storage_mock @@ -208,8 +217,16 @@ async def test_state_save_with_force(self): await self.user_state.load(self.context) # Use a mock storage to verify write is called even without changes - storage_mock = MagicMock(spec=Storage) - storage_mock.write = AsyncMock() + storage_mock = MagicMock(spec=StorageV2) + storage_key = self.user_state.get_storage_key(self.context) + write_results = StorageWriteResults( + { + storage_key: StorageWriteResult( + key=storage_key, status=StorageOperationStatus.SUCCEEDED + ) + } + ) + storage_mock.write = AsyncMock(return_value=write_results) self.user_state._storage = storage_mock await self.user_state.save(self.context, force=True) @@ -477,6 +494,64 @@ async def test_memory_storage_integration(self): assert storage_key in stored_data assert stored_data[storage_key] is not None + @pytest.mark.asyncio + async def test_memory_storage_v2_integration(self): + memory_storage = MemoryStorageV2() + user_state = UserState(memory_storage) + + await user_state.load(self.context) + property_accessor = user_state.create_property("memory_test") + await property_accessor.set(self.context, _MockTestDataItem("memory_value")) + await user_state.save(self.context) + + storage_key = user_state.get_storage_key(self.context) + stored_data = await memory_storage.read( + [storage_key], target_cls=CachedAgentState + ) + + assert stored_data[storage_key].value is not None + + @pytest.mark.asyncio + async def test_stale_state_save_does_not_overwrite_newer_state(self): + """Reject a stale turn after a newer turn saves the same state record.""" + storage = MemoryStorageV2() + + seed_state = UserState(storage) + await seed_state.load(self.context) + seed_property = seed_state.create_property("writer") + await seed_property.set(self.context, _MockTestDataItem("seed")) + await seed_state.save(self.context) + + slow_context = TurnContext(self.adapter, self.activity) + fast_context = TurnContext(self.adapter, self.activity) + slow_state = UserState(storage) + fast_state = UserState(storage) + slow_property = slow_state.create_property("writer") + fast_property = fast_state.create_property("writer") + + await slow_state.load(slow_context) + await fast_state.load(fast_context) + await slow_property.set(slow_context, _MockTestDataItem("slow")) + await fast_property.set(fast_context, _MockTestDataItem("fast")) + + await fast_state.save(fast_context) + expected_error = ( + "AgentState 'Internal.UserState' could not save key " + "'test-channel/users/test-user' because another turn updated the state " + "first (status: ConditionNotMet). This turn's state changes were not saved." + ) + with pytest.raises(RuntimeError) as error: + await slow_state.save(slow_context) + assert str(error.value) == expected_error + + final_context = TurnContext(self.adapter, self.activity) + final_state = UserState(storage) + final_property = final_state.create_property("writer") + await final_state.load(final_context) + value = await final_property.get(final_context, target_cls=_MockTestDataItem) + + assert value.value == "fast" + @pytest.mark.asyncio async def test_state_property_accessor_error_conditions(self): """Test StatePropertyAccessor error conditions.""" diff --git a/tests/hosting_core/storage/test_async_storage_base_v2.py b/tests/hosting_core/storage/test_async_storage_base_v2.py new file mode 100644 index 000000000..08e4224cb --- /dev/null +++ b/tests/hosting_core/storage/test_async_storage_base_v2.py @@ -0,0 +1,119 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +import pytest + +from microsoft_agents.hosting.core.storage import ( + AsyncStorageBaseV2, + StorageDeleteOptions, + StorageDeleteResult, + StorageDeleteResults, + StorageOperationStatus, + StorageReadResult, + StorageReadResults, + StorageWriteMode, + StorageWriteOptions, + StorageWriteResult, + StorageWriteResults, +) +from tests._common.storage.utils import MockStoreItem + + +class _RecordingStorageV2(AsyncStorageBaseV2): + def __init__(self) -> None: + self.initialize_count = 0 + self.read_calls = [] + self.write_calls = [] + self.delete_calls = [] + + async def initialize(self) -> None: + self.initialize_count += 1 + + async def _read_item(self, key, *, target_cls, **kwargs): + self.read_calls.append((key, target_cls, kwargs)) + return StorageReadResult( + key=key, + status=StorageOperationStatus.SUCCEEDED, + value=target_cls({"key": key}), + version=f"version-{key}", + ) + + async def _write_item(self, key, value, options): + self.write_calls.append((key, value, options)) + return StorageWriteResult( + key=key, + status=StorageOperationStatus.SUCCEEDED, + version=f"version-{key}", + ) + + async def _delete_item(self, key, options): + self.delete_calls.append((key, options)) + return StorageDeleteResult( + key=key, + status=StorageOperationStatus.SUCCEEDED, + ) + + +@pytest.mark.asyncio +async def test_base_builds_v2_bulk_operations_from_single_item_hooks(): + storage = _RecordingStorageV2() + write_options = StorageWriteOptions(mode=StorageWriteMode.CREATE_ONLY) + delete_options = StorageDeleteOptions(expected_version="version-a") + values = {"a": MockStoreItem(), "b": MockStoreItem()} + + reads = await storage.read( + ["a", "b"], target_cls=MockStoreItem, custom_argument=True + ) + writes = await storage.write(values, write_options) + deletes = await storage.delete(["a", "b"], delete_options) + + assert isinstance(reads, StorageReadResults) + assert isinstance(writes, StorageWriteResults) + assert isinstance(deletes, StorageDeleteResults) + assert list(reads) == ["a", "b"] + assert list(writes) == ["a", "b"] + assert list(deletes) == ["a", "b"] + assert storage.initialize_count == 3 + assert storage.read_calls == [ + ("a", MockStoreItem, {"custom_argument": True}), + ("b", MockStoreItem, {"custom_argument": True}), + ] + assert storage.write_calls == [ + ("a", values["a"], write_options), + ("b", values["b"], write_options), + ] + assert storage.delete_calls == [ + ("a", delete_options), + ("b", delete_options), + ] + + +@pytest.mark.asyncio +async def test_base_returns_empty_results_without_initializing(): + storage = _RecordingStorageV2() + + reads = await storage.read([], target_cls=MockStoreItem) + writes = await storage.write({}) + deletes = await storage.delete([]) + + assert isinstance(reads, StorageReadResults) + assert isinstance(writes, StorageWriteResults) + assert isinstance(deletes, StorageDeleteResults) + assert not reads + assert not writes + assert not deletes + assert storage.initialize_count == 0 + + +@pytest.mark.asyncio +async def test_base_validates_v2_inputs_before_initializing(): + storage = _RecordingStorageV2() + + with pytest.raises(ValueError, match="keys must be non-empty"): + await storage.read([" "], target_cls=MockStoreItem) + with pytest.raises(ValueError, match="values must implement"): + await storage.write({"key": object()}) # type: ignore[dict-item] + with pytest.raises(ValueError, match="expected_version cannot be empty"): + await storage.delete(["key"], StorageDeleteOptions(expected_version="")) + + assert storage.initialize_count == 0 diff --git a/tests/hosting_core/storage/test_memory_storage.py b/tests/hosting_core/storage/test_memory_storage.py index e09427938..79e0e1545 100644 --- a/tests/hosting_core/storage/test_memory_storage.py +++ b/tests/hosting_core/storage/test_memory_storage.py @@ -1,7 +1,31 @@ from contextlib import asynccontextmanager -from microsoft_agents.hosting.core.storage.memory_storage import MemoryStorage +import pytest + +from microsoft_agents.hosting.core.storage import ( + StorageDeleteOptions, + StorageOperationStatus, + StorageWriteMode, + StorageWriteOptions, +) +from microsoft_agents.hosting.core.storage.memory_storage import ( + MemoryStorage, + MemoryStorageV2, +) from tests._common.storage.utils import CRUDStorageTests +from tests._common.storage.utils import MockStoreItem + + +class _StoreItemShape: + def __init__(self, data=None): + self.data = data or {} + + def store_item_to_json(self): + return self.data + + @classmethod + def from_json_to_store_item(cls, data): + return cls(data) class TestMemoryStorage(CRUDStorageTests): @@ -13,3 +37,125 @@ async def storage(self, initial_data=None): for key, value in (initial_data or {}).items() } yield MemoryStorage(data) + + +@pytest.mark.asyncio +async def test_v2_returns_a_result_for_each_read_key(): + storage = MemoryStorageV2() + await storage.write({"existing": MockStoreItem({"value": 1})}) + + results = await storage.read(["existing", "missing"], target_cls=MockStoreItem) + + assert results["existing"].status == StorageOperationStatus.SUCCEEDED + assert results["existing"].value == MockStoreItem({"value": 1}) + assert results["existing"].version is not None + assert results["missing"].status == StorageOperationStatus.NOT_FOUND + + +@pytest.mark.asyncio +async def test_v2_create_replace_and_conditional_delete(): + storage = MemoryStorageV2() + created = await storage.write( + {"key": MockStoreItem({"value": 1})}, + StorageWriteOptions(mode=StorageWriteMode.CREATE_ONLY), + ) + duplicate = await storage.write( + {"key": MockStoreItem({"value": 2})}, + StorageWriteOptions(mode=StorageWriteMode.CREATE_ONLY), + ) + replaced = await storage.write( + {"key": MockStoreItem({"value": 2})}, + StorageWriteOptions( + mode=StorageWriteMode.REPLACE, + expected_version=created["key"].version, + ), + ) + stale = await storage.write( + {"key": MockStoreItem({"value": 3})}, + StorageWriteOptions( + mode=StorageWriteMode.REPLACE, + expected_version=created["key"].version, + ), + ) + deleted = await storage.delete( + ["key"], + StorageDeleteOptions(expected_version=replaced["key"].version), + ) + + assert created["key"].status == StorageOperationStatus.SUCCEEDED + assert duplicate["key"].status == StorageOperationStatus.CONFLICT + assert replaced["key"].status == StorageOperationStatus.SUCCEEDED + assert stale["key"].status == StorageOperationStatus.CONDITION_NOT_MET + assert deleted["key"].status == StorageOperationStatus.SUCCEEDED + + +@pytest.mark.asyncio +async def test_v2_does_not_mutate_or_share_store_item_data(): + storage = MemoryStorageV2() + value = MockStoreItem({"nested": {"value": 1}}) + await storage.write({"key": value}) + value.data["nested"]["value"] = 2 + + first = await storage.read(["key"], target_cls=MockStoreItem) + first["key"].value.data["nested"]["value"] = 3 + second = await storage.read(["key"], target_cls=MockStoreItem) + + assert second["key"].value == MockStoreItem({"nested": {"value": 1}}) + + +@pytest.mark.asyncio +async def test_v2_accepts_existing_store_item_shape_models(): + storage = MemoryStorageV2() + + await storage.write({"key": _StoreItemShape({"value": 1})}) + result = await storage.read(["key"], target_cls=_StoreItemShape) + + assert result["key"].value.data == {"value": 1} + + +@pytest.mark.asyncio +async def test_v2_accepts_empty_batches_and_rejects_empty_version(): + storage = MemoryStorageV2() + + assert await storage.read([], target_cls=MockStoreItem) == {} + assert await storage.write({}) == {} + assert await storage.delete([]) == {} + with pytest.raises(ValueError, match="expected_version cannot be empty"): + await storage.write( + {"key": MockStoreItem()}, + StorageWriteOptions(expected_version=""), + ) + + +@pytest.mark.asyncio +async def test_v1_interface_does_not_accept_v2_options(): + storage = MemoryStorage() + + with pytest.raises(TypeError, match="positional argument"): + await storage.write({"key": MockStoreItem()}, StorageWriteOptions()) + with pytest.raises(TypeError, match="positional argument"): + await storage.delete(["key"], StorageDeleteOptions()) + + +@pytest.mark.asyncio +async def test_v2_conditional_operations_report_missing_as_condition_not_met(): + storage = MemoryStorageV2() + + written = await storage.write( + {"missing": MockStoreItem()}, + StorageWriteOptions(mode=StorageWriteMode.REPLACE, expected_version="stale"), + ) + deleted = await storage.delete( + ["missing"], StorageDeleteOptions(expected_version="stale") + ) + + assert written["missing"].status == StorageOperationStatus.CONDITION_NOT_MET + assert deleted["missing"].status == StorageOperationStatus.CONDITION_NOT_MET + + +@pytest.mark.asyncio +async def test_v1_empty_write_reports_empty_batch(): + storage = MemoryStorage() + + with pytest.raises(ValueError, match="changes cannot be empty"): + await storage.write({}) diff --git a/tests/hosting_core/storage/test_storage_compatibility.py b/tests/hosting_core/storage/test_storage_compatibility.py new file mode 100644 index 000000000..656410a80 --- /dev/null +++ b/tests/hosting_core/storage/test_storage_compatibility.py @@ -0,0 +1,149 @@ +import pytest + +from microsoft_agents.activity import AgentsModel +from microsoft_agents.hosting.core.storage import ( + Storage, + StorageDeleteOptions, + StorageDeleteResult, + StorageOperationStatus, + StorageReadResult, + StorageReadResults, + StorageWriteResults, + StorageDeleteResults, + StorageWriteMode, + StorageWriteOptions, +) +from microsoft_agents.hosting.core.storage.storage_compatibility import ( + as_storage_v2, + assert_storage_delete_succeeded, + assert_storage_write_succeeded, + get_storage_read_value, +) +from microsoft_agents.hosting.core.client.conversation_id_factory import ( + _implement_store_item_for_agents_model_cls, +) +from tests._common.storage.utils import MockStoreItem + + +class _LegacyStorage(Storage): + def __init__(self): + self.items = {} + + async def read(self, keys, *, target_cls, **kwargs): + return {key: self.items[key] for key in keys if key in self.items} + + async def write(self, changes): + self.items.update(changes) + + async def delete(self, keys): + for key in keys: + self.items.pop(key, None) + + +class _ModelItem(AgentsModel): + value: str + + +def test_write_options_reject_invalid_mode(): + with pytest.raises(ValueError, match='mode "bogus"'): + StorageWriteOptions(mode="bogus") # type: ignore[arg-type] + + +def test_agents_model_store_item_serialization_uses_current_instance(): + first = _ModelItem(value="one") + second = _ModelItem(value="two") + _implement_store_item_for_agents_model_cls(first) + + assert first.store_item_to_json() == {"value": "one"} + assert second.store_item_to_json() == {"value": "two"} + + +@pytest.mark.asyncio +async def test_v1_adapter_returns_explicit_v2_results(): + storage = _LegacyStorage() + await storage.write({"existing": MockStoreItem({"value": 1})}) + + results = await as_storage_v2(storage).read( + ["existing", "missing"], target_cls=MockStoreItem + ) + + assert results["existing"].status == StorageOperationStatus.SUCCEEDED + assert results["missing"].status == StorageOperationStatus.NOT_FOUND + + +@pytest.mark.asyncio +async def test_v1_adapter_rejects_unsupported_conditions(): + storage = as_storage_v2(_LegacyStorage()) + + with pytest.raises(NotImplementedError, match='option "mode"'): + await storage.write( + {"key": MockStoreItem()}, + StorageWriteOptions(mode=StorageWriteMode.CREATE_ONLY), + ) + with pytest.raises(NotImplementedError, match='option "expected_version"'): + await storage.delete(["key"], StorageDeleteOptions(expected_version="1")) + + +def test_result_helpers_reject_missing_or_failed_results(): + assert ( + get_storage_read_value( + StorageReadResults( + { + "key": StorageReadResult( + key="key", status=StorageOperationStatus.NOT_FOUND + ) + } + ), + "key", + ) + is None + ) + with pytest.raises(RuntimeError, match='status "missing"'): + assert_storage_write_succeeded(StorageWriteResults(), ["key"]) + with pytest.raises(RuntimeError, match='status "conditionNotMet"'): + assert_storage_delete_succeeded( + StorageDeleteResults( + { + "key": StorageDeleteResult( + key="key", status=StorageOperationStatus.CONDITION_NOT_MET + ) + } + ), + ["key"], + ) + + +@pytest.mark.asyncio +async def test_v2_accepts_agents_model_store_item_shape(): + value = _ModelItem(value="one") + _implement_store_item_for_agents_model_cls(value) + storage = as_storage_v2(_LegacyStorage()) + + await storage.write({"key": value}) + result = await storage.read(["key"], target_cls=_ModelItem) + + assert get_storage_read_value(result, "key") == value + + +def test_result_collections_expose_result_handling_behavior(): + reads = StorageReadResults( + { + "missing": StorageReadResult( + key="missing", status=StorageOperationStatus.NOT_FOUND + ) + } + ) + writes = StorageWriteResults() + deletes = StorageDeleteResults( + { + "missing": StorageDeleteResult( + key="missing", status=StorageOperationStatus.NOT_FOUND + ) + } + ) + + assert reads.get_value("missing") is None + assert not writes + deletes.assert_succeeded(["missing"], allow_not_found=True) + with pytest.raises(RuntimeError, match='status "missing"'): + writes.assert_succeeded(["missing"]) diff --git a/tests/storage_blob/test_blob_storage.py b/tests/storage_blob/test_blob_storage.py index acc18debd..41fdbfb30 100644 --- a/tests/storage_blob/test_blob_storage.py +++ b/tests/storage_blob/test_blob_storage.py @@ -1,3 +1,4 @@ +import asyncio import json import gc import os @@ -8,7 +9,14 @@ import pytest_asyncio from dotenv import load_dotenv -from microsoft_agents.storage.blob import BlobStorage, BlobStorageConfig +from microsoft_agents.storage.blob import BlobStorage, BlobStorageConfig, BlobStorageV2 +from microsoft_agents.storage.blob.blob_storage import _BlobStorageBackend +from microsoft_agents.hosting.core.storage import ( + StorageDeleteOptions, + StorageOperationStatus, + StorageWriteOptions, + StorageWriteMode, +) from azure.storage.blob.aio import BlobServiceClient, ContainerClient from azure.core.exceptions import ResourceNotFoundError from azure.identity.aio import DefaultAzureCredential @@ -26,6 +34,278 @@ # TEST_BLOB_STORAGE_ACCOUNT_URL set +def test_public_blob_adapters_have_separate_implementations(): + assert not issubclass(BlobStorage, BlobStorageV2) + assert not issubclass(BlobStorageV2, BlobStorage) + + +@pytest.mark.asyncio +async def test_v1_rejects_v2_options_instead_of_ignoring_them(): + storage = object.__new__(BlobStorage) + with pytest.raises(TypeError, match="positional argument"): + await storage.write({"key": MockStoreItem()}, StorageWriteOptions()) + with pytest.raises(TypeError, match="positional argument"): + await storage.delete(["key"], StorageDeleteOptions()) + + +class _ConcurrentCallBarrier: + def __init__(self, expected: int): + self._expected = expected + self._active = 0 + self.max_active = 0 + self._ready = asyncio.Event() + + async def wait(self): + self._active += 1 + self.max_active = max(self.max_active, self._active) + if self._active == self._expected: + self._ready.set() + await asyncio.wait_for(self._ready.wait(), timeout=1) + self._active -= 1 + + +class _BlobDownloader: + properties = {"etag": "v1"} + + def __init__(self, key: str): + self._key = key + + async def readall(self): + return json.dumps({"value": self._key}).encode() + + +class _ConcurrentBlobClient: + def __init__(self, key: str, barrier: _ConcurrentCallBarrier): + self._key = key + self._barrier = barrier + + async def download_blob(self, **_kwargs): + await self._barrier.wait() + return _BlobDownloader(self._key) + + async def get_blob_properties(self): + return {"etag": "v1"} + + async def upload_blob(self, *_args, **_kwargs): + await self._barrier.wait() + return {"etag": "v2"} + + async def delete_blob(self, **_kwargs): + await self._barrier.wait() + return None + + +class _ConcurrentBlobContainer: + def __init__(self, barrier: _ConcurrentCallBarrier): + self._barrier = barrier + + def get_blob_client(self, key: str): + return _ConcurrentBlobClient(key, self._barrier) + + +def _create_v2_blob_storage(barrier: _ConcurrentCallBarrier): + storage = object.__new__(BlobStorageV2) + backend = object.__new__(_BlobStorageBackend) + backend._initialized = True + backend._container_client = _ConcurrentBlobContainer(barrier) + storage._backend = backend + return storage + + +@pytest.mark.asyncio +async def test_v2_batches_run_independent_blob_operations_concurrently(): + barrier = _ConcurrentCallBarrier(2) + storage = _create_v2_blob_storage(barrier) + read = await storage.read(["one", "two"], target_cls=MockStoreItem) + assert all( + result.status == StorageOperationStatus.SUCCEEDED for result in read.values() + ) + assert barrier.max_active == 2 + + barrier = _ConcurrentCallBarrier(2) + storage = _create_v2_blob_storage(barrier) + write = await storage.write({"one": MockStoreItem(), "two": MockStoreItem()}) + assert all( + result.status == StorageOperationStatus.SUCCEEDED for result in write.values() + ) + assert barrier.max_active == 2 + + barrier = _ConcurrentCallBarrier(2) + storage = _create_v2_blob_storage(barrier) + delete = await storage.delete(["one", "two"]) + assert all( + result.status == StorageOperationStatus.SUCCEEDED for result in delete.values() + ) + assert barrier.max_active == 2 + + +class _StatusError(Exception): + def __init__(self, status_code: int): + self.status_code = status_code + + +class _RecordingBlobClient: + def __init__(self): + self.property_calls = 0 + self.download_calls = [] + self.upload_calls = [] + self.delete_calls = [] + self.upload_error = None + self.delete_error = None + + async def get_blob_properties(self): + self.property_calls += 1 + return {"etag": "current"} + + async def download_blob(self, **kwargs): + self.download_calls.append(kwargs) + return _BlobDownloader("key") + + async def upload_blob(self, *_args, **kwargs): + self.upload_calls.append(kwargs) + if self.upload_error: + raise self.upload_error + return {"etag": "next"} + + async def delete_blob(self, **kwargs): + self.delete_calls.append(kwargs) + if self.delete_error: + raise self.delete_error + + +class _RecordingBlobContainer: + def __init__(self, client): + self.client = client + + def get_blob_client(self, _key): + return self.client + + +def _recording_v2_blob_storage(client): + storage = object.__new__(BlobStorageV2) + backend = object.__new__(_BlobStorageBackend) + backend._initialized = True + backend._container_client = _RecordingBlobContainer(client) + storage._backend = backend + return storage + + +@pytest.mark.asyncio +async def test_v2_blob_write_conditions_are_atomic_and_skip_prereads(): + client = _RecordingBlobClient() + storage = _recording_v2_blob_storage(client) + + await storage.write({"key": MockStoreItem()}) + await storage.write( + {"key": MockStoreItem()}, + StorageWriteOptions(mode=StorageWriteMode.REPLACE), + ) + await storage.write( + {"key": MockStoreItem()}, StorageWriteOptions(expected_version="expected") + ) + + assert client.property_calls == 0 + assert "etag" not in client.upload_calls[0] + assert client.upload_calls[1]["etag"] == "*" + assert client.upload_calls[2]["etag"] == "expected" + + +@pytest.mark.asyncio +async def test_v2_blob_replace_reports_missing_without_expected_version(): + client = _RecordingBlobClient() + client.upload_error = _StatusError(412) + storage = _recording_v2_blob_storage(client) + + result = await storage.write( + {"key": MockStoreItem()}, + StorageWriteOptions(mode=StorageWriteMode.REPLACE), + ) + + assert result["key"].status == StorageOperationStatus.NOT_FOUND + + +@pytest.mark.asyncio +async def test_v2_blob_read_forwards_provider_options(): + client = _RecordingBlobClient() + storage = _recording_v2_blob_storage(client) + + await storage.read(["key"], target_cls=MockStoreItem, timeout=9) + + assert client.download_calls == [{"timeout": 9}] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("storage_cls", [BlobStorage, BlobStorageV2]) +async def test_blob_adapters_expose_public_close(storage_cls): + class _ClosableClient: + closed = False + + async def close(self): + self.closed = True + + storage = object.__new__(storage_cls) + backend = object.__new__(_BlobStorageBackend) + backend._container_client = _ClosableClient() + backend._blob_service_client = _ClosableClient() + storage._backend = backend + + await storage.close() + + assert storage._backend._container_client.closed + assert storage._backend._blob_service_client.closed + + +@pytest.mark.asyncio +async def test_v2_blob_create_only_gives_conflict_precedence_for_existing_item(): + client = _RecordingBlobClient() + client.upload_error = _StatusError(412) + storage = _recording_v2_blob_storage(client) + + result = await storage.write( + {"key": MockStoreItem()}, + StorageWriteOptions( + mode=StorageWriteMode.CREATE_ONLY, expected_version="expected" + ), + ) + + assert result["key"].status == StorageOperationStatus.CONFLICT + + client.upload_error = _StatusError(409) + stale = await storage.write( + {"key": MockStoreItem()}, + StorageWriteOptions( + mode=StorageWriteMode.CREATE_ONLY, expected_version="stale" + ), + ) + matching = await storage.write( + {"key": MockStoreItem()}, + StorageWriteOptions( + mode=StorageWriteMode.CREATE_ONLY, expected_version="current" + ), + ) + + assert stale["key"].status == StorageOperationStatus.CONFLICT + assert matching["key"].status == StorageOperationStatus.CONFLICT + + +@pytest.mark.asyncio +async def test_v2_blob_delete_is_unconditional_without_expected_version(): + client = _RecordingBlobClient() + storage = _recording_v2_blob_storage(client) + + deleted = await storage.delete(["key"]) + client.delete_error = _StatusError(404) + conditional = await storage.delete( + ["key"], StorageDeleteOptions(expected_version="expected") + ) + + assert deleted["key"].status == StorageOperationStatus.SUCCEEDED + assert client.delete_calls[0] == {} + assert client.delete_calls[1]["etag"] == "expected" + assert conditional["key"].status == StorageOperationStatus.CONDITION_NOT_MET + assert client.property_calls == 0 + + async def reset_container(container_client: ContainerClient): blobs = container_client.list_blobs(timeout=5) @@ -80,8 +360,7 @@ async def blob_storage_instance(existing=False): yield storage, container_client - await storage._container_client.close() - await storage._blob_service_client.close() + await storage.close() await container_client.close() await blob_service_client.close() diff --git a/tests/storage_cosmos/test_cosmos_db_storage.py b/tests/storage_cosmos/test_cosmos_db_storage.py index d71da51b2..65820cb3d 100644 --- a/tests/storage_cosmos/test_cosmos_db_storage.py +++ b/tests/storage_cosmos/test_cosmos_db_storage.py @@ -1,6 +1,7 @@ # Copyright (c) Microsoft Corporation. All rights reserved. # Licensed under the MIT License. +import asyncio import os import gc from contextlib import asynccontextmanager @@ -12,10 +13,22 @@ from azure.cosmos import documents from azure.cosmos.aio import CosmosClient from azure.cosmos.exceptions import CosmosResourceNotFoundError +from azure.core import MatchConditions from azure.identity.aio import DefaultAzureCredential -from microsoft_agents.storage.cosmos import CosmosDBStorage, CosmosDBStorageConfig +from microsoft_agents.storage.cosmos import ( + CosmosDBStorage, + CosmosDBStorageConfig, + CosmosDBStorageV2, +) +from microsoft_agents.storage.cosmos.cosmos_db_storage import _CosmosStorageBackend from microsoft_agents.storage.cosmos.key_ops import sanitize_key +from microsoft_agents.hosting.core.storage import ( + StorageDeleteOptions, + StorageOperationStatus, + StorageWriteOptions, + StorageWriteMode, +) from tests._common.storage.utils import ( QuickCRUDStorageTests, @@ -50,6 +63,302 @@ def config(): return create_config(compat_mode=False) +def test_public_cosmos_adapters_have_separate_implementations(): + assert not issubclass(CosmosDBStorage, CosmosDBStorageV2) + assert not issubclass(CosmosDBStorageV2, CosmosDBStorage) + + +@pytest.mark.asyncio +async def test_v1_rejects_v2_options_instead_of_ignoring_them(): + storage = object.__new__(CosmosDBStorage) + with pytest.raises(TypeError, match="positional argument"): + await storage.write({"key": MockStoreItem()}, StorageWriteOptions()) + with pytest.raises(TypeError, match="positional argument"): + await storage.delete(["key"], StorageDeleteOptions()) + + +class _ConcurrentCallBarrier: + def __init__(self, expected: int): + self._expected = expected + self._active = 0 + self.max_active = 0 + self._ready = asyncio.Event() + + async def wait(self): + self._active += 1 + self.max_active = max(self.max_active, self._active) + if self._active == self._expected: + self._ready.set() + await asyncio.wait_for(self._ready.wait(), timeout=1) + self._active -= 1 + + +class _ConcurrentCosmosContainer: + def __init__(self, barrier: _ConcurrentCallBarrier): + self._barrier = barrier + + async def read_item(self, key, _partition_key): + await self._barrier.wait() + return {"id": key, "document": {"value": key}, "_etag": "v1"} + + async def upsert_item(self, **_kwargs): + await self._barrier.wait() + return {"_etag": "v2"} + + async def delete_item(self, *_args, **_kwargs): + await self._barrier.wait() + return None + + +def _create_v2_cosmos_storage(barrier: _ConcurrentCallBarrier): + storage = object.__new__(CosmosDBStorageV2) + backend = object.__new__(_CosmosStorageBackend) + backend._container = _ConcurrentCosmosContainer(barrier) + backend._sanitize = lambda key: key + backend._get_partition_key = lambda key: key + storage._backend = backend + return storage + + +@pytest.mark.asyncio +async def test_v2_batches_run_independent_cosmos_operations_concurrently(): + barrier = _ConcurrentCallBarrier(2) + storage = _create_v2_cosmos_storage(barrier) + read = await storage.read(["one", "two"], target_cls=MockStoreItem) + assert all( + result.status == StorageOperationStatus.SUCCEEDED for result in read.values() + ) + assert barrier.max_active == 2 + + barrier = _ConcurrentCallBarrier(2) + storage = _create_v2_cosmos_storage(barrier) + write = await storage.write({"one": MockStoreItem(), "two": MockStoreItem()}) + assert all( + result.status == StorageOperationStatus.SUCCEEDED for result in write.values() + ) + assert barrier.max_active == 2 + + barrier = _ConcurrentCallBarrier(2) + storage = _create_v2_cosmos_storage(barrier) + delete = await storage.delete(["one", "two"]) + assert all( + result.status == StorageOperationStatus.SUCCEEDED for result in delete.values() + ) + assert barrier.max_active == 2 + + +class _StatusError(Exception): + def __init__(self, status_code: int): + self.status_code = status_code + + +class _RecordingCosmosContainer: + def __init__(self): + self.read_calls = [] + self.create_calls = [] + self.upsert_calls = [] + self.replace_calls = [] + self.delete_calls = [] + self.replace_error = None + self.delete_error = None + + async def read_item(self, *args, **kwargs): + self.read_calls.append((args, kwargs)) + return {"_etag": "current", "document": {"value": "key"}} + + async def create_item(self, **kwargs): + self.create_calls.append(kwargs) + return {"_etag": "created"} + + async def upsert_item(self, **kwargs): + self.upsert_calls.append(kwargs) + return {"_etag": "upserted"} + + async def replace_item(self, **kwargs): + self.replace_calls.append(kwargs) + if self.replace_error: + raise self.replace_error + return {"_etag": "replaced"} + + async def delete_item(self, *args, **kwargs): + self.delete_calls.append((args, kwargs)) + if self.delete_error: + raise self.delete_error + + +class _CurrentCosmosReplaceContainer: + """Model the current azure-cosmos replace_item call signature.""" + + def __init__(self): + self.replace_calls = [] + + async def replace_item(self, *, item, body, etag=None, match_condition=None): + self.replace_calls.append( + { + "item": item, + "body": body, + "etag": etag, + "match_condition": match_condition, + } + ) + return {"_etag": "replaced"} + + +def _recording_v2_cosmos_storage(container): + storage = object.__new__(CosmosDBStorageV2) + backend = object.__new__(_CosmosStorageBackend) + backend._container = container + backend._sanitize = lambda key: key + backend._get_partition_key = lambda key: f"partition:{key}" + storage._backend = backend + return storage + + +def _recording_v1_cosmos_storage(container): + storage = object.__new__(CosmosDBStorage) + backend = object.__new__(_CosmosStorageBackend) + backend._container = container + backend._sanitize = lambda key: key + backend._get_partition_key = lambda key: f"partition:{key}" + storage._backend = backend + return storage + + +@pytest.mark.asyncio +async def test_v1_cosmos_write_still_rejects_an_empty_key(): + storage = _recording_v1_cosmos_storage(_RecordingCosmosContainer()) + + with pytest.raises(ValueError, match="Key cannot be empty"): + await storage.write({"": MockStoreItem()}) + + +@pytest.mark.asyncio +async def test_v2_cosmos_write_uses_atomic_operation_for_each_mode(): + container = _RecordingCosmosContainer() + storage = _recording_v2_cosmos_storage(container) + + await storage.write({"key": MockStoreItem()}) + await storage.write( + {"key": MockStoreItem()}, + StorageWriteOptions(mode=StorageWriteMode.REPLACE), + ) + await storage.write( + {"key": MockStoreItem()}, StorageWriteOptions(expected_version="expected") + ) + + assert container.read_calls == [] + assert len(container.upsert_calls) == 1 + assert "partition_key" not in container.replace_calls[0] + assert "etag" not in container.replace_calls[0] + assert container.replace_calls[1]["etag"] == "expected" + + +@pytest.mark.asyncio +async def test_v2_cosmos_conditional_write_uses_current_replace_signature(): + container = _CurrentCosmosReplaceContainer() + storage = _recording_v2_cosmos_storage(container) + + result = await storage.write( + {"key": MockStoreItem()}, StorageWriteOptions(expected_version="expected") + ) + + assert result["key"].status == StorageOperationStatus.SUCCEEDED + assert container.replace_calls == [ + { + "item": "key", + "body": {"id": "key", "realId": "key", "document": {}}, + "etag": "expected", + "match_condition": MatchConditions.IfNotModified, + } + ] + + +@pytest.mark.asyncio +async def test_v2_cosmos_read_forwards_provider_options(): + container = _RecordingCosmosContainer() + storage = _recording_v2_cosmos_storage(container) + + await storage.read(["key"], target_cls=MockStoreItem, consistency_level="Session") + + assert container.read_calls == [ + (("key", "partition:key"), {"consistency_level": "Session"}) + ] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("storage_cls", [CosmosDBStorage, CosmosDBStorageV2]) +async def test_cosmos_adapters_expose_public_close(storage_cls): + class _ClosableClient: + closed = False + + async def close(self): + self.closed = True + + storage = object.__new__(storage_cls) + backend = object.__new__(_CosmosStorageBackend) + backend._client = _ClosableClient() + storage._backend = backend + + await storage.close() + + assert storage._backend._client.closed + + +@pytest.mark.asyncio +async def test_v2_cosmos_conditional_missing_does_not_recreate_item(): + container = _RecordingCosmosContainer() + container.replace_error = _StatusError(404) + storage = _recording_v2_cosmos_storage(container) + + result = await storage.write( + {"key": MockStoreItem()}, StorageWriteOptions(expected_version="expected") + ) + + assert result["key"].status == StorageOperationStatus.CONDITION_NOT_MET + assert container.upsert_calls == [] + + +@pytest.mark.asyncio +async def test_v2_cosmos_create_only_honors_expected_version_without_writing(): + container = _RecordingCosmosContainer() + storage = _recording_v2_cosmos_storage(container) + + matching = await storage.write( + {"key": MockStoreItem()}, + StorageWriteOptions( + mode=StorageWriteMode.CREATE_ONLY, expected_version="current" + ), + ) + stale = await storage.write( + {"key": MockStoreItem()}, + StorageWriteOptions( + mode=StorageWriteMode.CREATE_ONLY, expected_version="stale" + ), + ) + + assert matching["key"].status == StorageOperationStatus.CONFLICT + assert stale["key"].status == StorageOperationStatus.CONFLICT + assert container.create_calls == [] + + +@pytest.mark.asyncio +async def test_v2_cosmos_delete_only_uses_condition_when_requested(): + container = _RecordingCosmosContainer() + storage = _recording_v2_cosmos_storage(container) + + deleted = await storage.delete(["key"]) + container.delete_error = _StatusError(404) + conditional = await storage.delete( + ["key"], StorageDeleteOptions(expected_version="expected") + ) + + assert deleted["key"].status == StorageOperationStatus.SUCCEEDED + assert container.delete_calls[0][1] == {} + assert container.delete_calls[1][1]["etag"] == "expected" + assert conditional["key"].status == StorageOperationStatus.CONDITION_NOT_MET + assert container.read_calls == [] + + async def reset_container(container_client): try: @@ -115,7 +424,7 @@ async def cosmos_db_storage_instance(compat_mode=False, existing=False): ) as container_client: storage = CosmosDBStorage(config) yield storage, container_client - await storage._close() + await storage.close() @pytest.mark.asyncio @@ -177,16 +486,16 @@ async def test_cosmos_db_storage_flow_existing_container_and_persistence( gc.collect() storage = CosmosDBStorage(config) - escaped_key = storage._sanitize("?test") + escaped_key = storage._backend._sanitize("?test") with pytest.raises(CosmosResourceNotFoundError): await container_client.read_item( - escaped_key, storage._get_partition_key(escaped_key) + escaped_key, storage._backend._get_partition_key(escaped_key) ) - escaped_key = storage._sanitize("1230") + escaped_key = storage._backend._sanitize("1230") item = ( await container_client.read_item( - escaped_key, storage._get_partition_key(escaped_key) + escaped_key, storage._backend._get_partition_key(escaped_key) ) ).get("document") assert MockStoreItemB.from_json_to_store_item(item) == initial_data["1230"] @@ -339,5 +648,5 @@ async def test_raises_error_different_partition_key(self, compat_mode): ) storage = CosmosDBStorage(config) await storage.initialize() - await storage._close() + await storage.close() await cosmos_client.close()