mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: add RedisContextProvider (#716)
* Setting up * Readme * Add redis tests path to all-tests * First pass integration * Keep provider convention * First pass integration * add redis integration tests * update README.md * Add basic sample for redis integration * Add partitioning, add partition-aware tests, improve sample script * Fix code quality check * Try to resolve pytest check * Try to identify if pytest is the cause of failed checks * Re-enable tests * Rename redis test file * Removing some tests to narrow down issue * Revert, no difference * Delete temp files * Starting refactor of RedisProvider * Build dynamic schema builder, still need to do dynamic embedding model config * Add scope control * Complete first pass functionality with OpenAI + HF vectors -> Tests, Samples, Demo to follow * Fix code quality * attempt to identify rootcause of failed test * attempt to identify rootcause of failed test * Attempt to resolve code quality fail * Update pyproject.toml for foundry to pin azure-ai-projects == 1.1.0b3,azure-ai-agents == 1.2.0b3 * Add tests for redisprovider * Remove invalid tests * Add API key handling for openai vectorizer * Update uv.locl * Use master uv.lock * Begin sample file, add lazy index creation, fix faulty override * Index drop and reinit depends on drop_redis_index not overwrite * Add samples, threading included, escaped queries, verify threading works, sample README.md * Refactor filters * Opinionated vars * Allow filter expression combination * Try inline stubs for mypy * Address mypy errors * Better docstrings, tweaks for feedback * Tweak example 3 in redis_threads.py sample * adjust confusing name * Enrich docstrings * Add descriptions and comments to samples, externalize vectorizer choice, remove nltk and sentencetransformers dependnecy * Add descriptions and comments to samples, externalize vectorizer choice, remove nltk and sentencetransformers dependnecy * Incorporate initial feedback from dmytrostruk * Fix uv.lock * Attempt to resolve conflict * Use remote .tomls * Sanity check * fix tests * Remove hardcoded API key from samples * Fix incorrect env var * Make add and redis_search private * Fix tests relying on private funcs * Expand tests * Explainer comments to each test * Add a 'get_conversation_history' function to RedisProvider - This just returns messages in sequential order. Added 'created_at_*' timestamps to facilitate sequential recovery. function has to be manually invoked by user * Add agent-framework-redis to python/pyproject.toml * Remove get_conversation_history * improve redis context provider with pydantic techniques and safe index handling patterns * add RedisChatMessageStore * remove integration test :( * fix mypy error * Remove unused params * Redo schema validation to be order-invariant, handle attrs (previously throwing errors due to strict ==) * Expand explanation * Add ChatMessageStore example * Fix comments in redis_conversation.py * Resolving uv.lock conflict, update to match main * Fix test in redis provider * Apply suggestion from @ekzhu * Update python/packages/main/pyproject.toml --------- Co-authored-by: Tyler Hutcherson <tyler.hutcherson@redis.com> Co-authored-by: Eric Zhu <ekzhu@users.noreply.github.com>
This commit is contained in:
@@ -6,6 +6,7 @@ from abc import ABC, abstractmethod
|
||||
from collections.abc import MutableSequence, Sequence
|
||||
from contextlib import AsyncExitStack
|
||||
from types import TracebackType
|
||||
from typing import ClassVar
|
||||
|
||||
from ._pydantic import AFBaseModel
|
||||
from ._types import ChatMessage, Contents
|
||||
@@ -47,6 +48,11 @@ class ContextProvider(AFBaseModel, ABC):
|
||||
just before invocation.
|
||||
"""
|
||||
|
||||
# Default prompt to be used by all context providers when assembling memories/instructions
|
||||
DEFAULT_CONTEXT_PROMPT: ClassVar[str] = (
|
||||
"## Memories\nConsider the following memories when answering user questions:"
|
||||
)
|
||||
|
||||
async def thread_created(self, thread_id: str | None) -> None:
|
||||
"""Called just after a new thread is created.
|
||||
|
||||
|
||||
@@ -0,0 +1,24 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import importlib
|
||||
from typing import Any
|
||||
|
||||
PACKAGE_NAME = "agent_framework_redis"
|
||||
PACKAGE_EXTRA = "redis"
|
||||
_IMPORTS = ["__version__", "RedisProvider", "RedisChatMessageStore"]
|
||||
|
||||
|
||||
def __getattr__(name: str) -> Any:
|
||||
if name in _IMPORTS:
|
||||
try:
|
||||
return getattr(importlib.import_module(PACKAGE_NAME), name)
|
||||
except ModuleNotFoundError as exc:
|
||||
raise ModuleNotFoundError(
|
||||
f"The '{PACKAGE_EXTRA}' extra is not installed, "
|
||||
f"please do `pip install agent-framework[{PACKAGE_EXTRA}]`"
|
||||
) from exc
|
||||
raise AttributeError(f"Module {PACKAGE_NAME} has no attribute {name}.")
|
||||
|
||||
|
||||
def __dir__() -> list[str]:
|
||||
return _IMPORTS
|
||||
@@ -0,0 +1,5 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from agent_framework_redis import RedisChatMessageStore, RedisProvider, __version__
|
||||
|
||||
__all__ = ["RedisChatMessageStore", "RedisProvider", "__version__"]
|
||||
@@ -44,6 +44,9 @@ azure = [
|
||||
foundry = [
|
||||
"agent-framework-foundry"
|
||||
]
|
||||
redis = [
|
||||
"agent-framework-redis"
|
||||
]
|
||||
viz = [
|
||||
"graphviz>=0.20.0",
|
||||
]
|
||||
@@ -61,6 +64,7 @@ all = [
|
||||
"agent-framework-foundry",
|
||||
"agent-framework-runtime",
|
||||
"agent-framework-mem0",
|
||||
"agent-framework-redis",
|
||||
"agent-framework-devui"
|
||||
]
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
import sys
|
||||
from collections.abc import MutableSequence, Sequence
|
||||
from typing import Any, Final
|
||||
from typing import Any
|
||||
|
||||
from agent_framework import ChatMessage, Context, ContextProvider, TextContent
|
||||
from agent_framework.exceptions import ServiceInitializationError
|
||||
@@ -15,9 +15,6 @@ else:
|
||||
from typing_extensions import Self # pragma: no cover
|
||||
|
||||
|
||||
DEFAULT_CONTEXT_PROMPT: Final[str] = "## Memories\nConsider the following memories when answering user questions:"
|
||||
|
||||
|
||||
class Mem0Provider(ContextProvider):
|
||||
mem0_client: AsyncMemoryClient
|
||||
api_key: str | None = None
|
||||
@@ -26,7 +23,7 @@ class Mem0Provider(ContextProvider):
|
||||
thread_id: str | None = None
|
||||
user_id: str | None = None
|
||||
scope_to_per_operation_thread_id: bool = False
|
||||
context_prompt: str = DEFAULT_CONTEXT_PROMPT
|
||||
context_prompt: str = ContextProvider.DEFAULT_CONTEXT_PROMPT
|
||||
|
||||
_should_close_client: bool = PrivateAttr(default=False) # Track whether we should close client connection
|
||||
|
||||
@@ -38,7 +35,7 @@ class Mem0Provider(ContextProvider):
|
||||
thread_id: str | None = None,
|
||||
user_id: str | None = None,
|
||||
scope_to_per_operation_thread_id: bool = False,
|
||||
context_prompt: str = DEFAULT_CONTEXT_PROMPT,
|
||||
context_prompt: str = ContextProvider.DEFAULT_CONTEXT_PROMPT,
|
||||
mem0_client: AsyncMemoryClient | None = None,
|
||||
) -> None:
|
||||
"""Initializes a new instance of the Mem0Provider class.
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) Microsoft Corporation.
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE
|
||||
@@ -0,0 +1,53 @@
|
||||
# Get Started with Microsoft Agent Framework Redis
|
||||
|
||||
Please install this package as the extra for `agent-framework`:
|
||||
|
||||
```bash
|
||||
pip install agent-framework[redis]
|
||||
```
|
||||
|
||||
## Components
|
||||
|
||||
### Memory Context Provider
|
||||
|
||||
The `RedisProvider` enables persistent context & memory capabilities for your agents, allowing them to remember user preferences and conversation context across sessions and threads.
|
||||
|
||||
#### Basic Usage Examples
|
||||
|
||||
Review the set of [getting started examples](../../samples/getting_started/context_providers/redis/README.md) for using the Redis context provider.
|
||||
|
||||
### Redis Chat Message Store
|
||||
|
||||
The `RedisChatMessageStore` provides persistent conversation storage using Redis Lists, enabling chat history to survive application restarts and support distributed applications.
|
||||
|
||||
#### Key Features
|
||||
|
||||
- **Persistent Storage**: Messages survive application restarts
|
||||
- **Thread Isolation**: Each conversation thread has its own Redis key
|
||||
- **Message Limits**: Configurable automatic trimming of old messages
|
||||
- **Serialization Support**: Full compatibility with Agent Framework thread serialization
|
||||
- **Production Ready**: Connection pooling, error handling, and performance optimized
|
||||
|
||||
#### Basic Usage Examples
|
||||
|
||||
See the complete [Redis chat message store examples](../../samples/getting_started/threads/redis_chat_message_store_thread.py) including:
|
||||
- User session management
|
||||
- Conversation persistence across restarts
|
||||
- Thread serialization and deserialization
|
||||
- Automatic message trimming
|
||||
- Error handling patterns
|
||||
|
||||
### Installing and running Redis
|
||||
|
||||
You have 3 options to set-up Redis:
|
||||
|
||||
#### Option A: Local Redis with Docker
|
||||
```bash
|
||||
docker run --name redis -p 6379:6379 -d redis:8.0.3
|
||||
```
|
||||
|
||||
#### Option B: Redis Cloud
|
||||
Get a free db at https://redis.io/cloud/
|
||||
|
||||
#### Option C: Azure Managed Redis
|
||||
Here's a quickstart guide to create **Azure Managed Redis** for as low as $12 monthly: https://learn.microsoft.com/en-us/azure/redis/quickstart-create-managed-redis
|
||||
@@ -0,0 +1,16 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
import importlib.metadata
|
||||
|
||||
from ._chat_message_store import RedisChatMessageStore
|
||||
from ._provider import RedisProvider
|
||||
|
||||
try:
|
||||
__version__ = importlib.metadata.version(__name__)
|
||||
except importlib.metadata.PackageNotFoundError:
|
||||
__version__ = "0.0.0" # Fallback for development mode
|
||||
|
||||
__all__ = [
|
||||
"RedisChatMessageStore",
|
||||
"RedisProvider",
|
||||
"__version__",
|
||||
]
|
||||
@@ -0,0 +1,507 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Sequence
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
import redis.asyncio as redis
|
||||
from agent_framework import ChatMessage
|
||||
from agent_framework._pydantic import AFBaseModel
|
||||
|
||||
|
||||
class RedisStoreState(AFBaseModel):
|
||||
"""State model for serializing and deserializing Redis chat message store data."""
|
||||
|
||||
thread_id: str
|
||||
redis_url: str | None = None
|
||||
key_prefix: str = "chat_messages"
|
||||
max_messages: int | None = None
|
||||
|
||||
|
||||
class RedisChatMessageStore:
|
||||
"""Redis-backed implementation of ChatMessageStore 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,
|
||||
thread_id: str | None = None,
|
||||
key_prefix: str = "chat_messages",
|
||||
max_messages: int | None = None,
|
||||
messages: Sequence[ChatMessage] | None = None,
|
||||
) -> None:
|
||||
"""Initialize the Redis chat message store.
|
||||
|
||||
Creates a Redis-backed chat message store for a specific conversation thread.
|
||||
The store will automatically create a Redis connection and manage message
|
||||
persistence using Redis List operations.
|
||||
|
||||
Args:
|
||||
redis_url: Redis connection URL (e.g., "redis://localhost:6379").
|
||||
Required for establishing Redis connection.
|
||||
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 redis_url is None (Redis connection is required).
|
||||
redis.ConnectionError: If unable to connect to Redis server.
|
||||
|
||||
|
||||
"""
|
||||
# Validate required parameters
|
||||
if redis_url is None:
|
||||
raise ValueError("redis_url is required for Redis connection")
|
||||
|
||||
# Store configuration
|
||||
self.redis_url = redis_url
|
||||
self.thread_id = thread_id or f"thread_{uuid4()}"
|
||||
self.key_prefix = key_prefix
|
||||
self.max_messages = max_messages
|
||||
|
||||
# Initialize Redis client with connection pooling and async support
|
||||
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[ChatMessage]) -> 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 ChatMessage 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[ChatMessage]) -> None:
|
||||
"""Add messages to the Redis store (ChatMessageStore protocol method).
|
||||
|
||||
This method implements the required ChatMessageStore 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 ChatMessage 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:
|
||||
```python
|
||||
messages = [ChatMessage(role=Role.USER, text="Hello"), ChatMessage(role=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[ChatMessage]:
|
||||
"""Get all messages from the store in chronological order (ChatMessageStore protocol method).
|
||||
|
||||
This method implements the required ChatMessageStore protocol for retrieving messages.
|
||||
Returns all messages stored in Redis, ordered from oldest (index 0) to newest (index -1).
|
||||
|
||||
Returns:
|
||||
List of ChatMessage objects in chronological order (oldest first).
|
||||
Returns empty list if no messages exist or if Redis connection fails.
|
||||
|
||||
Example:
|
||||
```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 ChatMessage
|
||||
message = self._deserialize_message(serialized_message)
|
||||
messages.append(message)
|
||||
|
||||
return messages
|
||||
|
||||
async def serialize_state(self, **kwargs: Any) -> Any:
|
||||
"""Serialize the current store state for persistence (ChatMessageStore protocol method).
|
||||
|
||||
This method implements the required ChatMessageStore protocol for state serialization.
|
||||
Captures the Redis connection configuration and thread information needed to
|
||||
reconstruct the store and reconnect to the same conversation data.
|
||||
|
||||
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.model_dump(**kwargs)
|
||||
|
||||
async def deserialize_state(self, serialized_store_state: Any, **kwargs: Any) -> None:
|
||||
"""Deserialize state data into this store instance (ChatMessageStore protocol method).
|
||||
|
||||
This method implements the required ChatMessageStore 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.
|
||||
**kwargs: Additional arguments passed to Pydantic model validation.
|
||||
"""
|
||||
if not serialized_store_state:
|
||||
return
|
||||
|
||||
# Validate and parse the serialized state using Pydantic
|
||||
state = RedisStoreState.model_validate(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:
|
||||
```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: ChatMessage) -> str:
|
||||
"""Serialize a ChatMessage to JSON string.
|
||||
|
||||
Args:
|
||||
message: ChatMessage to serialize.
|
||||
|
||||
Returns:
|
||||
JSON string representation of the message.
|
||||
"""
|
||||
# Convert ChatMessage to dictionary using Pydantic serialization
|
||||
message_dict = message.model_dump()
|
||||
# Serialize to compact JSON (no extra whitespace for Redis efficiency)
|
||||
return json.dumps(message_dict, separators=(",", ":"))
|
||||
|
||||
def _deserialize_message(self, serialized_message: str) -> ChatMessage:
|
||||
"""Deserialize a JSON string to ChatMessage.
|
||||
|
||||
Args:
|
||||
serialized_message: JSON string representation of a message.
|
||||
|
||||
Returns:
|
||||
ChatMessage object.
|
||||
"""
|
||||
# Parse JSON string back to dictionary
|
||||
message_dict = json.loads(serialized_message)
|
||||
# Reconstruct ChatMessage using Pydantic validation
|
||||
return ChatMessage.model_validate(message_dict)
|
||||
|
||||
# ============================================================================
|
||||
# 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) -> ChatMessage:
|
||||
"""Get a message by index using Redis LINDEX.
|
||||
|
||||
Args:
|
||||
index: The index of the message to retrieve.
|
||||
|
||||
Returns:
|
||||
The ChatMessage 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: ChatMessage) -> None:
|
||||
"""Set a message at the specified index using Redis LSET.
|
||||
|
||||
Args:
|
||||
index: The index at which to set the message.
|
||||
item: The ChatMessage 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: ChatMessage) -> None:
|
||||
"""Append a message to the end of the store.
|
||||
|
||||
Args:
|
||||
item: The ChatMessage 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: ChatMessage) -> 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 ChatMessage 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("ChatMessage not found in store")
|
||||
|
||||
async def remove(self, item: ChatMessage) -> 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 ChatMessage 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("ChatMessage not found in store")
|
||||
|
||||
async def extend(self, items: Sequence[ChatMessage]) -> None:
|
||||
"""Extend the store by appending all messages from the iterable.
|
||||
|
||||
Args:
|
||||
items: Sequence of ChatMessage 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})"
|
||||
)
|
||||
@@ -0,0 +1,547 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from collections.abc import MutableSequence, Sequence
|
||||
from functools import reduce
|
||||
from operator import and_
|
||||
from typing import Any, Literal, cast
|
||||
|
||||
from agent_framework import ChatMessage, Context, ContextProvider, Role, TextContent
|
||||
from agent_framework.exceptions import (
|
||||
ServiceInitializationError,
|
||||
ServiceInvalidRequestError,
|
||||
)
|
||||
|
||||
if sys.version_info >= (3, 11):
|
||||
from typing import Self # pragma: no cover
|
||||
else:
|
||||
from typing_extensions import Self # pragma: no cover
|
||||
|
||||
import json
|
||||
|
||||
import numpy as np
|
||||
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
|
||||
|
||||
|
||||
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.
|
||||
"""
|
||||
|
||||
# Connection and indexing
|
||||
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
|
||||
_per_operation_thread_id: str | None = None
|
||||
_token_escaper: TokenEscaper = TokenEscaper()
|
||||
_conversation_id: str | None = None
|
||||
_index_initialized: bool = False
|
||||
_schema_dict: dict[str, Any] | None = None
|
||||
|
||||
def model_post_init(self, __context: Any) -> None:
|
||||
"""Post-initialization hook to set up computed fields after Pydantic initialization.
|
||||
|
||||
This is called automatically by Pydantic after the model is initialized.
|
||||
"""
|
||||
# Create Redis index using the cached schema_dict property
|
||||
self.redis_index = 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.
|
||||
|
||||
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.
|
||||
|
||||
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.
|
||||
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
|
||||
|
||||
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
|
||||
|
||||
async def messages_adding(self, thread_id: str | None, new_messages: ChatMessage | Sequence[ChatMessage]) -> None:
|
||||
"""Called when a new message is being added to the thread.
|
||||
|
||||
Validates scope, normalizes allowed roles, and persists messages to Redis via add().
|
||||
|
||||
Args:
|
||||
thread_id: The ID of the thread or None.
|
||||
new_messages: New messages to add.
|
||||
"""
|
||||
self._validate_filters()
|
||||
self._validate_per_operation_thread_id(thread_id)
|
||||
self._per_operation_thread_id = self._per_operation_thread_id or thread_id
|
||||
|
||||
messages_list = [new_messages] if isinstance(new_messages, ChatMessage) else list(new_messages)
|
||||
|
||||
messages: list[dict[str, Any]] = []
|
||||
for message in messages_list:
|
||||
if (
|
||||
message.role.value in {Role.USER.value, Role.ASSISTANT.value, Role.SYSTEM.value}
|
||||
and message.text
|
||||
and message.text.strip()
|
||||
):
|
||||
shaped: dict[str, Any] = {
|
||||
"role": message.role.value,
|
||||
"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)
|
||||
|
||||
async def model_invoking(self, messages: ChatMessage | MutableSequence[ChatMessage]) -> 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.
|
||||
|
||||
Returns:
|
||||
Context: Context object containing instructions with memories.
|
||||
"""
|
||||
self._validate_filters()
|
||||
messages_list = [messages] if isinstance(messages, ChatMessage) 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")
|
||||
)
|
||||
content = TextContent(f"{self.context_prompt}\n{line_separated_memories}") if line_separated_memories else None
|
||||
return Context(contents=[content] if content 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."
|
||||
)
|
||||
@@ -0,0 +1,95 @@
|
||||
[project]
|
||||
name = "agent-framework-redis"
|
||||
description = "Redis integration for Microsoft Agent Framework."
|
||||
authors = [{ name = "Microsoft", email = "SK-Support@microsoft.com"}]
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
version = "0.0.0b1"
|
||||
license-files = ["LICENSE"]
|
||||
urls.homepage = "https://learn.microsoft.com/en-us/semantic-kernel/overview/"
|
||||
urls.source = "https://github.com/microsoft/agent-framework/tree/main/python"
|
||||
urls.release_notes = "https://github.com/microsoft/agent-framework/releases?q=tag%3Apython-1&expanded=true"
|
||||
urls.issues = "https://github.com/microsoft/agent-framework/issues"
|
||||
classifiers = [
|
||||
"License :: OSI Approved :: MIT License",
|
||||
"Development Status :: 5 - Production/Stable",
|
||||
"Intended Audience :: Developers",
|
||||
"Programming Language :: Python :: 3",
|
||||
"Programming Language :: Python :: 3.10",
|
||||
"Programming Language :: Python :: 3.11",
|
||||
"Programming Language :: Python :: 3.12",
|
||||
"Programming Language :: Python :: 3.13",
|
||||
"Framework :: Pydantic :: 2",
|
||||
"Typing :: Typed",
|
||||
]
|
||||
dependencies = [
|
||||
"agent-framework",
|
||||
"redis>=6.4.0",
|
||||
"redisvl>=0.8.2",
|
||||
"numpy>=2.2.6"
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
prerelease = "if-necessary-or-explicit"
|
||||
environments = [
|
||||
"sys_platform == 'darwin'",
|
||||
"sys_platform == 'linux'",
|
||||
"sys_platform == 'win32'"
|
||||
]
|
||||
|
||||
[tool.uv-dynamic-versioning]
|
||||
fallback-version = "0.0.0"
|
||||
[tool.pytest.ini_options]
|
||||
testpaths = 'tests'
|
||||
addopts = "-ra -q -r fEX"
|
||||
asyncio_mode = "auto"
|
||||
asyncio_default_fixture_loop_scope = "function"
|
||||
filterwarnings = [
|
||||
"ignore:Support for class-based `config` is deprecated:DeprecationWarning:pydantic.*"
|
||||
]
|
||||
timeout = 120
|
||||
|
||||
[tool.ruff]
|
||||
extend = "../../pyproject.toml"
|
||||
|
||||
[tool.coverage.run]
|
||||
omit = [
|
||||
"**/__init__.py"
|
||||
]
|
||||
|
||||
[tool.pyright]
|
||||
extend = "../../pyproject.toml"
|
||||
exclude = ['tests']
|
||||
|
||||
[tool.mypy]
|
||||
plugins = ['pydantic.mypy']
|
||||
strict = true
|
||||
python_version = "3.10"
|
||||
ignore_missing_imports = true
|
||||
disallow_untyped_defs = true
|
||||
no_implicit_optional = true
|
||||
check_untyped_defs = true
|
||||
warn_return_any = true
|
||||
show_error_codes = true
|
||||
warn_unused_ignores = false
|
||||
disallow_incomplete_defs = true
|
||||
disallow_untyped_decorators = true
|
||||
|
||||
[tool.bandit]
|
||||
targets = ["agent_framework_redis"]
|
||||
exclude_dirs = ["tests"]
|
||||
|
||||
[tool.poe]
|
||||
executor.type = "uv"
|
||||
include = "../../shared_tasks.toml"
|
||||
[tool.poe.tasks]
|
||||
mypy = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework_redis"
|
||||
test = "pytest --cov=agent_framework_redis --cov-report=term-missing:skip-covered tests"
|
||||
|
||||
[tool.uv.build-backend]
|
||||
module-name = "agent_framework_redis"
|
||||
module-root = ""
|
||||
|
||||
[build-system]
|
||||
requires = ["uv_build>=0.8.2,<0.9.0"]
|
||||
build-backend = "uv_build"
|
||||
@@ -0,0 +1,496 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from agent_framework import ChatMessage, Role, TextContent
|
||||
|
||||
from agent_framework_redis import RedisChatMessageStore
|
||||
|
||||
|
||||
class TestRedisChatMessageStore:
|
||||
"""Unit tests for RedisChatMessageStore using mocked Redis client.
|
||||
|
||||
These tests use mocked Redis operations to verify the logic and behavior
|
||||
of the RedisChatMessageStore without requiring a real Redis server.
|
||||
"""
|
||||
|
||||
@pytest.fixture
|
||||
def sample_messages(self):
|
||||
"""Sample chat messages for testing."""
|
||||
return [
|
||||
ChatMessage(role=Role.USER, text="Hello", message_id="msg1"),
|
||||
ChatMessage(role=Role.ASSISTANT, text="Hi there!", message_id="msg2"),
|
||||
ChatMessage(role=Role.USER, text="How are you?", message_id="msg3"),
|
||||
]
|
||||
|
||||
@pytest.fixture
|
||||
def mock_redis_client(self):
|
||||
"""Mock Redis client with all required methods."""
|
||||
client = MagicMock()
|
||||
# Core list operations
|
||||
client.lrange = AsyncMock(return_value=[])
|
||||
client.llen = AsyncMock(return_value=0)
|
||||
client.lindex = AsyncMock(return_value=None)
|
||||
client.lset = AsyncMock(return_value=True)
|
||||
client.lrem = AsyncMock(return_value=0)
|
||||
client.lpop = AsyncMock(return_value=None)
|
||||
client.rpop = AsyncMock(return_value=None)
|
||||
client.ltrim = AsyncMock(return_value=True)
|
||||
client.delete = AsyncMock(return_value=1)
|
||||
|
||||
# Pipeline operations
|
||||
mock_pipeline = AsyncMock()
|
||||
mock_pipeline.rpush = AsyncMock()
|
||||
mock_pipeline.execute = AsyncMock()
|
||||
client.pipeline.return_value.__aenter__.return_value = mock_pipeline
|
||||
|
||||
return client
|
||||
|
||||
@pytest.fixture
|
||||
def redis_store(self, mock_redis_client):
|
||||
"""Redis chat message store with mocked client."""
|
||||
with patch("agent_framework_redis._chat_message_store.redis.from_url") as mock_from_url:
|
||||
mock_from_url.return_value = mock_redis_client
|
||||
store = RedisChatMessageStore(redis_url="redis://localhost:6379", thread_id="test_thread_123")
|
||||
store._redis_client = mock_redis_client
|
||||
return store
|
||||
|
||||
def test_init_with_thread_id(self):
|
||||
"""Test initialization with explicit thread ID."""
|
||||
thread_id = "user123_session456"
|
||||
with patch("agent_framework_redis._chat_message_store.redis.from_url"):
|
||||
store = RedisChatMessageStore(redis_url="redis://localhost:6379", thread_id=thread_id)
|
||||
|
||||
assert store.thread_id == thread_id
|
||||
assert store.redis_url == "redis://localhost:6379"
|
||||
assert store.key_prefix == "chat_messages"
|
||||
assert store.redis_key == f"chat_messages:{thread_id}"
|
||||
|
||||
def test_init_auto_generate_thread_id(self):
|
||||
"""Test initialization with auto-generated thread ID."""
|
||||
with patch("agent_framework_redis._chat_message_store.redis.from_url"):
|
||||
store = RedisChatMessageStore(redis_url="redis://localhost:6379")
|
||||
|
||||
assert store.thread_id is not None
|
||||
assert store.thread_id.startswith("thread_")
|
||||
assert len(store.thread_id) > 10 # Should be a UUID
|
||||
|
||||
def test_init_with_custom_prefix(self):
|
||||
"""Test initialization with custom key prefix."""
|
||||
with patch("agent_framework_redis._chat_message_store.redis.from_url"):
|
||||
store = RedisChatMessageStore(
|
||||
redis_url="redis://localhost:6379", thread_id="test123", key_prefix="custom_messages"
|
||||
)
|
||||
|
||||
assert store.redis_key == "custom_messages:test123"
|
||||
|
||||
def test_init_with_max_messages(self):
|
||||
"""Test initialization with message limit."""
|
||||
with patch("agent_framework_redis._chat_message_store.redis.from_url"):
|
||||
store = RedisChatMessageStore(redis_url="redis://localhost:6379", thread_id="test123", max_messages=100)
|
||||
|
||||
assert store.max_messages == 100
|
||||
|
||||
def test_init_with_redis_url_required(self):
|
||||
"""Test that redis_url is required for initialization."""
|
||||
with pytest.raises(ValueError, match="redis_url is required for Redis connection"):
|
||||
# Should raise an exception since redis_url is required
|
||||
RedisChatMessageStore(thread_id="test123")
|
||||
|
||||
def test_init_with_initial_messages(self, sample_messages):
|
||||
"""Test initialization with initial messages."""
|
||||
with patch("agent_framework_redis._chat_message_store.redis.from_url"):
|
||||
store = RedisChatMessageStore(
|
||||
redis_url="redis://localhost:6379", thread_id="test123", messages=sample_messages
|
||||
)
|
||||
|
||||
assert store._initial_messages == sample_messages
|
||||
|
||||
async def test_add_messages_single(self, redis_store, mock_redis_client, sample_messages):
|
||||
"""Test adding a single message using pipeline operations."""
|
||||
message = sample_messages[0]
|
||||
|
||||
await redis_store.add_messages([message])
|
||||
|
||||
# Verify pipeline operations were called
|
||||
mock_redis_client.pipeline.assert_called_with(transaction=True)
|
||||
|
||||
# Get the pipeline mock and verify it was used correctly
|
||||
pipeline_mock = mock_redis_client.pipeline.return_value.__aenter__.return_value
|
||||
pipeline_mock.rpush.assert_called()
|
||||
pipeline_mock.execute.assert_called()
|
||||
|
||||
async def test_add_messages_multiple(self, redis_store, mock_redis_client, sample_messages):
|
||||
"""Test adding multiple messages using pipeline operations."""
|
||||
await redis_store.add_messages(sample_messages)
|
||||
|
||||
# Verify pipeline operations
|
||||
mock_redis_client.pipeline.assert_called_with(transaction=True)
|
||||
|
||||
# Verify rpush was called for each message
|
||||
pipeline_mock = mock_redis_client.pipeline.return_value.__aenter__.return_value
|
||||
assert pipeline_mock.rpush.call_count == len(sample_messages)
|
||||
|
||||
async def test_add_messages_with_max_limit(self, mock_redis_client):
|
||||
"""Test adding messages with max limit triggers trimming."""
|
||||
with patch("agent_framework_redis._chat_message_store.redis.from_url") as mock_from_url:
|
||||
mock_from_url.return_value = mock_redis_client
|
||||
|
||||
# Mock llen to return count that exceeds limit after adding
|
||||
mock_redis_client.llen.return_value = 5
|
||||
|
||||
store = RedisChatMessageStore(redis_url="redis://localhost:6379", thread_id="test123", max_messages=3)
|
||||
store._redis_client = mock_redis_client
|
||||
|
||||
message = ChatMessage(role=Role.USER, text="Test")
|
||||
await store.add_messages([message])
|
||||
|
||||
# Should trim after adding to keep only last 3 messages
|
||||
mock_redis_client.ltrim.assert_called_once_with("chat_messages:test123", -3, -1)
|
||||
|
||||
async def test_list_messages_empty(self, redis_store, mock_redis_client):
|
||||
"""Test listing messages when store is empty."""
|
||||
mock_redis_client.lrange.return_value = []
|
||||
|
||||
messages = await redis_store.list_messages()
|
||||
|
||||
assert messages == []
|
||||
mock_redis_client.lrange.assert_called_once_with("chat_messages:test_thread_123", 0, -1)
|
||||
|
||||
async def test_list_messages_with_data(self, redis_store, mock_redis_client, sample_messages):
|
||||
"""Test listing messages with data in Redis."""
|
||||
# Create proper serialized messages using the actual serialization method
|
||||
test_messages = [
|
||||
ChatMessage(role=Role.USER, text="Hello", message_id="msg1"),
|
||||
ChatMessage(role=Role.ASSISTANT, text="Hi there!", message_id="msg2"),
|
||||
]
|
||||
serialized_messages = [redis_store._serialize_message(msg) for msg in test_messages]
|
||||
mock_redis_client.lrange.return_value = serialized_messages
|
||||
|
||||
messages = await redis_store.list_messages()
|
||||
|
||||
assert len(messages) == 2
|
||||
assert messages[0].role == Role.USER
|
||||
assert messages[0].text == "Hello"
|
||||
assert messages[1].role == Role.ASSISTANT
|
||||
assert messages[1].text == "Hi there!"
|
||||
|
||||
async def test_list_messages_with_initial_messages(self, sample_messages):
|
||||
"""Test that initial messages are added to Redis and retrieved correctly."""
|
||||
with patch("agent_framework_redis._chat_message_store.redis.from_url") as mock_from_url:
|
||||
mock_redis_client = MagicMock()
|
||||
mock_redis_client.llen = AsyncMock(return_value=0) # Redis key is empty
|
||||
mock_redis_client.lrange = AsyncMock(return_value=[])
|
||||
|
||||
# Mock pipeline for adding initial messages
|
||||
mock_pipeline = AsyncMock()
|
||||
mock_pipeline.rpush = AsyncMock()
|
||||
mock_pipeline.execute = AsyncMock()
|
||||
mock_redis_client.pipeline.return_value.__aenter__.return_value = mock_pipeline
|
||||
|
||||
mock_from_url.return_value = mock_redis_client
|
||||
|
||||
store = RedisChatMessageStore(
|
||||
redis_url="redis://localhost:6379",
|
||||
thread_id="test123",
|
||||
messages=sample_messages[:1], # One initial message
|
||||
)
|
||||
store._redis_client = mock_redis_client
|
||||
|
||||
# Mock Redis to return the initial message after it's added
|
||||
initial_message_json = store._serialize_message(sample_messages[0])
|
||||
mock_redis_client.lrange.return_value = [initial_message_json]
|
||||
|
||||
messages = await store.list_messages()
|
||||
|
||||
assert len(messages) == 1
|
||||
assert messages[0].text == "Hello"
|
||||
# Verify initial message was added to Redis via pipeline
|
||||
mock_pipeline.rpush.assert_called()
|
||||
|
||||
async def test_initial_messages_not_added_if_key_exists(self, sample_messages):
|
||||
"""Test that initial messages are not added if Redis key already has data."""
|
||||
with patch("agent_framework_redis._chat_message_store.redis.from_url") as mock_from_url:
|
||||
mock_redis_client = MagicMock()
|
||||
mock_redis_client.llen = AsyncMock(return_value=5) # Key already has messages
|
||||
mock_redis_client.lrange = AsyncMock(return_value=[])
|
||||
|
||||
# Pipeline should not be called since key already exists
|
||||
mock_pipeline = AsyncMock()
|
||||
mock_pipeline.rpush = AsyncMock()
|
||||
mock_pipeline.execute = AsyncMock()
|
||||
mock_redis_client.pipeline.return_value.__aenter__.return_value = mock_pipeline
|
||||
|
||||
mock_from_url.return_value = mock_redis_client
|
||||
|
||||
store = RedisChatMessageStore(
|
||||
redis_url="redis://localhost:6379",
|
||||
thread_id="test123",
|
||||
messages=sample_messages[:1], # One initial message
|
||||
)
|
||||
store._redis_client = mock_redis_client
|
||||
|
||||
await store.list_messages()
|
||||
|
||||
# Should check length but not add messages since key exists
|
||||
mock_redis_client.llen.assert_called()
|
||||
mock_pipeline.rpush.assert_not_called()
|
||||
|
||||
async def test_serialize_state(self, redis_store):
|
||||
"""Test state serialization."""
|
||||
state = await redis_store.serialize_state()
|
||||
|
||||
expected_state = {
|
||||
"thread_id": "test_thread_123",
|
||||
"redis_url": "redis://localhost:6379",
|
||||
"key_prefix": "chat_messages",
|
||||
"max_messages": None,
|
||||
}
|
||||
|
||||
assert state == expected_state
|
||||
|
||||
async def test_deserialize_state(self, redis_store):
|
||||
"""Test state deserialization."""
|
||||
serialized_state = {
|
||||
"thread_id": "restored_thread_456",
|
||||
"redis_url": "redis://localhost:6380",
|
||||
"key_prefix": "restored_messages",
|
||||
"max_messages": 50,
|
||||
}
|
||||
|
||||
await redis_store.deserialize_state(serialized_state)
|
||||
|
||||
assert redis_store.thread_id == "restored_thread_456"
|
||||
assert redis_store.redis_url == "redis://localhost:6380"
|
||||
assert redis_store.key_prefix == "restored_messages"
|
||||
assert redis_store.max_messages == 50
|
||||
|
||||
async def test_deserialize_state_empty(self, redis_store):
|
||||
"""Test deserializing empty state doesn't change anything."""
|
||||
original_thread_id = redis_store.thread_id
|
||||
|
||||
await redis_store.deserialize_state(None)
|
||||
|
||||
assert redis_store.thread_id == original_thread_id
|
||||
|
||||
async def test_clear_messages(self, redis_store, mock_redis_client):
|
||||
"""Test clearing all messages."""
|
||||
await redis_store.clear()
|
||||
|
||||
mock_redis_client.delete.assert_called_once_with("chat_messages:test_thread_123")
|
||||
|
||||
async def test_message_serialization_roundtrip(self, sample_messages):
|
||||
"""Test message serialization and deserialization roundtrip."""
|
||||
with patch("agent_framework_redis._chat_message_store.redis.from_url"):
|
||||
store = RedisChatMessageStore(redis_url="redis://localhost:6379", thread_id="test123")
|
||||
|
||||
message = sample_messages[0]
|
||||
|
||||
# Test serialization
|
||||
serialized = store._serialize_message(message)
|
||||
assert isinstance(serialized, str)
|
||||
|
||||
# Test deserialization
|
||||
deserialized = store._deserialize_message(serialized)
|
||||
assert deserialized.role == message.role
|
||||
assert deserialized.text == message.text
|
||||
assert deserialized.message_id == message.message_id
|
||||
|
||||
async def test_message_serialization_with_complex_content(self):
|
||||
"""Test serialization of messages with complex content."""
|
||||
with patch("agent_framework_redis._chat_message_store.redis.from_url"):
|
||||
store = RedisChatMessageStore(redis_url="redis://localhost:6379", thread_id="test123")
|
||||
|
||||
# Message with multiple content types
|
||||
message = ChatMessage(
|
||||
role=Role.ASSISTANT,
|
||||
contents=[TextContent(text="Hello"), TextContent(text="World")],
|
||||
author_name="TestBot",
|
||||
message_id="complex_msg",
|
||||
additional_properties={"metadata": "test"},
|
||||
)
|
||||
|
||||
serialized = store._serialize_message(message)
|
||||
deserialized = store._deserialize_message(serialized)
|
||||
|
||||
assert deserialized.role == Role.ASSISTANT
|
||||
assert deserialized.text == "Hello World"
|
||||
assert deserialized.author_name == "TestBot"
|
||||
assert deserialized.message_id == "complex_msg"
|
||||
assert deserialized.additional_properties == {"metadata": "test"}
|
||||
|
||||
async def test_redis_connection_error_handling(self):
|
||||
"""Test handling Redis connection errors in add_messages."""
|
||||
with patch("agent_framework_redis._chat_message_store.redis.from_url") as mock_from_url:
|
||||
mock_client = MagicMock()
|
||||
|
||||
# Mock pipeline to raise exception during execution
|
||||
mock_pipeline = AsyncMock()
|
||||
mock_pipeline.rpush = AsyncMock()
|
||||
mock_pipeline.execute = AsyncMock(side_effect=Exception("Connection failed"))
|
||||
mock_client.pipeline.return_value.__aenter__.return_value = mock_pipeline
|
||||
|
||||
mock_from_url.return_value = mock_client
|
||||
|
||||
store = RedisChatMessageStore(redis_url="redis://localhost:6379", thread_id="test123")
|
||||
store._redis_client = mock_client
|
||||
|
||||
message = ChatMessage(role=Role.USER, text="Test")
|
||||
|
||||
# Should propagate Redis connection errors
|
||||
with pytest.raises(Exception, match="Connection failed"):
|
||||
await store.add_messages([message])
|
||||
|
||||
async def test_getitem(self, redis_store, mock_redis_client, sample_messages):
|
||||
"""Test getitem method using Redis LINDEX."""
|
||||
# Mock LINDEX to return specific messages
|
||||
serialized_msg0 = redis_store._serialize_message(sample_messages[0])
|
||||
serialized_msg1 = redis_store._serialize_message(sample_messages[1])
|
||||
|
||||
def mock_lindex(key, index):
|
||||
if index == 0:
|
||||
return serialized_msg0
|
||||
if index == -1 or index == 1:
|
||||
return serialized_msg1
|
||||
return None
|
||||
|
||||
mock_redis_client.lindex = AsyncMock(side_effect=mock_lindex)
|
||||
|
||||
# Test positive index
|
||||
message = await redis_store.getitem(0)
|
||||
assert message.text == "Hello"
|
||||
|
||||
# Test negative index
|
||||
message = await redis_store.getitem(-1)
|
||||
assert message.text == "Hi there!"
|
||||
|
||||
async def test_getitem_index_error(self, redis_store, mock_redis_client):
|
||||
"""Test getitem raises IndexError for invalid index."""
|
||||
mock_redis_client.lindex = AsyncMock(return_value=None)
|
||||
|
||||
with pytest.raises(IndexError):
|
||||
await redis_store.getitem(0)
|
||||
|
||||
async def test_setitem(self, redis_store, mock_redis_client, sample_messages):
|
||||
"""Test setitem method using Redis LSET."""
|
||||
mock_redis_client.llen.return_value = 2
|
||||
mock_redis_client.lset = AsyncMock()
|
||||
|
||||
new_message = ChatMessage(role=Role.USER, text="Updated message")
|
||||
await redis_store.setitem(0, new_message)
|
||||
|
||||
mock_redis_client.lset.assert_called_once()
|
||||
call_args = mock_redis_client.lset.call_args
|
||||
assert call_args[0][0] == "chat_messages:test_thread_123"
|
||||
assert call_args[0][1] == 0
|
||||
|
||||
async def test_setitem_index_error(self, redis_store, mock_redis_client):
|
||||
"""Test setitem raises IndexError for invalid index."""
|
||||
mock_redis_client.llen.return_value = 0
|
||||
|
||||
new_message = ChatMessage(role=Role.USER, text="Test")
|
||||
with pytest.raises(IndexError):
|
||||
await redis_store.setitem(0, new_message)
|
||||
|
||||
async def test_append(self, redis_store, mock_redis_client):
|
||||
"""Test append method delegates to add_messages."""
|
||||
message = ChatMessage(role=Role.USER, text="Appended message")
|
||||
await redis_store.append(message)
|
||||
|
||||
# Should call pipeline operations via add_messages
|
||||
mock_redis_client.pipeline.assert_called_with(transaction=True)
|
||||
|
||||
# Verify the message was added via pipeline
|
||||
pipeline_mock = mock_redis_client.pipeline.return_value.__aenter__.return_value
|
||||
pipeline_mock.rpush.assert_called()
|
||||
pipeline_mock.execute.assert_called()
|
||||
|
||||
async def test_count(self, redis_store, mock_redis_client):
|
||||
"""Test count method."""
|
||||
mock_redis_client.llen.return_value = 5
|
||||
|
||||
count = await redis_store.count()
|
||||
|
||||
assert count == 5
|
||||
mock_redis_client.llen.assert_called_with("chat_messages:test_thread_123")
|
||||
|
||||
async def test_len_method(self, redis_store, mock_redis_client):
|
||||
"""Test async __len__ method."""
|
||||
mock_redis_client.llen.return_value = 3
|
||||
|
||||
length = await redis_store.__len__()
|
||||
|
||||
assert length == 3
|
||||
mock_redis_client.llen.assert_called_with("chat_messages:test_thread_123")
|
||||
|
||||
def test_bool_method(self, redis_store):
|
||||
"""Test __bool__ method always returns True."""
|
||||
# Store should always be truthy
|
||||
assert bool(redis_store) is True
|
||||
assert redis_store.__bool__() is True
|
||||
|
||||
# Should work in if statements (this is what Agent Framework uses)
|
||||
if redis_store:
|
||||
assert True # Should reach this
|
||||
else:
|
||||
raise AssertionError("Store should be truthy")
|
||||
|
||||
async def test_index_found(self, redis_store, mock_redis_client, sample_messages):
|
||||
"""Test index method when message is found using Redis LINDEX."""
|
||||
mock_redis_client.llen.return_value = 2
|
||||
|
||||
# Mock LINDEX to return messages at each position
|
||||
serialized_msg0 = redis_store._serialize_message(sample_messages[0])
|
||||
serialized_msg1 = redis_store._serialize_message(sample_messages[1])
|
||||
|
||||
def mock_lindex(key, index):
|
||||
if index == 0:
|
||||
return serialized_msg0
|
||||
if index == 1:
|
||||
return serialized_msg1
|
||||
return None
|
||||
|
||||
mock_redis_client.lindex = AsyncMock(side_effect=mock_lindex)
|
||||
|
||||
index = await redis_store.index(sample_messages[1])
|
||||
assert index == 1
|
||||
|
||||
# Should have called lindex twice (index 0, then index 1)
|
||||
assert mock_redis_client.lindex.call_count == 2
|
||||
|
||||
async def test_index_not_found(self, redis_store, mock_redis_client, sample_messages):
|
||||
"""Test index method when message is not found."""
|
||||
mock_redis_client.llen.return_value = 1
|
||||
mock_redis_client.lindex = AsyncMock(return_value="different_message")
|
||||
|
||||
with pytest.raises(ValueError, match="ChatMessage not found in store"):
|
||||
await redis_store.index(sample_messages[0])
|
||||
|
||||
async def test_remove(self, redis_store, mock_redis_client, sample_messages):
|
||||
"""Test remove method using Redis LREM."""
|
||||
mock_redis_client.lrem = AsyncMock(return_value=1) # 1 element removed
|
||||
|
||||
await redis_store.remove(sample_messages[0])
|
||||
|
||||
# Should use LREM to remove the message
|
||||
expected_serialized = redis_store._serialize_message(sample_messages[0])
|
||||
mock_redis_client.lrem.assert_called_once_with("chat_messages:test_thread_123", 1, expected_serialized)
|
||||
|
||||
async def test_remove_not_found(self, redis_store, mock_redis_client, sample_messages):
|
||||
"""Test remove method when message is not found."""
|
||||
mock_redis_client.lrem = AsyncMock(return_value=0) # 0 elements removed
|
||||
|
||||
with pytest.raises(ValueError, match="ChatMessage not found in store"):
|
||||
await redis_store.remove(sample_messages[0])
|
||||
|
||||
async def test_extend(self, redis_store, mock_redis_client, sample_messages):
|
||||
"""Test extend method delegates to add_messages."""
|
||||
await redis_store.extend(sample_messages[:2])
|
||||
|
||||
# Should call pipeline operations via add_messages
|
||||
mock_redis_client.pipeline.assert_called_with(transaction=True)
|
||||
|
||||
# Verify rpush was called for each message
|
||||
pipeline_mock = mock_redis_client.pipeline.return_value.__aenter__.return_value
|
||||
assert pipeline_mock.rpush.call_count >= 2
|
||||
@@ -0,0 +1,28 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
|
||||
def test_self_through_main() -> None:
|
||||
try:
|
||||
from agent_framework.redis import __version__
|
||||
except ImportError:
|
||||
__version__ = None
|
||||
|
||||
assert __version__ is not None
|
||||
|
||||
|
||||
def test_self() -> None:
|
||||
try:
|
||||
from agent_framework_redis import __version__
|
||||
except ImportError:
|
||||
__version__ = None
|
||||
|
||||
assert __version__ is not None
|
||||
|
||||
|
||||
def test_agent_framework() -> None:
|
||||
try:
|
||||
from agent_framework import __version__
|
||||
except ImportError:
|
||||
__version__ = None
|
||||
|
||||
assert __version__ is not None
|
||||
@@ -0,0 +1,448 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
from agent_framework import ChatMessage, Role
|
||||
from agent_framework.exceptions import ServiceInitializationError
|
||||
from pydantic import ValidationError
|
||||
from redisvl.utils.vectorize import CustomTextVectorizer
|
||||
|
||||
from agent_framework_redis import RedisProvider
|
||||
|
||||
CUSTOM_VECTORIZER = CustomTextVectorizer(embed=lambda x: [1.0, 2.0, 3.0], dtype="float32")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_index() -> AsyncMock:
|
||||
idx = AsyncMock()
|
||||
idx.create = AsyncMock()
|
||||
idx.load = AsyncMock()
|
||||
idx.query = AsyncMock()
|
||||
idx.exists = AsyncMock(return_value=False)
|
||||
|
||||
async def _paginate_generator(*_args: Any, **_kwargs: Any):
|
||||
# Default empty generator; override per-test as needed
|
||||
if False: # pragma: no cover
|
||||
yield []
|
||||
return
|
||||
|
||||
idx.paginate = _paginate_generator
|
||||
return idx
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def patch_index_from_dict(mock_index: AsyncMock):
|
||||
with patch("agent_framework_redis._provider.AsyncSearchIndex") as mock_cls:
|
||||
mock_cls.from_dict = MagicMock(return_value=mock_index)
|
||||
|
||||
# Mock from_existing to return a mock with matching schema by default
|
||||
# This prevents schema validation errors in tests that don't specifically test schema validation
|
||||
async def mock_from_existing(index_name, redis_url):
|
||||
mock_existing = AsyncMock()
|
||||
# Return a schema that will match whatever the provider generates
|
||||
# This is a bit of a hack, but allows existing tests to continue working
|
||||
mock_existing.schema.to_dict = MagicMock(
|
||||
side_effect=lambda: mock_cls.from_dict.call_args[0][0] if mock_cls.from_dict.call_args else {}
|
||||
)
|
||||
return mock_existing
|
||||
|
||||
mock_cls.from_existing = AsyncMock(side_effect=mock_from_existing)
|
||||
|
||||
yield mock_cls
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def patch_queries():
|
||||
calls: dict[str, Any] = {"TextQuery": [], "HybridQuery": [], "FilterExpression": []}
|
||||
|
||||
def _mk_query(kind: str):
|
||||
class _Q: # simple marker object with captured kwargs
|
||||
def __init__(self, **kwargs):
|
||||
self.kind = kind
|
||||
self.kwargs = kwargs
|
||||
|
||||
return _Q
|
||||
|
||||
with (
|
||||
patch(
|
||||
"agent_framework_redis._provider.TextQuery",
|
||||
side_effect=lambda **k: calls["TextQuery"].append(k) or _mk_query("text")(**k),
|
||||
) as text_q,
|
||||
patch(
|
||||
"agent_framework_redis._provider.HybridQuery",
|
||||
side_effect=lambda **k: calls["HybridQuery"].append(k) or _mk_query("hybrid")(**k),
|
||||
) as hybrid_q,
|
||||
patch(
|
||||
"agent_framework_redis._provider.FilterExpression",
|
||||
side_effect=lambda s: calls["FilterExpression"].append(s) or ("FE", s),
|
||||
) as filt,
|
||||
):
|
||||
yield {"calls": calls, "TextQuery": text_q, "HybridQuery": hybrid_q, "FilterExpression": filt}
|
||||
|
||||
|
||||
class TestRedisProviderInitialization:
|
||||
# Verifies the provider can be imported from the package
|
||||
def test_import(self):
|
||||
from agent_framework_redis._provider import RedisProvider
|
||||
|
||||
assert RedisProvider is not None
|
||||
|
||||
# Constructing without filters should not raise; filters are enforced at call-time
|
||||
def test_init_without_filters_ok(self, patch_index_from_dict): # noqa: ARG002
|
||||
provider = RedisProvider()
|
||||
assert provider.user_id is None
|
||||
assert provider.agent_id is None
|
||||
assert provider.application_id is None
|
||||
assert provider.thread_id is None
|
||||
|
||||
# Schema should omit vector field when no vector configuration is provided
|
||||
def test_schema_without_vector_field(self, patch_index_from_dict):
|
||||
RedisProvider(user_id="u1")
|
||||
# Inspect schema passed to from_dict
|
||||
args, kwargs = patch_index_from_dict.from_dict.call_args
|
||||
schema = args[0]
|
||||
assert isinstance(schema, dict)
|
||||
names = [f["name"] for f in schema["fields"]]
|
||||
types = [f["type"] for f in schema["fields"]]
|
||||
assert "content" in names
|
||||
assert "text" in types
|
||||
assert "vector" not in types
|
||||
|
||||
|
||||
class TestRedisProviderMessages:
|
||||
@pytest.fixture
|
||||
def sample_messages(self) -> list[ChatMessage]:
|
||||
return [
|
||||
ChatMessage(role=Role.USER, text="Hello, how are you?"),
|
||||
ChatMessage(role=Role.ASSISTANT, text="I'm doing well, thank you!"),
|
||||
ChatMessage(role=Role.SYSTEM, text="You are a helpful assistant"),
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
# Writes require at least one scoping filter to avoid unbounded operations
|
||||
async def test_messages_adding_requires_filters(self, patch_index_from_dict): # noqa: ARG002
|
||||
provider = RedisProvider()
|
||||
with pytest.raises(ServiceInitializationError):
|
||||
await provider.messages_adding("thread123", ChatMessage(role=Role.USER, text="Hello"))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
# Captures the per-operation thread id when provided
|
||||
async def test_thread_created_sets_per_operation_id(self, patch_index_from_dict): # noqa: ARG002
|
||||
provider = RedisProvider(user_id="u1")
|
||||
await provider.thread_created("t1")
|
||||
assert provider._per_operation_thread_id == "t1"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
# Enforces single-thread usage when scope_to_per_operation_thread_id is True
|
||||
async def test_thread_created_conflict_when_scoped(self, patch_index_from_dict): # noqa: ARG002
|
||||
provider = RedisProvider(user_id="u1", scope_to_per_operation_thread_id=True)
|
||||
provider._per_operation_thread_id = "t1"
|
||||
with pytest.raises(ValueError) as exc:
|
||||
await provider.thread_created("t2")
|
||||
assert "only be used with one thread" in str(exc.value)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
# Aggregates all results from the async paginator into a flat list
|
||||
async def test_search_all_paginates(self, mock_index: AsyncMock, patch_index_from_dict): # noqa: ARG002
|
||||
async def gen(_q, page_size: int = 200): # noqa: ARG001, ANN001
|
||||
yield [{"id": 1}]
|
||||
yield [{"id": 2}, {"id": 3}]
|
||||
|
||||
mock_index.paginate = gen
|
||||
provider = RedisProvider(user_id="u1")
|
||||
res = await provider.search_all(page_size=2)
|
||||
assert res == [{"id": 1}, {"id": 2}, {"id": 3}]
|
||||
|
||||
|
||||
class TestRedisProviderModelInvoking:
|
||||
@pytest.mark.asyncio
|
||||
# Reads require at least one scoping filter to avoid unbounded operations
|
||||
async def test_model_invoking_requires_filters(self, patch_index_from_dict): # noqa: ARG002
|
||||
provider = RedisProvider()
|
||||
with pytest.raises(ServiceInitializationError):
|
||||
await provider.model_invoking(ChatMessage(role=Role.USER, text="Hi"))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
# Ensures text-only search path is used and context is composed from hits
|
||||
async def test_textquery_path_and_context_contents(
|
||||
self, mock_index: AsyncMock, patch_index_from_dict, patch_queries
|
||||
): # noqa: ARG002
|
||||
# Arrange: text-only search
|
||||
mock_index.query = AsyncMock(return_value=[{"content": "A"}, {"content": "B"}])
|
||||
provider = RedisProvider(user_id="u1")
|
||||
|
||||
# Act
|
||||
ctx = await provider.model_invoking([ChatMessage(role=Role.USER, text="q1")])
|
||||
|
||||
# Assert: TextQuery used (not HybridQuery), filter_expression included
|
||||
assert patch_queries["TextQuery"].call_count == 1
|
||||
assert patch_queries["HybridQuery"].call_count == 0
|
||||
kwargs = patch_queries["calls"]["TextQuery"][0]
|
||||
assert kwargs["text"] == "q1"
|
||||
assert kwargs["text_field_name"] == "content"
|
||||
assert kwargs["num_results"] == 10
|
||||
assert "filter_expression" in kwargs
|
||||
|
||||
# Context contains memories joined after the default prompt
|
||||
assert ctx.contents is not None and len(ctx.contents) == 1
|
||||
text = ctx.contents[0].text
|
||||
assert text.endswith("A\nB")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
# When no results are returned, Context should have no contents
|
||||
async def test_model_invoking_empty_results_returns_empty_context(
|
||||
self, mock_index: AsyncMock, patch_index_from_dict, patch_queries
|
||||
): # noqa: ARG002
|
||||
mock_index.query = AsyncMock(return_value=[])
|
||||
provider = RedisProvider(user_id="u1")
|
||||
ctx = await provider.model_invoking([ChatMessage(role=Role.USER, text="any")])
|
||||
assert ctx.contents is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
# Ensures hybrid vector-text search is used when a vectorizer and vector field are configured
|
||||
async def test_hybridquery_path_with_vectorizer(self, mock_index: AsyncMock, patch_index_from_dict, patch_queries): # noqa: ARG002
|
||||
mock_index.query = AsyncMock(return_value=[{"content": "Hit"}])
|
||||
provider = RedisProvider(user_id="u1", redis_vectorizer=CUSTOM_VECTORIZER, vector_field_name="vec")
|
||||
|
||||
ctx = await provider.model_invoking([ChatMessage(role=Role.USER, text="hello")])
|
||||
|
||||
# Assert: HybridQuery used with vector and vector field
|
||||
assert patch_queries["HybridQuery"].call_count == 1
|
||||
k = patch_queries["calls"]["HybridQuery"][0]
|
||||
assert k["text"] == "hello"
|
||||
assert k["vector_field_name"] == "vec"
|
||||
assert k["vector"] == [1.0, 2.0, 3.0]
|
||||
assert k["dtype"] == "float32"
|
||||
assert k["num_results"] == 10
|
||||
assert "filter_expression" in k
|
||||
|
||||
# Context assembled from returned memories
|
||||
assert ctx.contents and "Hit" in ctx.contents[0].text
|
||||
|
||||
|
||||
class TestRedisProviderContextManager:
|
||||
@pytest.mark.asyncio
|
||||
# Verifies async context manager returns self for chaining
|
||||
async def test_async_context_manager_returns_self(self, patch_index_from_dict): # noqa: ARG002
|
||||
provider = RedisProvider(user_id="u1")
|
||||
async with provider as ctx:
|
||||
assert ctx is provider
|
||||
|
||||
@pytest.mark.asyncio
|
||||
# Exit should be a no-op and not raise
|
||||
async def test_aexit_noop(self, patch_index_from_dict): # noqa: ARG002
|
||||
provider = RedisProvider(user_id="u1")
|
||||
assert await provider.__aexit__(None, None, None) is None
|
||||
|
||||
|
||||
class TestMessagesAddingBehavior:
|
||||
@pytest.mark.asyncio
|
||||
# Adds messages while injecting partition defaults and preserving allowed roles
|
||||
async def test_messages_adding_adds_partition_defaults_and_roles(
|
||||
self, mock_index: AsyncMock, patch_index_from_dict
|
||||
): # noqa: ARG002
|
||||
provider = RedisProvider(
|
||||
application_id="app",
|
||||
agent_id="agent",
|
||||
user_id="u1",
|
||||
scope_to_per_operation_thread_id=True,
|
||||
)
|
||||
|
||||
msgs = [
|
||||
ChatMessage(role=Role.USER, text="u"),
|
||||
ChatMessage(role=Role.ASSISTANT, text="a"),
|
||||
ChatMessage(role=Role.SYSTEM, text="s"),
|
||||
]
|
||||
|
||||
await provider.messages_adding("t1", msgs)
|
||||
|
||||
# Ensure load invoked with shaped docs containing defaults
|
||||
assert mock_index.load.await_count == 1
|
||||
(loaded_args, _kwargs) = mock_index.load.call_args
|
||||
docs = loaded_args[0]
|
||||
assert isinstance(docs, list) and len(docs) == 3
|
||||
for d in docs:
|
||||
assert d["role"] in {"user", "assistant", "system"}
|
||||
assert d["content"] in {"u", "a", "s"}
|
||||
assert d["application_id"] == "app"
|
||||
assert d["agent_id"] == "agent"
|
||||
assert d["user_id"] == "u1"
|
||||
assert d["thread_id"] == "t1" # scoped via per-operation thread id
|
||||
|
||||
@pytest.mark.asyncio
|
||||
# Skips blank text and disallowed roles (e.g., TOOL) when adding messages
|
||||
async def test_messages_adding_ignores_blank_and_disallowed_roles(
|
||||
self, mock_index: AsyncMock, patch_index_from_dict
|
||||
): # noqa: ARG002
|
||||
provider = RedisProvider(user_id="u1", scope_to_per_operation_thread_id=True)
|
||||
msgs = [
|
||||
ChatMessage(role=Role.USER, text=" "),
|
||||
ChatMessage(role=Role.TOOL, text="tool output"),
|
||||
]
|
||||
await provider.messages_adding("tid", msgs)
|
||||
# No valid messages -> no load
|
||||
assert mock_index.load.await_count == 0
|
||||
|
||||
|
||||
class TestIndexCreationPublicCalls:
|
||||
@pytest.mark.asyncio
|
||||
# Ensures index is created only once when drop=True on first public write call
|
||||
async def test_messages_adding_triggers_index_create_once_when_drop_true(
|
||||
self, mock_index: AsyncMock, patch_index_from_dict
|
||||
): # noqa: ARG002
|
||||
provider = RedisProvider(user_id="u1", drop_redis_index=True)
|
||||
await provider.messages_adding("t1", ChatMessage(role=Role.USER, text="m1"))
|
||||
await provider.messages_adding("t1", ChatMessage(role=Role.USER, text="m2"))
|
||||
# create only on first call
|
||||
assert mock_index.create.await_count == 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
# Ensures index is created when drop=False and the index does not exist on first read
|
||||
async def test_model_invoking_triggers_create_when_drop_false_and_not_exists(
|
||||
self, mock_index: AsyncMock, patch_index_from_dict
|
||||
): # noqa: ARG002
|
||||
mock_index.exists = AsyncMock(return_value=False)
|
||||
provider = RedisProvider(user_id="u1", drop_redis_index=False)
|
||||
mock_index.query = AsyncMock(return_value=[{"content": "C"}])
|
||||
await provider.model_invoking([ChatMessage(role=Role.USER, text="q")])
|
||||
assert mock_index.create.await_count == 1
|
||||
|
||||
|
||||
class TestThreadCreatedAdditional:
|
||||
@pytest.mark.asyncio
|
||||
# Allows None or same thread id repeatedly; different id raises when scoped
|
||||
async def test_thread_created_allows_none_and_same_id(self, patch_index_from_dict): # noqa: ARG002
|
||||
provider = RedisProvider(user_id="u1", scope_to_per_operation_thread_id=True)
|
||||
# None is allowed
|
||||
await provider.thread_created(None)
|
||||
# Same id is allowed repeatedly
|
||||
await provider.thread_created("t1")
|
||||
await provider.thread_created("t1")
|
||||
# Different id should raise
|
||||
with pytest.raises(ValueError):
|
||||
await provider.thread_created("t2")
|
||||
|
||||
|
||||
class TestVectorPopulation:
|
||||
@pytest.mark.asyncio
|
||||
# When vectorizer configured, messages_adding should embed content and populate the vector field
|
||||
async def test_messages_adding_populates_vector_field_when_vectorizer_present(
|
||||
self, mock_index: AsyncMock, patch_index_from_dict
|
||||
): # noqa: ARG002
|
||||
provider = RedisProvider(
|
||||
user_id="u1",
|
||||
scope_to_per_operation_thread_id=True,
|
||||
redis_vectorizer=CUSTOM_VECTORIZER,
|
||||
vector_field_name="vec",
|
||||
)
|
||||
|
||||
await provider.messages_adding("t1", ChatMessage(role=Role.USER, text="hello"))
|
||||
assert mock_index.load.await_count == 1
|
||||
(loaded_args, _kwargs) = mock_index.load.call_args
|
||||
docs = loaded_args[0]
|
||||
assert isinstance(docs, list) and len(docs) == 1
|
||||
vec = docs[0].get("vec")
|
||||
assert isinstance(vec, (bytes, bytearray))
|
||||
assert len(vec) == 3 * np.dtype(np.float32).itemsize
|
||||
|
||||
|
||||
class TestRedisProviderSchemaVectors:
|
||||
# Adds a vector field when vectorizer supplies dims implicitly
|
||||
def test_schema_with_vector_field_and_dims_inferred(self, patch_index_from_dict): # noqa: ARG002
|
||||
RedisProvider(user_id="u1", redis_vectorizer=CUSTOM_VECTORIZER, vector_field_name="vec")
|
||||
args, _ = patch_index_from_dict.from_dict.call_args
|
||||
schema = args[0]
|
||||
names = [f["name"] for f in schema["fields"]]
|
||||
types = {f["name"]: f["type"] for f in schema["fields"]}
|
||||
assert "vec" in names
|
||||
assert types["vec"] == "vector"
|
||||
|
||||
# Raises when redis_vectorizer is not the correct type
|
||||
def test_init_invalid_vectorizer(self, patch_index_from_dict): # noqa: ARG002
|
||||
class DummyVectorizer:
|
||||
pass
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
RedisProvider(user_id="u1", redis_vectorizer=DummyVectorizer(), vector_field_name="vec")
|
||||
|
||||
|
||||
class TestEnsureIndex:
|
||||
@pytest.mark.asyncio
|
||||
# Creates index once and marks _index_initialized to prevent duplicate calls
|
||||
async def test_ensure_index_creates_once(self, mock_index: AsyncMock, patch_index_from_dict): # noqa: ARG002
|
||||
# Mock index doesn't exist, so it will be created
|
||||
mock_index.exists = AsyncMock(return_value=False)
|
||||
provider = RedisProvider(user_id="u1", overwrite_index=False)
|
||||
|
||||
assert provider._index_initialized is False
|
||||
await provider._ensure_index()
|
||||
assert mock_index.create.await_count == 1
|
||||
assert provider._index_initialized is True
|
||||
|
||||
# Second call should not create again due to _index_initialized flag
|
||||
await provider._ensure_index()
|
||||
assert mock_index.create.await_count == 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
# Creates index with overwrite=True when overwrite_index=True
|
||||
async def test_ensure_index_with_overwrite_true(self, mock_index: AsyncMock, patch_index_from_dict): # noqa: ARG002
|
||||
mock_index.exists = AsyncMock(return_value=True)
|
||||
provider = RedisProvider(user_id="u1", overwrite_index=True)
|
||||
|
||||
await provider._ensure_index()
|
||||
|
||||
# Should call create with overwrite=True, drop=False
|
||||
mock_index.create.assert_called_once_with(overwrite=True, drop=False)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
# Creates index with overwrite=False when index doesn't exist
|
||||
async def test_ensure_index_create_if_missing(self, mock_index: AsyncMock, patch_index_from_dict): # noqa: ARG002
|
||||
mock_index.exists = AsyncMock(return_value=False)
|
||||
provider = RedisProvider(user_id="u1", overwrite_index=False)
|
||||
|
||||
await provider._ensure_index()
|
||||
|
||||
# Should call create with overwrite=False, drop=False
|
||||
mock_index.create.assert_called_once_with(overwrite=False, drop=False)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
# Validates schema compatibility when index exists and overwrite=False
|
||||
async def test_ensure_index_schema_validation_success(self, mock_index: AsyncMock, patch_index_from_dict): # noqa: ARG002
|
||||
mock_index.exists = AsyncMock(return_value=True)
|
||||
provider = RedisProvider(user_id="u1", overwrite_index=False)
|
||||
|
||||
# Mock existing index with matching schema
|
||||
expected_schema = provider.schema_dict
|
||||
patch_index_from_dict.from_existing.return_value.schema.to_dict.return_value = expected_schema
|
||||
|
||||
await provider._ensure_index()
|
||||
|
||||
# Should validate schema and proceed to create
|
||||
patch_index_from_dict.from_existing.assert_called_once_with("context", redis_url="redis://localhost:6379")
|
||||
mock_index.create.assert_called_once_with(overwrite=False, drop=False)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
# Raises ServiceInitializationError when schemas don't match
|
||||
async def test_ensure_index_schema_validation_failure(self, mock_index: AsyncMock, patch_index_from_dict): # noqa: ARG002
|
||||
mock_index.exists = AsyncMock(return_value=True)
|
||||
provider = RedisProvider(user_id="u1", overwrite_index=False)
|
||||
|
||||
# Override the mock to return a different schema after provider is created
|
||||
async def mock_from_existing_different(index_name, redis_url):
|
||||
mock_existing = AsyncMock()
|
||||
mock_existing.schema.to_dict = MagicMock(return_value={"different": "schema"})
|
||||
return mock_existing
|
||||
|
||||
patch_index_from_dict.from_existing = AsyncMock(side_effect=mock_from_existing_different)
|
||||
|
||||
with pytest.raises(ServiceInitializationError) as exc:
|
||||
await provider._ensure_index()
|
||||
|
||||
assert "incompatible with the current configuration" in str(exc.value)
|
||||
assert "overwrite_index=True" in str(exc.value)
|
||||
|
||||
# Should not call create when schema validation fails
|
||||
mock_index.create.assert_not_called()
|
||||
Reference in New Issue
Block a user