Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 6 additions & 3 deletions python/packages/azure-cosmos-memory/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
from __future__ import annotations

import asyncio
import inspect
import logging
import sys
from collections.abc import Mapping, Sequence
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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"
Comment thread
eavanvalkenburg marked this conversation as resolved.
]
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(
Comment thread
eavanvalkenburg marked this conversation as resolved.
user_id=user_id,
task=query_text,
Comment thread
eavanvalkenburg marked this conversation as resolved.
)
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())
Comment thread
eavanvalkenburg marked this conversation as resolved.
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,
Expand Down Expand Up @@ -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)
Expand Down
2 changes: 1 addition & 1 deletion python/packages/azure-cosmos-memory/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
182 changes: 181 additions & 1 deletion python/packages/azure-cosmos-memory/tests/test_context_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)
Expand Down Expand Up @@ -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 ------------------------------------------------------


Expand Down Expand Up @@ -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
Expand All @@ -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 = []
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
Loading
Loading