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:
co-authored by
Evan Mattson
parent
a4c9e43afb
commit
0521f5bed8
+35
-49
@@ -5,7 +5,7 @@ from dataclasses import dataclass
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from agent_framework import ChatContext, ChatMessage, MiddlewareTermination
|
||||
from agent_framework import ChatContext, Message, MiddlewareTermination
|
||||
from azure.core.credentials import AccessToken
|
||||
|
||||
from agent_framework_purview import PurviewChatPolicyMiddleware, PurviewSettings
|
||||
@@ -34,12 +34,10 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
@pytest.fixture
|
||||
def chat_context(self) -> ChatContext:
|
||||
chat_client = DummyChatClient()
|
||||
client = DummyChatClient()
|
||||
chat_options = MagicMock()
|
||||
chat_options.model = "test-model"
|
||||
return ChatContext(
|
||||
chat_client=chat_client, messages=[ChatMessage(role="user", text="Hello")], options=chat_options
|
||||
)
|
||||
return ChatContext(client=client, messages=[Message(role="user", text="Hello")], options=chat_options)
|
||||
|
||||
async def test_initialization(self, middleware: PurviewChatPolicyMiddleware) -> None:
|
||||
assert middleware._client is not None
|
||||
@@ -57,7 +55,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
class Result:
|
||||
def __init__(self):
|
||||
self.messages = [ChatMessage(role="assistant", text="Hi there")]
|
||||
self.messages = [Message(role="assistant", text="Hi there")]
|
||||
|
||||
ctx.result = Result()
|
||||
|
||||
@@ -93,7 +91,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
class Result:
|
||||
def __init__(self):
|
||||
self.messages = [ChatMessage(role="assistant", text="Sensitive output")] # pragma: no cover
|
||||
self.messages = [Message(role="assistant", text="Sensitive output")] # pragma: no cover
|
||||
|
||||
ctx.result = Result()
|
||||
|
||||
@@ -105,12 +103,12 @@ class TestPurviewChatPolicyMiddleware:
|
||||
assert "blocked" in first_msg.text.lower()
|
||||
|
||||
async def test_streaming_skips_post_check(self, middleware: PurviewChatPolicyMiddleware) -> None:
|
||||
chat_client = DummyChatClient()
|
||||
client = DummyChatClient()
|
||||
chat_options = MagicMock()
|
||||
chat_options.model = "test-model"
|
||||
streaming_context = ChatContext(
|
||||
chat_client=chat_client,
|
||||
messages=[ChatMessage(role="user", text="Hello")],
|
||||
client=client,
|
||||
messages=[Message(role="user", text="Hello")],
|
||||
options=chat_options,
|
||||
stream=True,
|
||||
)
|
||||
@@ -142,7 +140,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [ChatMessage(role="assistant", text="Response")]
|
||||
result.messages = [Message(role="assistant", text="Response")]
|
||||
ctx.result = result
|
||||
|
||||
await middleware.process(chat_context, mock_next)
|
||||
@@ -166,7 +164,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [ChatMessage(role="assistant", text="Response")]
|
||||
result.messages = [Message(role="assistant", text="Response")]
|
||||
ctx.result = result
|
||||
|
||||
await middleware.process(chat_context, mock_next)
|
||||
@@ -186,12 +184,10 @@ class TestPurviewChatPolicyMiddleware:
|
||||
settings = PurviewSettings(app_name="Test App", ignore_payment_required=False)
|
||||
middleware = PurviewChatPolicyMiddleware(mock_credential, settings)
|
||||
|
||||
chat_client = DummyChatClient()
|
||||
client = DummyChatClient()
|
||||
chat_options = MagicMock()
|
||||
chat_options.model = "test-model"
|
||||
context = ChatContext(
|
||||
chat_client=chat_client, messages=[ChatMessage(role="user", text="Hello")], options=chat_options
|
||||
)
|
||||
context = ChatContext(client=client, messages=[Message(role="user", text="Hello")], options=chat_options)
|
||||
|
||||
async def mock_process_messages(*args, **kwargs):
|
||||
raise PurviewPaymentRequiredError("Payment required")
|
||||
@@ -212,12 +208,10 @@ class TestPurviewChatPolicyMiddleware:
|
||||
settings = PurviewSettings(app_name="Test App", ignore_payment_required=False)
|
||||
middleware = PurviewChatPolicyMiddleware(mock_credential, settings)
|
||||
|
||||
chat_client = DummyChatClient()
|
||||
client = DummyChatClient()
|
||||
chat_options = MagicMock()
|
||||
chat_options.model = "test-model"
|
||||
context = ChatContext(
|
||||
chat_client=chat_client, messages=[ChatMessage(role="user", text="Hello")], options=chat_options
|
||||
)
|
||||
context = ChatContext(client=client, messages=[Message(role="user", text="Hello")], options=chat_options)
|
||||
|
||||
call_count = 0
|
||||
|
||||
@@ -232,7 +226,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [ChatMessage(role="assistant", text="OK")]
|
||||
result.messages = [Message(role="assistant", text="OK")]
|
||||
ctx.result = result
|
||||
|
||||
with pytest.raises(PurviewPaymentRequiredError):
|
||||
@@ -245,12 +239,10 @@ class TestPurviewChatPolicyMiddleware:
|
||||
settings = PurviewSettings(app_name="Test App", ignore_payment_required=True)
|
||||
middleware = PurviewChatPolicyMiddleware(mock_credential, settings)
|
||||
|
||||
chat_client = DummyChatClient()
|
||||
client = DummyChatClient()
|
||||
chat_options = MagicMock()
|
||||
chat_options.model = "test-model"
|
||||
context = ChatContext(
|
||||
chat_client=chat_client, messages=[ChatMessage(role="user", text="Hello")], options=chat_options
|
||||
)
|
||||
context = ChatContext(client=client, messages=[Message(role="user", text="Hello")], options=chat_options)
|
||||
|
||||
async def mock_process_messages(*args, **kwargs):
|
||||
raise PurviewPaymentRequiredError("Payment required")
|
||||
@@ -259,7 +251,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [ChatMessage(role="assistant", text="Response")]
|
||||
result.messages = [Message(role="assistant", text="Response")]
|
||||
context.result = result
|
||||
|
||||
# Should not raise, just log
|
||||
@@ -287,12 +279,10 @@ class TestPurviewChatPolicyMiddleware:
|
||||
settings = PurviewSettings(app_name="Test App", ignore_exceptions=True)
|
||||
middleware = PurviewChatPolicyMiddleware(mock_credential, settings)
|
||||
|
||||
chat_client = DummyChatClient()
|
||||
client = DummyChatClient()
|
||||
chat_options = MagicMock()
|
||||
chat_options.model = "test-model"
|
||||
context = ChatContext(
|
||||
chat_client=chat_client, messages=[ChatMessage(role="user", text="Hello")], options=chat_options
|
||||
)
|
||||
context = ChatContext(client=client, messages=[Message(role="user", text="Hello")], options=chat_options)
|
||||
|
||||
async def mock_process_messages(*args, **kwargs):
|
||||
raise ValueError("Some error")
|
||||
@@ -301,7 +291,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [ChatMessage(role="assistant", text="Response")]
|
||||
result.messages = [Message(role="assistant", text="Response")]
|
||||
context.result = result
|
||||
|
||||
# Should not raise, just log
|
||||
@@ -316,12 +306,10 @@ class TestPurviewChatPolicyMiddleware:
|
||||
settings = PurviewSettings(app_name="Test App", ignore_exceptions=False)
|
||||
middleware = PurviewChatPolicyMiddleware(mock_credential, settings)
|
||||
|
||||
chat_client = DummyChatClient()
|
||||
client = DummyChatClient()
|
||||
chat_options = MagicMock()
|
||||
chat_options.model = "test-model"
|
||||
context = ChatContext(
|
||||
chat_client=chat_client, messages=[ChatMessage(role="user", text="Hello")], options=chat_options
|
||||
)
|
||||
context = ChatContext(client=client, messages=[Message(role="user", text="Hello")], options=chat_options)
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=ValueError("boom")):
|
||||
|
||||
@@ -338,12 +326,10 @@ class TestPurviewChatPolicyMiddleware:
|
||||
settings = PurviewSettings(app_name="Test App", ignore_exceptions=False)
|
||||
middleware = PurviewChatPolicyMiddleware(mock_credential, settings)
|
||||
|
||||
chat_client = DummyChatClient()
|
||||
client = DummyChatClient()
|
||||
chat_options = MagicMock()
|
||||
chat_options.model = "test-model"
|
||||
context = ChatContext(
|
||||
chat_client=chat_client, messages=[ChatMessage(role="user", text="Hello")], options=chat_options
|
||||
)
|
||||
context = ChatContext(client=client, messages=[Message(role="user", text="Hello")], options=chat_options)
|
||||
|
||||
call_count = 0
|
||||
|
||||
@@ -358,7 +344,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [ChatMessage(role="assistant", text="OK")]
|
||||
result.messages = [Message(role="assistant", text="OK")]
|
||||
ctx.result = result
|
||||
|
||||
with pytest.raises(ValueError, match="post"):
|
||||
@@ -369,15 +355,15 @@ class TestPurviewChatPolicyMiddleware:
|
||||
) -> None:
|
||||
"""Test that session_id is extracted from context.options['conversation_id']."""
|
||||
chat_client = DummyChatClient()
|
||||
messages = [ChatMessage(role="user", text="Hello")]
|
||||
messages = [Message(role="user", text="Hello")]
|
||||
options = {"conversation_id": "conv-123", "model": "test-model"}
|
||||
context = ChatContext(chat_client=chat_client, messages=messages, options=options)
|
||||
context = ChatContext(client=chat_client, messages=messages, options=options)
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [ChatMessage(role="assistant", text="Hi")]
|
||||
result.messages = [Message(role="assistant", text="Hi")]
|
||||
ctx.result = result
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
@@ -391,14 +377,14 @@ class TestPurviewChatPolicyMiddleware:
|
||||
) -> None:
|
||||
"""Test that session_id is None when options don't contain conversation_id."""
|
||||
chat_client = DummyChatClient()
|
||||
messages = [ChatMessage(role="user", text="Hello")]
|
||||
context = ChatContext(chat_client=chat_client, messages=messages, options=None)
|
||||
messages = [Message(role="user", text="Hello")]
|
||||
context = ChatContext(client=chat_client, messages=messages, options=None)
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [ChatMessage(role="assistant", text="Hi")]
|
||||
result.messages = [Message(role="assistant", text="Hi")]
|
||||
ctx.result = result
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
@@ -409,15 +395,15 @@ class TestPurviewChatPolicyMiddleware:
|
||||
async def test_chat_middleware_session_id_used_in_post_check(self, middleware: PurviewChatPolicyMiddleware) -> None:
|
||||
"""Test that session_id is passed to post-check process_messages call."""
|
||||
chat_client = DummyChatClient()
|
||||
messages = [ChatMessage(role="user", text="Hello")]
|
||||
messages = [Message(role="user", text="Hello")]
|
||||
options = {"conversation_id": "conv-999"}
|
||||
context = ChatContext(chat_client=chat_client, messages=messages, options=options)
|
||||
context = ChatContext(client=chat_client, messages=messages, options=options)
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [ChatMessage(role="assistant", text="Response")]
|
||||
result.messages = [Message(role="assistant", text="Response")]
|
||||
ctx.result = result
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
+33
-33
@@ -5,7 +5,7 @@
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from agent_framework import AgentContext, AgentResponse, AgentThread, ChatMessage, MiddlewareTermination
|
||||
from agent_framework import AgentContext, AgentResponse, AgentThread, Message, MiddlewareTermination
|
||||
from azure.core.credentials import AccessToken
|
||||
|
||||
from agent_framework_purview import PurviewPolicyMiddleware, PurviewSettings
|
||||
@@ -50,7 +50,7 @@ class TestPurviewPolicyMiddleware:
|
||||
self, middleware: PurviewPolicyMiddleware, mock_agent: MagicMock
|
||||
) -> None:
|
||||
"""Test middleware allows prompt that passes policy check."""
|
||||
context = AgentContext(agent=mock_agent, messages=[ChatMessage(role="user", text="Hello, how are you?")])
|
||||
context = AgentContext(agent=mock_agent, messages=[Message(role="user", text="Hello, how are you?")])
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")):
|
||||
next_called = False
|
||||
@@ -58,7 +58,7 @@ class TestPurviewPolicyMiddleware:
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
nonlocal next_called
|
||||
next_called = True
|
||||
ctx.result = AgentResponse(messages=[ChatMessage(role="assistant", text="I'm good, thanks!")])
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="I'm good, thanks!")])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -69,7 +69,7 @@ class TestPurviewPolicyMiddleware:
|
||||
self, middleware: PurviewPolicyMiddleware, mock_agent: MagicMock
|
||||
) -> None:
|
||||
"""Test middleware blocks prompt that violates policy."""
|
||||
context = AgentContext(agent=mock_agent, messages=[ChatMessage(role="user", text="Sensitive information")])
|
||||
context = AgentContext(agent=mock_agent, messages=[Message(role="user", text="Sensitive information")])
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(True, "user-123")):
|
||||
next_called = False
|
||||
@@ -89,7 +89,7 @@ class TestPurviewPolicyMiddleware:
|
||||
|
||||
async def test_middleware_checks_response(self, middleware: PurviewPolicyMiddleware, mock_agent: MagicMock) -> None:
|
||||
"""Test middleware checks agent response for policy violations."""
|
||||
context = AgentContext(agent=mock_agent, messages=[ChatMessage(role="user", text="Hello")])
|
||||
context = AgentContext(agent=mock_agent, messages=[Message(role="user", text="Hello")])
|
||||
|
||||
call_count = 0
|
||||
|
||||
@@ -103,7 +103,7 @@ class TestPurviewPolicyMiddleware:
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(
|
||||
messages=[ChatMessage(role="assistant", text="Here's some sensitive information")]
|
||||
messages=[Message(role="assistant", text="Here's some sensitive information")]
|
||||
)
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
@@ -121,7 +121,7 @@ class TestPurviewPolicyMiddleware:
|
||||
# Set ignore_exceptions to True so AttributeError is caught and logged
|
||||
middleware._settings.ignore_exceptions = True
|
||||
|
||||
context = AgentContext(agent=mock_agent, messages=[ChatMessage(role="user", text="Hello")])
|
||||
context = AgentContext(agent=mock_agent, messages=[Message(role="user", text="Hello")])
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")):
|
||||
|
||||
@@ -138,12 +138,12 @@ class TestPurviewPolicyMiddleware:
|
||||
"""Test middleware passes correct activity type to processor."""
|
||||
from agent_framework_purview._models import Activity
|
||||
|
||||
context = AgentContext(agent=mock_agent, messages=[ChatMessage(role="user", text="Test")])
|
||||
context = AgentContext(agent=mock_agent, messages=[Message(role="user", text="Test")])
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_process:
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[ChatMessage(role="assistant", text="Response")])
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -155,13 +155,13 @@ class TestPurviewPolicyMiddleware:
|
||||
self, middleware: PurviewPolicyMiddleware, mock_agent: MagicMock
|
||||
) -> None:
|
||||
"""Test that streaming results skip post-check evaluation."""
|
||||
context = AgentContext(agent=mock_agent, messages=[ChatMessage(role="user", text="Hello")])
|
||||
context = AgentContext(agent=mock_agent, messages=[Message(role="user", text="Hello")])
|
||||
context.stream = True
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[ChatMessage(role="assistant", text="streaming")])
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="streaming")])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -173,7 +173,7 @@ class TestPurviewPolicyMiddleware:
|
||||
"""Test that 402 in pre-check is raised when ignore_payment_required=False."""
|
||||
from agent_framework_purview._exceptions import PurviewPaymentRequiredError
|
||||
|
||||
context = AgentContext(agent=mock_agent, messages=[ChatMessage(role="user", text="Hello")])
|
||||
context = AgentContext(agent=mock_agent, messages=[Message(role="user", text="Hello")])
|
||||
|
||||
with patch.object(
|
||||
middleware._processor,
|
||||
@@ -193,7 +193,7 @@ class TestPurviewPolicyMiddleware:
|
||||
"""Test that 402 in post-check is raised when ignore_payment_required=False."""
|
||||
from agent_framework_purview._exceptions import PurviewPaymentRequiredError
|
||||
|
||||
context = AgentContext(agent=mock_agent, messages=[ChatMessage(role="user", text="Hello")])
|
||||
context = AgentContext(agent=mock_agent, messages=[Message(role="user", text="Hello")])
|
||||
|
||||
call_count = 0
|
||||
|
||||
@@ -207,7 +207,7 @@ class TestPurviewPolicyMiddleware:
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=side_effect):
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[ChatMessage(role="assistant", text="OK")])
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="OK")])
|
||||
|
||||
with pytest.raises(PurviewPaymentRequiredError):
|
||||
await middleware.process(context, mock_next)
|
||||
@@ -218,7 +218,7 @@ class TestPurviewPolicyMiddleware:
|
||||
"""Test that post-check exceptions are propagated when ignore_exceptions=False."""
|
||||
middleware._settings.ignore_exceptions = False
|
||||
|
||||
context = AgentContext(agent=mock_agent, messages=[ChatMessage(role="user", text="Hello")])
|
||||
context = AgentContext(agent=mock_agent, messages=[Message(role="user", text="Hello")])
|
||||
|
||||
call_count = 0
|
||||
|
||||
@@ -232,7 +232,7 @@ class TestPurviewPolicyMiddleware:
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=side_effect):
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[ChatMessage(role="assistant", text="OK")])
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="OK")])
|
||||
|
||||
with pytest.raises(ValueError, match="Post-check blew up"):
|
||||
await middleware.process(context, mock_next)
|
||||
@@ -244,14 +244,14 @@ class TestPurviewPolicyMiddleware:
|
||||
# Set ignore_exceptions to True
|
||||
middleware._settings.ignore_exceptions = True
|
||||
|
||||
context = AgentContext(agent=mock_agent, messages=[ChatMessage(role="user", text="Test")])
|
||||
context = AgentContext(agent=mock_agent, messages=[Message(role="user", text="Test")])
|
||||
|
||||
with patch.object(
|
||||
middleware._processor, "process_messages", side_effect=Exception("Pre-check error")
|
||||
) as mock_process:
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[ChatMessage(role="assistant", text="Response")])
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -267,7 +267,7 @@ class TestPurviewPolicyMiddleware:
|
||||
# Set ignore_exceptions to True
|
||||
middleware._settings.ignore_exceptions = True
|
||||
|
||||
context = AgentContext(agent=mock_agent, messages=[ChatMessage(role="user", text="Test")])
|
||||
context = AgentContext(agent=mock_agent, messages=[Message(role="user", text="Test")])
|
||||
|
||||
call_count = 0
|
||||
|
||||
@@ -281,7 +281,7 @@ class TestPurviewPolicyMiddleware:
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[ChatMessage(role="assistant", text="Response")])
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -298,7 +298,7 @@ class TestPurviewPolicyMiddleware:
|
||||
|
||||
mock_agent = MagicMock()
|
||||
mock_agent.name = "test-agent"
|
||||
context = AgentContext(agent=mock_agent, messages=[ChatMessage(role="user", text="Test")])
|
||||
context = AgentContext(agent=mock_agent, messages=[Message(role="user", text="Test")])
|
||||
|
||||
# Mock processor to raise an exception
|
||||
async def mock_process_messages(*args, **kwargs):
|
||||
@@ -307,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="assistant", text="Response")])
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
|
||||
|
||||
# Should not raise, just log
|
||||
await middleware.process(context, mock_next)
|
||||
@@ -322,7 +322,7 @@ class TestPurviewPolicyMiddleware:
|
||||
|
||||
mock_agent = MagicMock()
|
||||
mock_agent.name = "test-agent"
|
||||
context = AgentContext(agent=mock_agent, messages=[ChatMessage(role="user", text="Test")])
|
||||
context = AgentContext(agent=mock_agent, messages=[Message(role="user", text="Test")])
|
||||
|
||||
# Mock processor to raise an exception
|
||||
async def mock_process_messages(*args, **kwargs):
|
||||
@@ -342,12 +342,12 @@ class TestPurviewPolicyMiddleware:
|
||||
) -> None:
|
||||
"""Test that session_id is extracted from thread.service_thread_id."""
|
||||
thread = AgentThread(service_thread_id="thread-123")
|
||||
context = AgentContext(agent=mock_agent, messages=[ChatMessage(role="user", text="Hello")], thread=thread)
|
||||
context = AgentContext(agent=mock_agent, messages=[Message(role="user", text="Hello")], thread=thread)
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[ChatMessage(role="assistant", text="Hi")])
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Hi")])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -359,13 +359,13 @@ class TestPurviewPolicyMiddleware:
|
||||
self, middleware: PurviewPolicyMiddleware, mock_agent: MagicMock
|
||||
) -> None:
|
||||
"""Test that session_id is extracted from message.additional_properties['conversation_id']."""
|
||||
messages = [ChatMessage(role="user", text="Hello", additional_properties={"conversation_id": "conv-456"})]
|
||||
messages = [Message(role="user", text="Hello", additional_properties={"conversation_id": "conv-456"})]
|
||||
context = AgentContext(agent=mock_agent, messages=messages)
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[ChatMessage(role="assistant", text="Hi")])
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Hi")])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -378,13 +378,13 @@ class TestPurviewPolicyMiddleware:
|
||||
) -> None:
|
||||
"""Test that thread.service_thread_id takes precedence over message conversation_id."""
|
||||
thread = AgentThread(service_thread_id="thread-789")
|
||||
messages = [ChatMessage(role="user", text="Hello", additional_properties={"conversation_id": "conv-456"})]
|
||||
messages = [Message(role="user", text="Hello", additional_properties={"conversation_id": "conv-456"})]
|
||||
context = AgentContext(agent=mock_agent, messages=messages, thread=thread)
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[ChatMessage(role="assistant", text="Hi")])
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Hi")])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -395,12 +395,12 @@ class TestPurviewPolicyMiddleware:
|
||||
self, middleware: PurviewPolicyMiddleware, mock_agent: MagicMock
|
||||
) -> None:
|
||||
"""Test that session_id is None when no thread or conversation_id is available."""
|
||||
context = AgentContext(agent=mock_agent, messages=[ChatMessage(role="user", text="Hello")])
|
||||
context = AgentContext(agent=mock_agent, messages=[Message(role="user", text="Hello")])
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[ChatMessage(role="assistant", text="Hi")])
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Hi")])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -412,12 +412,12 @@ class TestPurviewPolicyMiddleware:
|
||||
) -> None:
|
||||
"""Test that session_id is passed to post-check process_messages call."""
|
||||
thread = AgentThread(service_thread_id="thread-999")
|
||||
context = AgentContext(agent=mock_agent, messages=[ChatMessage(role="user", text="Hello")], thread=thread)
|
||||
context = AgentContext(agent=mock_agent, messages=[Message(role="user", text="Hello")], thread=thread)
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[ChatMessage(role="assistant", text="Response")])
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
+24
-24
@@ -5,7 +5,7 @@
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from agent_framework import ChatMessage
|
||||
from agent_framework import Message
|
||||
|
||||
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="user", text="Hello"),
|
||||
ChatMessage(role="assistant", text="Hi there"),
|
||||
Message(role="user", text="Hello"),
|
||||
Message(role="assistant", text="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="user", text="Sensitive content")]
|
||||
messages = [Message(role="user", text="Sensitive content")]
|
||||
|
||||
mock_request = process_content_request_factory("Sensitive content")
|
||||
|
||||
@@ -120,7 +120,7 @@ class TestScopedContentProcessor:
|
||||
) -> None:
|
||||
"""Test _map_messages creates ProcessContentRequest objects."""
|
||||
messages = [
|
||||
ChatMessage(
|
||||
Message(
|
||||
role="user",
|
||||
text="Test message",
|
||||
message_id="msg-123",
|
||||
@@ -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="user", text="Test", message_id="msg-123")]
|
||||
messages = [Message(role="user", text="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="user", text="Test", message_id="msg-123")]
|
||||
messages = [Message(role="user", text="Test", message_id="msg-123")]
|
||||
|
||||
with pytest.raises(ValueError, match="Tenant id required"):
|
||||
await processor._map_messages(messages, Activity.UPLOAD_TEXT)
|
||||
@@ -331,7 +331,7 @@ class TestScopedContentProcessor:
|
||||
processor = ScopedContentProcessor(mock_client, settings)
|
||||
|
||||
messages = [
|
||||
ChatMessage(
|
||||
Message(
|
||||
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="user", text="Test message")]
|
||||
messages = [Message(role="user", text="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="user", text="Test message")]
|
||||
messages = [Message(role="user", text="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="user", text="Test")]
|
||||
messages = [Message(role="user", text="Test")]
|
||||
|
||||
requests, user_id = await processor._map_messages(messages, Activity.UPLOAD_TEXT)
|
||||
|
||||
@@ -493,7 +493,7 @@ class TestUserIdResolution:
|
||||
processor = ScopedContentProcessor(mock_client, settings)
|
||||
|
||||
messages = [
|
||||
ChatMessage(
|
||||
Message(
|
||||
role="user",
|
||||
text="Test",
|
||||
additional_properties={"user_id": "22222222-2222-2222-2222-222222222222"},
|
||||
@@ -513,7 +513,7 @@ class TestUserIdResolution:
|
||||
processor = ScopedContentProcessor(mock_client, settings)
|
||||
|
||||
messages = [
|
||||
ChatMessage(
|
||||
Message(
|
||||
role="user",
|
||||
text="Test",
|
||||
author_name="33333333-3333-3333-3333-333333333333",
|
||||
@@ -531,7 +531,7 @@ class TestUserIdResolution:
|
||||
processor = ScopedContentProcessor(mock_client, settings)
|
||||
|
||||
messages = [
|
||||
ChatMessage(
|
||||
Message(
|
||||
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="user", text="Test")]
|
||||
messages = [Message(role="user", text="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="user", text="Test")]
|
||||
messages = [Message(role="user", text="Test")]
|
||||
|
||||
requests, user_id = await processor._map_messages(messages, Activity.UPLOAD_TEXT, provided_user_id="not-a-guid")
|
||||
|
||||
@@ -574,11 +574,11 @@ class TestUserIdResolution:
|
||||
processor = ScopedContentProcessor(mock_client, settings)
|
||||
|
||||
messages = [
|
||||
ChatMessage(
|
||||
Message(
|
||||
role="user", text="First", additional_properties={"user_id": "55555555-5555-5555-5555-555555555555"}
|
||||
),
|
||||
ChatMessage(role="assistant", text="Response"),
|
||||
ChatMessage(role="user", text="Second"),
|
||||
Message(role="assistant", text="Response"),
|
||||
Message(role="user", text="Second"),
|
||||
]
|
||||
|
||||
requests, user_id = await processor._map_messages(messages, Activity.UPLOAD_TEXT)
|
||||
@@ -594,13 +594,13 @@ class TestUserIdResolution:
|
||||
processor = ScopedContentProcessor(mock_client, settings)
|
||||
|
||||
messages = [
|
||||
ChatMessage(role="user", text="First", author_name="Not a GUID"),
|
||||
ChatMessage(
|
||||
Message(role="user", text="First", author_name="Not a GUID"),
|
||||
Message(
|
||||
role="assistant",
|
||||
text="Response",
|
||||
additional_properties={"user_id": "66666666-6666-6666-6666-666666666666"},
|
||||
),
|
||||
ChatMessage(
|
||||
Message(
|
||||
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="user", text="Test")]
|
||||
messages = [Message(role="user", text="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="user", text="Test")]
|
||||
messages = [Message(role="user", text="Test")]
|
||||
|
||||
with pytest.raises(PurviewPaymentRequiredError):
|
||||
await processor.process_messages(
|
||||
Reference in New Issue
Block a user