mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: [BREAKING] Types API Review improvements (#3647)
* Replace Role and FinishReason classes with NewType + Literal
- Remove EnumLike metaclass from _types.py
- Replace Role class with NewType('Role', str) + RoleLiteral
- Replace FinishReason class with NewType('FinishReason', str) + FinishReasonLiteral
- Update all usages across codebase to use string literals
- Remove .value access patterns (direct string comparison now works)
- Add backward compatibility for legacy dict serialization format
- Update tests to reflect new string-based types
Addresses #3591, #3615
* Simplify ChatResponse and AgentResponse type hints (#3592)
- Remove overloads from ChatResponse.__init__
- Remove text parameter from ChatResponse.__init__
- Remove | dict[str, Any] from finish_reason and usage_details params
- Remove **kwargs from AgentResponse.__init__
- Both now accept ChatMessage | Sequence[ChatMessage] | None for messages
- Update docstrings and examples to reflect changes
- Fix tests that were using removed kwargs
- Fix Role type hint usage in ag-ui utils
* Remove text parameter from ChatResponseUpdate and AgentResponseUpdate (#3597)
- Remove text parameter from ChatResponseUpdate.__init__
- Remove text parameter from AgentResponseUpdate.__init__
- Remove **kwargs from both update classes
- Simplify contents parameter type to Sequence[Content] | None
- Update all usages to use contents=[Content.from_text(...)] pattern
- Fix imports in test files
- Update docstrings and examples
* Rename from_chat_response_updates to from_updates (#3593)
- ChatResponse.from_chat_response_updates → ChatResponse.from_updates
- ChatResponse.from_chat_response_generator → ChatResponse.from_update_generator
- AgentResponse.from_agent_run_response_updates → AgentResponse.from_updates
* Remove try_parse_value method from ChatResponse and AgentResponse (#3595)
- Remove try_parse_value method from ChatResponse
- Remove try_parse_value method from AgentResponse
- Remove try_parse_value calls from from_updates and from_update_generator methods
- Update samples to use try/except with response.value instead
- Update tests to use response.value pattern
- Users should now use response.value with try/except for safe parsing
* Add agent_id to AgentResponse and clarify author_name documentation (#3596)
- Add agent_id parameter to AgentResponse class
- Document that author_name is on ChatMessage objects, not responses
- Update ChatResponse docstring with author_name note
- Update AgentResponse docstring with author_name note
* Simplify ChatMessage.__init__ signature (#3618)
- Make contents a positional argument accepting Sequence[Content | str]
- Auto-convert strings in contents to TextContent
- Remove overloads, keep text kwarg for backward compatibility with serialization
- Update _parse_content_list to handle string items
- Update all usages across codebase to use new format: ChatMessage("role", ["text"])
* Allow Content as input on run and get_response
- Update prepare_messages and normalize_messages to accept Content
- Update type signatures in _agents.py and _clients.py
- Add tests for Content input handling
* Fix ChatMessage usage across packages and samples
Update all remaining ChatMessage(role=..., text=...) to use new
ChatMessage('role', ['text']) signature.
* Fix Role string usage and response format parsing
- Fix redis provider: remove .value access on string literals
- Fix durabletask ensure_response_format: set _response_format before accessing .value
* Fix ollama .value and ai_model_id issues, handle None in content list
- Fix ollama _chat_client: remove .value on string literals
- Fix ollama _chat_client: rename ai_model_id to model_id
- Fix _parse_content_list: skip None values gracefully
* Fix A2AAgent type signature to include Content
* Fix Role/FinishReason NewType dict annotations and improve test coverage to 95%
* Fix mypy errors for Role/FinishReason NewType usage
* Fix Role.TOOL and Role.ASSISTANT usage in _orchestrator_helpers.py
* Fix Role NewType usage in durabletask _models.py
This commit is contained in:
@@ -3,7 +3,7 @@
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from agent_framework import ChatMessage, Content, Role
|
||||
from agent_framework import ChatMessage, Content
|
||||
|
||||
from agent_framework_redis import RedisChatMessageStore
|
||||
|
||||
@@ -19,9 +19,9 @@ class TestRedisChatMessageStore:
|
||||
def sample_messages(self):
|
||||
"""Sample chat messages for testing."""
|
||||
return [
|
||||
ChatMessage(role=Role.USER, text="Hello", message_id="msg1"),
|
||||
ChatMessage(role=Role.ASSISTANT, text="Hi there!", message_id="msg2"),
|
||||
ChatMessage(role=Role.USER, text="How are you?", message_id="msg3"),
|
||||
ChatMessage("user", ["Hello"], message_id="msg1"),
|
||||
ChatMessage("assistant", ["Hi there!"], message_id="msg2"),
|
||||
ChatMessage("user", ["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=Role.USER, text="Test")
|
||||
message = ChatMessage("user", ["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=Role.USER, text="Hello", message_id="msg1"),
|
||||
ChatMessage(role=Role.ASSISTANT, text="Hi there!", message_id="msg2"),
|
||||
ChatMessage("user", ["Hello"], message_id="msg1"),
|
||||
ChatMessage("assistant", ["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
|
||||
@@ -278,9 +278,9 @@ class TestRedisChatMessageStore:
|
||||
messages = await redis_store.list_messages()
|
||||
|
||||
assert len(messages) == 2
|
||||
assert messages[0].role == Role.USER
|
||||
assert messages[0].role == "user"
|
||||
assert messages[0].text == "Hello"
|
||||
assert messages[1].role == Role.ASSISTANT
|
||||
assert messages[1].role == "assistant"
|
||||
assert messages[1].text == "Hi there!"
|
||||
|
||||
async def test_list_messages_with_initial_messages(self, sample_messages):
|
||||
@@ -412,7 +412,7 @@ class TestRedisChatMessageStore:
|
||||
|
||||
# Message with multiple content types
|
||||
message = ChatMessage(
|
||||
role=Role.ASSISTANT,
|
||||
role="assistant",
|
||||
contents=[Content.from_text(text="Hello"), Content.from_text(text="World")],
|
||||
author_name="TestBot",
|
||||
message_id="complex_msg",
|
||||
@@ -422,7 +422,7 @@ class TestRedisChatMessageStore:
|
||||
serialized = store._serialize_message(message)
|
||||
deserialized = store._deserialize_message(serialized)
|
||||
|
||||
assert deserialized.role == Role.ASSISTANT
|
||||
assert deserialized.role == "assistant"
|
||||
assert deserialized.text == "Hello World"
|
||||
assert deserialized.author_name == "TestBot"
|
||||
assert deserialized.message_id == "complex_msg"
|
||||
@@ -444,7 +444,7 @@ class TestRedisChatMessageStore:
|
||||
store = RedisChatMessageStore(redis_url="redis://localhost:6379", thread_id="test123")
|
||||
store._redis_client = mock_client
|
||||
|
||||
message = ChatMessage(role=Role.USER, text="Test")
|
||||
message = ChatMessage("user", ["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=Role.USER, text="Updated message")
|
||||
new_message = ChatMessage("user", ["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=Role.USER, text="Test")
|
||||
new_message = ChatMessage("user", ["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=Role.USER, text="Appended message")
|
||||
message = ChatMessage("user", ["Appended message"])
|
||||
await redis_store.append(message)
|
||||
|
||||
# Should call pipeline operations via add_messages
|
||||
|
||||
@@ -5,7 +5,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
from agent_framework import ChatMessage, Role
|
||||
from agent_framework import ChatMessage
|
||||
from agent_framework.exceptions import AgentException, ServiceInitializationError
|
||||
from redisvl.utils.vectorize import CustomTextVectorizer
|
||||
|
||||
@@ -115,16 +115,16 @@ class TestRedisProviderMessages:
|
||||
@pytest.fixture
|
||||
def sample_messages(self) -> list[ChatMessage]:
|
||||
return [
|
||||
ChatMessage(role=Role.USER, text="Hello, how are you?"),
|
||||
ChatMessage(role=Role.ASSISTANT, text="I'm doing well, thank you!"),
|
||||
ChatMessage(role=Role.SYSTEM, text="You are a helpful assistant"),
|
||||
ChatMessage("user", ["Hello, how are you?"]),
|
||||
ChatMessage("assistant", ["I'm doing well, thank you!"]),
|
||||
ChatMessage("system", ["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=Role.USER, text="Hello"))
|
||||
await provider.invoked("thread123", ChatMessage("user", ["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=Role.USER, text="Hi"))
|
||||
await provider.invoking(ChatMessage("user", ["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=Role.USER, text="q1")])
|
||||
ctx = await provider.invoking([ChatMessage("user", ["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=Role.USER, text="any")])
|
||||
ctx = await provider.invoking([ChatMessage("user", ["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=Role.USER, text="hello")])
|
||||
ctx = await provider.invoking([ChatMessage("user", ["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=Role.USER, text="u"),
|
||||
ChatMessage(role=Role.ASSISTANT, text="a"),
|
||||
ChatMessage(role=Role.SYSTEM, text="s"),
|
||||
ChatMessage("user", ["u"]),
|
||||
ChatMessage("assistant", ["a"]),
|
||||
ChatMessage("system", ["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=Role.USER, text=" "),
|
||||
ChatMessage(role=Role.TOOL, text="tool output"),
|
||||
ChatMessage("user", [" "]),
|
||||
ChatMessage("tool", ["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=Role.USER, text="m1"))
|
||||
await provider.invoked(ChatMessage(role=Role.USER, text="m2"))
|
||||
await provider.invoked(ChatMessage("user", ["m1"]))
|
||||
await provider.invoked(ChatMessage("user", ["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=Role.USER, text="q")])
|
||||
await provider.invoking([ChatMessage("user", ["q"])])
|
||||
assert mock_index.create.await_count == 1
|
||||
|
||||
|
||||
@@ -321,7 +321,7 @@ class TestVectorPopulation:
|
||||
vector_field_name="vec",
|
||||
)
|
||||
|
||||
await provider.invoked(ChatMessage(role=Role.USER, text="hello"))
|
||||
await provider.invoked(ChatMessage("user", ["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