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:
Eduard van Valkenburg
2026-02-11 00:04:32 +01:00
committed by GitHub
Unverified
parent a4c9e43afb
commit 0521f5bed8
418 changed files with 5385 additions and 5389 deletions
@@ -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]