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
@@ -45,6 +45,25 @@ class PurviewPolicyMiddleware(AgentMiddleware):
self._processor = ScopedContentProcessor(self._client, settings, cache_provider)
self._settings = settings
@staticmethod
def _get_agent_session_id(context: AgentContext) -> str | None:
"""Resolve a session/conversation id from the agent run context.
Resolution order:
1. thread.service_thread_id
2. First message whose additional_properties contains 'conversation_id'
3. None: the downstream processor will generate a new UUID
"""
if context.thread and context.thread.service_thread_id:
return context.thread.service_thread_id
for message in context.messages:
conversation_id = message.additional_properties.get("conversation_id")
if conversation_id is not None:
return str(conversation_id)
return None
async def process(
self,
context: AgentContext,
@@ -53,8 +72,9 @@ class PurviewPolicyMiddleware(AgentMiddleware):
resolved_user_id: str | None = None
try:
# Pre (prompt) check
session_id = self._get_agent_session_id(context)
should_block_prompt, resolved_user_id = await self._processor.process_messages(
context.messages, Activity.UPLOAD_TEXT
context.messages, Activity.UPLOAD_TEXT, session_id=session_id
)
if should_block_prompt:
from agent_framework import AgentResponse, ChatMessage
@@ -79,10 +99,14 @@ class PurviewPolicyMiddleware(AgentMiddleware):
try:
# Post (response) check only if we have a normal AgentResponse
# Use the same user_id from the request for the response evaluation
session_id_response = self._get_agent_session_id(context)
if session_id_response is None:
session_id_response = session_id
if context.result and not context.stream:
should_block_response, _ = await self._processor.process_messages(
context.result.messages, # type: ignore[union-attr]
Activity.UPLOAD_TEXT,
session_id=session_id,
user_id=resolved_user_id,
)
if should_block_response:
@@ -144,8 +168,9 @@ class PurviewChatPolicyMiddleware(ChatMiddleware):
) -> None: # type: ignore[override]
resolved_user_id: str | None = None
try:
session_id = context.options.get("conversation_id") if context.options else None
should_block_prompt, resolved_user_id = await self._processor.process_messages(
context.messages, Activity.UPLOAD_TEXT
context.messages, Activity.UPLOAD_TEXT, session_id=session_id
)
if should_block_prompt:
from agent_framework import ChatMessage, ChatResponse
@@ -169,12 +194,15 @@ class PurviewChatPolicyMiddleware(ChatMiddleware):
try:
# Post (response) evaluation only if non-streaming and we have messages result shape
# Use the same user_id from the request for the response evaluation
session_id_response = context.options.get("conversation_id") if context.options else None
if session_id_response is None:
session_id_response = session_id
if context.result and not context.stream:
result_obj = context.result
messages = getattr(result_obj, "messages", None)
if messages:
should_block_response, _ = await self._processor.process_messages(
messages, Activity.UPLOAD_TEXT, user_id=resolved_user_id
messages, Activity.UPLOAD_TEXT, session_id=session_id_response, user_id=resolved_user_id
)
if should_block_response:
from agent_framework import ChatMessage, ChatResponse
@@ -1,6 +1,7 @@
# Copyright (c) Microsoft. All rights reserved.
import asyncio
import time
import uuid
from collections.abc import Iterable, MutableMapping
from typing import Any
@@ -62,13 +63,18 @@ class ScopedContentProcessor:
self._background_tasks: set[asyncio.Task[Any]] = set()
async def process_messages(
self, messages: Iterable[ChatMessage], activity: Activity, user_id: str | None = None
self,
messages: Iterable[ChatMessage],
activity: Activity,
session_id: str | None = None,
user_id: str | None = None,
) -> tuple[bool, str | None]:
"""Process messages for policy evaluation.
Args:
messages: The messages to process
activity: The activity type (e.g., UPLOAD_TEXT)
session_id: Optional session/conversation id. Else, a new GUID is generated.
user_id: Optional user_id to use for all messages. If provided, this is the fallback.
Returns:
@@ -76,7 +82,7 @@ class ScopedContentProcessor:
The resolved_user_id can be stored and passed back when processing the response
to ensure the same user context is maintained throughout the request/response cycle.
"""
pc_requests, resolved_user_id = await self._map_messages(messages, activity, user_id)
pc_requests, resolved_user_id = await self._map_messages(messages, activity, session_id, user_id)
should_block = False
for req in pc_requests:
resp = await self._process_with_scopes(req)
@@ -90,13 +96,18 @@ class ScopedContentProcessor:
return should_block, resolved_user_id
async def _map_messages(
self, messages: Iterable[ChatMessage], activity: Activity, provided_user_id: str | None = None
self,
messages: Iterable[ChatMessage],
activity: Activity,
session_id: str | None = None,
provided_user_id: str | None = None,
) -> tuple[list[ProcessContentRequest], str | None]:
"""Map messages to ProcessContentRequests.
Args:
messages: The messages to map
activity: The activity type
session_id: Optional session/conversation id to use for correlation
provided_user_id: Optional user_id to use. If provided, this is the fallback.
Returns:
@@ -137,12 +148,14 @@ class ScopedContentProcessor:
for m in messages:
message_id = m.message_id or str(uuid.uuid4())
content = PurviewTextContent(data=m.text or "")
correlation_id = (session_id or str(uuid.uuid4())) + "@AF"
meta = ProcessConversationMetadata(
identifier=message_id,
content=content,
name=f"Agent Framework Message {message_id}",
is_truncated=False,
correlation_id=str(uuid.uuid4()),
correlation_id=correlation_id,
sequence_number=time.time_ns(),
)
activity_meta = ActivityMetadata(activity=activity)
@@ -159,12 +172,13 @@ class ScopedContentProcessor:
else:
raise ValueError("App location not provided or inferable")
app_version = self._settings.app_version or "Unknown"
protected_app = ProtectedAppMetadata(
name=self._settings.app_name,
version="1.0",
version=app_version,
application_location=policy_location,
)
integrated_app = IntegratedAppMetadata(name=self._settings.app_name, version="1.0")
integrated_app = IntegratedAppMetadata(name=self._settings.app_name, version=app_version)
device_meta = DeviceMetadata(
operating_system_specifications=OperatingSystemSpecifications(
operating_system_platform="Unknown", operating_system_version="Unknown"
@@ -35,7 +35,7 @@ class PurviewAppLocation(BaseModel):
class PurviewSettings(AFBaseSettings):
"""Settings for Purview integration mirroring .NET PurviewSettings.
"""Settings for Purview integration.
Attributes:
app_name: Public app name.