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
@@ -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