Python: [Purview] Update CorrelationId (#3745)

This commit is contained in:
Rishabh Chawla
2026-02-10 19:19:14 +00:00
committed by GitHub
parent f106a1a2b1
commit 7e7d72275d
13 changed files with 245 additions and 29 deletions
@@ -9,6 +9,7 @@ from agent_framework import ChatContext, ChatMessage, MiddlewareTermination
from azure.core.credentials import AccessToken
from agent_framework_purview import PurviewChatPolicyMiddleware, PurviewSettings
from agent_framework_purview._models import Activity
@dataclass
@@ -82,7 +83,7 @@ class TestPurviewChatPolicyMiddleware:
async def test_blocks_response(self, middleware: PurviewChatPolicyMiddleware, chat_context: ChatContext) -> None:
call_state = {"count": 0}
async def side_effect(messages, activity, user_id=None):
async def side_effect(messages, activity, session_id=None, user_id=None):
call_state["count"] += 1
should_block = call_state["count"] == 2
return (should_block, "user-123")
@@ -157,7 +158,7 @@ class TestPurviewChatPolicyMiddleware:
"""Test that the same user_id from pre-check is used in post-check."""
captured_user_ids = []
async def mock_process_messages(messages, activity, user_id=None):
async def mock_process_messages(messages, activity, session_id=None, user_id=None):
captured_user_ids.append(user_id)
return (False, "resolved-user-123")
@@ -362,3 +363,67 @@ class TestPurviewChatPolicyMiddleware:
with pytest.raises(ValueError, match="post"):
await middleware.process(context, mock_next)
async def test_chat_middleware_uses_conversation_id_from_options(
self, middleware: PurviewChatPolicyMiddleware
) -> None:
"""Test that session_id is extracted from context.options['conversation_id']."""
chat_client = DummyChatClient()
messages = [ChatMessage(role="user", text="Hello")]
options = {"conversation_id": "conv-123", "model": "test-model"}
context = ChatContext(chat_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")]
ctx.result = result
await middleware.process(context, mock_next)
# Verify session_id is passed to both pre-check and post-check
assert mock_proc.call_count == 2
mock_proc.assert_any_call(messages, Activity.UPLOAD_TEXT, session_id="conv-123")
async def test_chat_middleware_passes_none_session_id_when_options_missing(
self, middleware: PurviewChatPolicyMiddleware
) -> 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)
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")]
ctx.result = result
await middleware.process(context, mock_next)
# Verify session_id=None is passed
mock_proc.assert_any_call(messages, Activity.UPLOAD_TEXT, session_id=None)
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")]
options = {"conversation_id": "conv-999"}
context = ChatContext(chat_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")]
ctx.result = result
await middleware.process(context, mock_next)
# Verify both calls include session_id
assert mock_proc.call_count == 2
# Check post-check call includes session_id
post_check_call = mock_proc.call_args_list[1]
assert post_check_call[1]["session_id"] == "conv-999"
@@ -5,10 +5,11 @@
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from agent_framework import AgentContext, AgentResponse, ChatMessage, MiddlewareTermination
from agent_framework import AgentContext, AgentResponse, AgentThread, ChatMessage, MiddlewareTermination
from azure.core.credentials import AccessToken
from agent_framework_purview import PurviewPolicyMiddleware, PurviewSettings
from agent_framework_purview._models import Activity
class TestPurviewPolicyMiddleware:
@@ -92,7 +93,7 @@ class TestPurviewPolicyMiddleware:
call_count = 0
async def mock_process_messages(messages, activity, user_id=None):
async def mock_process_messages(messages, activity, session_id=None, user_id=None):
nonlocal call_count
call_count += 1
should_block = call_count != 1
@@ -335,3 +336,93 @@ class TestPurviewPolicyMiddleware:
# Should raise the exception
with pytest.raises(ValueError, match="Test error"):
await middleware.process(context, mock_next)
async def test_middleware_uses_thread_service_thread_id_as_session_id(
self, middleware: PurviewPolicyMiddleware, mock_agent: MagicMock
) -> 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)
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")])
await middleware.process(context, mock_next)
# Verify session_id is passed to both pre-check and post-check
assert mock_proc.call_count == 2
mock_proc.assert_any_call(context.messages, Activity.UPLOAD_TEXT, session_id="thread-123")
async def test_middleware_uses_message_conversation_id_as_session_id(
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"})]
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")])
await middleware.process(context, mock_next)
# Verify session_id is passed to both pre-check and post-check
assert mock_proc.call_count == 2
mock_proc.assert_any_call(messages, Activity.UPLOAD_TEXT, session_id="conv-456")
async def test_middleware_thread_id_takes_precedence_over_message_conversation_id(
self, middleware: PurviewPolicyMiddleware, mock_agent: MagicMock
) -> 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"})]
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")])
await middleware.process(context, mock_next)
# Verify thread ID is used, not message conversation_id
mock_proc.assert_any_call(messages, Activity.UPLOAD_TEXT, session_id="thread-789")
async def test_middleware_passes_none_session_id_when_not_available(
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")])
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")])
await middleware.process(context, mock_next)
# Verify session_id=None is passed
mock_proc.assert_any_call(context.messages, Activity.UPLOAD_TEXT, session_id=None)
async def test_middleware_session_id_used_in_post_check(
self, middleware: PurviewPolicyMiddleware, mock_agent: MagicMock
) -> 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)
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")])
await middleware.process(context, mock_next)
# Verify both calls include session_id
assert mock_proc.call_count == 2
# Check post-check call includes session_id
post_check_call = mock_proc.call_args_list[1]
assert post_check_call[1]["session_id"] == "thread-999"
@@ -92,7 +92,7 @@ class TestScopedContentProcessor:
assert should_block is False
assert user_id is None
mock_map.assert_called_once_with(messages, Activity.UPLOAD_TEXT, None)
mock_map.assert_called_once_with(messages, Activity.UPLOAD_TEXT, None, None)
async def test_process_messages_blocks_content(
self, processor: ScopedContentProcessor, process_content_request_factory