mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: [BREAKING] Simplify API: ChatAgent -> Agent, ChatMessage -> Message (#3747)
* [BREAKING] Rename ChatAgent -> Agent, ChatMessage -> Message, ChatClientProtocol -> SupportsChatGetResponse Simplify the public API by removing redundant 'Chat' prefix from core types: - ChatAgent -> Agent - RawChatAgent -> RawAgent - ChatMessage -> Message - ChatClientProtocol -> SupportsChatGetResponse Also renamed internal WorkflowMessage (was Message in _runner_context) to avoid collision. No backward compatibility aliases - this is a clean breaking change. * [BREAKING] Rename Agent chat_client parameter to client * Fix rebase issues: WorkflowMessage references and broken markdown links * Fix formatting and lint issues from code quality checks * Fix import ordering in workflow sample files * fixed rebase * Fix test failures: use WorkflowMessage and A2AMessage after ChatMessage→Message rename - Replace Message(data=..., source_id=...) with WorkflowMessage(...) in workflow tests - Fix isinstance check in A2A agent to use A2AMessage instead of Message - Fix import in test_workflow_observability.py (Message→WorkflowMessage) * Fix lint, fmt, and sample errors after ChatMessage→Message rename - Auto-fix 70+ ruff lint issues across samples (ChatMessage→Message refs) - Fix HostedVectorStoreContent→Content.from_hosted_vector_store in file search sample - Fix _normalize_messages→normalize_messages in custom agent sample - Fix context.terminate→raise MiddlewareTermination in middleware samples - Fix with_update_hook→with_transform_hook in override middleware sample - Add TOptions_co import back to custom_chat_client sample - Add noqa for FastAPI File() default in chatkit sample - Fix B023 loop variable capture in weather agent sample * fix: update Agent constructor calls from chat_client to client in declaration-only tool tests * fix: add register_cleanup to devui lazy-loading proxy and type stub * fixed tests and updated new pieces * fix agui typevar * fix merge errors * fix merge conflicts * fiux merge * Remove unused links --------- Co-authored-by: Evan Mattson <evan.mattson@microsoft.com>
This commit is contained in:
committed by
GitHub
Unverified
parent
a4c9e43afb
commit
0521f5bed8
@@ -13,7 +13,7 @@ Redis-based storage for agent threads and context.
|
||||
from agent_framework.redis import RedisChatMessageStore
|
||||
|
||||
store = RedisChatMessageStore(redis_url="redis://localhost:6379")
|
||||
agent = ChatAgent(..., chat_message_store_factory=lambda: store)
|
||||
agent = Agent(..., chat_message_store_factory=lambda: store)
|
||||
```
|
||||
|
||||
## Import Path
|
||||
|
||||
@@ -7,7 +7,7 @@ from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
import redis.asyncio as redis
|
||||
from agent_framework import ChatMessage
|
||||
from agent_framework import Message
|
||||
from agent_framework._serialization import SerializationMixin
|
||||
from redis.credentials import CredentialProvider
|
||||
|
||||
@@ -64,7 +64,7 @@ class RedisChatMessageStore:
|
||||
thread_id: str | None = None,
|
||||
key_prefix: str = "chat_messages",
|
||||
max_messages: int | None = None,
|
||||
messages: Sequence[ChatMessage] | None = None,
|
||||
messages: Sequence[Message] | None = None,
|
||||
) -> None:
|
||||
"""Initialize the Redis chat message store.
|
||||
|
||||
@@ -186,14 +186,14 @@ class RedisChatMessageStore:
|
||||
self._initial_messages_added = True
|
||||
self._initial_messages.clear()
|
||||
|
||||
async def _add_redis_messages(self, messages: Sequence[ChatMessage]) -> None:
|
||||
async def _add_redis_messages(self, messages: Sequence[Message]) -> None:
|
||||
"""Add multiple messages to Redis using atomic pipeline operation.
|
||||
|
||||
This internal method efficiently adds multiple messages to the Redis list
|
||||
using a single atomic transaction to ensure consistency.
|
||||
|
||||
Args:
|
||||
messages: Sequence of ChatMessage objects to add to Redis.
|
||||
messages: Sequence of Message objects to add to Redis.
|
||||
"""
|
||||
if not messages:
|
||||
return
|
||||
@@ -207,7 +207,7 @@ class RedisChatMessageStore:
|
||||
await pipe.rpush(self.redis_key, serialized_message) # type: ignore[misc]
|
||||
await pipe.execute()
|
||||
|
||||
async def add_messages(self, messages: Sequence[ChatMessage]) -> None:
|
||||
async def add_messages(self, messages: Sequence[Message]) -> None:
|
||||
"""Add messages to the Redis store (ChatMessageStoreProtocol protocol method).
|
||||
|
||||
This method implements the required ChatMessageStoreProtocol protocol for adding messages.
|
||||
@@ -215,7 +215,7 @@ class RedisChatMessageStore:
|
||||
trimming if message limits are configured.
|
||||
|
||||
Args:
|
||||
messages: Sequence of ChatMessage objects to add to the store.
|
||||
messages: Sequence of Message objects to add to the store.
|
||||
Can be empty (no-op) or contain multiple messages.
|
||||
|
||||
Thread Safety:
|
||||
@@ -225,7 +225,7 @@ class RedisChatMessageStore:
|
||||
Example:
|
||||
.. code-block:: python
|
||||
|
||||
messages = [ChatMessage(role="user", text="Hello"), ChatMessage(role="assistant", text="Hi there!")]
|
||||
messages = [Message(role="user", text="Hello"), Message(role="assistant", text="Hi there!")]
|
||||
await store.add_messages(messages)
|
||||
"""
|
||||
if not messages:
|
||||
@@ -244,14 +244,14 @@ class RedisChatMessageStore:
|
||||
# Keep only the most recent max_messages using LTRIM
|
||||
await self._redis_client.ltrim(self.redis_key, -self.max_messages, -1) # type: ignore[misc]
|
||||
|
||||
async def list_messages(self) -> list[ChatMessage]:
|
||||
async def list_messages(self) -> list[Message]:
|
||||
"""Get all messages from the store in chronological order (ChatMessageStoreProtocol protocol method).
|
||||
|
||||
This method implements the required ChatMessageStoreProtocol protocol for retrieving messages.
|
||||
Returns all messages stored in Redis, ordered from oldest (index 0) to newest (index -1).
|
||||
|
||||
Returns:
|
||||
List of ChatMessage objects in chronological order (oldest first).
|
||||
List of Message objects in chronological order (oldest first).
|
||||
Returns empty list if no messages exist or if Redis connection fails.
|
||||
|
||||
Example:
|
||||
@@ -269,7 +269,7 @@ class RedisChatMessageStore:
|
||||
|
||||
if redis_messages:
|
||||
for serialized_message in redis_messages:
|
||||
# Deserialize each JSON message back to ChatMessage
|
||||
# Deserialize each JSON message back to Message
|
||||
message = self._deserialize_message(serialized_message)
|
||||
messages.append(message)
|
||||
|
||||
@@ -390,11 +390,11 @@ class RedisChatMessageStore:
|
||||
"""
|
||||
await self._redis_client.delete(self.redis_key)
|
||||
|
||||
def _serialize_message(self, message: ChatMessage) -> str:
|
||||
"""Serialize a ChatMessage to JSON string.
|
||||
def _serialize_message(self, message: Message) -> str:
|
||||
"""Serialize a Message to JSON string.
|
||||
|
||||
Args:
|
||||
message: ChatMessage to serialize.
|
||||
message: Message to serialize.
|
||||
|
||||
Returns:
|
||||
JSON string representation of the message.
|
||||
@@ -402,17 +402,17 @@ class RedisChatMessageStore:
|
||||
# Serialize to compact JSON (no extra whitespace for Redis efficiency)
|
||||
return message.to_json(separators=(",", ":"))
|
||||
|
||||
def _deserialize_message(self, serialized_message: str) -> ChatMessage:
|
||||
"""Deserialize a JSON string to ChatMessage.
|
||||
def _deserialize_message(self, serialized_message: str) -> Message:
|
||||
"""Deserialize a JSON string to Message.
|
||||
|
||||
Args:
|
||||
serialized_message: JSON string representation of a message.
|
||||
|
||||
Returns:
|
||||
ChatMessage object.
|
||||
Message object.
|
||||
"""
|
||||
# Reconstruct ChatMessage using custom deserialization
|
||||
return ChatMessage.from_json(serialized_message)
|
||||
# Reconstruct Message using custom deserialization
|
||||
return Message.from_json(serialized_message)
|
||||
|
||||
# ============================================================================
|
||||
# List-like Convenience Methods (Redis-optimized async versions)
|
||||
@@ -446,14 +446,14 @@ class RedisChatMessageStore:
|
||||
await self._ensure_initial_messages_added()
|
||||
return await self._redis_client.llen(self.redis_key) # type: ignore[misc,no-any-return]
|
||||
|
||||
async def getitem(self, index: int) -> ChatMessage:
|
||||
async def getitem(self, index: int) -> Message:
|
||||
"""Get a message by index using Redis LINDEX.
|
||||
|
||||
Args:
|
||||
index: The index of the message to retrieve.
|
||||
|
||||
Returns:
|
||||
The ChatMessage at the specified index.
|
||||
The Message at the specified index.
|
||||
|
||||
Raises:
|
||||
IndexError: If the index is out of range.
|
||||
@@ -467,12 +467,12 @@ class RedisChatMessageStore:
|
||||
|
||||
return self._deserialize_message(serialized_message)
|
||||
|
||||
async def setitem(self, index: int, item: ChatMessage) -> None:
|
||||
async def setitem(self, index: int, item: Message) -> None:
|
||||
"""Set a message at the specified index using Redis LSET.
|
||||
|
||||
Args:
|
||||
index: The index at which to set the message.
|
||||
item: The ChatMessage to set at the specified index.
|
||||
item: The Message to set at the specified index.
|
||||
|
||||
Raises:
|
||||
IndexError: If the index is out of range.
|
||||
@@ -490,11 +490,11 @@ class RedisChatMessageStore:
|
||||
serialized_message = self._serialize_message(item)
|
||||
await self._redis_client.lset(self.redis_key, index, serialized_message) # type: ignore[misc]
|
||||
|
||||
async def append(self, item: ChatMessage) -> None:
|
||||
async def append(self, item: Message) -> None:
|
||||
"""Append a message to the end of the store.
|
||||
|
||||
Args:
|
||||
item: The ChatMessage to append.
|
||||
item: The Message to append.
|
||||
"""
|
||||
await self.add_messages([item])
|
||||
|
||||
@@ -507,14 +507,14 @@ class RedisChatMessageStore:
|
||||
await self._ensure_initial_messages_added()
|
||||
return await self._redis_client.llen(self.redis_key) # type: ignore[misc,no-any-return]
|
||||
|
||||
async def index(self, item: ChatMessage) -> int:
|
||||
async def index(self, item: Message) -> int:
|
||||
"""Return the index of the first occurrence of the specified message.
|
||||
|
||||
Uses Redis LINDEX to iterate through the list without loading all messages.
|
||||
Still O(N) but more memory efficient for large lists.
|
||||
|
||||
Args:
|
||||
item: The ChatMessage to find.
|
||||
item: The Message to find.
|
||||
|
||||
Returns:
|
||||
The index of the first occurrence of the message.
|
||||
@@ -533,16 +533,16 @@ class RedisChatMessageStore:
|
||||
if redis_message == target_serialized:
|
||||
return i
|
||||
|
||||
raise ValueError("ChatMessage not found in store")
|
||||
raise ValueError("Message not found in store")
|
||||
|
||||
async def remove(self, item: ChatMessage) -> None:
|
||||
async def remove(self, item: Message) -> None:
|
||||
"""Remove the first occurrence of the specified message from the store.
|
||||
|
||||
Uses Redis LREM command for efficient removal by value.
|
||||
O(N) but performed natively in Redis without data transfer.
|
||||
|
||||
Args:
|
||||
item: The ChatMessage to remove.
|
||||
item: The Message to remove.
|
||||
|
||||
Raises:
|
||||
ValueError: If the message is not found in the store.
|
||||
@@ -556,13 +556,13 @@ class RedisChatMessageStore:
|
||||
removed_count = await self._redis_client.lrem(self.redis_key, 1, target_serialized) # type: ignore[misc]
|
||||
|
||||
if removed_count == 0:
|
||||
raise ValueError("ChatMessage not found in store")
|
||||
raise ValueError("Message not found in store")
|
||||
|
||||
async def extend(self, items: Sequence[ChatMessage]) -> None:
|
||||
async def extend(self, items: Sequence[Message]) -> None:
|
||||
"""Extend the store by appending all messages from the iterable.
|
||||
|
||||
Args:
|
||||
items: Sequence of ChatMessage objects to append.
|
||||
items: Sequence of Message objects to append.
|
||||
"""
|
||||
await self.add_messages(items)
|
||||
|
||||
|
||||
@@ -16,7 +16,7 @@ from operator import and_
|
||||
from typing import TYPE_CHECKING, Any, Literal, cast
|
||||
|
||||
import numpy as np
|
||||
from agent_framework import ChatMessage
|
||||
from agent_framework import Message
|
||||
from agent_framework._sessions import AgentSession, BaseContextProvider, SessionContext
|
||||
from agent_framework.exceptions import (
|
||||
AgentException,
|
||||
@@ -142,7 +142,7 @@ class _RedisContextProvider(BaseContextProvider):
|
||||
if line_separated_memories:
|
||||
context.extend_messages(
|
||||
self.source_id,
|
||||
[ChatMessage(role="user", text=f"{self.context_prompt}\n{line_separated_memories}")],
|
||||
[Message(role="user", text=f"{self.context_prompt}\n{line_separated_memories}")],
|
||||
)
|
||||
|
||||
@override
|
||||
@@ -157,7 +157,7 @@ class _RedisContextProvider(BaseContextProvider):
|
||||
"""Store request/response messages to Redis for future retrieval."""
|
||||
self._validate_filters()
|
||||
|
||||
messages_to_store: list[ChatMessage] = list(context.input_messages)
|
||||
messages_to_store: list[Message] = list(context.input_messages)
|
||||
if context.response and context.response.messages:
|
||||
messages_to_store.extend(context.response.messages)
|
||||
|
||||
|
||||
@@ -3,34 +3,31 @@
|
||||
"""New-pattern Redis history provider using BaseHistoryProvider.
|
||||
|
||||
This module provides ``_RedisHistoryProvider``, a side-by-side implementation of
|
||||
:class:`RedisChatMessageStore` built on the new :class:`BaseHistoryProvider` hooks pattern.
|
||||
:class:`RedisMessageStore` built on the new :class:`BaseHistoryProvider` hooks pattern.
|
||||
It will be renamed to ``RedisHistoryProvider`` in PR2 when the old class is removed.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import Any
|
||||
|
||||
import redis.asyncio as redis
|
||||
from agent_framework import ChatMessage
|
||||
from agent_framework import Message
|
||||
from agent_framework._sessions import BaseHistoryProvider
|
||||
from redis.credentials import CredentialProvider
|
||||
|
||||
if TYPE_CHECKING:
|
||||
pass
|
||||
|
||||
|
||||
class _RedisHistoryProvider(BaseHistoryProvider):
|
||||
"""Redis-backed history provider using the new BaseHistoryProvider hooks pattern.
|
||||
|
||||
Stores conversation history in Redis Lists, with each session isolated by a
|
||||
unique Redis key. This is the new-pattern equivalent of
|
||||
:class:`RedisChatMessageStore`.
|
||||
:class:`RedisMessageStore`.
|
||||
|
||||
Note:
|
||||
This class uses a temporary ``_`` prefix to coexist with the existing
|
||||
:class:`RedisChatMessageStore`. It will be renamed to ``RedisHistoryProvider``
|
||||
:class:`RedisMessageStore`. It will be renamed to ``RedisHistoryProvider``
|
||||
in PR2.
|
||||
"""
|
||||
|
||||
@@ -115,7 +112,7 @@ class _RedisHistoryProvider(BaseHistoryProvider):
|
||||
"""Get the Redis key for a given session's messages."""
|
||||
return f"{self.key_prefix}:{session_id or 'default'}"
|
||||
|
||||
async def get_messages(self, session_id: str | None, **kwargs: Any) -> list[ChatMessage]:
|
||||
async def get_messages(self, session_id: str | None, **kwargs: Any) -> list[Message]:
|
||||
"""Retrieve stored messages for this session from Redis.
|
||||
|
||||
Args:
|
||||
@@ -123,17 +120,17 @@ class _RedisHistoryProvider(BaseHistoryProvider):
|
||||
**kwargs: Additional arguments (unused).
|
||||
|
||||
Returns:
|
||||
List of stored ChatMessage objects in chronological order.
|
||||
List of stored Message objects in chronological order.
|
||||
"""
|
||||
key = self._redis_key(session_id)
|
||||
redis_messages = await self._redis_client.lrange(key, 0, -1) # type: ignore[misc]
|
||||
messages: list[ChatMessage] = []
|
||||
messages: list[Message] = []
|
||||
if redis_messages:
|
||||
for serialized in redis_messages:
|
||||
messages.append(ChatMessage.from_dict(self._deserialize_json(serialized)))
|
||||
messages.append(Message.from_dict(self._deserialize_json(serialized)))
|
||||
return messages
|
||||
|
||||
async def save_messages(self, session_id: str | None, messages: Sequence[ChatMessage], **kwargs: Any) -> None:
|
||||
async def save_messages(self, session_id: str | None, messages: Sequence[Message], **kwargs: Any) -> None:
|
||||
"""Persist messages for this session to Redis.
|
||||
|
||||
Args:
|
||||
@@ -158,8 +155,8 @@ class _RedisHistoryProvider(BaseHistoryProvider):
|
||||
await self._redis_client.ltrim(key, -self.max_messages, -1) # type: ignore[misc]
|
||||
|
||||
@staticmethod
|
||||
def _serialize_json(message: ChatMessage) -> str:
|
||||
"""Serialize a ChatMessage to a JSON string for Redis storage."""
|
||||
def _serialize_json(message: Message) -> str:
|
||||
"""Serialize a Message to a JSON string for Redis storage."""
|
||||
import json
|
||||
|
||||
return json.dumps(message.to_dict())
|
||||
|
||||
@@ -10,7 +10,7 @@ from operator import and_
|
||||
from typing import Any, Literal, cast
|
||||
|
||||
import numpy as np
|
||||
from agent_framework import ChatMessage, Context, ContextProvider
|
||||
from agent_framework import Context, ContextProvider, Message
|
||||
from agent_framework.exceptions import (
|
||||
AgentException,
|
||||
ServiceInitializationError,
|
||||
@@ -484,19 +484,17 @@ class RedisProvider(ContextProvider):
|
||||
@override
|
||||
async def invoked(
|
||||
self,
|
||||
request_messages: ChatMessage | Sequence[ChatMessage],
|
||||
response_messages: ChatMessage | Sequence[ChatMessage] | None = None,
|
||||
request_messages: Message | Sequence[Message],
|
||||
response_messages: Message | Sequence[Message] | None = None,
|
||||
invoke_exception: Exception | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
self._validate_filters()
|
||||
|
||||
request_messages_list = (
|
||||
[request_messages] if isinstance(request_messages, ChatMessage) else list(request_messages)
|
||||
)
|
||||
request_messages_list = [request_messages] if isinstance(request_messages, Message) else list(request_messages)
|
||||
response_messages_list = (
|
||||
[response_messages]
|
||||
if isinstance(response_messages, ChatMessage)
|
||||
if isinstance(response_messages, Message)
|
||||
else list(response_messages)
|
||||
if response_messages
|
||||
else []
|
||||
@@ -518,7 +516,7 @@ class RedisProvider(ContextProvider):
|
||||
await self._add(data=messages)
|
||||
|
||||
@override
|
||||
async def invoking(self, messages: ChatMessage | MutableSequence[ChatMessage], **kwargs: Any) -> Context:
|
||||
async def invoking(self, messages: Message | MutableSequence[Message], **kwargs: Any) -> Context:
|
||||
"""Called before invoking the model to provide scoped context.
|
||||
|
||||
Concatenates recent messages into a query, fetches matching memories from Redis.
|
||||
@@ -534,7 +532,7 @@ class RedisProvider(ContextProvider):
|
||||
Context: Context object containing instructions with memories.
|
||||
"""
|
||||
self._validate_filters()
|
||||
messages_list = [messages] if isinstance(messages, ChatMessage) else list(messages)
|
||||
messages_list = [messages] if isinstance(messages, Message) else list(messages)
|
||||
input_text = "\n".join(msg.text for msg in messages_list if msg and msg.text and msg.text.strip())
|
||||
|
||||
memories = await self._redis_search(text=input_text)
|
||||
@@ -543,7 +541,7 @@ class RedisProvider(ContextProvider):
|
||||
)
|
||||
|
||||
return Context(
|
||||
messages=[ChatMessage(role="user", text=f"{self.context_prompt}\n{line_separated_memories}")]
|
||||
messages=[Message(role="user", text=f"{self.context_prompt}\n{line_separated_memories}")]
|
||||
if line_separated_memories
|
||||
else None
|
||||
)
|
||||
|
||||
@@ -8,7 +8,7 @@ import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from agent_framework import AgentResponse, ChatMessage
|
||||
from agent_framework import AgentResponse, Message
|
||||
from agent_framework._sessions import AgentSession, SessionContext
|
||||
from agent_framework.exceptions import ServiceInitializationError
|
||||
|
||||
@@ -142,7 +142,7 @@ class TestRedisContextProviderBeforeRun:
|
||||
mock_index.query = AsyncMock(return_value=[{"content": "Memory A"}, {"content": "Memory B"}])
|
||||
provider = _RedisContextProvider(source_id="ctx", user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", contents=["test query"])], session_id="s1")
|
||||
ctx = SessionContext(input_messages=[Message(role="user", contents=["test query"])], session_id="s1")
|
||||
|
||||
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
|
||||
@@ -159,7 +159,7 @@ class TestRedisContextProviderBeforeRun:
|
||||
):
|
||||
provider = _RedisContextProvider(source_id="ctx", user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", contents=[" "])], session_id="s1")
|
||||
ctx = SessionContext(input_messages=[Message(role="user", contents=[" "])], session_id="s1")
|
||||
|
||||
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
|
||||
@@ -174,7 +174,7 @@ class TestRedisContextProviderBeforeRun:
|
||||
mock_index.query = AsyncMock(return_value=[])
|
||||
provider = _RedisContextProvider(source_id="ctx", user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", contents=["hello"])], session_id="s1")
|
||||
ctx = SessionContext(input_messages=[Message(role="user", contents=["hello"])], session_id="s1")
|
||||
|
||||
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
|
||||
@@ -189,8 +189,8 @@ class TestRedisContextProviderAfterRun:
|
||||
):
|
||||
provider = _RedisContextProvider(source_id="ctx", user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
response = AgentResponse(messages=[ChatMessage(role="assistant", contents=["response text"])])
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", contents=["user input"])], session_id="s1")
|
||||
response = AgentResponse(messages=[Message(role="assistant", contents=["response text"])])
|
||||
ctx = SessionContext(input_messages=[Message(role="user", contents=["user input"])], session_id="s1")
|
||||
ctx._response = response
|
||||
|
||||
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
@@ -208,7 +208,7 @@ class TestRedisContextProviderAfterRun:
|
||||
):
|
||||
provider = _RedisContextProvider(source_id="ctx", user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", contents=[" "])], session_id="s1")
|
||||
ctx = SessionContext(input_messages=[Message(role="user", contents=[" "])], session_id="s1")
|
||||
|
||||
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
|
||||
@@ -221,7 +221,7 @@ class TestRedisContextProviderAfterRun:
|
||||
):
|
||||
provider = _RedisContextProvider(source_id="ctx", application_id="app", agent_id="ag", user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", contents=["hello"])], session_id="s1")
|
||||
ctx = SessionContext(input_messages=[Message(role="user", contents=["hello"])], session_id="s1")
|
||||
|
||||
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
|
||||
@@ -325,8 +325,8 @@ class TestRedisHistoryProviderRedisKey:
|
||||
|
||||
class TestRedisHistoryProviderGetMessages:
|
||||
async def test_returns_deserialized_messages(self, mock_redis_client: MagicMock):
|
||||
msg1 = ChatMessage(role="user", contents=["Hello"])
|
||||
msg2 = ChatMessage(role="assistant", contents=["Hi!"])
|
||||
msg1 = Message(role="user", contents=["Hello"])
|
||||
msg2 = Message(role="assistant", contents=["Hi!"])
|
||||
mock_redis_client.lrange = AsyncMock(return_value=[json.dumps(msg1.to_dict()), json.dumps(msg2.to_dict())])
|
||||
|
||||
with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
|
||||
@@ -357,7 +357,7 @@ class TestRedisHistoryProviderSaveMessages:
|
||||
mock_from_url.return_value = mock_redis_client
|
||||
provider = _RedisHistoryProvider("mem", redis_url="redis://localhost:6379")
|
||||
|
||||
msgs = [ChatMessage(role="user", contents=["Hello"]), ChatMessage(role="assistant", contents=["Hi"])]
|
||||
msgs = [Message(role="user", contents=["Hello"]), Message(role="assistant", contents=["Hi"])]
|
||||
await provider.save_messages("s1", msgs)
|
||||
|
||||
pipeline = mock_redis_client.pipeline.return_value.__aenter__.return_value
|
||||
@@ -379,7 +379,7 @@ class TestRedisHistoryProviderSaveMessages:
|
||||
mock_from_url.return_value = mock_redis_client
|
||||
provider = _RedisHistoryProvider("mem", redis_url="redis://localhost:6379", max_messages=10)
|
||||
|
||||
await provider.save_messages("s1", [ChatMessage(role="user", contents=["msg"])])
|
||||
await provider.save_messages("s1", [Message(role="user", contents=["msg"])])
|
||||
|
||||
mock_redis_client.ltrim.assert_called_once_with("chat_messages:s1", -10, -1)
|
||||
|
||||
@@ -390,7 +390,7 @@ class TestRedisHistoryProviderSaveMessages:
|
||||
mock_from_url.return_value = mock_redis_client
|
||||
provider = _RedisHistoryProvider("mem", redis_url="redis://localhost:6379", max_messages=10)
|
||||
|
||||
await provider.save_messages("s1", [ChatMessage(role="user", contents=["msg"])])
|
||||
await provider.save_messages("s1", [Message(role="user", contents=["msg"])])
|
||||
|
||||
mock_redis_client.ltrim.assert_not_called()
|
||||
|
||||
@@ -409,7 +409,7 @@ class TestRedisHistoryProviderBeforeAfterRun:
|
||||
"""Test before_run/after_run integration via BaseHistoryProvider defaults."""
|
||||
|
||||
async def test_before_run_loads_history(self, mock_redis_client: MagicMock):
|
||||
msg = ChatMessage(role="user", contents=["old msg"])
|
||||
msg = Message(role="user", contents=["old msg"])
|
||||
mock_redis_client.lrange = AsyncMock(return_value=[json.dumps(msg.to_dict())])
|
||||
|
||||
with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
|
||||
@@ -417,7 +417,7 @@ class TestRedisHistoryProviderBeforeAfterRun:
|
||||
provider = _RedisHistoryProvider("mem", redis_url="redis://localhost:6379")
|
||||
|
||||
session = AgentSession(session_id="test")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", contents=["new msg"])], session_id="s1")
|
||||
ctx = SessionContext(input_messages=[Message(role="user", contents=["new msg"])], session_id="s1")
|
||||
|
||||
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
|
||||
@@ -431,8 +431,8 @@ class TestRedisHistoryProviderBeforeAfterRun:
|
||||
provider = _RedisHistoryProvider("mem", redis_url="redis://localhost:6379")
|
||||
|
||||
session = AgentSession(session_id="test")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", contents=["hi"])], session_id="s1")
|
||||
ctx._response = AgentResponse(messages=[ChatMessage(role="assistant", contents=["hello"])])
|
||||
ctx = SessionContext(input_messages=[Message(role="user", contents=["hi"])], session_id="s1")
|
||||
ctx._response = AgentResponse(messages=[Message(role="assistant", contents=["hello"])])
|
||||
|
||||
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
|
||||
@@ -448,7 +448,7 @@ class TestRedisHistoryProviderBeforeAfterRun:
|
||||
)
|
||||
|
||||
session = AgentSession(session_id="test")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", contents=["hi"])], session_id="s1")
|
||||
ctx = SessionContext(input_messages=[Message(role="user", contents=["hi"])], session_id="s1")
|
||||
|
||||
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from agent_framework import ChatMessage, Content
|
||||
from agent_framework import Content, Message
|
||||
|
||||
from agent_framework_redis import RedisChatMessageStore
|
||||
|
||||
@@ -19,9 +19,9 @@ class TestRedisChatMessageStore:
|
||||
def sample_messages(self):
|
||||
"""Sample chat messages for testing."""
|
||||
return [
|
||||
ChatMessage(role="user", text="Hello", message_id="msg1"),
|
||||
ChatMessage(role="assistant", text="Hi there!", message_id="msg2"),
|
||||
ChatMessage(role="user", text="How are you?", message_id="msg3"),
|
||||
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
|
||||
@@ -250,7 +250,7 @@ class TestRedisChatMessageStore:
|
||||
store = RedisChatMessageStore(redis_url="redis://localhost:6379", thread_id="test123", max_messages=3)
|
||||
store._redis_client = mock_redis_client
|
||||
|
||||
message = ChatMessage(role="user", text="Test")
|
||||
message = Message(role="user", text="Test")
|
||||
await store.add_messages([message])
|
||||
|
||||
# Should trim after adding to keep only last 3 messages
|
||||
@@ -269,8 +269,8 @@ class TestRedisChatMessageStore:
|
||||
"""Test listing messages with data in Redis."""
|
||||
# Create proper serialized messages using the actual serialization method
|
||||
test_messages = [
|
||||
ChatMessage(role="user", text="Hello", message_id="msg1"),
|
||||
ChatMessage(role="assistant", text="Hi there!", message_id="msg2"),
|
||||
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
|
||||
@@ -411,7 +411,7 @@ class TestRedisChatMessageStore:
|
||||
store = RedisChatMessageStore(redis_url="redis://localhost:6379", thread_id="test123")
|
||||
|
||||
# Message with multiple content types
|
||||
message = ChatMessage(
|
||||
message = Message(
|
||||
role="assistant",
|
||||
contents=[Content.from_text(text="Hello"), Content.from_text(text="World")],
|
||||
author_name="TestBot",
|
||||
@@ -444,7 +444,7 @@ class TestRedisChatMessageStore:
|
||||
store = RedisChatMessageStore(redis_url="redis://localhost:6379", thread_id="test123")
|
||||
store._redis_client = mock_client
|
||||
|
||||
message = ChatMessage(role="user", text="Test")
|
||||
message = Message(role="user", text="Test")
|
||||
|
||||
# Should propagate Redis connection errors
|
||||
with pytest.raises(Exception, match="Connection failed"):
|
||||
@@ -485,7 +485,7 @@ class TestRedisChatMessageStore:
|
||||
mock_redis_client.llen.return_value = 2
|
||||
mock_redis_client.lset = AsyncMock()
|
||||
|
||||
new_message = ChatMessage(role="user", text="Updated message")
|
||||
new_message = Message(role="user", text="Updated message")
|
||||
await redis_store.setitem(0, new_message)
|
||||
|
||||
mock_redis_client.lset.assert_called_once()
|
||||
@@ -497,13 +497,13 @@ class TestRedisChatMessageStore:
|
||||
"""Test setitem raises IndexError for invalid index."""
|
||||
mock_redis_client.llen.return_value = 0
|
||||
|
||||
new_message = ChatMessage(role="user", text="Test")
|
||||
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 = ChatMessage(role="user", text="Appended message")
|
||||
message = Message(role="user", text="Appended message")
|
||||
await redis_store.append(message)
|
||||
|
||||
# Should call pipeline operations via add_messages
|
||||
@@ -572,7 +572,7 @@ class TestRedisChatMessageStore:
|
||||
mock_redis_client.llen.return_value = 1
|
||||
mock_redis_client.lindex = AsyncMock(return_value="different_message")
|
||||
|
||||
with pytest.raises(ValueError, match="ChatMessage not found in store"):
|
||||
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):
|
||||
@@ -589,7 +589,7 @@ class TestRedisChatMessageStore:
|
||||
"""Test remove method when message is not found."""
|
||||
mock_redis_client.lrem = AsyncMock(return_value=0) # 0 elements removed
|
||||
|
||||
with pytest.raises(ValueError, match="ChatMessage not found in store"):
|
||||
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):
|
||||
|
||||
@@ -5,7 +5,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
from agent_framework import ChatMessage
|
||||
from agent_framework import Message
|
||||
from agent_framework.exceptions import AgentException, ServiceInitializationError
|
||||
from redisvl.utils.vectorize import CustomTextVectorizer
|
||||
|
||||
@@ -113,18 +113,18 @@ class TestRedisProviderInitialization:
|
||||
|
||||
class TestRedisProviderMessages:
|
||||
@pytest.fixture
|
||||
def sample_messages(self) -> list[ChatMessage]:
|
||||
def sample_messages(self) -> list[Message]:
|
||||
return [
|
||||
ChatMessage(role="user", text="Hello, how are you?"),
|
||||
ChatMessage(role="assistant", text="I'm doing well, thank you!"),
|
||||
ChatMessage(role="system", text="You are a helpful assistant"),
|
||||
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", ChatMessage(role="user", text="Hello"))
|
||||
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
|
||||
@@ -157,7 +157,7 @@ class TestRedisProviderModelInvoking:
|
||||
async def test_model_invoking_requires_filters(self, patch_index_from_dict): # noqa: ARG002
|
||||
provider = RedisProvider()
|
||||
with pytest.raises(ServiceInitializationError):
|
||||
await provider.invoking(ChatMessage(role="user", text="Hi"))
|
||||
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(
|
||||
@@ -168,7 +168,7 @@ class TestRedisProviderModelInvoking:
|
||||
provider = RedisProvider(user_id="u1")
|
||||
|
||||
# Act
|
||||
ctx = await provider.invoking([ChatMessage(role="user", text="q1")])
|
||||
ctx = await provider.invoking([Message(role="user", text="q1")])
|
||||
|
||||
# Assert: TextQuery used (not HybridQuery), filter_expression included
|
||||
assert patch_queries["TextQuery"].call_count == 1
|
||||
@@ -190,7 +190,7 @@ class TestRedisProviderModelInvoking:
|
||||
): # noqa: ARG002
|
||||
mock_index.query = AsyncMock(return_value=[])
|
||||
provider = RedisProvider(user_id="u1")
|
||||
ctx = await provider.invoking([ChatMessage(role="user", text="any")])
|
||||
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
|
||||
@@ -198,7 +198,7 @@ class TestRedisProviderModelInvoking:
|
||||
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([ChatMessage(role="user", text="hello")])
|
||||
ctx = await provider.invoking([Message(role="user", text="hello")])
|
||||
|
||||
# Assert: HybridQuery used with vector and vector field
|
||||
assert patch_queries["HybridQuery"].call_count == 1
|
||||
@@ -240,9 +240,9 @@ class TestMessagesAddingBehavior:
|
||||
)
|
||||
|
||||
msgs = [
|
||||
ChatMessage(role="user", text="u"),
|
||||
ChatMessage(role="assistant", text="a"),
|
||||
ChatMessage(role="system", text="s"),
|
||||
Message(role="user", text="u"),
|
||||
Message(role="assistant", text="a"),
|
||||
Message(role="system", text="s"),
|
||||
]
|
||||
|
||||
await provider.invoked(msgs)
|
||||
@@ -265,8 +265,8 @@ class TestMessagesAddingBehavior:
|
||||
): # noqa: ARG002
|
||||
provider = RedisProvider(user_id="u1", scope_to_per_operation_thread_id=True)
|
||||
msgs = [
|
||||
ChatMessage(role="user", text=" "),
|
||||
ChatMessage(role="tool", text="tool output"),
|
||||
Message(role="user", text=" "),
|
||||
Message(role="tool", text="tool output"),
|
||||
]
|
||||
await provider.invoked(msgs)
|
||||
# No valid messages -> no load
|
||||
@@ -279,8 +279,8 @@ class TestIndexCreationPublicCalls:
|
||||
self, mock_index: AsyncMock, patch_index_from_dict
|
||||
): # noqa: ARG002
|
||||
provider = RedisProvider(user_id="u1")
|
||||
await provider.invoked(ChatMessage(role="user", text="m1"))
|
||||
await provider.invoked(ChatMessage(role="user", text="m2"))
|
||||
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
|
||||
|
||||
@@ -291,7 +291,7 @@ class TestIndexCreationPublicCalls:
|
||||
mock_index.exists = AsyncMock(return_value=False)
|
||||
provider = RedisProvider(user_id="u1")
|
||||
mock_index.query = AsyncMock(return_value=[{"content": "C"}])
|
||||
await provider.invoking([ChatMessage(role="user", text="q")])
|
||||
await provider.invoking([Message(role="user", text="q")])
|
||||
assert mock_index.create.await_count == 1
|
||||
|
||||
|
||||
@@ -321,7 +321,7 @@ class TestVectorPopulation:
|
||||
vector_field_name="vec",
|
||||
)
|
||||
|
||||
await provider.invoked(ChatMessage(role="user", text="hello"))
|
||||
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]
|
||||
|
||||
Reference in New Issue
Block a user