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:
@@ -72,7 +72,7 @@ async def main():
|
||||
middleware=[purview_middleware]
|
||||
)
|
||||
|
||||
response = await agent.run(ChatMessage(role=Role.USER, text="Summarize zero trust in one sentence."))
|
||||
response = await agent.run(ChatMessage("user", ["Summarize zero trust in one sentence."]))
|
||||
print(response)
|
||||
|
||||
asyncio.run(main())
|
||||
|
||||
@@ -57,10 +57,10 @@ class PurviewPolicyMiddleware(AgentMiddleware):
|
||||
context.messages, Activity.UPLOAD_TEXT
|
||||
)
|
||||
if should_block_prompt:
|
||||
from agent_framework import AgentResponse, ChatMessage, Role
|
||||
from agent_framework import AgentResponse, ChatMessage
|
||||
|
||||
context.result = AgentResponse(
|
||||
messages=[ChatMessage(role=Role.SYSTEM, text=self._settings.blocked_prompt_message)]
|
||||
messages=[ChatMessage("system", [self._settings.blocked_prompt_message])]
|
||||
)
|
||||
context.terminate = True
|
||||
return
|
||||
@@ -85,10 +85,10 @@ class PurviewPolicyMiddleware(AgentMiddleware):
|
||||
user_id=resolved_user_id,
|
||||
)
|
||||
if should_block_response:
|
||||
from agent_framework import AgentResponse, ChatMessage, Role
|
||||
from agent_framework import AgentResponse, ChatMessage
|
||||
|
||||
context.result = AgentResponse(
|
||||
messages=[ChatMessage(role=Role.SYSTEM, text=self._settings.blocked_response_message)]
|
||||
messages=[ChatMessage("system", [self._settings.blocked_response_message])]
|
||||
)
|
||||
else:
|
||||
# Streaming responses are not supported for post-checks
|
||||
@@ -149,7 +149,7 @@ class PurviewChatPolicyMiddleware(ChatMiddleware):
|
||||
if should_block_prompt:
|
||||
from agent_framework import ChatMessage, ChatResponse
|
||||
|
||||
blocked_message = ChatMessage(role="system", text=self._settings.blocked_prompt_message)
|
||||
blocked_message = ChatMessage("system", [self._settings.blocked_prompt_message])
|
||||
context.result = ChatResponse(messages=[blocked_message])
|
||||
context.terminate = True
|
||||
return
|
||||
@@ -177,7 +177,7 @@ class PurviewChatPolicyMiddleware(ChatMiddleware):
|
||||
if should_block_response:
|
||||
from agent_framework import ChatMessage, ChatResponse
|
||||
|
||||
blocked_message = ChatMessage(role="system", text=self._settings.blocked_response_message)
|
||||
blocked_message = ChatMessage("system", [self._settings.blocked_response_message])
|
||||
context.result = ChatResponse(messages=[blocked_message])
|
||||
else:
|
||||
logger.debug("Streaming responses are not supported for Purview policy post-checks")
|
||||
|
||||
@@ -5,7 +5,7 @@ from dataclasses import dataclass
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from agent_framework import ChatContext, ChatMessage, Role
|
||||
from agent_framework import ChatContext, ChatMessage
|
||||
from azure.core.credentials import AccessToken
|
||||
|
||||
from agent_framework_purview import PurviewChatPolicyMiddleware, PurviewSettings
|
||||
@@ -36,9 +36,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
chat_client = DummyChatClient()
|
||||
chat_options = MagicMock()
|
||||
chat_options.model = "test-model"
|
||||
return ChatContext(
|
||||
chat_client=chat_client, messages=[ChatMessage(role=Role.USER, text="Hello")], options=chat_options
|
||||
)
|
||||
return ChatContext(chat_client=chat_client, messages=[ChatMessage("user", ["Hello"])], options=chat_options)
|
||||
|
||||
async def test_initialization(self, middleware: PurviewChatPolicyMiddleware) -> None:
|
||||
assert middleware._client is not None
|
||||
@@ -56,14 +54,14 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
class Result:
|
||||
def __init__(self):
|
||||
self.messages = [ChatMessage(role=Role.ASSISTANT, text="Hi there")]
|
||||
self.messages = [ChatMessage("assistant", ["Hi there"])]
|
||||
|
||||
ctx.result = Result()
|
||||
|
||||
await middleware.process(chat_context, mock_next)
|
||||
assert next_called
|
||||
assert mock_proc.call_count == 2
|
||||
assert chat_context.result.messages[0].role == Role.ASSISTANT
|
||||
assert chat_context.result.messages[0].role == "assistant"
|
||||
|
||||
async def test_blocks_prompt(self, middleware: PurviewChatPolicyMiddleware, chat_context: ChatContext) -> None:
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(True, "user-123")):
|
||||
@@ -76,7 +74,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
assert chat_context.result
|
||||
assert hasattr(chat_context.result, "messages")
|
||||
msg = chat_context.result.messages[0]
|
||||
assert msg.role in ("system", Role.SYSTEM)
|
||||
assert msg.role in ("system", "system")
|
||||
assert "blocked" in msg.text.lower()
|
||||
|
||||
async def test_blocks_response(self, middleware: PurviewChatPolicyMiddleware, chat_context: ChatContext) -> None:
|
||||
@@ -92,7 +90,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
class Result:
|
||||
def __init__(self):
|
||||
self.messages = [ChatMessage(role=Role.ASSISTANT, text="Sensitive output")] # pragma: no cover
|
||||
self.messages = [ChatMessage("assistant", ["Sensitive output"])] # pragma: no cover
|
||||
|
||||
ctx.result = Result()
|
||||
|
||||
@@ -100,7 +98,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
assert call_state["count"] == 2
|
||||
msgs = getattr(chat_context.result, "messages", None) or chat_context.result
|
||||
first_msg = msgs[0]
|
||||
assert first_msg.role in ("system", Role.SYSTEM)
|
||||
assert first_msg.role in ("system", "system")
|
||||
assert "blocked" in first_msg.text.lower()
|
||||
|
||||
async def test_streaming_skips_post_check(self, middleware: PurviewChatPolicyMiddleware) -> None:
|
||||
@@ -109,7 +107,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
chat_options.model = "test-model"
|
||||
streaming_context = ChatContext(
|
||||
chat_client=chat_client,
|
||||
messages=[ChatMessage(role=Role.USER, text="Hello")],
|
||||
messages=[ChatMessage("user", ["Hello"])],
|
||||
options=chat_options,
|
||||
is_streaming=True,
|
||||
)
|
||||
@@ -141,7 +139,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [ChatMessage(role=Role.ASSISTANT, text="Response")]
|
||||
result.messages = [ChatMessage("assistant", ["Response"])]
|
||||
ctx.result = result
|
||||
|
||||
await middleware.process(chat_context, mock_next)
|
||||
@@ -165,7 +163,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [ChatMessage(role=Role.ASSISTANT, text="Response")]
|
||||
result.messages = [ChatMessage("assistant", ["Response"])]
|
||||
ctx.result = result
|
||||
|
||||
await middleware.process(chat_context, mock_next)
|
||||
@@ -188,9 +186,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
chat_client = DummyChatClient()
|
||||
chat_options = MagicMock()
|
||||
chat_options.model = "test-model"
|
||||
context = ChatContext(
|
||||
chat_client=chat_client, messages=[ChatMessage(role=Role.USER, text="Hello")], options=chat_options
|
||||
)
|
||||
context = ChatContext(chat_client=chat_client, messages=[ChatMessage("user", ["Hello"])], options=chat_options)
|
||||
|
||||
async def mock_process_messages(*args, **kwargs):
|
||||
raise PurviewPaymentRequiredError("Payment required")
|
||||
@@ -214,9 +210,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
chat_client = DummyChatClient()
|
||||
chat_options = MagicMock()
|
||||
chat_options.model = "test-model"
|
||||
context = ChatContext(
|
||||
chat_client=chat_client, messages=[ChatMessage(role=Role.USER, text="Hello")], options=chat_options
|
||||
)
|
||||
context = ChatContext(chat_client=chat_client, messages=[ChatMessage("user", ["Hello"])], options=chat_options)
|
||||
|
||||
call_count = 0
|
||||
|
||||
@@ -231,7 +225,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [ChatMessage(role=Role.ASSISTANT, text="OK")]
|
||||
result.messages = [ChatMessage("assistant", ["OK"])]
|
||||
ctx.result = result
|
||||
|
||||
with pytest.raises(PurviewPaymentRequiredError):
|
||||
@@ -247,9 +241,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
chat_client = DummyChatClient()
|
||||
chat_options = MagicMock()
|
||||
chat_options.model = "test-model"
|
||||
context = ChatContext(
|
||||
chat_client=chat_client, messages=[ChatMessage(role=Role.USER, text="Hello")], options=chat_options
|
||||
)
|
||||
context = ChatContext(chat_client=chat_client, messages=[ChatMessage("user", ["Hello"])], options=chat_options)
|
||||
|
||||
async def mock_process_messages(*args, **kwargs):
|
||||
raise PurviewPaymentRequiredError("Payment required")
|
||||
@@ -258,7 +250,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [ChatMessage(role=Role.ASSISTANT, text="Response")]
|
||||
result.messages = [ChatMessage("assistant", ["Response"])]
|
||||
context.result = result
|
||||
|
||||
# Should not raise, just log
|
||||
@@ -289,9 +281,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
chat_client = DummyChatClient()
|
||||
chat_options = MagicMock()
|
||||
chat_options.model = "test-model"
|
||||
context = ChatContext(
|
||||
chat_client=chat_client, messages=[ChatMessage(role=Role.USER, text="Hello")], options=chat_options
|
||||
)
|
||||
context = ChatContext(chat_client=chat_client, messages=[ChatMessage("user", ["Hello"])], options=chat_options)
|
||||
|
||||
async def mock_process_messages(*args, **kwargs):
|
||||
raise ValueError("Some error")
|
||||
@@ -300,7 +290,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [ChatMessage(role=Role.ASSISTANT, text="Response")]
|
||||
result.messages = [ChatMessage("assistant", ["Response"])]
|
||||
context.result = result
|
||||
|
||||
# Should not raise, just log
|
||||
@@ -318,9 +308,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
chat_client = DummyChatClient()
|
||||
chat_options = MagicMock()
|
||||
chat_options.model = "test-model"
|
||||
context = ChatContext(
|
||||
chat_client=chat_client, messages=[ChatMessage(role=Role.USER, text="Hello")], options=chat_options
|
||||
)
|
||||
context = ChatContext(chat_client=chat_client, messages=[ChatMessage("user", ["Hello"])], options=chat_options)
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=ValueError("boom")):
|
||||
|
||||
@@ -340,9 +328,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
chat_client = DummyChatClient()
|
||||
chat_options = MagicMock()
|
||||
chat_options.model = "test-model"
|
||||
context = ChatContext(
|
||||
chat_client=chat_client, messages=[ChatMessage(role=Role.USER, text="Hello")], options=chat_options
|
||||
)
|
||||
context = ChatContext(chat_client=chat_client, messages=[ChatMessage("user", ["Hello"])], options=chat_options)
|
||||
|
||||
call_count = 0
|
||||
|
||||
@@ -357,7 +343,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [ChatMessage(role=Role.ASSISTANT, text="OK")]
|
||||
result.messages = [ChatMessage("assistant", ["OK"])]
|
||||
ctx.result = result
|
||||
|
||||
with pytest.raises(ValueError, match="post"):
|
||||
|
||||
@@ -5,7 +5,7 @@
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from agent_framework import AgentResponse, AgentRunContext, ChatMessage, Role
|
||||
from agent_framework import AgentResponse, AgentRunContext, ChatMessage
|
||||
from azure.core.credentials import AccessToken
|
||||
|
||||
from agent_framework_purview import PurviewPolicyMiddleware, PurviewSettings
|
||||
@@ -49,7 +49,7 @@ class TestPurviewPolicyMiddleware:
|
||||
self, middleware: PurviewPolicyMiddleware, mock_agent: MagicMock
|
||||
) -> None:
|
||||
"""Test middleware allows prompt that passes policy check."""
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage(role=Role.USER, text="Hello, how are you?")])
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage("user", ["Hello, how are you?"])])
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")):
|
||||
next_called = False
|
||||
@@ -57,7 +57,7 @@ class TestPurviewPolicyMiddleware:
|
||||
async def mock_next(ctx: AgentRunContext) -> None:
|
||||
nonlocal next_called
|
||||
next_called = True
|
||||
ctx.result = AgentResponse(messages=[ChatMessage(role=Role.ASSISTANT, text="I'm good, thanks!")])
|
||||
ctx.result = AgentResponse(messages=[ChatMessage("assistant", ["I'm good, thanks!"])])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -69,9 +69,7 @@ class TestPurviewPolicyMiddleware:
|
||||
self, middleware: PurviewPolicyMiddleware, mock_agent: MagicMock
|
||||
) -> None:
|
||||
"""Test middleware blocks prompt that violates policy."""
|
||||
context = AgentRunContext(
|
||||
agent=mock_agent, messages=[ChatMessage(role=Role.USER, text="Sensitive information")]
|
||||
)
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage("user", ["Sensitive information"])])
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(True, "user-123")):
|
||||
next_called = False
|
||||
@@ -86,12 +84,12 @@ class TestPurviewPolicyMiddleware:
|
||||
assert context.result is not None
|
||||
assert context.terminate
|
||||
assert len(context.result.messages) == 1
|
||||
assert context.result.messages[0].role == Role.SYSTEM
|
||||
assert context.result.messages[0].role == "system"
|
||||
assert "blocked by policy" in context.result.messages[0].text.lower()
|
||||
|
||||
async def test_middleware_checks_response(self, middleware: PurviewPolicyMiddleware, mock_agent: MagicMock) -> None:
|
||||
"""Test middleware checks agent response for policy violations."""
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage(role=Role.USER, text="Hello")])
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage("user", ["Hello"])])
|
||||
|
||||
call_count = 0
|
||||
|
||||
@@ -104,16 +102,14 @@ class TestPurviewPolicyMiddleware:
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
|
||||
|
||||
async def mock_next(ctx: AgentRunContext) -> None:
|
||||
ctx.result = AgentResponse(
|
||||
messages=[ChatMessage(role=Role.ASSISTANT, text="Here's some sensitive information")]
|
||||
)
|
||||
ctx.result = AgentResponse(messages=[ChatMessage("assistant", ["Here's some sensitive information"])])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
assert call_count == 2
|
||||
assert context.result is not None
|
||||
assert len(context.result.messages) == 1
|
||||
assert context.result.messages[0].role == Role.SYSTEM
|
||||
assert context.result.messages[0].role == "system"
|
||||
assert "blocked by policy" in context.result.messages[0].text.lower()
|
||||
|
||||
async def test_middleware_handles_result_without_messages(
|
||||
@@ -123,7 +119,7 @@ class TestPurviewPolicyMiddleware:
|
||||
# Set ignore_exceptions to True so AttributeError is caught and logged
|
||||
middleware._settings.ignore_exceptions = True
|
||||
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage(role=Role.USER, text="Hello")])
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage("user", ["Hello"])])
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")):
|
||||
|
||||
@@ -140,12 +136,12 @@ class TestPurviewPolicyMiddleware:
|
||||
"""Test middleware passes correct activity type to processor."""
|
||||
from agent_framework_purview._models import Activity
|
||||
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage(role=Role.USER, text="Test")])
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage("user", ["Test"])])
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_process:
|
||||
|
||||
async def mock_next(ctx: AgentRunContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[ChatMessage(role=Role.ASSISTANT, text="Response")])
|
||||
ctx.result = AgentResponse(messages=[ChatMessage("assistant", ["Response"])])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -157,13 +153,13 @@ class TestPurviewPolicyMiddleware:
|
||||
self, middleware: PurviewPolicyMiddleware, mock_agent: MagicMock
|
||||
) -> None:
|
||||
"""Test that streaming results skip post-check evaluation."""
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage(role=Role.USER, text="Hello")])
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage("user", ["Hello"])])
|
||||
context.is_streaming = True
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
|
||||
async def mock_next(ctx: AgentRunContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[ChatMessage(role=Role.ASSISTANT, text="streaming")])
|
||||
ctx.result = AgentResponse(messages=[ChatMessage("assistant", ["streaming"])])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -175,7 +171,7 @@ class TestPurviewPolicyMiddleware:
|
||||
"""Test that 402 in pre-check is raised when ignore_payment_required=False."""
|
||||
from agent_framework_purview._exceptions import PurviewPaymentRequiredError
|
||||
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage(role=Role.USER, text="Hello")])
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage("user", ["Hello"])])
|
||||
|
||||
with patch.object(
|
||||
middleware._processor,
|
||||
@@ -195,7 +191,7 @@ class TestPurviewPolicyMiddleware:
|
||||
"""Test that 402 in post-check is raised when ignore_payment_required=False."""
|
||||
from agent_framework_purview._exceptions import PurviewPaymentRequiredError
|
||||
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage(role=Role.USER, text="Hello")])
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage("user", ["Hello"])])
|
||||
|
||||
call_count = 0
|
||||
|
||||
@@ -209,7 +205,7 @@ class TestPurviewPolicyMiddleware:
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=side_effect):
|
||||
|
||||
async def mock_next(ctx: AgentRunContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[ChatMessage(role=Role.ASSISTANT, text="OK")])
|
||||
ctx.result = AgentResponse(messages=[ChatMessage("assistant", ["OK"])])
|
||||
|
||||
with pytest.raises(PurviewPaymentRequiredError):
|
||||
await middleware.process(context, mock_next)
|
||||
@@ -220,7 +216,7 @@ class TestPurviewPolicyMiddleware:
|
||||
"""Test that post-check exceptions are propagated when ignore_exceptions=False."""
|
||||
middleware._settings.ignore_exceptions = False
|
||||
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage(role=Role.USER, text="Hello")])
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage("user", ["Hello"])])
|
||||
|
||||
call_count = 0
|
||||
|
||||
@@ -234,7 +230,7 @@ class TestPurviewPolicyMiddleware:
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=side_effect):
|
||||
|
||||
async def mock_next(ctx: AgentRunContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[ChatMessage(role=Role.ASSISTANT, text="OK")])
|
||||
ctx.result = AgentResponse(messages=[ChatMessage("assistant", ["OK"])])
|
||||
|
||||
with pytest.raises(ValueError, match="Post-check blew up"):
|
||||
await middleware.process(context, mock_next)
|
||||
@@ -246,14 +242,14 @@ class TestPurviewPolicyMiddleware:
|
||||
# Set ignore_exceptions to True
|
||||
middleware._settings.ignore_exceptions = True
|
||||
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage(role=Role.USER, text="Test")])
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage("user", ["Test"])])
|
||||
|
||||
with patch.object(
|
||||
middleware._processor, "process_messages", side_effect=Exception("Pre-check error")
|
||||
) as mock_process:
|
||||
|
||||
async def mock_next(ctx: AgentRunContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[ChatMessage(role=Role.ASSISTANT, text="Response")])
|
||||
ctx.result = AgentResponse(messages=[ChatMessage("assistant", ["Response"])])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -271,7 +267,7 @@ class TestPurviewPolicyMiddleware:
|
||||
# Set ignore_exceptions to True
|
||||
middleware._settings.ignore_exceptions = True
|
||||
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage(role=Role.USER, text="Test")])
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage("user", ["Test"])])
|
||||
|
||||
call_count = 0
|
||||
|
||||
@@ -285,7 +281,7 @@ class TestPurviewPolicyMiddleware:
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
|
||||
|
||||
async def mock_next(ctx: AgentRunContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[ChatMessage(role=Role.ASSISTANT, text="Response")])
|
||||
ctx.result = AgentResponse(messages=[ChatMessage("assistant", ["Response"])])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -302,7 +298,7 @@ class TestPurviewPolicyMiddleware:
|
||||
|
||||
mock_agent = MagicMock()
|
||||
mock_agent.name = "test-agent"
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage(role=Role.USER, text="Test")])
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage("user", ["Test"])])
|
||||
|
||||
# Mock processor to raise an exception
|
||||
async def mock_process_messages(*args, **kwargs):
|
||||
@@ -311,7 +307,7 @@ class TestPurviewPolicyMiddleware:
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
|
||||
|
||||
async def mock_next(ctx):
|
||||
ctx.result = AgentResponse(messages=[ChatMessage(role=Role.ASSISTANT, text="Response")])
|
||||
ctx.result = AgentResponse(messages=[ChatMessage("assistant", ["Response"])])
|
||||
|
||||
# Should not raise, just log
|
||||
await middleware.process(context, mock_next)
|
||||
@@ -326,7 +322,7 @@ class TestPurviewPolicyMiddleware:
|
||||
|
||||
mock_agent = MagicMock()
|
||||
mock_agent.name = "test-agent"
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage(role=Role.USER, text="Test")])
|
||||
context = AgentRunContext(agent=mock_agent, messages=[ChatMessage("user", ["Test"])])
|
||||
|
||||
# Mock processor to raise an exception
|
||||
async def mock_process_messages(*args, **kwargs):
|
||||
|
||||
@@ -5,7 +5,7 @@
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from agent_framework import ChatMessage, Role
|
||||
from agent_framework import ChatMessage
|
||||
|
||||
from agent_framework_purview import PurviewAppLocation, PurviewLocationType, PurviewSettings
|
||||
from agent_framework_purview._models import (
|
||||
@@ -83,8 +83,8 @@ class TestScopedContentProcessor:
|
||||
async def test_process_messages_with_defaults(self, processor: ScopedContentProcessor) -> None:
|
||||
"""Test process_messages with settings that have defaults."""
|
||||
messages = [
|
||||
ChatMessage(role=Role.USER, text="Hello"),
|
||||
ChatMessage(role=Role.ASSISTANT, text="Hi there"),
|
||||
ChatMessage("user", ["Hello"]),
|
||||
ChatMessage("assistant", ["Hi there"]),
|
||||
]
|
||||
|
||||
with patch.object(processor, "_map_messages", return_value=([], None)) as mock_map:
|
||||
@@ -98,7 +98,7 @@ class TestScopedContentProcessor:
|
||||
self, processor: ScopedContentProcessor, process_content_request_factory
|
||||
) -> None:
|
||||
"""Test process_messages returns True when content should be blocked."""
|
||||
messages = [ChatMessage(role=Role.USER, text="Sensitive content")]
|
||||
messages = [ChatMessage("user", ["Sensitive content"])]
|
||||
|
||||
mock_request = process_content_request_factory("Sensitive content")
|
||||
|
||||
@@ -121,7 +121,7 @@ class TestScopedContentProcessor:
|
||||
"""Test _map_messages creates ProcessContentRequest objects."""
|
||||
messages = [
|
||||
ChatMessage(
|
||||
role=Role.USER,
|
||||
role="user",
|
||||
text="Test message",
|
||||
message_id="msg-123",
|
||||
author_name="12345678-1234-1234-1234-123456789012",
|
||||
@@ -139,7 +139,7 @@ class TestScopedContentProcessor:
|
||||
"""Test _map_messages gets token info when settings lack some defaults."""
|
||||
settings = PurviewSettings(app_name="Test App", tenant_id="12345678-1234-1234-1234-123456789012")
|
||||
processor = ScopedContentProcessor(mock_client, settings)
|
||||
messages = [ChatMessage(role=Role.USER, text="Test", message_id="msg-123")]
|
||||
messages = [ChatMessage("user", ["Test"], message_id="msg-123")]
|
||||
|
||||
requests, user_id = await processor._map_messages(messages, Activity.UPLOAD_TEXT)
|
||||
|
||||
@@ -156,7 +156,7 @@ class TestScopedContentProcessor:
|
||||
return_value={"user_id": "test-user", "client_id": "test-client"}
|
||||
)
|
||||
|
||||
messages = [ChatMessage(role=Role.USER, text="Test", message_id="msg-123")]
|
||||
messages = [ChatMessage("user", ["Test"], message_id="msg-123")]
|
||||
|
||||
with pytest.raises(ValueError, match="Tenant id required"):
|
||||
await processor._map_messages(messages, Activity.UPLOAD_TEXT)
|
||||
@@ -332,7 +332,7 @@ class TestScopedContentProcessor:
|
||||
|
||||
messages = [
|
||||
ChatMessage(
|
||||
role=Role.USER,
|
||||
role="user",
|
||||
text="Test message",
|
||||
additional_properties={"user_id": "22345678-1234-1234-1234-123456789012"},
|
||||
),
|
||||
@@ -355,7 +355,7 @@ class TestScopedContentProcessor:
|
||||
)
|
||||
processor = ScopedContentProcessor(mock_client, settings)
|
||||
|
||||
messages = [ChatMessage(role=Role.USER, text="Test message")]
|
||||
messages = [ChatMessage("user", ["Test message"])]
|
||||
|
||||
requests, user_id = await processor._map_messages(
|
||||
messages, Activity.UPLOAD_TEXT, provided_user_id="32345678-1234-1234-1234-123456789012"
|
||||
@@ -376,7 +376,7 @@ class TestScopedContentProcessor:
|
||||
)
|
||||
processor = ScopedContentProcessor(mock_client, settings)
|
||||
|
||||
messages = [ChatMessage(role=Role.USER, text="Test message")]
|
||||
messages = [ChatMessage("user", ["Test message"])]
|
||||
|
||||
requests, user_id = await processor._map_messages(messages, Activity.UPLOAD_TEXT)
|
||||
|
||||
@@ -479,7 +479,7 @@ class TestUserIdResolution:
|
||||
settings = PurviewSettings(app_name="Test App") # No tenant_id or app_location
|
||||
processor = ScopedContentProcessor(mock_client, settings)
|
||||
|
||||
messages = [ChatMessage(role=Role.USER, text="Test")]
|
||||
messages = [ChatMessage("user", ["Test"])]
|
||||
|
||||
requests, user_id = await processor._map_messages(messages, Activity.UPLOAD_TEXT)
|
||||
|
||||
@@ -494,7 +494,7 @@ class TestUserIdResolution:
|
||||
|
||||
messages = [
|
||||
ChatMessage(
|
||||
role=Role.USER,
|
||||
role="user",
|
||||
text="Test",
|
||||
additional_properties={"user_id": "22222222-2222-2222-2222-222222222222"},
|
||||
)
|
||||
@@ -514,7 +514,7 @@ class TestUserIdResolution:
|
||||
|
||||
messages = [
|
||||
ChatMessage(
|
||||
role=Role.USER,
|
||||
role="user",
|
||||
text="Test",
|
||||
author_name="33333333-3333-3333-3333-333333333333",
|
||||
)
|
||||
@@ -532,7 +532,7 @@ class TestUserIdResolution:
|
||||
|
||||
messages = [
|
||||
ChatMessage(
|
||||
role=Role.USER,
|
||||
role="user",
|
||||
text="Test",
|
||||
author_name="John Doe", # Not a GUID
|
||||
)
|
||||
@@ -550,7 +550,7 @@ class TestUserIdResolution:
|
||||
"""Test provided_user_id parameter is used as last resort."""
|
||||
processor = ScopedContentProcessor(mock_client, settings)
|
||||
|
||||
messages = [ChatMessage(role=Role.USER, text="Test")]
|
||||
messages = [ChatMessage("user", ["Test"])]
|
||||
|
||||
requests, user_id = await processor._map_messages(
|
||||
messages, Activity.UPLOAD_TEXT, provided_user_id="44444444-4444-4444-4444-444444444444"
|
||||
@@ -562,7 +562,7 @@ class TestUserIdResolution:
|
||||
"""Test invalid provided_user_id is ignored."""
|
||||
processor = ScopedContentProcessor(mock_client, settings)
|
||||
|
||||
messages = [ChatMessage(role=Role.USER, text="Test")]
|
||||
messages = [ChatMessage("user", ["Test"])]
|
||||
|
||||
requests, user_id = await processor._map_messages(messages, Activity.UPLOAD_TEXT, provided_user_id="not-a-guid")
|
||||
|
||||
@@ -575,10 +575,10 @@ class TestUserIdResolution:
|
||||
|
||||
messages = [
|
||||
ChatMessage(
|
||||
role=Role.USER, text="First", additional_properties={"user_id": "55555555-5555-5555-5555-555555555555"}
|
||||
role="user", text="First", additional_properties={"user_id": "55555555-5555-5555-5555-555555555555"}
|
||||
),
|
||||
ChatMessage(role=Role.ASSISTANT, text="Response"),
|
||||
ChatMessage(role=Role.USER, text="Second"),
|
||||
ChatMessage("assistant", ["Response"]),
|
||||
ChatMessage("user", ["Second"]),
|
||||
]
|
||||
|
||||
requests, user_id = await processor._map_messages(messages, Activity.UPLOAD_TEXT)
|
||||
@@ -594,14 +594,14 @@ class TestUserIdResolution:
|
||||
processor = ScopedContentProcessor(mock_client, settings)
|
||||
|
||||
messages = [
|
||||
ChatMessage(role=Role.USER, text="First", author_name="Not a GUID"),
|
||||
ChatMessage("user", ["First"], author_name="Not a GUID"),
|
||||
ChatMessage(
|
||||
role=Role.ASSISTANT,
|
||||
role="assistant",
|
||||
text="Response",
|
||||
additional_properties={"user_id": "66666666-6666-6666-6666-666666666666"},
|
||||
),
|
||||
ChatMessage(
|
||||
role=Role.USER, text="Third", additional_properties={"user_id": "77777777-7777-7777-7777-777777777777"}
|
||||
role="user", text="Third", additional_properties={"user_id": "77777777-7777-7777-7777-777777777777"}
|
||||
),
|
||||
]
|
||||
|
||||
@@ -654,7 +654,7 @@ class TestScopedContentProcessorCaching:
|
||||
scope_identifier="scope-123", scopes=[]
|
||||
)
|
||||
|
||||
messages = [ChatMessage(role=Role.USER, text="Test")]
|
||||
messages = [ChatMessage("user", ["Test"])]
|
||||
|
||||
await processor.process_messages(messages, Activity.UPLOAD_TEXT, user_id="12345678-1234-1234-1234-123456789012")
|
||||
|
||||
@@ -676,7 +676,7 @@ class TestScopedContentProcessorCaching:
|
||||
|
||||
mock_client.get_protection_scopes.side_effect = PurviewPaymentRequiredError("Payment required")
|
||||
|
||||
messages = [ChatMessage(role=Role.USER, text="Test")]
|
||||
messages = [ChatMessage("user", ["Test"])]
|
||||
|
||||
with pytest.raises(PurviewPaymentRequiredError):
|
||||
await processor.process_messages(
|
||||
|
||||
Reference in New Issue
Block a user