diff --git a/python/packages/azure-cosmos-memory/README.md b/python/packages/azure-cosmos-memory/README.md index db466b5ed33..e5516ae8f0c 100644 --- a/python/packages/azure-cosmos-memory/README.md +++ b/python/packages/azure-cosmos-memory/README.md @@ -311,7 +311,7 @@ memory_provider = CosmosMemoryContextProvider( credential=DefaultAzureCredential(), # Azure credential # Memory retrieval options - top_k=5, # Number of memories to retrieve + top_k=5, # Shared fact/episode result limit min_confidence=0.7, # Minimum confidence score (0.0-1.0) memory_types=["fact", "procedural"], # Types to retrieve @@ -326,16 +326,19 @@ memory_provider = CosmosMemoryContextProvider( ### Memory Types -The provider retrieves four types of memories: +The provider retrieves three types of memories: | Type | Description | Default TTL | |------|-------------|-------------| | **fact** | Declarative knowledge ("user prefers dark mode") | None | | **procedural** | Behavioral rules ("always confirm before deleting") | None | | **episodic** | Past experiences with context and outcomes | 90 days | -| **unclassified** | Memories that couldn't be confidently classified | None | Each memory has a confidence score (0.0-1.0). Use `min_confidence` to filter low-quality extractions. +With Agent Memory Toolkit 0.3.0b2 or later, facts and selected episodes are ranked together under +the single `top_k` limit. Selected procedures are compiled separately for the current task and do +not consume that ranking budget. Older supported Toolkit versions retain their generic retrieval +behavior. ### Processing Pipeline diff --git a/python/packages/azure-cosmos-memory/agent_framework_azure_cosmos_memory/_context_provider.py b/python/packages/azure-cosmos-memory/agent_framework_azure_cosmos_memory/_context_provider.py index 9fa93db0598..4c2412e97ef 100644 --- a/python/packages/azure-cosmos-memory/agent_framework_azure_cosmos_memory/_context_provider.py +++ b/python/packages/azure-cosmos-memory/agent_framework_azure_cosmos_memory/_context_provider.py @@ -9,6 +9,7 @@ from __future__ import annotations import asyncio +import inspect import logging import sys from collections.abc import Mapping, Sequence @@ -238,6 +239,12 @@ def __init__( self.memory_client = memory_client self._cosmos_endpoint = cosmos_endpoint self._foundry_endpoint = foundry_endpoint + procedural_builder = getattr(self.memory_client, "build_procedural_context", None) + self._supports_toolkit_03_retrieval = ( + "include_episodes" in inspect.signature(self.memory_client.search_cosmos).parameters + and procedural_builder is not None + and "task" in inspect.signature(procedural_builder).parameters + ) def _resolve_user_id(self, state: dict[str, Any], session: AgentSession) -> str: """Resolve the user id for memory scoping. @@ -358,26 +365,51 @@ async def before_run( # Get user_id from state or session (warns once if no stable user_id was provided) user_id = self._resolve_user_id(state, session) - # Memory search and user-summary retrieval are independent: the user summary - # provides baseline context even when no memories match the query, so a failure - # in one must not suppress the other. They get separate error handling. - try: - results = await self.memory_client.search_cosmos( - search_terms=query_text, - user_id=user_id, - top_k=self.top_k, - memory_types=[str(t) for t in self.memory_types], - min_confidence=self.min_confidence, - ) - - if results: - # Format and inject memories - memory_content = self._format_memories(results) - context.extend_messages( - self.source_id, [Message(role="user", contents=[f"{self.context_prompt}\n{memory_content}"])] + memory_sections: list[str] = [] + + # Toolkit 0.3.0b2 compiles task-aware procedures separately. Older supported + # clients keep their existing generic procedural search behavior. + search_memory_types = [ + str(memory_type) + for memory_type in self.memory_types + if not self._supports_toolkit_03_retrieval or memory_type != "procedural" + ] + if search_memory_types: + try: + search_kwargs: dict[str, Any] = { + "search_terms": query_text, + "user_id": user_id, + "top_k": self.top_k, + "memory_types": search_memory_types, + "min_confidence": self.min_confidence, + } + if self._supports_toolkit_03_retrieval: + search_kwargs["include_episodes"] = "episodic" in self.memory_types + + results = await self.memory_client.search_cosmos(**search_kwargs) + if results: + memory_sections.append(self._format_memories(results)) + except Exception as e: + logger.warning("Failed to retrieve memories: %s", e, exc_info=True) + + if self._supports_toolkit_03_retrieval and "procedural" in self.memory_types: + try: + procedural_builder = cast("Any", self.memory_client.build_procedural_context) + procedural_context = await procedural_builder( + user_id=user_id, + task=query_text, ) - except Exception as e: - logger.warning("Failed to retrieve memories: %s", e, exc_info=True) + if procedural_context and procedural_context.strip(): + memory_sections.append(procedural_context.strip()) + except Exception as e: + logger.warning("Failed to retrieve procedural context: %s", e, exc_info=True) + + if memory_sections: + memory_content = "\n".join(memory_sections) + context.extend_messages( + self.source_id, + [Message(role="user", contents=[f"{self.context_prompt}\n{memory_content}"])], + ) # Retrieve and inject user summary as untrusted context. # This is INDEPENDENT of search results - even if no memories match the query, @@ -437,7 +469,9 @@ async def after_run( # TODO(atty57): The toolkit renamed add_cosmos -> upsert_memory (same kwargs); accept either # until the declared azure-cosmos-agent-memory floor is past the rename, then inline it. - write_turn = getattr(self.memory_client, "upsert_memory", None) or self.memory_client.add_cosmos + write_turn = getattr(self.memory_client, "upsert_memory", None) + if write_turn is None: + write_turn = getattr(self.memory_client, "add_cosmos") # ruff: ignore[get-attr-with-constant] try: # Store input messages (skip empty/whitespace-only content to avoid junk turns) diff --git a/python/packages/azure-cosmos-memory/pyproject.toml b/python/packages/azure-cosmos-memory/pyproject.toml index 07e99fe3eea..2e5f8182ae8 100644 --- a/python/packages/azure-cosmos-memory/pyproject.toml +++ b/python/packages/azure-cosmos-memory/pyproject.toml @@ -23,7 +23,7 @@ classifiers = [ ] dependencies = [ "agent-framework-core>=1.13.0,<2", - "azure-cosmos-agent-memory>=0.2.0b3", + "azure-cosmos-agent-memory>=0.2.0b3,<0.4", # azure-cosmos-agent-memory depends transitively on a prompty pre-release # (prompty>=2.0.0a9, which has no stable 2.x release yet). Declaring it here as a # direct dependency makes the pre-release "explicit" so the workspace's diff --git a/python/packages/azure-cosmos-memory/tests/test_context_provider.py b/python/packages/azure-cosmos-memory/tests/test_context_provider.py index ebee3232348..3ed16fd0a09 100644 --- a/python/packages/azure-cosmos-memory/tests/test_context_provider.py +++ b/python/packages/azure-cosmos-memory/tests/test_context_provider.py @@ -13,7 +13,7 @@ pytest.importorskip("azure.cosmos.agent_memory") from typing import Any -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, create_autospec, patch from agent_framework import AgentResponse, Message from agent_framework._sessions import AgentSession, SessionContext @@ -30,6 +30,28 @@ _STUB_AGENT: Any = None +class Toolkit03MemoryClient: + """Signature-bearing test double for the Toolkit 0.3.0b2 retrieval contract.""" + + async def search_cosmos( + self, + search_terms: str, + *, + user_id: str, + top_k: int, + memory_types: list[str], + min_confidence: float, + include_episodes: bool = False, + ) -> list[dict[str, Any]]: + return [] + + async def build_procedural_context(self, user_id: str, task: str | None = None) -> str: + return "" + + async def get_user_summary(self, user_id: str) -> dict[str, Any] | None: + return None + + async def test_before_run_marks_cosmos_memory_used_before_empty_return() -> None: provider = object.__new__(CosmosMemoryContextProvider) context = MagicMock(spec=SessionContext) @@ -59,6 +81,16 @@ def mock_memory_client() -> AsyncMock: return mock_client +@pytest.fixture +def toolkit_03_memory_client() -> Any: + """Create a mock client that exposes the Toolkit 0.3.0b2 retrieval signatures.""" + mock_client = create_autospec(Toolkit03MemoryClient, instance=True) + mock_client.search_cosmos.return_value = [] + mock_client.build_procedural_context.return_value = "" + mock_client.get_user_summary.return_value = None + return mock_client + + # -- Initialization tests ------------------------------------------------------ @@ -270,6 +302,7 @@ async def test_retrieves_and_injects_memories(self, mock_memory_client: AsyncMoc assert call_kwargs["top_k"] == 5 assert call_kwargs["memory_types"] == ["fact", "procedural"] assert call_kwargs["min_confidence"] == 0.7 + assert "include_episodes" not in call_kwargs # Verify memories added to context assert "cosmos_memory" in ctx.context_messages @@ -280,6 +313,153 @@ async def test_retrieves_and_injects_memories(self, mock_memory_client: AsyncMoc assert "0.95" in added[0].text # type: ignore assert "0.85" in added[0].text # type: ignore + async def test_older_toolkit_preserves_generic_mixed_search(self, mock_memory_client: AsyncMock) -> None: + """Older Toolkit clients receive all selected types through generic search.""" + provider = CosmosMemoryContextProvider( + memory_client=mock_memory_client, + memory_types=["fact", "episodic", "procedural"], + ) + session = AgentSession(session_id="test-session") + ctx = SessionContext(input_messages=[Message(role="user", contents=["test"])], session_id="s1") + + await provider.before_run( + agent=_STUB_AGENT, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {}) + ) + + call_kwargs = mock_memory_client.search_cosmos.call_args.kwargs + assert call_kwargs["memory_types"] == ["fact", "episodic", "procedural"] + assert "include_episodes" not in call_kwargs + mock_memory_client.build_procedural_context.assert_not_awaited() + + async def test_toolkit_03_retrieves_mixed_memory_types(self, toolkit_03_memory_client: Any) -> None: + """Toolkit 0.3 uses one fact/episode search plus task-aware procedural projection.""" + toolkit_03_memory_client.search_cosmos.return_value = [ + {"content": "User prefers Python", "memory_type": "fact", "confidence": 0.95}, + {"content": "User completed ML course", "memory_type": "episodic", "confidence": 0.85}, + ] + toolkit_03_memory_client.build_procedural_context.return_value = "Always provide runnable examples." + + provider = CosmosMemoryContextProvider( + memory_client=toolkit_03_memory_client, + top_k=7, + memory_types=["fact", "episodic", "procedural"], + ) + session = AgentSession(session_id="test-session") + ctx = SessionContext( + input_messages=[Message(role="user", contents=["Help me write Python"])], + session_id="s1", + ) + + await provider.before_run( + agent=_STUB_AGENT, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {}) + ) + + toolkit_03_memory_client.search_cosmos.assert_awaited_once_with( + search_terms="Help me write Python", + user_id="test-session", + top_k=7, + memory_types=["fact", "episodic"], + min_confidence=0.7, + include_episodes=True, + ) + toolkit_03_memory_client.build_procedural_context.assert_awaited_once_with( + user_id="test-session", + task="Help me write Python", + ) + added = ctx.context_messages["cosmos_memory"] + assert len(added) == 1 + assert "User prefers Python" in added[0].text # type: ignore + assert "User completed ML course" in added[0].text # type: ignore + assert "Always provide runnable examples." in added[0].text # type: ignore + + async def test_toolkit_03_fact_only_disables_episodes(self, toolkit_03_memory_client: Any) -> None: + """Fact-only retrieval explicitly disables episodes and skips procedures.""" + provider = CosmosMemoryContextProvider(memory_client=toolkit_03_memory_client, memory_types=["fact"]) + session = AgentSession(session_id="test-session") + ctx = SessionContext(input_messages=[Message(role="user", contents=["test"])], session_id="s1") + + await provider.before_run( + agent=_STUB_AGENT, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {}) + ) + + assert toolkit_03_memory_client.search_cosmos.call_args.kwargs["memory_types"] == ["fact"] + assert toolkit_03_memory_client.search_cosmos.call_args.kwargs["include_episodes"] is False + toolkit_03_memory_client.build_procedural_context.assert_not_awaited() + + async def test_toolkit_03_procedural_only_skips_generic_search(self, toolkit_03_memory_client: Any) -> None: + """Procedural-only retrieval compiles the task context without running generic search.""" + toolkit_03_memory_client.build_procedural_context.return_value = "Use the deployment runbook." + provider = CosmosMemoryContextProvider(memory_client=toolkit_03_memory_client, memory_types=["procedural"]) + session = AgentSession(session_id="test-session") + ctx = SessionContext(input_messages=[Message(role="user", contents=["Deploy the app"])], session_id="s1") + + await provider.before_run( + agent=_STUB_AGENT, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {}) + ) + + toolkit_03_memory_client.search_cosmos.assert_not_awaited() + toolkit_03_memory_client.build_procedural_context.assert_awaited_once_with( + user_id="test-session", + task="Deploy the app", + ) + assert "Use the deployment runbook." in ctx.context_messages["cosmos_memory"][0].text # type: ignore + + async def test_toolkit_03_empty_procedural_context_is_not_injected(self, toolkit_03_memory_client: Any) -> None: + """An empty compiled procedure does not add a context message.""" + provider = CosmosMemoryContextProvider(memory_client=toolkit_03_memory_client, memory_types=["procedural"]) + session = AgentSession(session_id="test-session") + ctx = SessionContext(input_messages=[Message(role="user", contents=["test"])], session_id="s1") + + await provider.before_run( + agent=_STUB_AGENT, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {}) + ) + + assert "cosmos_memory" not in ctx.context_messages + + async def test_toolkit_03_search_failure_does_not_block_other_context(self, toolkit_03_memory_client: Any) -> None: + """Generic search failure does not suppress procedures or the user summary.""" + toolkit_03_memory_client.search_cosmos.side_effect = Exception("search boom") + toolkit_03_memory_client.build_procedural_context.return_value = "Use the recovery procedure." + toolkit_03_memory_client.get_user_summary.return_value = {"content": "Prefers concise answers"} + provider = CosmosMemoryContextProvider( + memory_client=toolkit_03_memory_client, + memory_types=["fact", "procedural"], + ) + session = AgentSession(session_id="test-session") + ctx = SessionContext(input_messages=[Message(role="user", contents=["test"])], session_id="s1") + + await provider.before_run( + agent=_STUB_AGENT, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {}) + ) + + added = ctx.context_messages["cosmos_memory"] + assert any("Use the recovery procedure." in message.text for message in added) # type: ignore + assert any("Prefers concise answers" in message.text for message in added) # type: ignore + + async def test_toolkit_03_procedural_failure_does_not_block_other_context( + self, toolkit_03_memory_client: Any + ) -> None: + """Procedural projection failure does not suppress facts or the user summary.""" + toolkit_03_memory_client.search_cosmos.return_value = [ + {"content": "User likes hiking", "memory_type": "fact", "confidence": 0.9} + ] + toolkit_03_memory_client.build_procedural_context.side_effect = Exception("procedure boom") + toolkit_03_memory_client.get_user_summary.return_value = {"content": "Prefers concise answers"} + provider = CosmosMemoryContextProvider( + memory_client=toolkit_03_memory_client, + memory_types=["fact", "procedural"], + ) + session = AgentSession(session_id="test-session") + ctx = SessionContext(input_messages=[Message(role="user", contents=["test"])], session_id="s1") + + await provider.before_run( + agent=_STUB_AGENT, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {}) + ) + + added = ctx.context_messages["cosmos_memory"] + assert any("User likes hiking" in message.text for message in added) # type: ignore + assert any("Prefers concise answers" in message.text for message in added) # type: ignore + async def test_user_summary_injected_as_untrusted_message(self, mock_memory_client: AsyncMock) -> None: """User summary is injected as an untrusted context message, not as agent instructions.""" mock_memory_client.search_cosmos.return_value = [] diff --git a/python/packages/azure-cosmos-memory/tests/test_emulator.py b/python/packages/azure-cosmos-memory/tests/test_emulator.py index caad815db4d..771a82ad802 100644 --- a/python/packages/azure-cosmos-memory/tests/test_emulator.py +++ b/python/packages/azure-cosmos-memory/tests/test_emulator.py @@ -183,7 +183,9 @@ async def test_before_run_retrieves_seeded_fact(self, emulator_provider: CosmosM # embeddings client). This lands in the memories container under the quantizedFlat # vector index, without needing LLM extraction. assert provider.memory_client is not None - seed = getattr(provider.memory_client, "upsert_memory", None) or provider.memory_client.add_cosmos + seed = getattr(provider.memory_client, "upsert_memory", None) + if seed is None: + seed = getattr(provider.memory_client, "add_cosmos") # noqa: B009 await seed( user_id=user_id, thread_id=thread_id, diff --git a/python/uv.lock b/python/uv.lock index 2ae5de419bd..047cbf03701 100644 --- a/python/uv.lock +++ b/python/uv.lock @@ -333,7 +333,7 @@ dev = [ [package.metadata] requires-dist = [ { name = "agent-framework-core", editable = "packages/core" }, - { name = "azure-cosmos-agent-memory", specifier = ">=0.2.0b3" }, + { name = "azure-cosmos-agent-memory", specifier = ">=0.2.0b3,<0.4" }, { name = "prompty", specifier = ">=2.0.0a9" }, ] @@ -1469,7 +1469,7 @@ wheels = [ [[package]] name = "azure-cosmos-agent-memory" -version = "0.3.0b1" +version = "0.3.0b2" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "aiohttp", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, @@ -1482,9 +1482,9 @@ dependencies = [ { name = "tiktoken", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "typing-extensions", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/71/85/f66f7e5989125428eba72917224fbdb4fb22671de3a39ec8d24446f93b2b/azure_cosmos_agent_memory-0.3.0b1.tar.gz", hash = "sha256:4dd017bd4c69895f3e053511c549ee8b215e5eb98f684152cf852ebbb69199b3", size = 169559, upload-time = "2026-07-24T21:37:19.433Z" } +sdist = { url = "https://files.pythonhosted.org/packages/59/98/1bba425448cfe620e27818c5cc770a5c5ca69d76b8e5989c0ff69daef01d/azure_cosmos_agent_memory-0.3.0b2.tar.gz", hash = "sha256:5f3d4590b74fd15da2973d09217d9eefc2a3787c03d283210b875b05f06e68e3", size = 189529, upload-time = "2026-08-17T23:09:33.819Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/f6/f5/18be23d17aee288562adb9ef66826829524077e7024e2b5535615c82f621/azure_cosmos_agent_memory-0.3.0b1-py3-none-any.whl", hash = "sha256:7d36e28cba2adfc1567a94d8b057c018e7195bf3d5cec14e92747cb23b27465b", size = 187471, upload-time = "2026-07-24T21:37:18.268Z" }, + { url = "https://files.pythonhosted.org/packages/e0/b4/d18966a12e3ba2c99dca4a3888bd0b35c1d1fd2c341690bd926e80550018/azure_cosmos_agent_memory-0.3.0b2-py3-none-any.whl", hash = "sha256:2af5d6e30b485052f52abc8980be980632e2ace5750eb6329e2af8ca66c5aaa9", size = 216786, upload-time = "2026-08-17T23:09:32.666Z" }, ] [[package]]