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

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

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

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

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

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

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

* refactor: rename _RunContext.thread_messages to session_messages

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

* rename: remove _new_ prefix from test files

* refactor: rewrite SlidingWindowChatMessageStore as SlidingWindowHistoryProvider(InMemoryHistoryProvider)

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

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

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

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

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

* feat: add Pydantic BaseModel support to session state serialization

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

Also export register_state_type as a public API.

* fix mem0

* Update sample README links and descriptions for session terminology

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

* Fix broken Redis README link to renamed sample

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

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

Port of fix from #3844 to new Mem0ContextProvider.

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

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

* Add STORES_BY_DEFAULT ClassVar to skip redundant InMemoryHistoryProvider injection

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

* Fix broken markdown links in azure_ai and redis READMEs

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

* updates to workflow as agent

* fix group chat import

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

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

* Fix broken markdown links after thread→session file renames

* fix azure ai test
This commit is contained in:
Eduard van Valkenburg
2026-02-12 22:00:32 +01:00
committed by GitHub
Unverified
parent 0c67dbbce5
commit 1e350ea22f
312 changed files with 6669 additions and 11423 deletions
@@ -1,6 +1,6 @@
# Copyright (c) Microsoft. All rights reserved.
"""Tests for _RedisContextProvider and _RedisHistoryProvider."""
"""Tests for RedisContextProvider and RedisHistoryProvider."""
from __future__ import annotations
@@ -12,8 +12,8 @@ from agent_framework import AgentResponse, Message
from agent_framework._sessions import AgentSession, SessionContext
from agent_framework.exceptions import ServiceInitializationError
from agent_framework_redis._context_provider import _RedisContextProvider
from agent_framework_redis._history_provider import _RedisHistoryProvider
from agent_framework_redis._context_provider import RedisContextProvider
from agent_framework_redis._history_provider import RedisHistoryProvider
# ---------------------------------------------------------------------------
# Shared fixtures
@@ -63,13 +63,13 @@ def mock_redis_client():
# ===========================================================================
# _RedisContextProvider tests
# RedisContextProvider tests
# ===========================================================================
class TestRedisContextProviderInit:
def test_basic_construction(self, patch_index_from_dict: MagicMock): # noqa: ARG002
provider = _RedisContextProvider(source_id="ctx", user_id="u1")
provider = RedisContextProvider(source_id="ctx", user_id="u1")
assert provider.source_id == "ctx"
assert provider.user_id == "u1"
assert provider.redis_url == "redis://localhost:6379"
@@ -77,7 +77,7 @@ class TestRedisContextProviderInit:
assert provider.prefix == "context"
def test_custom_params(self, patch_index_from_dict: MagicMock): # noqa: ARG002
provider = _RedisContextProvider(
provider = RedisContextProvider(
source_id="ctx",
redis_url="redis://custom:6380",
index_name="my_idx",
@@ -95,31 +95,31 @@ class TestRedisContextProviderInit:
assert provider.context_prompt == "Custom prompt"
def test_default_context_prompt(self, patch_index_from_dict: MagicMock): # noqa: ARG002
provider = _RedisContextProvider(source_id="ctx", user_id="u1")
provider = RedisContextProvider(source_id="ctx", user_id="u1")
assert "Memories" in provider.context_prompt
def test_invalid_vectorizer_raises(self, patch_index_from_dict: MagicMock): # noqa: ARG002
from agent_framework.exceptions import AgentException
with pytest.raises(AgentException, match="not a valid type"):
_RedisContextProvider(source_id="ctx", user_id="u1", redis_vectorizer="bad") # type: ignore[arg-type]
RedisContextProvider(source_id="ctx", user_id="u1", redis_vectorizer="bad") # type: ignore[arg-type]
class TestRedisContextProviderValidateFilters:
def test_no_filters_raises(self, patch_index_from_dict: MagicMock): # noqa: ARG002
provider = _RedisContextProvider(source_id="ctx")
provider = RedisContextProvider(source_id="ctx")
with pytest.raises(ServiceInitializationError, match="(?i)at least one"):
provider._validate_filters()
def test_any_single_filter_ok(self, patch_index_from_dict: MagicMock): # noqa: ARG002
for kwargs in [{"user_id": "u"}, {"agent_id": "a"}, {"application_id": "app"}]:
provider = _RedisContextProvider(source_id="ctx", **kwargs)
provider = RedisContextProvider(source_id="ctx", **kwargs)
provider._validate_filters() # should not raise
class TestRedisContextProviderSchema:
def test_schema_has_expected_fields(self, patch_index_from_dict: MagicMock): # noqa: ARG002
provider = _RedisContextProvider(source_id="ctx", user_id="u1")
provider = RedisContextProvider(source_id="ctx", user_id="u1")
schema = provider.schema_dict
field_names = [f["name"] for f in schema["fields"]]
for expected in ("role", "content", "conversation_id", "message_id", "application_id", "agent_id", "user_id"):
@@ -128,7 +128,7 @@ class TestRedisContextProviderSchema:
assert schema["index"]["prefix"] == "context"
def test_schema_no_vector_without_vectorizer(self, patch_index_from_dict: MagicMock): # noqa: ARG002
provider = _RedisContextProvider(source_id="ctx", user_id="u1")
provider = RedisContextProvider(source_id="ctx", user_id="u1")
field_types = [f["type"] for f in provider.schema_dict["fields"]]
assert "vector" not in field_types
@@ -140,7 +140,7 @@ class TestRedisContextProviderBeforeRun:
patch_index_from_dict: MagicMock, # noqa: ARG002
):
mock_index.query = AsyncMock(return_value=[{"content": "Memory A"}, {"content": "Memory B"}])
provider = _RedisContextProvider(source_id="ctx", user_id="u1")
provider = RedisContextProvider(source_id="ctx", user_id="u1")
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[Message(role="user", contents=["test query"])], session_id="s1")
@@ -157,7 +157,7 @@ class TestRedisContextProviderBeforeRun:
mock_index: AsyncMock,
patch_index_from_dict: MagicMock, # noqa: ARG002
):
provider = _RedisContextProvider(source_id="ctx", user_id="u1")
provider = RedisContextProvider(source_id="ctx", user_id="u1")
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[Message(role="user", contents=[" "])], session_id="s1")
@@ -172,7 +172,7 @@ class TestRedisContextProviderBeforeRun:
patch_index_from_dict: MagicMock, # noqa: ARG002
):
mock_index.query = AsyncMock(return_value=[])
provider = _RedisContextProvider(source_id="ctx", user_id="u1")
provider = RedisContextProvider(source_id="ctx", user_id="u1")
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[Message(role="user", contents=["hello"])], session_id="s1")
@@ -187,7 +187,7 @@ class TestRedisContextProviderAfterRun:
mock_index: AsyncMock,
patch_index_from_dict: MagicMock, # noqa: ARG002
):
provider = _RedisContextProvider(source_id="ctx", user_id="u1")
provider = RedisContextProvider(source_id="ctx", user_id="u1")
session = AgentSession(session_id="test-session")
response = AgentResponse(messages=[Message(role="assistant", contents=["response text"])])
ctx = SessionContext(input_messages=[Message(role="user", contents=["user input"])], session_id="s1")
@@ -206,7 +206,7 @@ class TestRedisContextProviderAfterRun:
mock_index: AsyncMock,
patch_index_from_dict: MagicMock, # noqa: ARG002
):
provider = _RedisContextProvider(source_id="ctx", user_id="u1")
provider = RedisContextProvider(source_id="ctx", user_id="u1")
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[Message(role="user", contents=[" "])], session_id="s1")
@@ -219,7 +219,7 @@ class TestRedisContextProviderAfterRun:
mock_index: AsyncMock,
patch_index_from_dict: MagicMock, # noqa: ARG002
):
provider = _RedisContextProvider(source_id="ctx", application_id="app", agent_id="ag", user_id="u1")
provider = RedisContextProvider(source_id="ctx", application_id="app", agent_id="ag", user_id="u1")
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[Message(role="user", contents=["hello"])], session_id="s1")
@@ -235,13 +235,13 @@ class TestRedisContextProviderAfterRun:
class TestRedisContextProviderContextManager:
async def test_aenter_returns_self(self, patch_index_from_dict: MagicMock): # noqa: ARG002
provider = _RedisContextProvider(source_id="ctx", user_id="u1")
provider = RedisContextProvider(source_id="ctx", user_id="u1")
async with provider as p:
assert p is provider
# ===========================================================================
# _RedisHistoryProvider tests
# RedisHistoryProvider tests
# ===========================================================================
@@ -249,7 +249,7 @@ class TestRedisHistoryProviderInit:
def test_basic_construction(self, mock_redis_client: MagicMock):
with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
provider = _RedisHistoryProvider("memory", redis_url="redis://localhost:6379")
provider = RedisHistoryProvider("memory", redis_url="redis://localhost:6379")
assert provider.source_id == "memory"
assert provider.key_prefix == "chat_messages"
@@ -261,7 +261,7 @@ class TestRedisHistoryProviderInit:
def test_custom_params(self, mock_redis_client: MagicMock):
with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
provider = _RedisHistoryProvider(
provider = RedisHistoryProvider(
"mem",
redis_url="redis://localhost:6379",
key_prefix="custom",
@@ -279,12 +279,12 @@ class TestRedisHistoryProviderInit:
def test_no_redis_url_or_credential_raises(self):
with pytest.raises(ValueError, match="Either redis_url or credential_provider must be provided"):
_RedisHistoryProvider("mem")
RedisHistoryProvider("mem")
def test_both_url_and_credential_raises(self):
mock_cred = MagicMock()
with pytest.raises(ValueError, match="mutually exclusive"):
_RedisHistoryProvider(
RedisHistoryProvider(
"mem",
redis_url="redis://localhost:6379",
credential_provider=mock_cred,
@@ -294,13 +294,13 @@ class TestRedisHistoryProviderInit:
def test_credential_provider_without_host_raises(self):
mock_cred = MagicMock()
with pytest.raises(ValueError, match="host is required"):
_RedisHistoryProvider("mem", credential_provider=mock_cred)
RedisHistoryProvider("mem", credential_provider=mock_cred)
def test_credential_provider_with_host(self):
mock_cred = MagicMock()
with patch("agent_framework_redis._history_provider.redis.Redis") as mock_redis_cls:
mock_redis_cls.return_value = MagicMock()
provider = _RedisHistoryProvider("mem", credential_provider=mock_cred, host="myhost")
provider = RedisHistoryProvider("mem", credential_provider=mock_cred, host="myhost")
mock_redis_cls.assert_called_once_with(
host="myhost",
@@ -317,7 +317,7 @@ class TestRedisHistoryProviderRedisKey:
def test_key_format(self, mock_redis_client: MagicMock):
with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
provider = _RedisHistoryProvider("mem", redis_url="redis://localhost:6379", key_prefix="msgs")
provider = RedisHistoryProvider("mem", redis_url="redis://localhost:6379", key_prefix="msgs")
assert provider._redis_key("session-123") == "msgs:session-123"
assert provider._redis_key(None) == "msgs:default"
@@ -331,7 +331,7 @@ class TestRedisHistoryProviderGetMessages:
with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
provider = _RedisHistoryProvider("mem", redis_url="redis://localhost:6379")
provider = RedisHistoryProvider("mem", redis_url="redis://localhost:6379")
messages = await provider.get_messages("s1")
assert len(messages) == 2
@@ -345,7 +345,7 @@ class TestRedisHistoryProviderGetMessages:
with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
provider = _RedisHistoryProvider("mem", redis_url="redis://localhost:6379")
provider = RedisHistoryProvider("mem", redis_url="redis://localhost:6379")
messages = await provider.get_messages("s1")
assert messages == []
@@ -355,7 +355,7 @@ class TestRedisHistoryProviderSaveMessages:
async def test_saves_serialized_messages(self, mock_redis_client: MagicMock):
with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
provider = _RedisHistoryProvider("mem", redis_url="redis://localhost:6379")
provider = RedisHistoryProvider("mem", redis_url="redis://localhost:6379")
msgs = [Message(role="user", contents=["Hello"]), Message(role="assistant", contents=["Hi"])]
await provider.save_messages("s1", msgs)
@@ -367,7 +367,7 @@ class TestRedisHistoryProviderSaveMessages:
async def test_empty_messages_noop(self, mock_redis_client: MagicMock):
with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
provider = _RedisHistoryProvider("mem", redis_url="redis://localhost:6379")
provider = RedisHistoryProvider("mem", redis_url="redis://localhost:6379")
await provider.save_messages("s1", [])
mock_redis_client.pipeline.assert_not_called()
@@ -377,7 +377,7 @@ class TestRedisHistoryProviderSaveMessages:
with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
provider = _RedisHistoryProvider("mem", redis_url="redis://localhost:6379", max_messages=10)
provider = RedisHistoryProvider("mem", redis_url="redis://localhost:6379", max_messages=10)
await provider.save_messages("s1", [Message(role="user", contents=["msg"])])
@@ -388,7 +388,7 @@ class TestRedisHistoryProviderSaveMessages:
with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
provider = _RedisHistoryProvider("mem", redis_url="redis://localhost:6379", max_messages=10)
provider = RedisHistoryProvider("mem", redis_url="redis://localhost:6379", max_messages=10)
await provider.save_messages("s1", [Message(role="user", contents=["msg"])])
@@ -399,7 +399,7 @@ class TestRedisHistoryProviderClear:
async def test_clear_calls_delete(self, mock_redis_client: MagicMock):
with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
provider = _RedisHistoryProvider("mem", redis_url="redis://localhost:6379")
provider = RedisHistoryProvider("mem", redis_url="redis://localhost:6379")
await provider.clear("session-1")
mock_redis_client.delete.assert_called_once_with("chat_messages:session-1")
@@ -414,7 +414,7 @@ class TestRedisHistoryProviderBeforeAfterRun:
with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
provider = _RedisHistoryProvider("mem", redis_url="redis://localhost:6379")
provider = RedisHistoryProvider("mem", redis_url="redis://localhost:6379")
session = AgentSession(session_id="test")
ctx = SessionContext(input_messages=[Message(role="user", contents=["new msg"])], session_id="s1")
@@ -428,7 +428,7 @@ class TestRedisHistoryProviderBeforeAfterRun:
async def test_after_run_stores_input_and_response(self, mock_redis_client: MagicMock):
with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
provider = _RedisHistoryProvider("mem", redis_url="redis://localhost:6379")
provider = RedisHistoryProvider("mem", redis_url="redis://localhost:6379")
session = AgentSession(session_id="test")
ctx = SessionContext(input_messages=[Message(role="user", contents=["hi"])], session_id="s1")
@@ -443,7 +443,7 @@ class TestRedisHistoryProviderBeforeAfterRun:
async def test_after_run_skips_when_no_messages(self, mock_redis_client: MagicMock):
with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
provider = _RedisHistoryProvider(
provider = RedisHistoryProvider(
"mem", redis_url="redis://localhost:6379", store_inputs=False, store_outputs=False
)
@@ -1,621 +0,0 @@
# Copyright (c) Microsoft. All rights reserved.
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from agent_framework import Content, Message
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 [
Message(role="user", text="Hello", message_id="msg1"),
Message(role="assistant", text="Hi there!", message_id="msg2"),
Message(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 either redis_url or credential_provider is required."""
with pytest.raises(ValueError, match="Either redis_url or credential_provider must be provided"):
RedisChatMessageStore(thread_id="test123")
def test_init_with_credential_provider(self):
"""Test initialization with credential_provider."""
mock_credential_provider = MagicMock()
with patch("agent_framework_redis._chat_message_store.redis.Redis") as mock_redis_class:
mock_redis_instance = MagicMock()
mock_redis_class.return_value = mock_redis_instance
store = RedisChatMessageStore(
credential_provider=mock_credential_provider,
host="myredis.redis.cache.windows.net",
thread_id="test123",
)
# Verify Redis.Redis was called with correct parameters
mock_redis_class.assert_called_once_with(
host="myredis.redis.cache.windows.net",
port=6380,
ssl=True,
username=None,
credential_provider=mock_credential_provider,
decode_responses=True,
)
# Verify store instance is properly initialized
assert store.thread_id == "test123"
assert store.redis_url is None # Should be None for credential provider auth
assert store.key_prefix == "chat_messages"
assert store.max_messages is None
def test_init_with_credential_provider_custom_port(self):
"""Test initialization with credential_provider and custom port."""
mock_credential_provider = MagicMock()
with patch("agent_framework_redis._chat_message_store.redis.Redis") as mock_redis_class:
mock_redis_instance = MagicMock()
mock_redis_class.return_value = mock_redis_instance
store = RedisChatMessageStore(
credential_provider=mock_credential_provider,
host="myredis.redis.cache.windows.net",
port=6379,
ssl=False,
username="admin",
thread_id="test123",
)
# Verify custom parameters were passed
mock_redis_class.assert_called_once_with(
host="myredis.redis.cache.windows.net",
port=6379,
ssl=False,
username="admin",
credential_provider=mock_credential_provider,
decode_responses=True,
)
# Verify store instance is properly initialized
assert store.thread_id == "test123"
assert store.redis_url is None # Should be None for credential provider auth
assert store.key_prefix == "chat_messages"
def test_init_credential_provider_requires_host(self):
"""Test that credential_provider requires host parameter."""
mock_credential_provider = MagicMock()
with pytest.raises(ValueError, match="host is required when using credential_provider"):
RedisChatMessageStore(
credential_provider=mock_credential_provider,
thread_id="test123",
)
def test_init_mutually_exclusive_params(self):
"""Test that redis_url and credential_provider are mutually exclusive."""
mock_credential_provider = MagicMock()
with pytest.raises(ValueError, match="redis_url and credential_provider are mutually exclusive"):
RedisChatMessageStore(
redis_url="redis://localhost:6379",
credential_provider=mock_credential_provider,
host="myredis.redis.cache.windows.net",
thread_id="test123",
)
async def test_serialize_with_credential_provider(self):
"""Test that serialization works correctly with credential provider authentication."""
mock_credential_provider = MagicMock()
with patch("agent_framework_redis._chat_message_store.redis.Redis") as mock_redis_class:
mock_redis_instance = MagicMock()
mock_redis_class.return_value = mock_redis_instance
store = RedisChatMessageStore(
credential_provider=mock_credential_provider,
host="myredis.redis.cache.windows.net",
thread_id="test123",
key_prefix="custom_prefix",
max_messages=100,
)
# Serialize the store state
state = await store.serialize()
# Verify serialization includes correct values
assert state["thread_id"] == "test123"
assert state["redis_url"] is None # Should be None for credential provider auth
assert state["key_prefix"] == "custom_prefix"
assert state["max_messages"] == 100
assert state["type"] == "redis_store_state"
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 = Message(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 = [
Message(role="user", text="Hello", message_id="msg1"),
Message(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 == "user"
assert messages[0].text == "Hello"
assert messages[1].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()
expected_state = {
"type": "redis_store_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.update_from_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.update_from_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 = Message(
role="assistant",
contents=[Content.from_text(text="Hello"), Content.from_text(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 == "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 = Message(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 = Message(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 = Message(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 = Message(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="Message 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="Message 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
async def test_serialize_with_agent_thread(self, redis_store, sample_messages):
"""Test that RedisChatMessageStore can be serialized within an AgentThread.
This test verifies the fix for issue #1991 where calling thread.serialize()
with a RedisChatMessageStore would fail with "Messages should be a list" error.
"""
from agent_framework import AgentThread
thread = AgentThread(message_store=redis_store)
await thread.on_new_messages(sample_messages)
serialized = await thread.serialize()
assert serialized is not None
assert "chat_message_store_state" in serialized
assert serialized["chat_message_store_state"] is not None
@@ -1,425 +0,0 @@
# 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 Message
from agent_framework.exceptions import AgentException, ServiceInitializationError
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[Message]:
return [
Message(role="user", text="Hello, how are you?"),
Message(role="assistant", text="I'm doing well, thank you!"),
Message(role="system", text="You are a helpful assistant"),
]
# 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.invoked("thread123", Message(role="user", text="Hello"))
# 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"
# 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)
# 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:
# 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.invoking(Message(role="user", text="Hi"))
# 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.invoking([Message(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.messages is not None and len(ctx.messages) == 1
text = ctx.messages[0].text
assert text.endswith("A\nB")
# 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.invoking([Message(role="user", text="any")])
assert ctx.messages == []
# 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.invoking([Message(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.messages and "Hit" in ctx.messages[0].text
class TestRedisProviderContextManager:
# 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
# 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:
# 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 = [
Message(role="user", text="u"),
Message(role="assistant", text="a"),
Message(role="system", text="s"),
]
await provider.invoked(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"
# 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 = [
Message(role="user", text=" "),
Message(role="tool", text="tool output"),
]
await provider.invoked(msgs)
# No valid messages -> no load
assert mock_index.load.await_count == 0
class TestIndexCreationPublicCalls:
# 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")
await provider.invoked(Message(role="user", text="m1"))
await provider.invoked(Message(role="user", text="m2"))
# create only on first call
assert mock_index.create.await_count == 1
# 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")
mock_index.query = AsyncMock(return_value=[{"content": "C"}])
await provider.invoking([Message(role="user", text="q")])
assert mock_index.create.await_count == 1
class TestThreadCreatedAdditional:
# 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:
# When vectorizer configured, invoked 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.invoked(Message(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(AgentException):
RedisProvider(user_id="u1", redis_vectorizer=DummyVectorizer(), vector_field_name="vec")
class TestEnsureIndex:
# 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
# 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)
# 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)
# 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)
# 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()