Python: [BREAKING] PR2 — Wire context provider pipeline, remove old types, update all consumers (#3850)

* PR2: Wire context provider pipeline and update all internal consumers

- Replace AgentThread with AgentSession across all packages
- Replace ContextProvider with BaseContextProvider across all packages
- Replace context_provider param with context_providers (Sequence)
- Replace thread= with session= in run() signatures
- Replace get_new_thread() with create_session()
- Add get_session(service_session_id) to agent interface
- DurableAgentThread -> DurableAgentSession
- Remove _notify_thread_of_new_messages from WorkflowAgent
- Wire before_run/after_run context provider pipeline in RawAgent
- Auto-inject InMemoryHistoryProvider when no providers configured

* fix: update all tests for context provider pipeline, fix lazy-loaders, remove old test files

* refactor: update all sample files for context provider pipeline (AgentThread→AgentSession, ContextProvider→BaseContextProvider)

* fix: update remaining ag-ui references (client docstring, getting_started sample)

* fix: make get_session service_session_id keyword-only to avoid confusion with session_id

* refactor: rename _RunContext.thread_messages to session_messages

* refactor: remove _threads.py, _memory.py, and old provider files; migrate devui to use plain message lists

* rename: remove _new_ prefix from test files

* refactor: rewrite SlidingWindowChatMessageStore as SlidingWindowHistoryProvider(InMemoryHistoryProvider)

* fix: read full history from session state directly instead of reaching into provider internals

* fix: update stale .pyi stubs, sample imports, and README references for new provider types

* fix: remove stale message_store, _notify_thread_of_new_messages, and session_id.key references in samples

* refactor: merge context_providers and sessions sample folders into sessions, remove aggregate_context_provider

* refactor: UserInfoMemory stores state in session.state instead of instance attributes

* feat: add Pydantic BaseModel support to session state serialization

Pydantic models stored in session.state are now automatically serialized
via model_dump() and restored via model_validate() during to_dict()/from_dict()
round-trips. Models are auto-registered on first serialization; use
register_state_type() for cold-start deserialization.

Also export register_state_type as a public API.

* fix mem0

* Update sample README links and descriptions for session terminology

- Replace 'thread' with 'session' in sample descriptions across all READMEs
- Update file links for renamed samples (mem0_sessions, redis_sessions, etc.)
- Fix Threads section → Sessions section in main samples/README.md
- Update tools, middleware, workflows, durabletask, azure_functions READMEs
- Update architecture diagrams in concepts/tools/README.md
- Update migration guides (autogen, semantic-kernel)

* Fix broken Redis README link to renamed sample

* Fix Mem0 OSS client search: pass scoping params as direct kwargs

AsyncMemory (OSS) expects user_id/agent_id/run_id as direct kwargs,
while AsyncMemoryClient (Platform) expects them in a filters dict.
Adds tests for both client types.

Port of fix from #3844 to new Mem0ContextProvider.

* Fix rebase issues: restore missing _conversation_state.py and checkpoint decode logic

- Add back _conversation_state.py (encode/decode_chat_messages) lost in rebase
- Fix on_checkpoint_restore to decode cache/conversation with decode_chat_messages
- Fix on_checkpoint_restore to use decode_checkpoint_value for pending requests
- Add tests/workflow/__init__.py for relative import support
- Fix test_agent_executor checkpoint selection (checkpoints[1] not superstep)

* Add STORES_BY_DEFAULT ClassVar to skip redundant InMemoryHistoryProvider injection

Chat clients that store history server-side by default (OpenAI Responses API,
Azure AI Agent) now declare STORES_BY_DEFAULT = True. The agent checks this
during auto-injection and skips InMemoryHistoryProvider unless the user
explicitly sets store=False.

* Fix broken markdown links in azure_ai and redis READMEs

* Fix getting-started samples to use session API instead of removed thread/ContextProvider API

* updates to workflow as agent

* fix group chat import

* Rename Thread→Session throughout, fix service_session_id propagation, remove stale AGUIThread

- Fix: Propagate conversation_id from ChatResponse back to session.service_session_id
  in both streaming and non-streaming paths in _agents.py
- Rename AgentThreadException → AgentSessionException
- Remove stale AGUIThread from ag_ui lazy-loader
- Rename use_service_thread → use_service_session in ag-ui package
- Rename test functions from *_thread_* to *_session_*
- Rename sample files from *_thread* to *_session*
- Update docstrings and comments: thread → session
- Update _mcp.py kwargs filter: add 'session' alongside 'thread'
- Fix ContinuationToken docstring example: thread=thread → session=session
- Fix _clients.py docstring: 'Agent threads' → 'Agent sessions'

* Fix broken markdown links after thread→session file renames

* fix azure ai test
This commit is contained in:
Eduard van Valkenburg
2026-02-12 22:00:32 +01:00
committed by GitHub
Unverified
parent 0c67dbbce5
commit 1e350ea22f
312 changed files with 6669 additions and 11423 deletions
@@ -1,10 +1,8 @@
# Copyright (c) Microsoft. All rights reserved.
import importlib.metadata
from ._chat_message_store import RedisChatMessageStore
from ._context_provider import _RedisContextProvider
from ._history_provider import _RedisHistoryProvider
from ._provider import RedisProvider
from ._context_provider import RedisContextProvider
from ._history_provider import RedisHistoryProvider
try:
__version__ = importlib.metadata.version(__name__)
@@ -12,9 +10,7 @@ except importlib.metadata.PackageNotFoundError:
__version__ = "0.0.0" # Fallback for development mode
__all__ = [
"RedisChatMessageStore",
"RedisProvider",
"_RedisContextProvider",
"_RedisHistoryProvider",
"RedisContextProvider",
"RedisHistoryProvider",
"__version__",
]
@@ -1,595 +0,0 @@
# Copyright (c) Microsoft. All rights reserved.
from __future__ import annotations
from collections.abc import Sequence
from typing import Any
from uuid import uuid4
import redis.asyncio as redis
from agent_framework import Message
from agent_framework._serialization import SerializationMixin
from redis.credentials import CredentialProvider
class RedisStoreState(SerializationMixin):
"""State model for serializing and deserializing Redis chat message store data."""
def __init__(
self,
thread_id: str,
redis_url: str | None = None,
key_prefix: str = "chat_messages",
max_messages: int | None = None,
) -> None:
"""State model for serializing and deserializing Redis chat message store data."""
self.thread_id = thread_id
self.redis_url = redis_url
self.key_prefix = key_prefix
self.max_messages = max_messages
class RedisChatMessageStore:
"""Redis-backed implementation of ChatMessageStoreProtocol using Redis Lists.
This implementation provides persistent, thread-safe chat message storage using Redis Lists.
Messages are stored as JSON-serialized strings in chronological order, with each conversation
thread isolated by a unique Redis key.
Key Features:
============
- **Persistent Storage**: Messages survive application restarts and crashes
- **Thread Isolation**: Each conversation thread has its own Redis key namespace
- **Auto Message Limits**: Configurable automatic trimming of old messages using LTRIM
- **Performance Optimized**: Uses native Redis operations for efficiency
- **State Serialization**: Full compatibility with Agent Framework thread serialization
- **Initial Message Support**: Pre-load conversations with existing message history
- **Production Ready**: Atomic operations, error handling, connection pooling
Redis Operations:
- RPUSH: Add messages to the end of the list (chronological order)
- LRANGE: Retrieve messages in chronological order
- LTRIM: Maintain message limits by trimming old messages
- DELETE: Clear all messages for a thread
"""
def __init__(
self,
redis_url: str | None = None,
credential_provider: CredentialProvider | None = None,
host: str | None = None,
port: int = 6380,
ssl: bool = True,
username: str | None = None,
thread_id: str | None = None,
key_prefix: str = "chat_messages",
max_messages: int | None = None,
messages: Sequence[Message] | None = None,
) -> None:
"""Initialize the Redis chat message store.
Creates a Redis-backed chat message store for a specific conversation thread.
Supports both traditional URL-based authentication and Azure Managed Redis
with credential provider.
Args:
redis_url: Redis connection URL (e.g., "redis://localhost:6379").
Used for traditional authentication. Mutually exclusive with credential_provider.
credential_provider: Redis credential provider (redis.credentials.CredentialProvider) for
Azure AD authentication. Requires host parameter. Mutually exclusive with redis_url.
host: Redis host name (e.g., "myredis.redis.cache.windows.net").
Required when using credential_provider.
port: Redis port number. Defaults to 6380 (Azure Redis SSL port).
ssl: Enable SSL/TLS connection. Defaults to True.
username: Redis username. Defaults to None.
thread_id: Unique identifier for this conversation thread.
If not provided, a UUID will be auto-generated.
This becomes part of the Redis key: {key_prefix}:{thread_id}
key_prefix: Prefix for Redis keys to namespace different applications.
Defaults to 'chat_messages'. Useful for multi-tenant scenarios.
max_messages: Maximum number of messages to retain in Redis.
When exceeded, oldest messages are automatically trimmed using LTRIM.
None means unlimited storage.
messages: Initial messages to pre-populate the conversation.
These are added to Redis on first access if the Redis key is empty.
Useful for resuming conversations or seeding with context.
Raises:
ValueError: If neither redis_url nor credential_provider is provided.
ValueError: If both redis_url and credential_provider are provided.
ValueError: If credential_provider is used without host parameter.
Examples:
Traditional connection:
store = RedisChatMessageStore(
redis_url="redis://localhost:6379",
thread_id="conversation_123"
)
Azure Managed Redis with credential provider:
from redis.credentials import CredentialProvider
from azure.identity.aio import DefaultAzureCredential
store = RedisChatMessageStore(
credential_provider=CredentialProvider(DefaultAzureCredential()),
host="myredis.redis.cache.windows.net",
thread_id="conversation_123"
)
"""
# Validate connection parameters
if redis_url is None and credential_provider is None:
raise ValueError("Either redis_url or credential_provider must be provided")
if redis_url is not None and credential_provider is not None:
raise ValueError("redis_url and credential_provider are mutually exclusive")
if credential_provider is not None and host is None:
raise ValueError("host is required when using credential_provider")
# Store configuration
self.thread_id = thread_id or f"thread_{uuid4()}"
self.key_prefix = key_prefix
self.max_messages = max_messages
# Initialize Redis client based on authentication method
if credential_provider is not None and host is not None:
# Azure AD authentication with credential provider
self.redis_url = None # Not using URL-based auth
self._redis_client = redis.Redis(
host=host,
port=port,
ssl=ssl,
username=username,
credential_provider=credential_provider,
decode_responses=True,
)
else:
# Traditional URL-based authentication
self.redis_url = redis_url
self._redis_client = redis.from_url(redis_url, decode_responses=True) # type: ignore[no-untyped-call]
# Handle initial messages (will be moved to Redis on first access)
self._initial_messages = list(messages) if messages else []
self._initial_messages_added = False
@property
def redis_key(self) -> str:
"""Get the Redis key for this thread's messages.
The key format is: {key_prefix}:{thread_id}
Returns:
Redis key string used for storing this thread's messages.
Example:
For key_prefix="chat_messages" and thread_id="user_123_session_456":
Returns "chat_messages:user_123_session_456"
"""
return f"{self.key_prefix}:{self.thread_id}"
async def _ensure_initial_messages_added(self) -> None:
"""Ensure initial messages are added to Redis if not already present.
This method is called before any Redis operations to guarantee that
initial messages provided during construction are persisted to Redis.
"""
if not self._initial_messages or self._initial_messages_added:
return
# Check if Redis key already has messages (prevents duplicate additions)
existing_count = await self._redis_client.llen(self.redis_key) # type: ignore[misc] # type: ignore[misc]
if existing_count == 0:
# Add initial messages using atomic pipeline operation
await self._add_redis_messages(self._initial_messages)
# Mark as completed and free memory
self._initial_messages_added = True
self._initial_messages.clear()
async def _add_redis_messages(self, messages: Sequence[Message]) -> None:
"""Add multiple messages to Redis using atomic pipeline operation.
This internal method efficiently adds multiple messages to the Redis list
using a single atomic transaction to ensure consistency.
Args:
messages: Sequence of Message objects to add to Redis.
"""
if not messages:
return
# Pre-serialize all messages for efficient pipeline operation
serialized_messages = [self._serialize_message(message) for message in messages]
# Use Redis pipeline for atomic batch operation
async with self._redis_client.pipeline(transaction=True) as pipe:
for serialized_message in serialized_messages:
await pipe.rpush(self.redis_key, serialized_message) # type: ignore[misc]
await pipe.execute()
async def add_messages(self, messages: Sequence[Message]) -> None:
"""Add messages to the Redis store (ChatMessageStoreProtocol protocol method).
This method implements the required ChatMessageStoreProtocol protocol for adding messages.
Messages are appended to the Redis list in chronological order, with automatic
trimming if message limits are configured.
Args:
messages: Sequence of Message objects to add to the store.
Can be empty (no-op) or contain multiple messages.
Thread Safety:
- Atomic pipeline ensures all messages are added together
- LTRIM operation is atomic for consistent message limits
Example:
.. code-block:: python
messages = [Message(role="user", text="Hello"), Message(role="assistant", text="Hi there!")]
await store.add_messages(messages)
"""
if not messages:
return
# Ensure any initial messages are persisted first
await self._ensure_initial_messages_added()
# Add new messages using atomic pipeline operation
await self._add_redis_messages(messages)
# Apply message limit if configured (automatic cleanup)
if self.max_messages is not None:
current_count = await self._redis_client.llen(self.redis_key) # type: ignore[misc]
if current_count > self.max_messages:
# Keep only the most recent max_messages using LTRIM
await self._redis_client.ltrim(self.redis_key, -self.max_messages, -1) # type: ignore[misc]
async def list_messages(self) -> list[Message]:
"""Get all messages from the store in chronological order (ChatMessageStoreProtocol protocol method).
This method implements the required ChatMessageStoreProtocol protocol for retrieving messages.
Returns all messages stored in Redis, ordered from oldest (index 0) to newest (index -1).
Returns:
List of Message objects in chronological order (oldest first).
Returns empty list if no messages exist or if Redis connection fails.
Example:
.. code-block:: python
# Get all conversation history
messages = await store.list_messages()
"""
# Ensure any initial messages are persisted to Redis first
await self._ensure_initial_messages_added()
messages = []
# Retrieve all messages from Redis list (oldest to newest)
redis_messages = await self._redis_client.lrange(self.redis_key, 0, -1) # type: ignore[misc]
if redis_messages:
for serialized_message in redis_messages:
# Deserialize each JSON message back to Message
message = self._deserialize_message(serialized_message)
messages.append(message)
return messages
async def serialize(self, **kwargs: Any) -> Any:
"""Serialize the current store state for persistence (ChatMessageStoreProtocol protocol method).
This method implements the required ChatMessageStoreProtocol protocol for state serialization.
Captures the Redis connection configuration and thread information needed to
reconstruct the store and reconnect to the same conversation data.
Keyword Args:
**kwargs: Additional arguments passed to Pydantic model_dump() for serialization.
Common options: exclude_none=True, by_alias=True
Returns:
Dictionary containing serialized store configuration that can be persisted
to databases, files, or other storage mechanisms.
"""
state = RedisStoreState(
thread_id=self.thread_id,
redis_url=self.redis_url,
key_prefix=self.key_prefix,
max_messages=self.max_messages,
)
return state.to_dict(exclude_none=False, **kwargs)
@classmethod
async def deserialize(cls, serialized_store_state: Any, **kwargs: Any) -> RedisChatMessageStore:
"""Deserialize state data into a new store instance (ChatMessageStoreProtocol protocol method).
This method implements the required ChatMessageStoreProtocol protocol for state deserialization.
Creates a new RedisChatMessageStore instance from previously serialized data,
allowing the store to reconnect to the same conversation data in Redis.
Args:
serialized_store_state: Previously serialized state data from serialize_state().
Should be a dictionary with thread_id, redis_url, etc.
Keyword Args:
**kwargs: Additional arguments passed to Pydantic model validation.
Returns:
A new RedisChatMessageStore instance configured from the serialized state.
Raises:
ValueError: If required fields are missing or invalid in the serialized state.
"""
if not serialized_store_state:
raise ValueError("serialized_store_state is required for deserialization")
# Validate and parse the serialized state using Pydantic
state = RedisStoreState.from_dict(serialized_store_state, **kwargs)
# Create and return a new store instance with the deserialized configuration
return cls(
redis_url=state.redis_url,
thread_id=state.thread_id,
key_prefix=state.key_prefix,
max_messages=state.max_messages,
)
async def update_from_state(self, serialized_store_state: Any, **kwargs: Any) -> None:
"""Deserialize state data into this store instance (ChatMessageStoreProtocol protocol method).
This method implements the required ChatMessageStoreProtocol protocol for state deserialization.
Restores the store configuration from previously serialized data, allowing the store
to reconnect to the same conversation data in Redis.
Args:
serialized_store_state: Previously serialized state data from serialize_state().
Should be a dictionary with thread_id, redis_url, etc.
Keyword Args:
**kwargs: Additional arguments passed to Pydantic model validation.
"""
if not serialized_store_state:
return
# Validate and parse the serialized state using Pydantic
state = RedisStoreState.from_dict(serialized_store_state, **kwargs)
# Update store configuration from deserialized state
self.thread_id = state.thread_id
if state.redis_url is not None:
self.redis_url = state.redis_url
self.key_prefix = state.key_prefix
self.max_messages = state.max_messages
# Recreate Redis client if the URL changed
if state.redis_url and state.redis_url != getattr(self, "_last_redis_url", None):
self._redis_client = redis.from_url(state.redis_url, decode_responses=True) # type: ignore[no-untyped-call]
self._last_redis_url = state.redis_url
# Reset initial message state since we're connecting to existing data
self._initial_messages_added = False
async def clear(self) -> None:
"""Remove all messages from the store.
Permanently deletes all messages for this conversation thread by removing
the Redis key. This operation cannot be undone.
Warning:
- This permanently deletes all conversation history
- Consider exporting messages before clearing if backup is needed
Example:
.. code-block:: python
# Clear conversation history
await store.clear()
# Verify messages are gone
messages = await store.list_messages()
assert len(messages) == 0
"""
await self._redis_client.delete(self.redis_key)
def _serialize_message(self, message: Message) -> str:
"""Serialize a Message to JSON string.
Args:
message: Message to serialize.
Returns:
JSON string representation of the message.
"""
# Serialize to compact JSON (no extra whitespace for Redis efficiency)
return message.to_json(separators=(",", ":"))
def _deserialize_message(self, serialized_message: str) -> Message:
"""Deserialize a JSON string to Message.
Args:
serialized_message: JSON string representation of a message.
Returns:
Message object.
"""
# Reconstruct Message using custom deserialization
return Message.from_json(serialized_message)
# ============================================================================
# List-like Convenience Methods (Redis-optimized async versions)
# ============================================================================
def __bool__(self) -> bool:
"""Return True since the store always exists once created.
This method is called by Python's truthiness checks (if store:).
Since a RedisChatMessageStore instance always represents a valid store,
this always returns True.
Returns:
Always True - the store exists and is ready for operations.
Note:
This is used by the Agent Framework to check if a message store
is configured: `if thread.message_store:`
"""
return True
async def __len__(self) -> int:
"""Return the number of messages in the Redis store.
Provides efficient message counting using Redis LLEN command.
This is the async equivalent of Python's built-in len() function.
Returns:
The count of messages currently stored in Redis.
"""
await self._ensure_initial_messages_added()
return await self._redis_client.llen(self.redis_key) # type: ignore[misc,no-any-return]
async def getitem(self, index: int) -> Message:
"""Get a message by index using Redis LINDEX.
Args:
index: The index of the message to retrieve.
Returns:
The Message at the specified index.
Raises:
IndexError: If the index is out of range.
"""
await self._ensure_initial_messages_added()
# Use Redis LINDEX for efficient single-item access
serialized_message = await self._redis_client.lindex(self.redis_key, index) # type: ignore[misc]
if serialized_message is None:
raise IndexError("list index out of range")
return self._deserialize_message(serialized_message)
async def setitem(self, index: int, item: Message) -> None:
"""Set a message at the specified index using Redis LSET.
Args:
index: The index at which to set the message.
item: The Message to set at the specified index.
Raises:
IndexError: If the index is out of range.
"""
await self._ensure_initial_messages_added()
# Validate index exists using LLEN
current_count = await self._redis_client.llen(self.redis_key) # type: ignore[misc]
if index < 0:
index = current_count + index
if index < 0 or index >= current_count:
raise IndexError("list index out of range")
# Use Redis LSET for efficient single-item update
serialized_message = self._serialize_message(item)
await self._redis_client.lset(self.redis_key, index, serialized_message) # type: ignore[misc]
async def append(self, item: Message) -> None:
"""Append a message to the end of the store.
Args:
item: The Message to append.
"""
await self.add_messages([item])
async def count(self) -> int:
"""Return the number of messages in the Redis store.
Returns:
The count of messages currently stored in Redis.
"""
await self._ensure_initial_messages_added()
return await self._redis_client.llen(self.redis_key) # type: ignore[misc,no-any-return]
async def index(self, item: Message) -> int:
"""Return the index of the first occurrence of the specified message.
Uses Redis LINDEX to iterate through the list without loading all messages.
Still O(N) but more memory efficient for large lists.
Args:
item: The Message to find.
Returns:
The index of the first occurrence of the message.
Raises:
ValueError: If the message is not found in the store.
"""
await self._ensure_initial_messages_added()
target_serialized = self._serialize_message(item)
list_length = await self._redis_client.llen(self.redis_key) # type: ignore[misc]
# Iterate through Redis list using LINDEX
for i in range(list_length):
redis_message = await self._redis_client.lindex(self.redis_key, i) # type: ignore[misc]
if redis_message == target_serialized:
return i
raise ValueError("Message not found in store")
async def remove(self, item: Message) -> None:
"""Remove the first occurrence of the specified message from the store.
Uses Redis LREM command for efficient removal by value.
O(N) but performed natively in Redis without data transfer.
Args:
item: The Message to remove.
Raises:
ValueError: If the message is not found in the store.
"""
await self._ensure_initial_messages_added()
# Serialize the message to match Redis storage format
target_serialized = self._serialize_message(item)
# Use LREM to remove first occurrence (count=1)
removed_count = await self._redis_client.lrem(self.redis_key, 1, target_serialized) # type: ignore[misc]
if removed_count == 0:
raise ValueError("Message not found in store")
async def extend(self, items: Sequence[Message]) -> None:
"""Extend the store by appending all messages from the iterable.
Args:
items: Sequence of Message objects to append.
"""
await self.add_messages(items)
async def ping(self) -> bool:
"""Test the Redis connection.
Returns:
True if the connection is successful, False otherwise.
"""
try:
await self._redis_client.ping() # type: ignore[misc]
return True
except Exception:
return False
async def aclose(self) -> None:
"""Close the Redis connection.
This method provides a clean way to close the underlying Redis connection
when the store is no longer needed. This is particularly useful in samples
and applications where explicit resource cleanup is desired.
"""
await self._redis_client.aclose() # type: ignore[misc]
def __repr__(self) -> str:
"""String representation of the store."""
return (
f"RedisChatMessageStore(thread_id='{self.thread_id}', "
f"redis_key='{self.redis_key}', max_messages={self.max_messages})"
)
@@ -2,9 +2,8 @@
"""New-pattern Redis context provider using BaseContextProvider.
This module provides ``_RedisContextProvider``, a side-by-side implementation of
:class:`RedisProvider` built on the new :class:`BaseContextProvider` hooks pattern.
It will be renamed to ``RedisContextProvider`` in PR2 when the old class is removed.
This module provides ``RedisContextProvider``, built on the new
:class:`BaseContextProvider` hooks pattern.
"""
from __future__ import annotations
@@ -43,17 +42,11 @@ if TYPE_CHECKING:
from agent_framework._agents import SupportsAgentRun
class _RedisContextProvider(BaseContextProvider):
class RedisContextProvider(BaseContextProvider):
"""Redis context provider using the new BaseContextProvider hooks pattern.
Stores context in Redis and retrieves scoped context via full-text or
optional hybrid vector search. This is the new-pattern equivalent of
:class:`RedisProvider`.
Note:
This class uses a temporary ``_`` prefix to coexist with the existing
:class:`RedisProvider`. It will be renamed to ``RedisContextProvider``
in PR2.
optional hybrid vector search.
"""
DEFAULT_CONTEXT_PROMPT = "## Memories\nConsider the following memories when answering user questions:"
@@ -429,4 +422,4 @@ class _RedisContextProvider(BaseContextProvider):
"""Async context manager exit."""
__all__ = ["_RedisContextProvider"]
__all__ = ["RedisContextProvider"]
@@ -2,9 +2,8 @@
"""New-pattern Redis history provider using BaseHistoryProvider.
This module provides ``_RedisHistoryProvider``, a side-by-side implementation of
:class:`RedisMessageStore` built on the new :class:`BaseHistoryProvider` hooks pattern.
It will be renamed to ``RedisHistoryProvider`` in PR2 when the old class is removed.
This module provides ``RedisHistoryProvider``, built on the new
:class:`BaseHistoryProvider` hooks pattern.
"""
from __future__ import annotations
@@ -18,17 +17,11 @@ from agent_framework._sessions import BaseHistoryProvider
from redis.credentials import CredentialProvider
class _RedisHistoryProvider(BaseHistoryProvider):
class RedisHistoryProvider(BaseHistoryProvider):
"""Redis-backed history provider using the new BaseHistoryProvider hooks pattern.
Stores conversation history in Redis Lists, with each session isolated by a
unique Redis key. This is the new-pattern equivalent of
:class:`RedisMessageStore`.
Note:
This class uses a temporary ``_`` prefix to coexist with the existing
:class:`RedisMessageStore`. It will be renamed to ``RedisHistoryProvider``
in PR2.
unique Redis key.
"""
def __init__(
@@ -181,4 +174,4 @@ class _RedisHistoryProvider(BaseHistoryProvider):
await self._redis_client.aclose() # type: ignore[misc]
__all__ = ["_RedisHistoryProvider"]
__all__ = ["RedisHistoryProvider"]
@@ -1,595 +0,0 @@
# Copyright (c) Microsoft. All rights reserved.
from __future__ import annotations
import json
import sys
from collections.abc import MutableSequence, Sequence
from functools import reduce
from operator import and_
from typing import Any, Literal, cast
import numpy as np
from agent_framework import Context, ContextProvider, Message
from agent_framework.exceptions import (
AgentException,
ServiceInitializationError,
ServiceInvalidRequestError,
)
from redisvl.index import AsyncSearchIndex
from redisvl.query import FilterQuery, HybridQuery, TextQuery
from redisvl.query.filter import FilterExpression, Tag
from redisvl.utils.token_escaper import TokenEscaper
from redisvl.utils.vectorize import BaseVectorizer
if sys.version_info >= (3, 11):
from typing import Self # pragma: no cover
else:
from typing_extensions import Self # pragma: no cover
if sys.version_info >= (3, 12):
from typing import override # type: ignore # pragma: no cover
else:
from typing_extensions import override # type: ignore[import] # pragma: no cover
class RedisProvider(ContextProvider):
"""Redis context provider with dynamic, filterable schema.
Stores context in Redis and retrieves scoped context.
Uses full-text or optional hybrid vector search to ground model responses.
"""
def __init__(
self,
redis_url: str = "redis://localhost:6379",
index_name: str = "context",
prefix: str = "context",
# Redis vectorizer configuration (optional, injected by client)
redis_vectorizer: BaseVectorizer | None = None,
vector_field_name: str | None = None,
vector_algorithm: Literal["flat", "hnsw"] | None = None,
vector_distance_metric: Literal["cosine", "ip", "l2"] | None = None,
# Partition fields (indexed for filtering)
application_id: str | None = None,
agent_id: str | None = None,
user_id: str | None = None,
thread_id: str | None = None,
scope_to_per_operation_thread_id: bool = False,
# Prompt and runtime
context_prompt: str = ContextProvider.DEFAULT_CONTEXT_PROMPT,
redis_index: Any = None,
overwrite_index: bool = False,
):
"""Create a Redis Context Provider.
Args:
redis_url: The Redis server URL.
index_name: The name of the Redis index.
prefix: The prefix for all keys in the Redis database.
redis_vectorizer: The vectorizer to use for Redis.
vector_field_name: The name of the vector field in Redis.
vector_algorithm: The algorithm to use for vector search.
vector_distance_metric: The distance metric to use for vector search.
application_id: The application ID to scope the context.
agent_id: The agent ID to scope the context.
user_id: The user ID to scope the context.
thread_id: The thread ID to scope the context.
scope_to_per_operation_thread_id: Whether to scope to the per-operation thread ID.
context_prompt: The context prompt to use for the provider.
redis_index: The Redis index to use for the provider.
overwrite_index: Whether to overwrite the existing Redis index.
"""
self.redis_url = redis_url
self.index_name = index_name
self.prefix = prefix
if redis_vectorizer is not None and not isinstance(redis_vectorizer, BaseVectorizer):
raise AgentException(
f"The redis vectorizer is not a valid type, got: {type(redis_vectorizer)}, expected: BaseVectorizer."
)
self.redis_vectorizer = redis_vectorizer
self.vector_field_name = vector_field_name
self.vector_algorithm: Literal["flat", "hnsw"] | None = vector_algorithm
self.vector_distance_metric: Literal["cosine", "ip", "l2"] | None = vector_distance_metric
self.application_id = application_id
self.agent_id = agent_id
self.user_id = user_id
self.thread_id = thread_id
self.scope_to_per_operation_thread_id = scope_to_per_operation_thread_id
self.context_prompt = context_prompt
self.overwrite_index = overwrite_index
self._per_operation_thread_id: str | None = None
self._token_escaper: TokenEscaper = TokenEscaper()
self._conversation_id: str | None = None
self._index_initialized: bool = False
self._schema_dict: dict[str, Any] | None = None
self.redis_index = redis_index or AsyncSearchIndex.from_dict(
self.schema_dict, redis_url=self.redis_url, validate_on_load=True
)
@property
def schema_dict(self) -> dict[str, Any]:
"""Get the Redis schema dictionary, computing and caching it on first access."""
if self._schema_dict is None:
# Get vector configuration from vectorizer if available
vector_dims = self.redis_vectorizer.dims if self.redis_vectorizer is not None else None
vector_datatype = self.redis_vectorizer.dtype if self.redis_vectorizer is not None else None
self._schema_dict = self._build_schema_dict(
index_name=self.index_name,
prefix=self.prefix,
vector_field_name=self.vector_field_name,
vector_dims=vector_dims,
vector_datatype=vector_datatype,
vector_algorithm=self.vector_algorithm,
vector_distance_metric=self.vector_distance_metric,
)
return self._schema_dict
def _build_filter_from_dict(self, filters: dict[str, str | None]) -> Any | None:
"""Builds a combined filter expression from simple equality tags.
This ANDs non-empty tag filters and is used to scope all operations to app/agent/user/thread partitions.
Args:
filters: Mapping of field name to value; falsy values are ignored.
Returns:
A combined filter expression or None if no filters are provided.
"""
parts = [Tag(k) == v for k, v in filters.items() if v]
return reduce(and_, parts) if parts else None
def _build_schema_dict(
self,
*,
index_name: str,
prefix: str,
vector_field_name: str | None,
vector_dims: int | None,
vector_datatype: str | None,
vector_algorithm: Literal["flat", "hnsw"] | None,
vector_distance_metric: Literal["cosine", "ip", "l2"] | None,
) -> dict[str, Any]:
"""Builds the RediSearch schema configuration dictionary.
Defines text and tag fields for messages plus an optional vector field enabling KNN/hybrid search.
Keyword Args:
index_name: Index name.
prefix: Key prefix.
vector_field_name: Vector field name or None.
vector_dims: Vector dimensionality or None.
vector_datatype: Vector datatype or None.
vector_algorithm: Vector index algorithm or None.
vector_distance_metric: Vector distance metric or None.
Returns:
Dict representing the index and fields configuration.
"""
fields: list[dict[str, Any]] = [
{"name": "role", "type": "tag"},
{"name": "mime_type", "type": "tag"},
{"name": "content", "type": "text"},
# Conversation tracking
{"name": "conversation_id", "type": "tag"},
{"name": "message_id", "type": "tag"},
{"name": "author_name", "type": "tag"},
# Partition fields (TAG for fast filtering)
{"name": "application_id", "type": "tag"},
{"name": "agent_id", "type": "tag"},
{"name": "user_id", "type": "tag"},
{"name": "thread_id", "type": "tag"},
]
# Add vector field only if configured (keeps provider runnable with no params)
if vector_field_name is not None and vector_dims is not None:
fields.append({
"name": vector_field_name,
"type": "vector",
"attrs": {
"algorithm": (vector_algorithm or "hnsw"),
"dims": int(vector_dims),
"distance_metric": (vector_distance_metric or "cosine"),
"datatype": (vector_datatype or "float32"),
},
})
return {
"index": {
"name": index_name,
"prefix": prefix,
"key_separator": ":",
"storage_type": "hash",
},
"fields": fields,
}
async def _ensure_index(self) -> None:
"""Initialize the search index.
- Connect to existing index if it exists and schema matches
- Create new index if it doesn't exist
- Overwrite if requested via overwrite_index=True
- Validate schema compatibility to prevent accidental data loss
"""
if self._index_initialized:
return
# Check if index already exists
index_exists = await self.redis_index.exists()
if not self.overwrite_index and index_exists:
# Validate schema compatibility before connecting
await self._validate_schema_compatibility()
# Create the index (will connect to existing or create new)
await self.redis_index.create(overwrite=self.overwrite_index, drop=False)
self._index_initialized = True
async def _validate_schema_compatibility(self) -> None:
"""Validate that existing index schema matches current configuration.
Raises ServiceInitializationError if schemas don't match, with helpful guidance.
self._build_schema_dict returns a minimal schema while Redis returns an expanded
schema with all defaults filled in. To compare for incompatibilities, compare
significant parts of the schema by creating signatures with normalized default values.
"""
# Defaults for attr normalization
TAG_DEFAULTS = {"separator": ",", "case_sensitive": False, "withsuffixtrie": False}
TEXT_DEFAULTS = {"weight": 1.0, "no_stem": False}
def _significant_index(i: dict[str, Any]) -> dict[str, Any]:
return {k: i.get(k) for k in ("name", "prefix", "key_separator", "storage_type")}
def _sig_tag(attrs: dict[str, Any] | None) -> dict[str, Any]:
a = {**TAG_DEFAULTS, **(attrs or {})}
return {k: a[k] for k in ("separator", "case_sensitive", "withsuffixtrie")}
def _sig_text(attrs: dict[str, Any] | None) -> dict[str, Any]:
a = {**TEXT_DEFAULTS, **(attrs or {})}
return {k: a[k] for k in ("weight", "no_stem")}
def _sig_vector(attrs: dict[str, Any] | None) -> dict[str, Any]:
a = {**(attrs or {})}
# Require these to exist if vector field is present
return {k: a.get(k) for k in ("algorithm", "dims", "distance_metric", "datatype")}
def _schema_signature(schema: dict[str, Any]) -> dict[str, Any]:
# Order-independent, minimal signature
sig: dict[str, Any] = {"index": _significant_index(schema.get("index", {})), "fields": {}}
for f in schema.get("fields", []):
name, ftype = f.get("name"), f.get("type")
if not name:
continue
if ftype == "tag":
sig["fields"][name] = {"type": "tag", "attrs": _sig_tag(f.get("attrs"))}
elif ftype == "text":
sig["fields"][name] = {"type": "text", "attrs": _sig_text(f.get("attrs"))}
elif ftype == "vector":
sig["fields"][name] = {"type": "vector", "attrs": _sig_vector(f.get("attrs"))}
else:
# Unknown field types: compare by type only
sig["fields"][name] = {"type": ftype}
return sig
existing_index = await AsyncSearchIndex.from_existing(self.index_name, redis_url=self.redis_url)
existing_schema = existing_index.schema.to_dict()
current_schema = self.schema_dict
existing_sig = _schema_signature(existing_schema)
current_sig = _schema_signature(current_schema)
if existing_sig != current_sig:
# Add sigs to error message
raise ServiceInitializationError(
"Existing Redis index schema is incompatible with the current configuration.\n"
f"Existing (significant): {json.dumps(existing_sig, indent=2, sort_keys=True)}\n"
f"Current (significant): {json.dumps(current_sig, indent=2, sort_keys=True)}\n"
"Set overwrite_index=True to rebuild if this change is intentional."
)
async def _add(
self,
*,
data: dict[str, Any] | list[dict[str, Any]],
metadata: dict[str, Any] | None = None,
) -> None:
"""Inserts one or many documents with partition fields populated.
Fills default partition fields, optionally embeds content when configured, and loads documents in a batch.
Keyword Args:
data: Single document or list of documents to insert.
metadata: Optional metadata dictionary (unused placeholder).
Raises:
ServiceInvalidRequestError: If required fields are missing or invalid.
"""
# Ensure provider has at least one scope set (symmetry with Mem0Provider)
self._validate_filters()
await self._ensure_index()
docs = data if isinstance(data, list) else [data]
prepared: list[dict[str, Any]] = []
for doc in docs:
d = dict(doc) # shallow copy
# Partition defaults
d.setdefault("application_id", self.application_id)
d.setdefault("agent_id", self.agent_id)
d.setdefault("user_id", self.user_id)
d.setdefault("thread_id", self._effective_thread_id)
# Conversation defaults
d.setdefault("conversation_id", self._conversation_id)
# Logical requirement
if "content" not in d:
raise ServiceInvalidRequestError("add() requires a 'content' field in data")
# Vector field requirement (only if schema has one)
if self.vector_field_name:
d.setdefault(self.vector_field_name, None)
prepared.append(d)
# Batch embed contents for every message
if self.redis_vectorizer and self.vector_field_name:
text_list = [d["content"] for d in prepared]
embeddings = await self.redis_vectorizer.aembed_many(text_list, batch_size=len(text_list))
for i, d in enumerate(prepared):
vec = np.asarray(embeddings[i], dtype=np.float32).tobytes()
field_name: str = self.vector_field_name
d[field_name] = vec
# Load all at once if supported
await self.redis_index.load(prepared)
return
async def _redis_search(
self,
text: str,
*,
text_scorer: str = "BM25STD",
filter_expression: Any | None = None,
return_fields: list[str] | None = None,
num_results: int = 10,
alpha: float = 0.7,
) -> list[dict[str, Any]]:
"""Runs a text or hybrid vector-text search with optional filters.
Builds a TextQuery or HybridQuery and automatically ANDs partition filters to keep results scoped and safe.
Args:
text: Query text.
Keyword Args:
text_scorer: Scorer to use for text ranking.
filter_expression: Additional filter expression to AND with partition filters.
return_fields: Fields to return in results.
num_results: Maximum number of results.
alpha: Hybrid balancing parameter when vectors are enabled.
Returns:
List of result dictionaries.
Raises:
ServiceInvalidRequestError: If input is invalid or the query fails.
"""
# Enforce presence of at least one provider-level filter (symmetry with Mem0Provider)
await self._ensure_index()
self._validate_filters()
q = (text or "").strip()
if not q:
raise ServiceInvalidRequestError("text_search() requires non-empty text")
num_results = max(int(num_results or 10), 1)
combined_filter = self._build_filter_from_dict({
"application_id": self.application_id,
"agent_id": self.agent_id,
"user_id": self.user_id,
"thread_id": self._effective_thread_id,
"conversation_id": self._conversation_id,
})
if filter_expression is not None:
combined_filter = (combined_filter & filter_expression) if combined_filter else filter_expression
# Choose return fields
return_fields = (
return_fields
if return_fields is not None
else ["content", "role", "application_id", "agent_id", "user_id", "thread_id"]
)
try:
if self.redis_vectorizer and self.vector_field_name:
# Build hybrid query: combine full-text and vector similarity
vector = await self.redis_vectorizer.aembed(q)
query = HybridQuery(
text=q,
text_field_name="content",
vector=vector,
vector_field_name=self.vector_field_name,
text_scorer=text_scorer,
filter_expression=combined_filter,
alpha=alpha,
dtype=self.redis_vectorizer.dtype,
num_results=num_results,
return_fields=return_fields,
stopwords=None,
)
hybrid_results = await self.redis_index.query(query)
return cast(list[dict[str, Any]], hybrid_results)
# Text-only search
query = TextQuery(
text=q,
text_field_name="content",
text_scorer=text_scorer,
filter_expression=combined_filter,
num_results=num_results,
return_fields=return_fields,
stopwords=None,
)
text_results = await self.redis_index.query(query)
return cast(list[dict[str, Any]], text_results)
except Exception as exc: # pragma: no cover - surface as framework error
raise ServiceInvalidRequestError(f"Redis text search failed: {exc}") from exc
async def search_all(self, page_size: int = 200) -> list[dict[str, Any]]:
"""Returns all documents in the index.
Streams results via pagination to avoid excessive memory and response sizes.
Args:
page_size: Page size used for pagination under the hood.
Returns:
List of all documents.
"""
out: list[dict[str, Any]] = []
async for batch in self.redis_index.paginate(
FilterQuery(FilterExpression("*"), return_fields=[], num_results=page_size),
page_size=page_size,
):
out.extend(batch)
return out
@property
def _effective_thread_id(self) -> str | None:
"""Resolves the active thread id.
Returns per-operation thread id when scoping is enabled; otherwise the provider's thread id.
"""
return self._per_operation_thread_id if self.scope_to_per_operation_thread_id else self.thread_id
@override
async def thread_created(self, thread_id: str | None) -> None:
"""Called when a new thread is created.
Captures the per-operation thread id when scoping is enabled to enforce single-thread usage.
Args:
thread_id: The ID of the thread or None.
"""
self._validate_per_operation_thread_id(thread_id)
self._per_operation_thread_id = self._per_operation_thread_id or thread_id
# Track current conversation id (Agent passes conversation_id here)
self._conversation_id = thread_id or self._conversation_id
@override
async def invoked(
self,
request_messages: Message | Sequence[Message],
response_messages: Message | Sequence[Message] | None = None,
invoke_exception: Exception | None = None,
**kwargs: Any,
) -> None:
self._validate_filters()
request_messages_list = [request_messages] if isinstance(request_messages, Message) else list(request_messages)
response_messages_list = (
[response_messages]
if isinstance(response_messages, Message)
else list(response_messages)
if response_messages
else []
)
messages_list = [*request_messages_list, *response_messages_list]
messages: list[dict[str, Any]] = []
for message in messages_list:
if message.role in {"user", "assistant", "system"} and message.text and message.text.strip():
shaped: dict[str, Any] = {
"role": message.role,
"content": message.text,
"conversation_id": self._conversation_id,
"message_id": message.message_id,
"author_name": message.author_name,
}
messages.append(shaped)
if messages:
await self._add(data=messages)
@override
async def invoking(self, messages: Message | MutableSequence[Message], **kwargs: Any) -> Context:
"""Called before invoking the model to provide scoped context.
Concatenates recent messages into a query, fetches matching memories from Redis.
Prepends them as instructions.
Args:
messages: List of new messages in the thread.
Keyword Args:
**kwargs: not used at present at present.
Returns:
Context: Context object containing instructions with memories.
"""
self._validate_filters()
messages_list = [messages] if isinstance(messages, Message) else list(messages)
input_text = "\n".join(msg.text for msg in messages_list if msg and msg.text and msg.text.strip())
memories = await self._redis_search(text=input_text)
line_separated_memories = "\n".join(
str(memory.get("content", "")) for memory in memories if memory.get("content")
)
return Context(
messages=[Message(role="user", text=f"{self.context_prompt}\n{line_separated_memories}")]
if line_separated_memories
else None
)
async def __aenter__(self) -> Self:
"""Async context manager entry.
No special setup is required; provided for symmetry with the Mem0 provider.
"""
return self
async def __aexit__(self, exc_type: type[BaseException] | None, exc_val: BaseException | None, exc_tb: Any) -> None:
"""Async context manager exit.
No cleanup is required; indexes/keys remain unless explicitly cleared.
"""
return
def _validate_filters(self) -> None:
"""Validates that at least one filter is provided.
Prevents unbounded operations by requiring a partition filter before reads or writes.
Raises:
ServiceInitializationError: If no filters are provided.
"""
if not self.agent_id and not self.user_id and not self.application_id and not self.thread_id:
raise ServiceInitializationError(
"At least one of the filters: agent_id, user_id, application_id, or thread_id is required."
)
def _validate_per_operation_thread_id(self, thread_id: str | None) -> None:
"""Validates that a new thread ID doesn't conflict when scoped.
Prevents cross-thread data leakage by enforcing single-thread usage when per-operation scoping is enabled.
Args:
thread_id: The new thread ID or None.
Raises:
ValueError: If a new thread ID conflicts with the existing one.
"""
if (
self.scope_to_per_operation_thread_id
and thread_id
and self._per_operation_thread_id
and thread_id != self._per_operation_thread_id
):
raise ValueError(
"RedisProvider can only be used with one thread, when scope_to_per_operation_thread_id is True."
)