mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
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:
committed by
GitHub
Unverified
parent
0c67dbbce5
commit
1e350ea22f
@@ -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."
|
||||
)
|
||||
Reference in New Issue
Block a user