mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: [BREAKING] PR2 — Wire context provider pipeline, remove old types, update all consumers (#3850)
* PR2: Wire context provider pipeline and update all internal consumers - Replace AgentThread with AgentSession across all packages - Replace ContextProvider with BaseContextProvider across all packages - Replace context_provider param with context_providers (Sequence) - Replace thread= with session= in run() signatures - Replace get_new_thread() with create_session() - Add get_session(service_session_id) to agent interface - DurableAgentThread -> DurableAgentSession - Remove _notify_thread_of_new_messages from WorkflowAgent - Wire before_run/after_run context provider pipeline in RawAgent - Auto-inject InMemoryHistoryProvider when no providers configured * fix: update all tests for context provider pipeline, fix lazy-loaders, remove old test files * refactor: update all sample files for context provider pipeline (AgentThread→AgentSession, ContextProvider→BaseContextProvider) * fix: update remaining ag-ui references (client docstring, getting_started sample) * fix: make get_session service_session_id keyword-only to avoid confusion with session_id * refactor: rename _RunContext.thread_messages to session_messages * refactor: remove _threads.py, _memory.py, and old provider files; migrate devui to use plain message lists * rename: remove _new_ prefix from test files * refactor: rewrite SlidingWindowChatMessageStore as SlidingWindowHistoryProvider(InMemoryHistoryProvider) * fix: read full history from session state directly instead of reaching into provider internals * fix: update stale .pyi stubs, sample imports, and README references for new provider types * fix: remove stale message_store, _notify_thread_of_new_messages, and session_id.key references in samples * refactor: merge context_providers and sessions sample folders into sessions, remove aggregate_context_provider * refactor: UserInfoMemory stores state in session.state instead of instance attributes * feat: add Pydantic BaseModel support to session state serialization Pydantic models stored in session.state are now automatically serialized via model_dump() and restored via model_validate() during to_dict()/from_dict() round-trips. Models are auto-registered on first serialization; use register_state_type() for cold-start deserialization. Also export register_state_type as a public API. * fix mem0 * Update sample README links and descriptions for session terminology - Replace 'thread' with 'session' in sample descriptions across all READMEs - Update file links for renamed samples (mem0_sessions, redis_sessions, etc.) - Fix Threads section → Sessions section in main samples/README.md - Update tools, middleware, workflows, durabletask, azure_functions READMEs - Update architecture diagrams in concepts/tools/README.md - Update migration guides (autogen, semantic-kernel) * Fix broken Redis README link to renamed sample * Fix Mem0 OSS client search: pass scoping params as direct kwargs AsyncMemory (OSS) expects user_id/agent_id/run_id as direct kwargs, while AsyncMemoryClient (Platform) expects them in a filters dict. Adds tests for both client types. Port of fix from #3844 to new Mem0ContextProvider. * Fix rebase issues: restore missing _conversation_state.py and checkpoint decode logic - Add back _conversation_state.py (encode/decode_chat_messages) lost in rebase - Fix on_checkpoint_restore to decode cache/conversation with decode_chat_messages - Fix on_checkpoint_restore to use decode_checkpoint_value for pending requests - Add tests/workflow/__init__.py for relative import support - Fix test_agent_executor checkpoint selection (checkpoints[1] not superstep) * Add STORES_BY_DEFAULT ClassVar to skip redundant InMemoryHistoryProvider injection Chat clients that store history server-side by default (OpenAI Responses API, Azure AI Agent) now declare STORES_BY_DEFAULT = True. The agent checks this during auto-injection and skips InMemoryHistoryProvider unless the user explicitly sets store=False. * Fix broken markdown links in azure_ai and redis READMEs * Fix getting-started samples to use session API instead of removed thread/ContextProvider API * updates to workflow as agent * fix group chat import * Rename Thread→Session throughout, fix service_session_id propagation, remove stale AGUIThread - Fix: Propagate conversation_id from ChatResponse back to session.service_session_id in both streaming and non-streaming paths in _agents.py - Rename AgentThreadException → AgentSessionException - Remove stale AGUIThread from ag_ui lazy-loader - Rename use_service_thread → use_service_session in ag-ui package - Rename test functions from *_thread_* to *_session_* - Rename sample files from *_thread* to *_session* - Update docstrings and comments: thread → session - Update _mcp.py kwargs filter: add 'session' alongside 'thread' - Fix ContinuationToken docstring example: thread=thread → session=session - Fix _clients.py docstring: 'Agent threads' → 'Agent sessions' * Fix broken markdown links after thread→session file renames * fix azure ai test
This commit is contained in:
committed by
GitHub
Unverified
parent
0c67dbbce5
commit
1e350ea22f
+37
-37
@@ -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()
|
||||
Reference in New Issue
Block a user