Python: [BREAKING] Python: Intro group chat and refactor orchestrations. Fix as_agent(). Standardize orchestration start msg types. (#1538)

* Intro group chat and refactor magentic. Fix as_agent()

* Cleanup and improvements

* Add as_agent docstring clarification

* Standardize orchestration messages to use agent-style inputs.

* Simplify group chat constructs

* Further cleanup

* Add sk to af group chat migration sample. Update README.

* Improvements and simplifications

* consolidating shared orchestration logic

* Further clean up

* Add group chat sample

* Improve typing

* Fix test imports

* Fix readme links

* Cleanup per PR Feedback
This commit is contained in:
Evan Mattson
2025-10-25 09:14:06 +09:00
committed by GitHub
Unverified
parent 899d8ff775
commit e3aad8e4e0
38 changed files with 5024 additions and 814 deletions
@@ -52,29 +52,29 @@ from ._executor import (
handler,
)
from ._function_executor import FunctionExecutor, executor
from ._group_chat import (
DEFAULT_MANAGER_INSTRUCTIONS,
DEFAULT_MANAGER_STRUCTURED_OUTPUT_PROMPT,
GroupChatBuilder,
GroupChatDirective,
GroupChatStateSnapshot,
ManagerDirectiveModel,
)
from ._handoff import HandoffBuilder, HandoffUserInputRequest
from ._magentic import (
MagenticAgentDeltaEvent,
MagenticAgentExecutor,
MagenticAgentMessageEvent,
MagenticBuilder,
MagenticCallbackEvent,
MagenticCallbackMode,
MagenticContext,
MagenticFinalResultEvent,
MagenticManagerBase,
MagenticOrchestratorExecutor,
MagenticOrchestratorMessageEvent,
MagenticPlanReviewDecision,
MagenticPlanReviewReply,
MagenticPlanReviewRequest,
MagenticProgressLedger,
MagenticProgressLedgerItem,
MagenticRequestMessage,
MagenticResponseMessage,
MagenticStartMessage,
StandardMagenticManager,
)
from ._orchestration_state import OrchestrationState
from ._request_info_executor import (
PendingRequestDetails,
RequestInfoExecutor,
@@ -105,6 +105,8 @@ from ._workflow_context import WorkflowContext
from ._workflow_executor import WorkflowExecutor
__all__ = [
"DEFAULT_MANAGER_INSTRUCTIONS",
"DEFAULT_MANAGER_STRUCTURED_OUTPUT_PROMPT",
"DEFAULT_MAX_ITERATIONS",
"AgentExecutor",
"AgentExecutorRequest",
@@ -128,30 +130,26 @@ __all__ = [
"FileCheckpointStorage",
"FunctionExecutor",
"GraphConnectivityError",
"GroupChatBuilder",
"GroupChatDirective",
"GroupChatStateSnapshot",
"HandoffBuilder",
"HandoffUserInputRequest",
"InMemoryCheckpointStorage",
"InProcRunnerContext",
"MagenticAgentDeltaEvent",
"MagenticAgentExecutor",
"MagenticAgentMessageEvent",
"MagenticBuilder",
"MagenticCallbackEvent",
"MagenticCallbackMode",
"MagenticContext",
"MagenticFinalResultEvent",
"MagenticManagerBase",
"MagenticOrchestratorExecutor",
"MagenticOrchestratorMessageEvent",
"MagenticPlanReviewDecision",
"MagenticPlanReviewReply",
"MagenticPlanReviewRequest",
"MagenticProgressLedger",
"MagenticProgressLedgerItem",
"MagenticRequestMessage",
"MagenticResponseMessage",
"MagenticStartMessage",
"ManagerDirectiveModel",
"Message",
"OrchestrationState",
"PendingRequestDetails",
"RequestInfoEvent",
"RequestInfoExecutor",
@@ -50,29 +50,28 @@ from ._executor import (
handler,
)
from ._function_executor import FunctionExecutor, executor
from ._group_chat import (
DEFAULT_MANAGER_INSTRUCTIONS,
DEFAULT_MANAGER_STRUCTURED_OUTPUT_PROMPT,
GroupChatBuilder,
GroupChatDirective,
GroupChatStateSnapshot,
)
from ._handoff import HandoffBuilder, HandoffUserInputRequest
from ._magentic import (
MagenticAgentDeltaEvent,
MagenticAgentExecutor,
MagenticAgentMessageEvent,
MagenticBuilder,
MagenticCallbackEvent,
MagenticCallbackMode,
MagenticContext,
MagenticFinalResultEvent,
MagenticManagerBase,
MagenticOrchestratorExecutor,
MagenticOrchestratorMessageEvent,
MagenticPlanReviewDecision,
MagenticPlanReviewReply,
MagenticPlanReviewRequest,
MagenticProgressLedger,
MagenticProgressLedgerItem,
MagenticRequestMessage,
MagenticResponseMessage,
MagenticStartMessage,
StandardMagenticManager,
)
from ._orchestration_state import OrchestrationState
from ._request_info_executor import (
PendingRequestDetails,
RequestInfoExecutor,
@@ -103,6 +102,8 @@ from ._workflow_context import WorkflowContext
from ._workflow_executor import WorkflowExecutor
__all__ = [
"DEFAULT_MANAGER_INSTRUCTIONS",
"DEFAULT_MANAGER_STRUCTURED_OUTPUT_PROMPT",
"DEFAULT_MAX_ITERATIONS",
"AgentExecutor",
"AgentExecutorRequest",
@@ -126,30 +127,25 @@ __all__ = [
"FileCheckpointStorage",
"FunctionExecutor",
"GraphConnectivityError",
"GroupChatBuilder",
"GroupChatDirective",
"GroupChatStateSnapshot",
"HandoffBuilder",
"HandoffUserInputRequest",
"InMemoryCheckpointStorage",
"InProcRunnerContext",
"MagenticAgentDeltaEvent",
"MagenticAgentExecutor",
"MagenticAgentMessageEvent",
"MagenticBuilder",
"MagenticCallbackEvent",
"MagenticCallbackMode",
"MagenticContext",
"MagenticFinalResultEvent",
"MagenticManagerBase",
"MagenticOrchestratorExecutor",
"MagenticOrchestratorMessageEvent",
"MagenticPlanReviewDecision",
"MagenticPlanReviewReply",
"MagenticPlanReviewRequest",
"MagenticProgressLedger",
"MagenticProgressLedgerItem",
"MagenticRequestMessage",
"MagenticResponseMessage",
"MagenticStartMessage",
"Message",
"OrchestrationState",
"PendingRequestDetails",
"RequestInfoEvent",
"RequestInfoExecutor",
@@ -3,7 +3,7 @@
import json
import logging
import uuid
from collections.abc import AsyncIterable, Sequence
from collections.abc import AsyncIterable
from dataclasses import dataclass
from datetime import datetime
from typing import TYPE_CHECKING, Any, ClassVar, TypedDict, cast
@@ -19,7 +19,6 @@ from agent_framework import (
FunctionCallContent,
FunctionResultContent,
Role,
TextContent,
UsageDetails,
)
@@ -29,6 +28,7 @@ from ._events import (
RequestInfoEvent,
WorkflowEvent,
)
from ._message_utils import normalize_messages_input
if TYPE_CHECKING:
from ._workflow import Workflow
@@ -131,7 +131,7 @@ class WorkflowAgent(BaseAgent):
"""
# Collect all streaming updates
response_updates: list[AgentRunResponseUpdate] = []
input_messages = self._normalize_messages(messages)
input_messages = normalize_messages_input(messages)
thread = thread or self.get_new_thread()
response_id = str(uuid.uuid4())
@@ -165,7 +165,7 @@ class WorkflowAgent(BaseAgent):
Yields:
AgentRunResponseUpdate objects representing the workflow execution progress.
"""
input_messages = self._normalize_messages(messages)
input_messages = normalize_messages_input(messages)
thread = thread or self.get_new_thread()
response_updates: list[AgentRunResponseUpdate] = []
response_id = str(uuid.uuid4())
@@ -225,28 +225,6 @@ class WorkflowAgent(BaseAgent):
if update:
yield update
def _normalize_messages(
self,
messages: str | ChatMessage | Sequence[str] | Sequence[ChatMessage] | None = None,
) -> list[ChatMessage]:
"""Normalize input messages to a list of ChatMessage objects."""
if messages is None:
return []
if isinstance(messages, str):
return [ChatMessage(role=Role.USER, contents=[TextContent(text=messages)])]
if isinstance(messages, ChatMessage):
return [messages]
normalized: list[ChatMessage] = []
for msg in messages:
if isinstance(msg, str):
normalized.append(ChatMessage(role=Role.USER, contents=[TextContent(text=msg)]))
elif isinstance(msg, ChatMessage):
normalized.append(msg)
return normalized
def _convert_workflow_event_to_agent_update(
self,
response_id: str,
@@ -12,6 +12,7 @@ from ._events import (
AgentRunUpdateEvent, # type: ignore[reportPrivateUsage]
)
from ._executor import Executor, handler
from ._message_utils import normalize_messages_input
from ._workflow_context import WorkflowContext
logger = logging.getLogger(__name__)
@@ -167,7 +168,7 @@ class AgentExecutor(Executor):
@handler
async def from_str(self, text: str, ctx: WorkflowContext[AgentExecutorResponse, AgentRunResponse]) -> None:
"""Accept a raw user prompt string and run the agent (one-shot)."""
self._cache = [ChatMessage(role="user", text=text)] # type: ignore[arg-type]
self._cache = normalize_messages_input(text)
await self._run_agent_and_emit(ctx)
@handler
@@ -177,15 +178,50 @@ class AgentExecutor(Executor):
ctx: WorkflowContext[AgentExecutorResponse, AgentRunResponse],
) -> None:
"""Accept a single ChatMessage as input."""
self._cache = [message]
self._cache = normalize_messages_input(message)
await self._run_agent_and_emit(ctx)
@handler
async def from_messages(
self,
messages: list[ChatMessage],
messages: list[str | ChatMessage],
ctx: WorkflowContext[AgentExecutorResponse, AgentRunResponse],
) -> None:
"""Accept a list of ChatMessage objects as conversation context."""
self._cache = list(messages)
"""Accept a list of chat inputs (strings or ChatMessage) as conversation context."""
self._cache = normalize_messages_input(messages)
await self._run_agent_and_emit(ctx)
def snapshot_state(self) -> dict[str, Any]:
"""Capture current executor state for checkpointing.
Returns:
Dict containing serialized cache state
"""
from ._conversation_state import encode_chat_messages
return {
"cache": encode_chat_messages(self._cache),
}
def restore_state(self, state: dict[str, Any]) -> None:
"""Restore executor state from checkpoint.
Args:
state: Checkpoint data dict
"""
from ._conversation_state import decode_chat_messages
cache_payload = state.get("cache")
if cache_payload:
try:
self._cache = decode_chat_messages(cache_payload)
except Exception as exc:
logger.warning("Failed to restore cache: %s", exc)
self._cache = []
else:
self._cache = []
def reset(self) -> None:
"""Reset the internal cache of the executor."""
logger.debug("AgentExecutor %s: Resetting cache", self.id)
self._cache.clear()
@@ -0,0 +1,265 @@
# Copyright (c) Microsoft. All rights reserved.
"""Base class for group chat orchestrators that manages conversation flow and participant selection."""
import inspect
import logging
from abc import ABC, abstractmethod
from collections.abc import Awaitable, Callable, Sequence
from typing import Any
from .._types import ChatMessage
from ._executor import Executor
from ._orchestrator_helpers import ParticipantRegistry
from ._workflow_context import WorkflowContext
logger = logging.getLogger(__name__)
class BaseGroupChatOrchestrator(Executor, ABC):
"""Abstract base class for group chat orchestrators.
Provides shared functionality for participant registration, routing,
and round limit checking that is common across all group chat patterns.
Subclasses must implement pattern-specific orchestration logic while
inheriting the common participant management infrastructure.
"""
def __init__(self, executor_id: str) -> None:
"""Initialize base orchestrator.
Args:
executor_id: Unique identifier for this orchestrator executor
"""
super().__init__(executor_id)
self._registry = ParticipantRegistry()
# Shared conversation state management
self._conversation: list[ChatMessage] = []
self._round_index: int = 0
self._max_rounds: int | None = None
self._termination_condition: Callable[[list[ChatMessage]], bool | Awaitable[bool]] | None = None
def register_participant_entry(self, name: str, *, entry_id: str, is_agent: bool) -> None:
"""Record routing details for a participant's entry executor.
This method provides a unified interface for registering participants
across all orchestrator patterns, whether they are agents or custom executors.
Args:
name: Participant name (used for selection and tracking)
entry_id: Executor ID for this participant's entry point
is_agent: Whether this is an AgentExecutor (True) or custom Executor (False)
"""
self._registry.register(name, entry_id=entry_id, is_agent=is_agent)
# Conversation state management (shared across all patterns)
def _append_messages(self, messages: Sequence[ChatMessage]) -> None:
"""Append messages to the conversation history.
Args:
messages: Messages to append
"""
self._conversation.extend(messages)
def _get_conversation(self) -> list[ChatMessage]:
"""Get a copy of the current conversation.
Returns:
Cloned conversation list
"""
return list(self._conversation)
def _clear_conversation(self) -> None:
"""Clear the conversation history."""
self._conversation.clear()
def _increment_round(self) -> None:
"""Increment the round counter."""
self._round_index += 1
async def _check_termination(self) -> bool:
"""Check if conversation should terminate based on termination condition.
Supports both synchronous and asynchronous termination conditions.
Returns:
True if termination condition met, False otherwise
"""
if self._termination_condition is None:
return False
result = self._termination_condition(self._get_conversation())
if inspect.iscoroutine(result) or inspect.isawaitable(result):
result = await result
return bool(result)
@abstractmethod
def _get_author_name(self) -> str:
"""Get the author name for orchestrator-generated messages.
Subclasses must implement this to provide a stable author name
for completion messages and other orchestrator-generated content.
Returns:
Author name to use for messages generated by this orchestrator
"""
...
def _create_completion_message(
self,
text: str | None = None,
reason: str = "completed",
) -> ChatMessage:
"""Create a standardized completion message.
Args:
text: Optional message text (auto-generated if None)
reason: Completion reason for default text
Returns:
ChatMessage with completion content
"""
from .._types import Role
message_text = text or f"Conversation {reason}."
return ChatMessage(
role=Role.ASSISTANT,
text=message_text,
author_name=self._get_author_name(),
)
# Participant routing (shared across all patterns)
async def _route_to_participant(
self,
participant_name: str,
conversation: list[ChatMessage],
ctx: WorkflowContext[Any, Any],
*,
instruction: str | None = None,
task: ChatMessage | None = None,
metadata: dict[str, Any] | None = None,
) -> None:
"""Route a conversation to a participant.
This method handles the dual envelope pattern:
- AgentExecutors receive AgentExecutorRequest (messages only)
- Custom executors receive GroupChatRequestMessage (full context)
Args:
participant_name: Name of the participant to route to
conversation: Conversation history to send
ctx: Workflow context for message routing
instruction: Optional instruction from manager/orchestrator
task: Optional task context
metadata: Optional metadata dict
Raises:
ValueError: If participant is not registered
"""
from ._agent_executor import AgentExecutorRequest
from ._orchestrator_helpers import prepare_participant_request
entry_id = self._registry.get_entry_id(participant_name)
if entry_id is None:
raise ValueError(f"No registered entry executor for participant '{participant_name}'.")
if self._registry.is_agent(participant_name):
# AgentExecutors receive simple message list
await ctx.send_message(
AgentExecutorRequest(messages=conversation, should_respond=True),
target_id=entry_id,
)
else:
# Custom executors receive full context envelope
request = prepare_participant_request(
participant_name=participant_name,
conversation=conversation,
instruction=instruction or "",
task=task,
metadata=metadata,
)
await ctx.send_message(request, target_id=entry_id)
# Round limit enforcement (shared across all patterns)
def _check_round_limit(self) -> bool:
"""Check if round limit has been reached.
Uses instance variables _round_index and _max_rounds.
Returns:
True if limit reached, False otherwise
"""
if self._max_rounds is None:
return False
if self._round_index >= self._max_rounds:
logger.warning(
"%s reached max_rounds=%s; forcing completion.",
self.__class__.__name__,
self._max_rounds,
)
return True
return False
# State persistence (shared across all patterns)
# State persistence (shared across all patterns)
def snapshot_state(self) -> dict[str, Any]:
"""Capture current orchestrator state for checkpointing.
Default implementation uses OrchestrationState to serialize common state.
Subclasses should override _snapshot_pattern_metadata() to add pattern-specific data.
Returns:
Serialized state dict
"""
from ._orchestration_state import OrchestrationState
state = OrchestrationState(
conversation=list(self._conversation),
round_index=self._round_index,
metadata=self._snapshot_pattern_metadata(),
)
return state.to_dict()
def _snapshot_pattern_metadata(self) -> dict[str, Any]:
"""Serialize pattern-specific state.
Override this method to add pattern-specific checkpoint data.
Returns:
Dict with pattern-specific state (empty by default)
"""
return {}
def restore_state(self, state: dict[str, Any]) -> None:
"""Restore orchestrator state from checkpoint.
Default implementation uses OrchestrationState to deserialize common state.
Subclasses should override _restore_pattern_metadata() to restore pattern-specific data.
Args:
state: Serialized state dict
"""
from ._orchestration_state import OrchestrationState
orch_state = OrchestrationState.from_dict(state)
self._conversation = list(orch_state.conversation)
self._round_index = orch_state.round_index
self._restore_pattern_metadata(orch_state.metadata)
def _restore_pattern_metadata(self, metadata: dict[str, Any]) -> None:
"""Restore pattern-specific state.
Override this method to restore pattern-specific checkpoint data.
Args:
metadata: Pattern-specific state dict
"""
pass
@@ -13,6 +13,7 @@ from agent_framework import AgentProtocol, ChatMessage, Role
from ._agent_executor import AgentExecutorRequest, AgentExecutorResponse
from ._checkpoint import CheckpointStorage
from ._executor import Executor, handler
from ._message_utils import normalize_messages_input
from ._workflow import Workflow
from ._workflow_builder import WorkflowBuilder
from ._workflow_context import WorkflowContext
@@ -50,17 +51,21 @@ class _DispatchToAllParticipants(Executor):
@handler
async def from_str(self, prompt: str, ctx: WorkflowContext[AgentExecutorRequest]) -> None:
request = AgentExecutorRequest(messages=[ChatMessage(Role.USER, text=prompt)], should_respond=True)
request = AgentExecutorRequest(messages=normalize_messages_input(prompt), should_respond=True)
await ctx.send_message(request)
@handler
async def from_message(self, message: ChatMessage, ctx: WorkflowContext[AgentExecutorRequest]) -> None: # type: ignore[name-defined]
request = AgentExecutorRequest(messages=[message], should_respond=True)
async def from_message(self, message: ChatMessage, ctx: WorkflowContext[AgentExecutorRequest]) -> None:
request = AgentExecutorRequest(messages=normalize_messages_input(message), should_respond=True)
await ctx.send_message(request)
@handler
async def from_messages(self, messages: list[ChatMessage], ctx: WorkflowContext[AgentExecutorRequest]) -> None: # type: ignore[name-defined]
request = AgentExecutorRequest(messages=list(messages), should_respond=True)
async def from_messages(
self,
messages: list[str | ChatMessage],
ctx: WorkflowContext[AgentExecutorRequest],
) -> None:
request = AgentExecutorRequest(messages=normalize_messages_input(messages), should_respond=True)
await ctx.send_message(request)
@@ -0,0 +1,50 @@
# Copyright (c) Microsoft. All rights reserved.
"""Helpers for managing chat conversation history.
These utilities operate on standard `list[ChatMessage]` collections and simple
dictionary snapshots so orchestrators can share logic without new mixins.
"""
import json
from collections.abc import Mapping, Sequence
from typing import Any
from .._types import ChatMessage
def latest_user_message(conversation: Sequence[ChatMessage]) -> ChatMessage:
"""Return the most recent user-authored message from `conversation`."""
for message in reversed(conversation):
role_value = getattr(message.role, "value", message.role)
if str(role_value).lower() == "user":
return message
raise ValueError("No user message in conversation")
def ensure_author(message: ChatMessage, fallback: str) -> ChatMessage:
"""Attach `fallback` author if message is missing `author_name`."""
message.author_name = message.author_name or fallback
return message
def snapshot_state(conversation: Sequence[ChatMessage]) -> dict[str, Any]:
"""Build an immutable snapshot for checkpoint storage."""
if hasattr(conversation, "to_dict"):
result = conversation.to_dict() # type: ignore[attr-defined]
if isinstance(result, dict):
return result # type: ignore[return-value]
if isinstance(result, Mapping):
return dict(result) # type: ignore[arg-type]
serialisable: list[dict[str, Any]] = []
for message in conversation:
if hasattr(message, "to_dict") and callable(message.to_dict): # type: ignore[attr-defined]
msg_dict = message.to_dict() # type: ignore[attr-defined]
serialisable.append(dict(msg_dict) if isinstance(msg_dict, Mapping) else msg_dict) # type: ignore[arg-type]
elif hasattr(message, "to_json") and callable(message.to_json): # type: ignore[attr-defined]
json_payload = message.to_json() # type: ignore[attr-defined]
parsed = json.loads(json_payload) if isinstance(json_payload, str) else json_payload
serialisable.append(dict(parsed) if isinstance(parsed, Mapping) else parsed) # type: ignore[arg-type]
else:
serialisable.append(dict(getattr(message, "__dict__", {}))) # type: ignore[arg-type]
return {"messages": serialisable}
@@ -450,13 +450,7 @@ ContextT = TypeVar("ContextT", bound="WorkflowContext[Any, Any]")
def handler(
func: Callable[[ExecutorT, Any, ContextT], Awaitable[Any]],
) -> (
Callable[[ExecutorT, Any, ContextT], Awaitable[Any]]
| Callable[
[Callable[[ExecutorT, Any, ContextT], Awaitable[Any]]],
Callable[[ExecutorT, Any, ContextT], Awaitable[Any]],
]
):
) -> Callable[[ExecutorT, Any, ContextT], Awaitable[Any]]:
"""Decorator to register a handler for an executor.
Args:
File diff suppressed because it is too large Load Diff
@@ -35,9 +35,16 @@ from agent_framework import (
from .._agents import ChatAgent
from .._middleware import FunctionInvocationContext, FunctionMiddleware
from ._agent_executor import AgentExecutor, AgentExecutorRequest, AgentExecutorResponse
from ._base_group_chat_orchestrator import BaseGroupChatOrchestrator
from ._checkpoint import CheckpointStorage
from ._conversation_state import decode_chat_messages, encode_chat_messages
from ._executor import Executor, handler
from ._group_chat import (
_default_participant_factory, # type: ignore[reportPrivateUsage]
_GroupChatConfig, # type: ignore[reportPrivateUsage]
assemble_group_chat_workflow,
)
from ._orchestrator_helpers import clean_conversation_for_handoff
from ._participant_utils import GroupChatParticipantSpec, prepare_participant_metadata, sanitize_identifier
from ._request_info_executor import RequestInfoExecutor, RequestInfoMessage, RequestResponse
from ._workflow import Workflow
from ._workflow_builder import WorkflowBuilder
@@ -49,19 +56,9 @@ logger = logging.getLogger(__name__)
_HANDOFF_TOOL_PATTERN = re.compile(r"(?:handoff|transfer)[_\s-]*to[_\s-]*(?P<target>[\w-]+)", re.IGNORECASE)
def _sanitize_alias(value: str) -> str:
"""Normalise an agent alias into a lowercase identifier-safe string."""
cleaned = re.sub(r"[^0-9a-zA-Z]+", "_", value).strip("_")
if not cleaned:
cleaned = "agent"
if cleaned[0].isdigit():
cleaned = f"agent_{cleaned}"
return cleaned.lower()
def _create_handoff_tool(alias: str, description: str | None = None) -> AIFunction[Any, Any]:
"""Construct the synthetic handoff tool that signals routing to `alias`."""
sanitized = _sanitize_alias(alias)
sanitized = sanitize_identifier(alias)
tool_name = f"handoff_to_{sanitized}"
doc = description or f"Handoff to the {alias} agent."
@@ -257,7 +254,7 @@ def _target_from_tool_name(name: str | None) -> str | None:
return None
class _HandoffCoordinator(Executor):
class _HandoffCoordinator(BaseGroupChatOrchestrator):
"""Coordinates agent-to-agent transfers and user turn requests."""
def __init__(
@@ -266,7 +263,7 @@ class _HandoffCoordinator(Executor):
starting_agent_id: str,
specialist_ids: Mapping[str, str],
input_gateway_id: str,
termination_condition: Callable[[list[ChatMessage]], bool],
termination_condition: Callable[[list[ChatMessage]], bool | Awaitable[bool]],
id: str,
handoff_tool_targets: Mapping[str, str] | None = None,
) -> None:
@@ -277,9 +274,12 @@ class _HandoffCoordinator(Executor):
self._specialist_ids = set(specialist_ids.values())
self._input_gateway_id = input_gateway_id
self._termination_condition = termination_condition
self._full_conversation: list[ChatMessage] = []
self._handoff_tool_targets = {k.lower(): v for k, v in (handoff_tool_targets or {}).items()}
def _get_author_name(self) -> str:
"""Get the coordinator name for orchestrator-generated messages."""
return "handoff_coordinator"
@handler
async def handle_agent_response(
self,
@@ -290,38 +290,39 @@ class _HandoffCoordinator(Executor):
# Hydrate coordinator state (and detect new run) using checkpointable executor state
state = await ctx.get_executor_state()
if not state:
self._full_conversation = []
elif not self._full_conversation:
self._clear_conversation()
elif not self._get_conversation():
restored = self._restore_conversation_from_state(state)
if restored:
self._full_conversation = restored
self._conversation = list(restored)
source = ctx.get_source_executor_id()
is_starting_agent = source == self._starting_agent_id
# On first turn of a run, full_conversation is empty
# On first turn of a run, conversation is empty
# Track new messages only, build authoritative history incrementally
if not self._full_conversation:
conversation_msgs = self._get_conversation()
if not conversation_msgs:
# First response from starting agent - initialize with authoritative conversation snapshot
# Keep the FULL conversation including tool calls (OpenAI SDK default behavior)
full_conv = self._conversation_from_response(response)
self._full_conversation = list(full_conv)
self._conversation = list(full_conv)
else:
# Subsequent responses - append only new messages from this agent
# Keep ALL messages including tool calls to maintain complete history
new_messages = list(response.agent_run_response.messages)
self._full_conversation.extend(new_messages)
new_messages = response.agent_run_response.messages or []
self._conversation.extend(new_messages)
self._apply_response_metadata(self._full_conversation, response.agent_run_response)
self._apply_response_metadata(self._conversation, response.agent_run_response)
conversation = list(self._full_conversation)
conversation = list(self._conversation)
# Check for handoff from ANY agent (starting agent or specialist)
target = self._resolve_specialist(response.agent_run_response, conversation)
if target is not None:
await self._persist_state(ctx)
# Clean tool-related content before sending to next agent
cleaned = self._get_cleaned_conversation(conversation)
cleaned = clean_conversation_for_handoff(conversation)
request = AgentExecutorRequest(messages=cleaned, should_respond=True)
await ctx.send_message(request, target_id=target)
return
@@ -332,7 +333,7 @@ class _HandoffCoordinator(Executor):
await self._persist_state(ctx)
if self._termination_condition(conversation):
if await self._check_termination():
logger.info("Handoff workflow termination condition met. Ending conversation.")
await ctx.yield_output(list(conversation))
return
@@ -346,18 +347,18 @@ class _HandoffCoordinator(Executor):
ctx: WorkflowContext[AgentExecutorRequest, list[ChatMessage]],
) -> None:
"""Receive full conversation with new user input from gateway, update history, trim for agent."""
# Update authoritative full conversation
self._full_conversation = list(message.full_conversation)
# Update authoritative conversation
self._conversation = list(message.full_conversation)
await self._persist_state(ctx)
# Check termination before sending to agent
if self._termination_condition(self._full_conversation):
if await self._check_termination():
logger.info("Handoff workflow termination condition met. Ending conversation.")
await ctx.yield_output(list(self._full_conversation))
await ctx.yield_output(list(self._conversation))
return
# Clean before sending to starting agent
cleaned = self._get_cleaned_conversation(self._full_conversation)
cleaned = clean_conversation_for_handoff(self._conversation)
request = AgentExecutorRequest(messages=cleaned, should_respond=True)
await ctx.send_message(request, target_id=self._starting_agent_id)
@@ -409,8 +410,8 @@ class _HandoffCoordinator(Executor):
author_name=function_call.name,
)
# Add tool acknowledgement to both the conversation being sent and the full history
conversation.append(tool_message)
self._full_conversation.append(tool_message)
conversation.extend((tool_message,))
self._append_messages((tool_message,))
def _conversation_from_response(self, response: AgentExecutorResponse) -> list[ChatMessage]:
"""Return the authoritative conversation snapshot from an executor response."""
@@ -421,78 +422,41 @@ class _HandoffCoordinator(Executor):
)
return list(conversation)
def _get_cleaned_conversation(self, conversation: list[ChatMessage]) -> list[ChatMessage]:
"""Create a cleaned copy of conversation with tool-related content removed.
This method creates a copy of the conversation and removes tool-related content
before passing it to agents. The original conversation is preserved for handoff
detection and state management.
During handoffs, tool calls (including handoff tools) cause OpenAI API errors. The OpenAI
API requires that:
1. Assistant messages with tool_calls must be followed by corresponding tool responses
2. Tool response messages must follow an assistant message with tool_calls
To avoid these errors, we remove ALL tool-related content from the conversation:
- FunctionApprovalRequestContent and FunctionCallContent from assistant messages
- Tool response messages (Role.TOOL)
This follows the pattern from OpenAI Agents SDK's `remove_all_tools` filter, which strips
all tool-related content from conversation history during handoffs.
Removes:
- FunctionApprovalRequestContent: Approval requests for tools
- FunctionCallContent: Tool calls made by the agent
- Tool response messages (Role.TOOL with FunctionResultContent)
- Messages with only tool calls and no text content
Preserves:
- User messages
- Assistant messages with text content (tool calls are stripped out)
"""
# Create a copy to avoid modifying the original
cleaned: list[ChatMessage] = []
for msg in conversation:
# Skip tool response messages - they must be paired with tool calls which we're removing
if msg.role == Role.TOOL:
continue
# Check if message has tool-related content
has_tool_content = False
if msg.contents:
has_tool_content = any(
isinstance(content, (FunctionApprovalRequestContent, FunctionCallContent))
for content in msg.contents
)
# If no tool content, keep the original message
if not has_tool_content:
cleaned.append(msg)
continue
# Message has tool content - only keep if it also has text
if msg.text and msg.text.strip():
# Create fresh text-only message to avoid tool_calls being regenerated
msg_copy = ChatMessage(
role=msg.role,
text=msg.text,
author_name=msg.author_name,
)
cleaned.append(msg_copy)
return cleaned
async def _persist_state(self, ctx: WorkflowContext[Any, Any]) -> None:
"""Store authoritative conversation snapshot without losing rich metadata."""
state_payload = {"full_conversation": encode_chat_messages(self._full_conversation)}
state_payload = self.snapshot_state()
await ctx.set_executor_state(state_payload)
def _snapshot_pattern_metadata(self) -> dict[str, Any]:
"""Serialize pattern-specific state.
Handoff has no additional metadata beyond base conversation state.
Returns:
Empty dict (no pattern-specific state)
"""
return {}
def _restore_pattern_metadata(self, metadata: dict[str, Any]) -> None:
"""Restore pattern-specific state.
Handoff has no additional metadata beyond base conversation state.
Args:
metadata: Pattern-specific state dict (ignored)
"""
pass
def _restore_conversation_from_state(self, state: Mapping[str, Any]) -> list[ChatMessage]:
"""Rehydrate the coordinator's conversation history from checkpointed state."""
raw_conv = state.get("full_conversation")
if not isinstance(raw_conv, list):
return []
return decode_chat_messages(raw_conv) # type: ignore[arg-type]
"""Rehydrate the coordinator's conversation history from checkpointed state.
DEPRECATED: Use restore_state() instead. Kept for backward compatibility.
"""
from ._orchestration_state import OrchestrationState
orch_state_dict = {"conversation": state.get("full_conversation", state.get("conversation", []))}
temp_state = OrchestrationState.from_dict(orch_state_dict)
return list(temp_state.conversation)
def _apply_response_metadata(self, conversation: list[ChatMessage], agent_response: AgentRunResponse) -> None:
"""Merge top-level response metadata into the latest assistant message."""
@@ -766,7 +730,10 @@ class HandoffBuilder:
self._starting_agent_id: str | None = None
self._checkpoint_storage: CheckpointStorage | None = None
self._request_prompt: str | None = None
self._termination_condition: Callable[[list[ChatMessage]], bool] = _default_termination_condition
# Termination condition
self._termination_condition: Callable[[list[ChatMessage]], bool | Awaitable[bool]] = (
_default_termination_condition
)
self._auto_register_handoff_tools: bool = True
self._handoff_config: dict[str, list[str]] = {} # Maps agent_id -> [target_agent_ids]
@@ -814,36 +781,41 @@ class HandoffBuilder:
if not participants:
raise ValueError("participants cannot be empty")
wrapped: list[Executor] = []
named: dict[str, AgentProtocol | Executor] = {}
for participant in participants:
identifier: str
if isinstance(participant, Executor):
identifier = participant.id
elif isinstance(participant, AgentProtocol):
name_attr = getattr(participant, "name", None)
if not name_attr:
raise ValueError(
"Agents used in handoff workflows must have a stable name "
"so they can be addressed during routing."
)
identifier = str(name_attr)
else:
raise TypeError(
f"Participants must be AgentProtocol or Executor instances. Got {type(participant).__name__}."
)
if identifier in named:
raise ValueError(f"Duplicate participant name '{identifier}' detected")
named[identifier] = participant
metadata = prepare_participant_metadata(
named,
description_factory=lambda name, participant: getattr(participant, "description", None) or name,
)
wrapped = metadata["executors"]
seen_ids: set[str] = set()
alias_map: dict[str, str] = {}
def _register_alias(alias: str | None, exec_id: str) -> None:
"""Record canonical and sanitised aliases that resolve to the executor id."""
if not alias:
return
alias_map[alias] = exec_id
sanitized = _sanitize_alias(alias)
if sanitized and sanitized not in alias_map:
alias_map[sanitized] = exec_id
for p in participants:
executor = self._wrap_participant(p)
for executor in wrapped.values():
if executor.id in seen_ids:
raise ValueError(f"Duplicate participant with id '{executor.id}' detected")
seen_ids.add(executor.id)
wrapped.append(executor)
_register_alias(executor.id, executor.id)
if isinstance(p, AgentProtocol):
name = getattr(p, "name", None)
_register_alias(name, executor.id)
display = getattr(p, "display_name", None)
if isinstance(display, str) and display:
_register_alias(display, executor.id)
self._executors = {executor.id: executor for executor in wrapped}
self._aliases = alias_map
self._executors = {executor.id: executor for executor in wrapped.values()}
self._aliases = metadata["aliases"]
self._starting_agent_id = None
return self
@@ -1023,7 +995,7 @@ class HandoffBuilder:
new_tools: list[Any] = []
for exec_id in specialists:
alias = exec_id
sanitized = _sanitize_alias(alias)
sanitized = sanitize_identifier(alias)
tool = _create_handoff_tool(alias)
if tool.name not in existing_names:
new_tools.append(tool)
@@ -1184,12 +1156,16 @@ class HandoffBuilder:
self._checkpoint_storage = checkpoint_storage
return self
def with_termination_condition(self, condition: Callable[[list[ChatMessage]], bool]) -> "HandoffBuilder":
def with_termination_condition(
self, condition: Callable[[list[ChatMessage]], bool | Awaitable[bool]]
) -> "HandoffBuilder":
"""Set a custom termination condition for the handoff workflow.
The condition can be either synchronous or asynchronous.
Args:
condition: Function that receives the full conversation and returns True
if the workflow should terminate (not request further user input).
(or awaitable True) if the workflow should terminate (not request further user input).
Returns:
Self for chaining.
@@ -1198,9 +1174,19 @@ class HandoffBuilder:
.. code-block:: python
# Synchronous condition
builder.with_termination_condition(
lambda conv: len(conv) > 20 or any("goodbye" in msg.text.lower() for msg in conv[-2:])
)
# Asynchronous condition
async def check_termination(conv: list[ChatMessage]) -> bool:
# Can perform async operations
return len(conv) > 20
builder.with_termination_condition(check_termination)
"""
self._termination_condition = condition
return self
@@ -1308,6 +1294,14 @@ class HandoffBuilder:
if not specialists:
logger.warning("Handoff workflow has no specialist agents; the coordinator will loop with the user.")
descriptions = {
exec_id: getattr(executor, "description", None) or exec_id for exec_id, executor in self._executors.items()
}
participant_specs = {
exec_id: GroupChatParticipantSpec(name=exec_id, participant=executor, description=descriptions[exec_id])
for exec_id, executor in self._executors.items()
}
input_node = _InputToConversation(id="input-conversation")
request_info = RequestInfoExecutor(id=f"{starting_executor.id}_handoff_requests")
user_gateway = _UserInputGateway(
@@ -1316,48 +1310,50 @@ class HandoffBuilder:
prompt=self._request_prompt,
id="handoff-user-input",
)
coordinator = _HandoffCoordinator(
starting_agent_id=starting_executor.id,
specialist_ids={alias: exec_id for alias, exec_id in self._aliases.items() if exec_id in specialists},
input_gateway_id=user_gateway.id,
termination_condition=self._termination_condition,
id="handoff-coordinator",
handoff_tool_targets=handoff_tool_targets,
specialist_aliases = {alias: exec_id for alias, exec_id in self._aliases.items() if exec_id in specialists}
def _handoff_orchestrator_factory(_: _GroupChatConfig) -> Executor:
return _HandoffCoordinator(
starting_agent_id=starting_executor.id,
specialist_ids=specialist_aliases,
input_gateway_id=user_gateway.id,
termination_condition=self._termination_condition,
id="handoff-coordinator",
handoff_tool_targets=handoff_tool_targets,
)
wiring = _GroupChatConfig(
manager=None,
manager_name=self._starting_agent_id,
participants=participant_specs,
max_rounds=None,
participant_aliases=self._aliases,
participant_executors=self._executors,
)
builder = WorkflowBuilder(name=self._name, description=self._description)
builder.set_start_executor(input_node)
builder.add_edge(input_node, starting_executor)
builder.add_edge(starting_executor, coordinator)
result = assemble_group_chat_workflow(
wiring=wiring,
participant_factory=_default_participant_factory,
orchestrator_factory=_handoff_orchestrator_factory,
interceptors=(),
checkpoint_storage=self._checkpoint_storage,
builder=WorkflowBuilder(name=self._name, description=self._description),
return_builder=True,
)
if not isinstance(result, tuple):
raise TypeError("Expected tuple from assemble_group_chat_workflow with return_builder=True")
builder, coordinator = result
for specialist in specialists.values():
builder.add_edge(coordinator, specialist)
builder.add_edge(specialist, coordinator)
builder.add_edge(coordinator, user_gateway)
builder.add_edge(user_gateway, request_info)
builder.add_edge(request_info, user_gateway)
builder.add_edge(user_gateway, coordinator) # Route back to coordinator, not directly to agent
builder.add_edge(coordinator, starting_executor) # Coordinator sends trimmed request to agent
if self._checkpoint_storage is not None:
builder = builder.with_checkpointing(self._checkpoint_storage)
builder = builder.set_start_executor(input_node)
builder = builder.add_edge(input_node, starting_executor)
builder = builder.add_edge(coordinator, user_gateway)
builder = builder.add_edge(user_gateway, request_info)
builder = builder.add_edge(request_info, user_gateway)
builder = builder.add_edge(user_gateway, coordinator)
return builder.build()
def _wrap_participant(self, participant: AgentProtocol | Executor) -> Executor:
"""Ensure every participant is represented as an Executor instance."""
if isinstance(participant, Executor):
return participant
if isinstance(participant, AgentProtocol):
name = getattr(participant, "name", None)
if not name:
raise ValueError(
"Agents used in handoff workflows must have a stable name so they can be addressed during routing."
)
return AgentExecutor(participant, id=name)
raise TypeError(f"Participants must be AgentProtocol or Executor instances. Got {type(participant).__name__}.")
def _resolve_to_id(self, candidate: str | AgentProtocol | Executor) -> str:
"""Resolve a participant reference into a concrete executor identifier."""
if isinstance(candidate, Executor):
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,43 @@
# Copyright (c) Microsoft. All rights reserved.
"""Shared helpers for normalizing workflow message inputs."""
from collections.abc import Sequence
from agent_framework import ChatMessage, Role
def normalize_messages_input(
messages: str | ChatMessage | Sequence[str | ChatMessage] | None = None,
) -> list[ChatMessage]:
"""Normalize heterogeneous message inputs to a list of ChatMessage objects.
Args:
messages: String, ChatMessage, or sequence of either. None yields empty list.
Returns:
List of ChatMessage instances suitable for workflow consumption.
"""
if messages is None:
return []
if isinstance(messages, str):
return [ChatMessage(role=Role.USER, text=messages)]
if isinstance(messages, ChatMessage):
return [messages]
normalized: list[ChatMessage] = []
for item in messages:
if isinstance(item, str):
normalized.append(ChatMessage(role=Role.USER, text=item))
elif isinstance(item, ChatMessage):
normalized.append(item)
else:
raise TypeError(
f"Messages sequence must contain only str or ChatMessage instances; found {type(item).__name__}."
)
return normalized
__all__ = ["normalize_messages_input"]
@@ -2,7 +2,7 @@
import copy
import sys
from typing import Any, TypeVar
from typing import Any, TypeVar, cast
if sys.version_info >= (3, 11):
from typing import Self # pragma: no cover
@@ -37,7 +37,7 @@ class DictConvertible:
data = json.loads(raw)
if not isinstance(data, dict):
raise ValueError("JSON payload must decode to a mapping")
return cls.from_dict(data)
return cls.from_dict(cast(dict[str, Any], data))
def encode_value(value: Any) -> Any:
@@ -0,0 +1,92 @@
# Copyright (c) Microsoft. All rights reserved.
"""Unified state management for group chat orchestrators.
Provides OrchestrationState dataclass for standardized checkpoint serialization
across GroupChat, Handoff, and Magentic patterns.
"""
from dataclasses import dataclass, field
from typing import Any
from .._types import ChatMessage
def _new_chat_message_list() -> list[ChatMessage]:
"""Factory function for typed empty ChatMessage list.
Satisfies the type checker.
"""
return []
def _new_metadata_dict() -> dict[str, Any]:
"""Factory function for typed empty metadata dict.
Satisfies the type checker.
"""
return {}
@dataclass
class OrchestrationState:
"""Unified state container for orchestrator checkpointing.
This dataclass standardizes checkpoint serialization across all three
group chat patterns while allowing pattern-specific extensions via metadata.
Common attributes cover shared orchestration concerns (task, conversation,
round tracking). Pattern-specific state goes in the metadata dict.
Attributes:
conversation: Full conversation history (all messages)
round_index: Number of coordination rounds completed (0 if not tracked)
metadata: Extensible dict for pattern-specific state
task: Optional primary task/question being orchestrated
"""
conversation: list[ChatMessage] = field(default_factory=_new_chat_message_list)
round_index: int = 0
metadata: dict[str, Any] = field(default_factory=_new_metadata_dict)
task: ChatMessage | None = None
def to_dict(self) -> dict[str, Any]:
"""Serialize to dict for checkpointing.
Returns:
Dict with encoded conversation and metadata for persistence
"""
from ._conversation_state import encode_chat_messages
result: dict[str, Any] = {
"conversation": encode_chat_messages(self.conversation),
"round_index": self.round_index,
"metadata": dict(self.metadata),
}
if self.task is not None:
result["task"] = encode_chat_messages([self.task])[0]
return result
@classmethod
def from_dict(cls, data: dict[str, Any]) -> "OrchestrationState":
"""Deserialize from checkpointed dict.
Args:
data: Checkpoint data with encoded conversation
Returns:
Restored OrchestrationState instance
"""
from ._conversation_state import decode_chat_messages
task = None
if "task" in data:
decoded_tasks = decode_chat_messages([data["task"]])
task = decoded_tasks[0] if decoded_tasks else None
return cls(
conversation=decode_chat_messages(data.get("conversation", [])),
round_index=data.get("round_index", 0),
metadata=dict(data.get("metadata", {})),
task=task,
)
@@ -0,0 +1,190 @@
# Copyright (c) Microsoft. All rights reserved.
"""Shared orchestrator utilities for group chat patterns.
This module provides simple, reusable functions for common orchestration tasks.
No inheritance required - just import and call.
"""
import logging
from typing import TYPE_CHECKING, Any
from .._types import ChatMessage, Role
if TYPE_CHECKING:
from ._group_chat import _GroupChatRequestMessage # type: ignore[reportPrivateUsage]
logger = logging.getLogger(__name__)
def clean_conversation_for_handoff(conversation: list[ChatMessage]) -> list[ChatMessage]:
"""Remove tool-related content from conversation for clean handoffs.
During handoffs, tool calls can cause API errors because:
1. Assistant messages with tool_calls must be followed by tool responses
2. Tool response messages must follow an assistant message with tool_calls
This creates a cleaned copy removing ALL tool-related content.
Removes:
- FunctionApprovalRequestContent and FunctionCallContent from assistant messages
- Tool response messages (Role.TOOL)
- Messages with only tool calls and no text
Preserves:
- User messages
- Assistant messages with text content
Args:
conversation: Original conversation with potential tool content
Returns:
Cleaned conversation safe for handoff routing
"""
from agent_framework import FunctionApprovalRequestContent, FunctionCallContent
cleaned: list[ChatMessage] = []
for msg in conversation:
# Skip tool response messages entirely
if msg.role == Role.TOOL:
continue
# Check for tool-related content
has_tool_content = False
if msg.contents:
has_tool_content = any(
isinstance(content, (FunctionApprovalRequestContent, FunctionCallContent)) for content in msg.contents
)
# If no tool content, keep original
if not has_tool_content:
cleaned.append(msg)
continue
# Has tool content - only keep if it also has text
if msg.text and msg.text.strip():
# Create fresh text-only message
msg_copy = ChatMessage(
role=msg.role,
text=msg.text,
author_name=msg.author_name,
)
cleaned.append(msg_copy)
return cleaned
def create_completion_message(
*,
text: str | None = None,
author_name: str,
reason: str = "completed",
) -> ChatMessage:
"""Create a standardized completion message.
Simple helper to avoid duplicating completion message creation.
Args:
text: Message text, or None to generate default
author_name: Author/orchestrator name
reason: Reason for completion (for default text generation)
Returns:
ChatMessage with ASSISTANT role
"""
message_text = text or f"Conversation {reason}."
return ChatMessage(
role=Role.ASSISTANT,
text=message_text,
author_name=author_name,
)
def prepare_participant_request(
*,
participant_name: str,
conversation: list[ChatMessage],
instruction: str | None = None,
task: ChatMessage | None = None,
metadata: dict[str, Any] | None = None,
) -> "_GroupChatRequestMessage":
"""Create a standardized participant request message.
Simple helper to avoid duplicating request construction.
Args:
participant_name: Name of the target participant
conversation: Conversation history to send
instruction: Optional instruction from manager/orchestrator
task: Optional task context
metadata: Optional metadata dict
Returns:
GroupChatRequestMessage ready to send
"""
# Import here to avoid circular dependency
from ._group_chat import _GroupChatRequestMessage # type: ignore[reportPrivateUsage]
return _GroupChatRequestMessage(
agent_name=participant_name,
conversation=list(conversation),
instruction=instruction or "",
task=task,
metadata=metadata,
)
class ParticipantRegistry:
"""Simple registry for tracking participant executor IDs and routing info.
Provides a clean interface for the common pattern of mapping participant names
to executor IDs and tracking which are agents vs custom executors.
"""
def __init__(self) -> None:
self._participant_entry_ids: dict[str, str] = {}
self._agent_executor_ids: dict[str, str] = {}
self._executor_id_to_participant: dict[str, str] = {}
self._non_agent_participants: set[str] = set()
def register(
self,
name: str,
*,
entry_id: str,
is_agent: bool,
) -> None:
"""Register a participant's routing information.
Args:
name: Participant name
entry_id: Executor ID for this participant's entry point
is_agent: Whether this is an AgentExecutor (True) or custom Executor (False)
"""
self._participant_entry_ids[name] = entry_id
if is_agent:
self._agent_executor_ids[name] = entry_id
self._executor_id_to_participant[entry_id] = name
else:
self._non_agent_participants.add(name)
def get_entry_id(self, name: str) -> str | None:
"""Get the entry executor ID for a participant name."""
return self._participant_entry_ids.get(name)
def get_participant_name(self, executor_id: str) -> str | None:
"""Get the participant name for an executor ID (agents only)."""
return self._executor_id_to_participant.get(executor_id)
def is_agent(self, name: str) -> bool:
"""Check if a participant is an agent (vs custom executor)."""
return name in self._agent_executor_ids
def is_registered(self, name: str) -> bool:
"""Check if a participant is registered."""
return name in self._participant_entry_ids
def all_participants(self) -> set[str]:
"""Get all registered participant names."""
return set(self._participant_entry_ids.keys())
@@ -0,0 +1,136 @@
# Copyright (c) Microsoft. All rights reserved.
"""Shared participant helpers for orchestration builders."""
import re
from collections.abc import Callable, Iterable, Mapping
from dataclasses import dataclass
from typing import Any
from .._agents import AgentProtocol
from ._agent_executor import AgentExecutor
from ._executor import Executor
@dataclass
class GroupChatParticipantSpec:
"""Metadata describing a single participant in group chat orchestrations.
Used by multiple orchestration patterns (GroupChat, Handoff, Magentic) to describe
participants with consistent structure across different workflow types.
Attributes:
name: Unique identifier for the participant used by managers for selection
participant: AgentProtocol or Executor instance representing the participant
description: Human-readable description provided to managers for selection context
"""
name: str
participant: AgentProtocol | Executor
description: str
_SANITIZE_PATTERN = re.compile(r"[^0-9a-zA-Z]+")
def sanitize_identifier(value: str, *, default: str = "agent") -> str:
"""Return a deterministic, lowercase identifier derived from `value`."""
cleaned = _SANITIZE_PATTERN.sub("_", value).strip("_")
if not cleaned:
cleaned = default
if cleaned[0].isdigit():
cleaned = f"{default}_{cleaned}"
return cleaned.lower()
def wrap_participant(participant: AgentProtocol | Executor, *, executor_id: str | None = None) -> Executor:
"""Represent `participant` as an `Executor`."""
if isinstance(participant, Executor):
return participant
if not isinstance(participant, AgentProtocol):
raise TypeError(
f"Participants must implement AgentProtocol or be Executor instances. Got {type(participant).__name__}."
)
name = getattr(participant, "name", None)
if executor_id is None:
if not name:
raise ValueError("Agent participants must expose a stable 'name' attribute.")
executor_id = str(name)
return AgentExecutor(participant, id=executor_id)
def participant_description(participant: AgentProtocol | Executor, fallback: str) -> str:
"""Produce a human-readable description for manager context."""
if isinstance(participant, Executor):
description = getattr(participant, "description", None)
if isinstance(description, str) and description.strip():
return description.strip()
return fallback
description = getattr(participant, "description", None)
if isinstance(description, str) and description.strip():
return description.strip()
return fallback
def build_alias_map(participant: AgentProtocol | Executor, executor: Executor) -> dict[str, str]:
"""Collect canonical and sanitised aliases that should resolve to `executor`."""
aliases: dict[str, str] = {}
def _register(values: Iterable[str | None]) -> None:
for value in values:
if not value:
continue
key = str(value)
if key not in aliases:
aliases[key] = executor.id
sanitized = sanitize_identifier(key)
if sanitized not in aliases:
aliases[sanitized] = executor.id
_register([executor.id])
if isinstance(participant, AgentProtocol):
name = getattr(participant, "name", None)
display = getattr(participant, "display_name", None)
_register([name, display])
else:
display = getattr(participant, "display_name", None)
_register([display])
return aliases
def merge_alias_maps(maps: Iterable[Mapping[str, str]]) -> dict[str, str]:
"""Merge alias mappings, preserving the first occurrence of each alias."""
merged: dict[str, str] = {}
for mapping in maps:
for key, value in mapping.items():
merged.setdefault(key, value)
return merged
def prepare_participant_metadata(
participants: Mapping[str, AgentProtocol | Executor],
*,
executor_id_factory: Callable[[str, AgentProtocol | Executor], str | None] | None = None,
description_factory: Callable[[str, AgentProtocol | Executor], str] | None = None,
) -> dict[str, dict[str, Any]]:
"""Return metadata dicts for participants keyed by participant name."""
executors: dict[str, Executor] = {}
descriptions: dict[str, str] = {}
alias_maps: list[Mapping[str, str]] = []
for name, participant in participants.items():
desired_id = executor_id_factory(name, participant) if executor_id_factory else None
executor = wrap_participant(participant, executor_id=desired_id)
fallback_description = description_factory(name, participant) if description_factory else executor.id
descriptions[name] = participant_description(participant, fallback_description)
executors[name] = executor
alias_maps.append(build_alias_map(participant, executor))
aliases = merge_alias_maps(alias_maps)
return {
"executors": executors,
"descriptions": descriptions,
"aliases": aliases,
}
@@ -40,7 +40,7 @@ import logging
from collections.abc import Sequence
from typing import Any
from agent_framework import AgentProtocol, ChatMessage, Role
from agent_framework import AgentProtocol, ChatMessage
from ._agent_executor import (
AgentExecutor,
@@ -51,6 +51,7 @@ from ._executor import (
Executor,
handler,
)
from ._message_utils import normalize_messages_input
from ._workflow import Workflow
from ._workflow_builder import WorkflowBuilder
from ._workflow_context import WorkflowContext
@@ -63,16 +64,21 @@ class _InputToConversation(Executor):
@handler
async def from_str(self, prompt: str, ctx: WorkflowContext[list[ChatMessage]]) -> None:
await ctx.send_message([ChatMessage(Role.USER, text=prompt)])
await ctx.send_message(normalize_messages_input(prompt))
@handler
async def from_message(self, message: ChatMessage, ctx: WorkflowContext[list[ChatMessage]]) -> None: # type: ignore[name-defined]
await ctx.send_message([message])
async def from_message(self, message: ChatMessage, ctx: WorkflowContext[list[ChatMessage]]) -> None:
await ctx.send_message(normalize_messages_input(message))
@handler
async def from_messages(self, messages: list[ChatMessage], ctx: WorkflowContext[list[ChatMessage]]) -> None: # type: ignore[name-defined]
async def from_messages(
self,
messages: list[str | ChatMessage],
ctx: WorkflowContext[list[ChatMessage]],
) -> None:
# Make a copy to avoid mutation downstream
await ctx.send_message(list(messages))
normalized = normalize_messages_input(messages)
await ctx.send_message(list(normalized))
class _ResponseToConversation(Executor):
@@ -4,56 +4,72 @@ import logging
from collections.abc import Mapping
from dataclasses import fields, is_dataclass
from types import UnionType
from typing import Any, Union, get_args, get_origin
from typing import Any, TypeVar, Union, cast, get_args, get_origin
logger = logging.getLogger(__name__)
T = TypeVar("T")
def _coerce_to_type(value: Any, target_type: type) -> Any | None:
"""Best-effort conversion of value into target_type."""
def _coerce_to_type(value: Any, target_type: type[T]) -> T | None:
"""Best-effort conversion of value into target_type.
Args:
value: The value to convert (can be dict, dataclass, or object with __dict__)
target_type: The target type to convert to
Returns:
Instance of target_type if conversion succeeds, None otherwise
"""
if isinstance(value, target_type):
return value
return value # type: ignore[return-value]
# Convert dataclass instances or objects with __dict__ into dict first
value_as_dict: dict[str, Any]
if not isinstance(value, dict):
if is_dataclass(value):
value = {f.name: getattr(value, f.name) for f in fields(value)}
value_as_dict = {f.name: getattr(value, f.name) for f in fields(value)}
else:
value_dict = getattr(value, "__dict__", None)
if isinstance(value_dict, dict):
value = dict(value_dict)
value_as_dict = cast(dict[str, Any], value_dict)
else:
return None
else:
value_as_dict = cast(dict[str, Any], value)
if isinstance(value, dict):
ctor_kwargs: dict[str, Any] = dict(value)
# Try to construct the target type from the dict
ctor_kwargs: dict[str, Any] = dict(value_as_dict)
if is_dataclass(target_type):
field_names = {f.name for f in fields(target_type)}
ctor_kwargs = {k: v for k, v in value.items() if k in field_names}
if is_dataclass(target_type):
field_names = {f.name for f in fields(target_type)}
ctor_kwargs = {k: v for k, v in value_as_dict.items() if k in field_names}
try:
return target_type(**ctor_kwargs) # type: ignore[call-arg,return-value]
except TypeError as exc:
logger.debug(f"_coerce_to_type could not call {target_type.__name__}(**..): {exc}")
except Exception as exc: # pragma: no cover - unexpected constructor failure
logger.warning(
f"_coerce_to_type encountered unexpected error calling {target_type.__name__} constructor: {exc}"
)
# Fallback: try to create instance without __init__ and set attributes
try:
instance = object.__new__(target_type)
except Exception as exc: # pragma: no cover - pathological type
logger.debug(f"_coerce_to_type could not allocate {target_type.__name__} without __init__: {exc}")
return None
for key, val in value_as_dict.items():
try:
return target_type(**ctor_kwargs) # type: ignore[arg-type]
except TypeError as exc:
logger.debug(f"_coerce_to_type could not call {target_type.__name__}(**..): {exc}")
except Exception as exc: # pragma: no cover - unexpected constructor failure
logger.warning(
f"_coerce_to_type encountered unexpected error calling {target_type.__name__} constructor: {exc}"
setattr(instance, key, val)
except Exception as exc:
logger.debug(
f"_coerce_to_type could not set {target_type.__name__}.{key} during fallback assignment: {exc}"
)
try:
instance: Any = object.__new__(target_type)
except Exception as exc: # pragma: no cover - pathological type
logger.debug(f"_coerce_to_type could not allocate {target_type.__name__} without __init__: {exc}")
return None
for key, val in value.items():
try:
setattr(instance, key, val)
except Exception as exc:
logger.debug(
f"_coerce_to_type could not set {target_type.__name__}.{key} during fallback assignment: {exc}"
)
continue
return instance
return None
continue
return instance # type: ignore[return-value]
def is_instance_of(data: Any, target_type: type | UnionType | Any) -> bool:
@@ -89,14 +105,14 @@ def is_instance_of(data: Any, target_type: type | UnionType | Any) -> bool:
# Case 3: target_type is a generic type
if origin in [list, set]:
return isinstance(data, origin) and (
not args or all(any(is_instance_of(item, arg) for arg in args) for item in data)
not args or all(any(is_instance_of(item, arg) for arg in args) for item in data) # type: ignore[misc]
) # type: ignore
# Case 4: target_type is a tuple
if origin is tuple:
if len(args) == 2 and args[1] is Ellipsis: # Tuple[T, ...] case
element_type = args[0]
return isinstance(data, tuple) and all(is_instance_of(item, element_type) for item in data)
return isinstance(data, tuple) and all(is_instance_of(item, element_type) for item in data) # type: ignore[misc]
if len(args) == 1 and args[0] is Ellipsis: # Tuple[...] case
return isinstance(data, tuple)
if len(args) == 0:
@@ -135,7 +151,7 @@ def is_instance_of(data: Any, target_type: type | UnionType | Any) -> bool:
# and validators still receive a fully typed RequestResponse instance.
original_request = data.original_request
if isinstance(original_request, Mapping):
coerced = _coerce_to_type(dict(original_request), request_type)
coerced = _coerce_to_type(dict(original_request), request_type) # type: ignore[arg-type]
if coerced is None or not isinstance(coerced, request_type):
return False
data.original_request = coerced
@@ -838,11 +838,24 @@ class Workflow(DictConvertible):
def as_agent(self, name: str | None = None) -> WorkflowAgent:
"""Create a WorkflowAgent that wraps this workflow.
The returned agent converts standard agent inputs (strings, ChatMessage, or lists of these)
into a list[ChatMessage] that is passed to the workflow's start executor. This conversion
happens in WorkflowAgent._normalize_messages() which transforms:
- str -> [ChatMessage(role=USER, text=str)]
- ChatMessage -> [ChatMessage]
- list[str | ChatMessage] -> list[ChatMessage] (with string elements converted)
The workflow's start executor must accept list[ChatMessage] as an input type, otherwise
initialization will fail with a ValueError.
Args:
name: Optional name for the agent. If None, a default name will be generated.
Returns:
A WorkflowAgent instance that wraps this workflow.
Raises:
ValueError: If the workflow's start executor cannot handle list[ChatMessage] input.
"""
# Import here to avoid circular imports
from ._agent import WorkflowAgent
@@ -21,7 +21,7 @@ from ._events import (
WorkflowStartedEvent,
WorkflowStatusEvent,
WorkflowWarningEvent,
_framework_event_origin,
_framework_event_origin, # type: ignore
)
from ._runner_context import Message, RunnerContext
from ._shared_state import SharedState
@@ -0,0 +1,744 @@
# Copyright (c) Microsoft. All rights reserved.
from collections.abc import AsyncIterable, Callable
from typing import Any
import pytest
from agent_framework import (
AgentRunResponse,
AgentRunResponseUpdate,
AgentThread,
BaseAgent,
ChatMessage,
GroupChatBuilder,
GroupChatDirective,
GroupChatStateSnapshot,
MagenticAgentMessageEvent,
MagenticBuilder,
MagenticContext,
MagenticManagerBase,
MagenticOrchestratorMessageEvent,
Role,
TextContent,
Workflow,
WorkflowOutputEvent,
)
from agent_framework._workflows._checkpoint import InMemoryCheckpointStorage
from agent_framework._workflows._group_chat import (
GroupChatOrchestratorExecutor,
_default_orchestrator_factory, # type: ignore
_GroupChatConfig, # type: ignore
_PromptBasedGroupChatManager, # type: ignore
_SpeakerSelectorAdapter, # type: ignore
)
from agent_framework._workflows._magentic import (
_MagenticProgressLedger, # type: ignore
_MagenticProgressLedgerItem, # type: ignore
_MagenticStartMessage, # type: ignore
)
class StubAgent(BaseAgent):
def __init__(self, agent_name: str, reply_text: str, **kwargs: Any) -> None:
super().__init__(name=agent_name, description=f"Stub agent {agent_name}", **kwargs)
self._reply_text = reply_text
async def run( # type: ignore[override]
self,
messages: str | ChatMessage | list[str] | list[ChatMessage] | None = None,
*,
thread: AgentThread | None = None,
**kwargs: Any,
) -> AgentRunResponse:
response = ChatMessage(role=Role.ASSISTANT, text=self._reply_text, author_name=self.name)
return AgentRunResponse(messages=[response])
def run_stream( # type: ignore[override]
self,
messages: str | ChatMessage | list[str] | list[ChatMessage] | None = None,
*,
thread: AgentThread | None = None,
**kwargs: Any,
) -> AsyncIterable[AgentRunResponseUpdate]:
async def _stream() -> AsyncIterable[AgentRunResponseUpdate]:
yield AgentRunResponseUpdate(
contents=[TextContent(text=self._reply_text)], role=Role.ASSISTANT, author_name=self.name
)
return _stream()
def make_sequence_selector() -> Callable[[GroupChatStateSnapshot], Any]:
state_counter = {"value": 0}
async def _selector(state: GroupChatStateSnapshot) -> str | None:
participants = list(state["participants"].keys())
step = state_counter["value"]
if step == 0:
state_counter["value"] = step + 1
return participants[0]
if step == 1 and len(participants) > 1:
state_counter["value"] = step + 1
return participants[1]
return None
_selector.name = "manager" # type: ignore[attr-defined]
return _selector
class StubMagenticManager(MagenticManagerBase):
def __init__(self) -> None:
super().__init__(max_stall_count=3, max_round_count=5)
self._round = 0
async def plan(self, magentic_context: MagenticContext) -> ChatMessage:
return ChatMessage(role=Role.ASSISTANT, text="plan", author_name="magentic_manager")
async def replan(self, magentic_context: MagenticContext) -> ChatMessage:
return await self.plan(magentic_context)
async def create_progress_ledger(self, magentic_context: MagenticContext) -> _MagenticProgressLedger:
participants = list(magentic_context.participant_descriptions.keys())
target = participants[0] if participants else "agent"
if self._round == 0:
self._round += 1
return _MagenticProgressLedger(
is_request_satisfied=_MagenticProgressLedgerItem(reason="", answer=False),
is_in_loop=_MagenticProgressLedgerItem(reason="", answer=False),
is_progress_being_made=_MagenticProgressLedgerItem(reason="", answer=True),
next_speaker=_MagenticProgressLedgerItem(reason="", answer=target),
instruction_or_question=_MagenticProgressLedgerItem(reason="", answer="respond"),
)
return _MagenticProgressLedger(
is_request_satisfied=_MagenticProgressLedgerItem(reason="", answer=True),
is_in_loop=_MagenticProgressLedgerItem(reason="", answer=False),
is_progress_being_made=_MagenticProgressLedgerItem(reason="", answer=True),
next_speaker=_MagenticProgressLedgerItem(reason="", answer=target),
instruction_or_question=_MagenticProgressLedgerItem(reason="", answer=""),
)
async def prepare_final_answer(self, magentic_context: MagenticContext) -> ChatMessage:
return ChatMessage(role=Role.ASSISTANT, text="final", author_name="magentic_manager")
async def test_group_chat_builder_basic_flow() -> None:
selector = make_sequence_selector()
alpha = StubAgent("alpha", "ack from alpha")
beta = StubAgent("beta", "ack from beta")
workflow = (
GroupChatBuilder()
.select_speakers(selector, display_name="manager", final_message="done")
.participants(alpha=alpha, beta=beta)
.build()
)
outputs: list[ChatMessage] = []
async for event in workflow.run_stream("coordinate task"):
if isinstance(event, WorkflowOutputEvent):
data = event.data
if isinstance(data, ChatMessage):
outputs.append(data)
assert len(outputs) == 1
assert outputs[0].text == "done"
assert outputs[0].author_name == "manager"
async def test_magentic_builder_returns_workflow_and_runs() -> None:
manager = StubMagenticManager()
agent = StubAgent("writer", "first draft")
workflow = MagenticBuilder().participants(writer=agent).with_standard_manager(manager=manager).build()
assert isinstance(workflow, Workflow)
outputs: list[ChatMessage] = []
orchestrator_events: list[MagenticOrchestratorMessageEvent] = []
agent_events: list[MagenticAgentMessageEvent] = []
start_message = _MagenticStartMessage.from_string("compose summary")
async for event in workflow.run_stream(start_message):
if isinstance(event, MagenticOrchestratorMessageEvent):
orchestrator_events.append(event)
if isinstance(event, MagenticAgentMessageEvent):
agent_events.append(event)
if isinstance(event, WorkflowOutputEvent):
msg = event.data
if isinstance(msg, ChatMessage):
outputs.append(msg)
assert outputs, "Expected a final output message"
final = outputs[-1]
assert final.text == "final"
assert final.author_name == "magentic_manager"
assert orchestrator_events, "Expected orchestrator events to be emitted"
assert agent_events, "Expected agent message events to be emitted"
async def test_group_chat_as_agent_accepts_conversation() -> None:
selector = make_sequence_selector()
alpha = StubAgent("alpha", "ack from alpha")
beta = StubAgent("beta", "ack from beta")
workflow = (
GroupChatBuilder()
.select_speakers(selector, display_name="manager", final_message="done")
.participants(alpha=alpha, beta=beta)
.build()
)
agent = workflow.as_agent(name="group-chat-agent")
conversation = [
ChatMessage(role=Role.USER, text="kickoff", author_name="user"),
ChatMessage(role=Role.ASSISTANT, text="noted", author_name="alpha"),
]
response = await agent.run(conversation)
assert response.messages, "Expected agent conversation output"
async def test_magentic_as_agent_accepts_conversation() -> None:
manager = StubMagenticManager()
writer = StubAgent("writer", "draft")
workflow = MagenticBuilder().participants(writer=writer).with_standard_manager(manager=manager).build()
agent = workflow.as_agent(name="magentic-agent")
conversation = [
ChatMessage(role=Role.SYSTEM, text="Guidelines", author_name="system"),
ChatMessage(role=Role.USER, text="Summarize the findings", author_name="requester"),
]
response = await agent.run(conversation)
assert isinstance(response, AgentRunResponse)
# Comprehensive tests for group chat functionality
class TestGroupChatBuilder:
"""Tests for GroupChatBuilder validation and configuration."""
def test_build_without_manager_raises_error(self) -> None:
"""Test that building without a manager raises ValueError."""
agent = StubAgent("test", "response")
builder = GroupChatBuilder().participants([agent])
with pytest.raises(ValueError, match="manager must be configured before build"):
builder.build()
def test_build_without_participants_raises_error(self) -> None:
"""Test that building without participants raises ValueError."""
def selector(state: GroupChatStateSnapshot) -> str | None:
return None
builder = GroupChatBuilder().select_speakers(selector)
with pytest.raises(ValueError, match="participants must be configured before build"):
builder.build()
def test_duplicate_manager_configuration_raises_error(self) -> None:
"""Test that configuring multiple managers raises ValueError."""
def selector(state: GroupChatStateSnapshot) -> str | None:
return None
builder = GroupChatBuilder().select_speakers(selector)
with pytest.raises(ValueError, match="already has a manager configured"):
builder.select_speakers(selector)
def test_empty_participants_raises_error(self) -> None:
"""Test that empty participants list raises ValueError."""
def selector(state: GroupChatStateSnapshot) -> str | None:
return None
builder = GroupChatBuilder().select_speakers(selector)
with pytest.raises(ValueError, match="participants cannot be empty"):
builder.participants([])
def test_duplicate_participant_names_raises_error(self) -> None:
"""Test that duplicate participant names raise ValueError."""
agent1 = StubAgent("test", "response1")
agent2 = StubAgent("test", "response2")
def selector(state: GroupChatStateSnapshot) -> str | None:
return None
builder = GroupChatBuilder().select_speakers(selector)
with pytest.raises(ValueError, match="Duplicate participant name 'test'"):
builder.participants([agent1, agent2])
def test_agent_without_name_raises_error(self) -> None:
"""Test that agent without name attribute raises ValueError."""
class AgentWithoutName(BaseAgent):
def __init__(self) -> None:
super().__init__(name="", description="test")
async def run(self, messages: Any = None, *, thread: Any = None, **kwargs: Any) -> AgentRunResponse:
return AgentRunResponse(messages=[])
def run_stream(
self, messages: Any = None, *, thread: Any = None, **kwargs: Any
) -> AsyncIterable[AgentRunResponseUpdate]:
async def _stream() -> AsyncIterable[AgentRunResponseUpdate]:
yield AgentRunResponseUpdate(contents=[])
return _stream()
agent = AgentWithoutName()
def selector(state: GroupChatStateSnapshot) -> str | None:
return None
builder = GroupChatBuilder().select_speakers(selector)
with pytest.raises(ValueError, match="must define a non-empty 'name' attribute"):
builder.participants([agent])
def test_empty_participant_name_raises_error(self) -> None:
"""Test that empty participant name raises ValueError."""
agent = StubAgent("test", "response")
def selector(state: GroupChatStateSnapshot) -> str | None:
return None
builder = GroupChatBuilder().select_speakers(selector)
with pytest.raises(ValueError, match="participant names must be non-empty strings"):
builder.participants({"": agent})
class TestGroupChatOrchestrator:
"""Tests for GroupChatOrchestratorExecutor core functionality."""
async def test_max_rounds_enforcement(self) -> None:
"""Test that max_rounds properly limits conversation rounds."""
call_count = {"value": 0}
def selector(state: GroupChatStateSnapshot) -> str | None:
call_count["value"] += 1
# Always return the agent name to try to continue indefinitely
return "agent"
agent = StubAgent("agent", "response")
workflow = (
GroupChatBuilder()
.select_speakers(selector)
.participants([agent])
.with_max_rounds(2) # Limit to 2 rounds
.build()
)
outputs: list[ChatMessage] = []
async for event in workflow.run_stream("test task"):
if isinstance(event, WorkflowOutputEvent):
data = event.data
if isinstance(data, ChatMessage):
outputs.append(data)
# Should have terminated due to max_rounds, expect at least one output
assert len(outputs) >= 1
# The final message should be about round limit
final_output = outputs[-1]
assert "round limit" in final_output.text.lower()
async def test_unknown_participant_error(self) -> None:
"""Test that _apply_directive raises error for unknown participants."""
def selector(state: GroupChatStateSnapshot) -> str | None:
return "unknown_agent" # Return non-existent participant
agent = StubAgent("agent", "response")
workflow = GroupChatBuilder().select_speakers(selector).participants([agent]).build()
with pytest.raises(ValueError, match="Manager selected unknown participant 'unknown_agent'"):
async for _ in workflow.run_stream("test task"):
pass
async def test_directive_without_agent_name_raises_error(self) -> None:
"""Test that directive without agent_name raises error when finish=False."""
def bad_selector(state: GroupChatStateSnapshot) -> GroupChatDirective:
# Return a GroupChatDirective object instead of string to trigger error
return GroupChatDirective(finish=False, agent_name=None) # type: ignore
agent = StubAgent("agent", "response")
# The _SpeakerSelectorAdapter will catch this and raise TypeError
workflow = GroupChatBuilder().select_speakers(bad_selector).participants([agent]).build() # type: ignore
# This should raise a TypeError because selector doesn't return str or None
with pytest.raises(TypeError, match="must return a participant name \\(str\\) or None"):
async for _ in workflow.run_stream("test"):
pass
async def test_handle_empty_conversation_raises_error(self) -> None:
"""Test that empty conversation list raises ValueError."""
def selector(state: GroupChatStateSnapshot) -> str | None:
return None
agent = StubAgent("agent", "response")
workflow = GroupChatBuilder().select_speakers(selector).participants([agent]).build()
with pytest.raises(ValueError, match="requires at least one chat message"):
async for _ in workflow.run_stream([]):
pass
async def test_unknown_participant_response_raises_error(self) -> None:
"""Test that responses from unknown participants raise errors."""
def selector(state: GroupChatStateSnapshot) -> str | None:
return "agent"
# Create orchestrator to test _ingest_participant_message directly
orchestrator = GroupChatOrchestratorExecutor(
manager=selector, # type: ignore
participants={"agent": "test agent"},
manager_name="test_manager", # type: ignore
)
# Mock the workflow context
class MockContext:
async def yield_output(self, message: ChatMessage) -> None:
pass
ctx = MockContext()
# Initialize orchestrator state
orchestrator._task_message = ChatMessage(role=Role.USER, text="test") # type: ignore
orchestrator._conversation = [orchestrator._task_message] # type: ignore
orchestrator._history = [] # type: ignore
orchestrator._pending_agent = None # type: ignore
orchestrator._round_index = 0 # type: ignore
# Test with unknown participant
message = ChatMessage(role=Role.ASSISTANT, text="response")
with pytest.raises(ValueError, match="Received response from unknown participant 'unknown'"):
await orchestrator._ingest_participant_message("unknown", message, ctx) # type: ignore
async def test_state_build_before_initialization_raises_error(self) -> None:
"""Test that _build_state raises error before task message initialization."""
def selector(state: GroupChatStateSnapshot) -> str | None:
return None
orchestrator = GroupChatOrchestratorExecutor(
manager=selector, # type: ignore
participants={"agent": "test agent"},
manager_name="test_manager", # type: ignore
)
with pytest.raises(RuntimeError, match="state not initialized with task message"):
orchestrator._build_state() # type: ignore
class TestSpeakerSelectorAdapter:
"""Tests for _SpeakerSelectorAdapter functionality."""
async def test_selector_returning_list_with_multiple_items_raises_error(self) -> None:
"""Test that selector returning list with multiple items raises error."""
def bad_selector(state: GroupChatStateSnapshot) -> list[str]:
return ["agent1", "agent2"] # Multiple items
adapter = _SpeakerSelectorAdapter(bad_selector, manager_name="manager")
state = {
"participants": {"agent1": "desc1", "agent2": "desc2"},
"task": ChatMessage(role=Role.USER, text="test"),
"conversation": (),
"history": (),
"round_index": 0,
"pending_agent": None,
}
with pytest.raises(ValueError, match="must return a single participant name"):
await adapter(state)
async def test_selector_returning_non_string_raises_error(self) -> None:
"""Test that selector returning non-string raises TypeError."""
def bad_selector(state: GroupChatStateSnapshot) -> int:
return 42 # Not a string
adapter = _SpeakerSelectorAdapter(bad_selector, manager_name="manager")
state = {
"participants": {"agent": "desc"},
"task": ChatMessage(role=Role.USER, text="test"),
"conversation": (),
"history": (),
"round_index": 0,
"pending_agent": None,
}
with pytest.raises(TypeError, match="must return a participant name \\(str\\) or None"):
await adapter(state)
async def test_selector_returning_empty_list_finishes(self) -> None:
"""Test that selector returning empty list finishes conversation."""
def empty_selector(state: GroupChatStateSnapshot) -> list[str]:
return [] # Empty list should finish
adapter = _SpeakerSelectorAdapter(empty_selector, manager_name="manager")
state = {
"participants": {"agent": "desc"},
"task": ChatMessage(role=Role.USER, text="test"),
"conversation": (),
"history": (),
"round_index": 0,
"pending_agent": None,
}
directive = await adapter(state)
assert directive.finish is True
assert directive.final_message is not None
class TestCheckpointing:
"""Tests for checkpointing functionality."""
async def test_workflow_with_checkpointing(self) -> None:
"""Test that workflow works with checkpointing enabled."""
def selector(state: GroupChatStateSnapshot) -> str | None:
if state["round_index"] >= 1:
return None
return "agent"
agent = StubAgent("agent", "response")
storage = InMemoryCheckpointStorage()
workflow = (
GroupChatBuilder().select_speakers(selector).participants([agent]).with_checkpointing(storage).build()
)
outputs: list[ChatMessage] = []
async for event in workflow.run_stream("test task"):
if isinstance(event, WorkflowOutputEvent):
data = event.data
if isinstance(data, ChatMessage):
outputs.append(data)
assert len(outputs) == 1 # Should complete normally
class TestPromptBasedManager:
"""Tests for _PromptBasedGroupChatManager."""
async def test_manager_with_missing_next_agent_raises_error(self) -> None:
"""Test that manager directive without next_agent raises RuntimeError."""
class MockChatClient:
async def get_response(self, messages: Any, response_format: Any = None) -> Any:
# Return response that has finish=False but no next_agent
class MockResponse:
def __init__(self) -> None:
self.value = {"finish": False, "next_agent": None}
self.messages: list[Any] = []
return MockResponse()
manager = _PromptBasedGroupChatManager(MockChatClient()) # type: ignore
state = {
"participants": {"agent": "desc"},
"task": ChatMessage(role=Role.USER, text="test"),
"conversation": (),
}
with pytest.raises(RuntimeError, match="missing next_agent while finish is False"):
await manager(state)
async def test_manager_with_unknown_participant_raises_error(self) -> None:
"""Test that manager selecting unknown participant raises RuntimeError."""
class MockChatClient:
async def get_response(self, messages: Any, response_format: Any = None) -> Any:
# Return response selecting unknown participant
class MockResponse:
def __init__(self) -> None:
self.value = {"finish": False, "next_agent": "unknown"}
self.messages: list[Any] = []
return MockResponse()
manager = _PromptBasedGroupChatManager(MockChatClient()) # type: ignore
state = {
"participants": {"agent": "desc"},
"task": ChatMessage(role=Role.USER, text="test"),
"conversation": (),
}
with pytest.raises(RuntimeError, match="Manager selected unknown participant 'unknown'"):
await manager(state)
class TestFactoryFunctions:
"""Tests for factory functions."""
def test_default_orchestrator_factory_without_manager_raises_error(self) -> None:
"""Test that default factory requires manager to be set."""
config = _GroupChatConfig(manager=None, manager_name="test", participants={})
with pytest.raises(RuntimeError, match="requires a manager to be set"):
_default_orchestrator_factory(config)
class TestConversationHandling:
"""Tests for different conversation input types."""
async def test_handle_string_input(self) -> None:
"""Test handling string input creates proper ChatMessage."""
def selector(state: GroupChatStateSnapshot) -> str | None:
# Verify the task was properly converted
assert state["task"].role == Role.USER
assert state["task"].text == "test string"
return None
agent = StubAgent("agent", "response")
workflow = GroupChatBuilder().select_speakers(selector).participants([agent]).build()
outputs: list[ChatMessage] = []
async for event in workflow.run_stream("test string"):
if isinstance(event, WorkflowOutputEvent):
data = event.data
if isinstance(data, ChatMessage):
outputs.append(data)
assert len(outputs) == 1
async def test_handle_chat_message_input(self) -> None:
"""Test handling ChatMessage input directly."""
task_message = ChatMessage(role=Role.USER, text="test message")
def selector(state: GroupChatStateSnapshot) -> str | None:
# Verify the task message was preserved
assert state["task"] == task_message
return None
agent = StubAgent("agent", "response")
workflow = GroupChatBuilder().select_speakers(selector).participants([agent]).build()
outputs: list[ChatMessage] = []
async for event in workflow.run_stream(task_message):
if isinstance(event, WorkflowOutputEvent):
data = event.data
if isinstance(data, ChatMessage):
outputs.append(data)
assert len(outputs) == 1
async def test_handle_conversation_list_input(self) -> None:
"""Test handling conversation list preserves context."""
conversation = [
ChatMessage(role=Role.SYSTEM, text="system message"),
ChatMessage(role=Role.USER, text="user message"),
]
def selector(state: GroupChatStateSnapshot) -> str | None:
# Verify conversation context is preserved
assert len(state["conversation"]) == 2
assert state["task"].text == "user message"
return None
agent = StubAgent("agent", "response")
workflow = GroupChatBuilder().select_speakers(selector).participants([agent]).build()
outputs: list[ChatMessage] = []
async for event in workflow.run_stream(conversation):
if isinstance(event, WorkflowOutputEvent):
data = event.data
if isinstance(data, ChatMessage):
outputs.append(data)
assert len(outputs) == 1
class TestRoundLimitEnforcement:
"""Tests for round limit checking functionality."""
async def test_round_limit_in_apply_directive(self) -> None:
"""Test round limit enforcement in _apply_directive."""
rounds_called = {"count": 0}
def selector(state: GroupChatStateSnapshot) -> str | None:
rounds_called["count"] += 1
# Keep trying to select agent to test limit enforcement
return "agent"
agent = StubAgent("agent", "response")
workflow = (
GroupChatBuilder()
.select_speakers(selector)
.participants([agent])
.with_max_rounds(1) # Very low limit
.build()
)
outputs: list[ChatMessage] = []
async for event in workflow.run_stream("test"):
if isinstance(event, WorkflowOutputEvent):
data = event.data
if isinstance(data, ChatMessage):
outputs.append(data)
# Should have at least one output (the round limit message)
assert len(outputs) >= 1
# The last message should be about round limit
final_output = outputs[-1]
assert "round limit" in final_output.text.lower()
async def test_round_limit_in_ingest_participant_message(self) -> None:
"""Test round limit enforcement after participant response."""
responses_received = {"count": 0}
def selector(state: GroupChatStateSnapshot) -> str | None:
responses_received["count"] += 1
if responses_received["count"] == 1:
return "agent" # First call selects agent
return "agent" # Try to continue, but should hit limit
agent = StubAgent("agent", "response from agent")
workflow = (
GroupChatBuilder()
.select_speakers(selector)
.participants([agent])
.with_max_rounds(1) # Hit limit after first response
.build()
)
outputs: list[ChatMessage] = []
async for event in workflow.run_stream("test"):
if isinstance(event, WorkflowOutputEvent):
data = event.data
if isinstance(data, ChatMessage):
outputs.append(data)
# Should have at least one output (the round limit message)
assert len(outputs) >= 1
# The last message should be about round limit
final_output = outputs[-1]
assert "round limit" in final_output.text.lower()
@@ -54,6 +54,7 @@ class _RecordingAgent(BaseAgent):
extra_properties: dict[str, object] | None = None,
) -> None:
super().__init__(id=name, name=name, display_name=name)
self._agent_name = name
self.handoff_to = handoff_to
self.calls: list[list[ChatMessage]] = []
self._text_handoff = text_handoff
@@ -72,7 +73,7 @@ class _RecordingAgent(BaseAgent):
additional_properties = _merge_additional_properties(
self.handoff_to, self._text_handoff, self._extra_properties
)
contents = _build_reply_contents(self.name, self.handoff_to, self._text_handoff, self._next_call_id())
contents = _build_reply_contents(self._agent_name, self.handoff_to, self._text_handoff, self._next_call_id())
reply = ChatMessage(
role=Role.ASSISTANT,
contents=contents,
@@ -91,7 +92,7 @@ class _RecordingAgent(BaseAgent):
conversation = _normalise(messages)
self.calls.append(conversation)
additional_props = _merge_additional_properties(self.handoff_to, self._text_handoff, self._extra_properties)
contents = _build_reply_contents(self.name, self.handoff_to, self._text_handoff, self._next_call_id())
contents = _build_reply_contents(self._agent_name, self.handoff_to, self._text_handoff, self._next_call_id())
yield AgentRunResponseUpdate(
contents=contents,
role=Role.ASSISTANT,
@@ -357,3 +358,38 @@ async def test_multiple_runs_dont_leak_conversation():
assert not any("First run message" in msg.text for msg in second_run_user_messages if msg.text), (
"Second run should NOT contain first run's messages"
)
async def test_handoff_async_termination_condition() -> None:
"""Test that async termination conditions work correctly."""
termination_call_count = 0
async def async_termination(conv: list[ChatMessage]) -> bool:
nonlocal termination_call_count
termination_call_count += 1
user_count = sum(1 for msg in conv if msg.role == Role.USER)
return user_count >= 2
coordinator = _RecordingAgent(name="coordinator")
workflow = (
HandoffBuilder(participants=[coordinator])
.set_coordinator(coordinator)
.with_termination_condition(async_termination)
.build()
)
events = await _drain(workflow.run_stream("First user message"))
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
assert requests
events = await _drain(workflow.send_responses_streaming({requests[-1].request_id: "Second user message"}))
outputs = [ev for ev in events if isinstance(ev, WorkflowOutputEvent)]
assert len(outputs) == 1
final_conversation = outputs[0].data
assert isinstance(final_conversation, list)
final_conv_list = cast(list[ChatMessage], final_conversation)
user_messages = [msg for msg in final_conv_list if msg.role == Role.USER]
assert len(user_messages) == 2
assert termination_call_count > 0
@@ -2,7 +2,7 @@
from collections.abc import AsyncIterable
from dataclasses import dataclass
from typing import Any
from typing import Any, cast
import pytest
@@ -15,13 +15,12 @@ from agent_framework import (
ChatResponse,
ChatResponseUpdate,
Executor,
MagenticAgentMessageEvent,
MagenticBuilder,
MagenticManagerBase,
MagenticPlanReviewDecision,
MagenticPlanReviewReply,
MagenticPlanReviewRequest,
MagenticProgressLedger,
MagenticProgressLedgerItem,
RequestInfoEvent,
Role,
TextContent,
@@ -34,17 +33,19 @@ from agent_framework import (
handler,
)
from agent_framework._workflows._checkpoint import InMemoryCheckpointStorage
from agent_framework._workflows._magentic import (
from agent_framework._workflows._magentic import ( # type: ignore[reportPrivateUsage]
MagenticAgentExecutor,
MagenticContext,
MagenticOrchestratorExecutor,
MagenticStartMessage,
_MagenticProgressLedger, # type: ignore
_MagenticProgressLedgerItem, # type: ignore
_MagenticStartMessage, # type: ignore
)
def test_magentic_start_message_from_string():
msg = MagenticStartMessage.from_string("Do the thing")
assert isinstance(msg, MagenticStartMessage)
msg = _MagenticStartMessage.from_string("Do the thing")
assert isinstance(msg, _MagenticStartMessage)
assert isinstance(msg.task, ChatMessage)
assert msg.task.role == Role.USER
assert msg.task.text == "Do the thing"
@@ -114,8 +115,9 @@ class FakeManager(MagenticManagerBase):
super().restore_state(state)
ledger_state = state.get("task_ledger")
if isinstance(ledger_state, dict):
facts_payload = ledger_state.get("facts") # type: ignore[reportUnknownMemberType]
plan_payload = ledger_state.get("plan") # type: ignore[reportUnknownMemberType]
ledger_dict = cast(dict[str, Any], ledger_state)
facts_payload = cast(dict[str, Any] | None, ledger_dict.get("facts"))
plan_payload = cast(dict[str, Any] | None, ledger_dict.get("plan"))
if facts_payload is not None and plan_payload is not None:
try:
facts = ChatMessage.from_dict(facts_payload)
@@ -138,14 +140,14 @@ class FakeManager(MagenticManagerBase):
combined = f"Task: {magentic_context.task.text}\n\nFacts:\n{facts.text}\n\nPlan:\n{plan.text}"
return ChatMessage(role=Role.ASSISTANT, text=combined, author_name="magentic_manager")
async def create_progress_ledger(self, magentic_context: MagenticContext) -> MagenticProgressLedger:
async def create_progress_ledger(self, magentic_context: MagenticContext) -> _MagenticProgressLedger:
is_satisfied = self.satisfied_after_signoff and len(magentic_context.chat_history) > 0
return MagenticProgressLedger(
is_request_satisfied=MagenticProgressLedgerItem(reason="test", answer=is_satisfied),
is_in_loop=MagenticProgressLedgerItem(reason="test", answer=False),
is_progress_being_made=MagenticProgressLedgerItem(reason="test", answer=True),
next_speaker=MagenticProgressLedgerItem(reason="test", answer=self.next_speaker_name),
instruction_or_question=MagenticProgressLedgerItem(reason="test", answer=self.instruction_text),
return _MagenticProgressLedger(
is_request_satisfied=_MagenticProgressLedgerItem(reason="test", answer=is_satisfied),
is_in_loop=_MagenticProgressLedgerItem(reason="test", answer=False),
is_progress_being_made=_MagenticProgressLedgerItem(reason="test", answer=True),
next_speaker=_MagenticProgressLedgerItem(reason="test", answer=self.next_speaker_name),
instruction_or_question=_MagenticProgressLedgerItem(reason="test", answer=self.instruction_text),
)
async def prepare_final_answer(self, magentic_context: MagenticContext) -> ChatMessage:
@@ -175,7 +177,7 @@ async def test_standard_manager_progress_ledger_and_fallback():
)
ledger = await manager.create_progress_ledger(ctx.clone())
assert isinstance(ledger, MagenticProgressLedger)
assert isinstance(ledger, _MagenticProgressLedger)
assert ledger.next_speaker.answer == "agentA"
manager.satisfied_after_signoff = False
@@ -328,13 +330,11 @@ async def test_magentic_checkpoint_resume_round_trip():
.build()
)
orchestrator = next(
exec for exec in wf_resume.workflow.executors.values() if isinstance(exec, MagenticOrchestratorExecutor)
)
orchestrator = next(exec for exec in wf_resume.executors.values() if isinstance(exec, MagenticOrchestratorExecutor))
reply = MagenticPlanReviewReply(decision=MagenticPlanReviewDecision.APPROVE)
completed: WorkflowOutputEvent | None = None
async for event in wf_resume.workflow.run_stream_from_checkpoint(
async for event in wf_resume.run_stream_from_checkpoint(
resume_checkpoint.checkpoint_id,
responses={req_event.request_id: reply},
):
@@ -346,8 +346,8 @@ async def test_magentic_checkpoint_resume_round_trip():
assert orchestrator._context.chat_history # type: ignore[reportPrivateUsage]
assert orchestrator._task_ledger is not None # type: ignore[reportPrivateUsage]
assert manager2.task_ledger is not None
# Initial message should be the task ledger plan
assert orchestrator._context.chat_history[0].text == orchestrator._task_ledger.text # type: ignore[reportPrivateUsage]
# Latest entry in chat history should be the task ledger plan
assert orchestrator._context.chat_history[-1].text == orchestrator._task_ledger.text # type: ignore[reportPrivateUsage]
class _DummyExec(Executor):
@@ -472,24 +472,24 @@ class InvokeOnceManager(MagenticManagerBase):
async def replan(self, magentic_context: MagenticContext) -> ChatMessage:
return ChatMessage(role=Role.ASSISTANT, text="re-ledger")
async def create_progress_ledger(self, magentic_context: MagenticContext) -> MagenticProgressLedger:
async def create_progress_ledger(self, magentic_context: MagenticContext) -> _MagenticProgressLedger:
if not self._invoked:
# First round: ask agentA to respond
self._invoked = True
return MagenticProgressLedger(
is_request_satisfied=MagenticProgressLedgerItem(reason="r", answer=False),
is_in_loop=MagenticProgressLedgerItem(reason="r", answer=False),
is_progress_being_made=MagenticProgressLedgerItem(reason="r", answer=True),
next_speaker=MagenticProgressLedgerItem(reason="r", answer="agentA"),
instruction_or_question=MagenticProgressLedgerItem(reason="r", answer="say hi"),
return _MagenticProgressLedger(
is_request_satisfied=_MagenticProgressLedgerItem(reason="r", answer=False),
is_in_loop=_MagenticProgressLedgerItem(reason="r", answer=False),
is_progress_being_made=_MagenticProgressLedgerItem(reason="r", answer=True),
next_speaker=_MagenticProgressLedgerItem(reason="r", answer="agentA"),
instruction_or_question=_MagenticProgressLedgerItem(reason="r", answer="say hi"),
)
# Next round: mark satisfied so run can conclude
return MagenticProgressLedger(
is_request_satisfied=MagenticProgressLedgerItem(reason="r", answer=True),
is_in_loop=MagenticProgressLedgerItem(reason="r", answer=False),
is_progress_being_made=MagenticProgressLedgerItem(reason="r", answer=True),
next_speaker=MagenticProgressLedgerItem(reason="r", answer="agentA"),
instruction_or_question=MagenticProgressLedgerItem(reason="r", answer="done"),
return _MagenticProgressLedger(
is_request_satisfied=_MagenticProgressLedgerItem(reason="r", answer=True),
is_in_loop=_MagenticProgressLedgerItem(reason="r", answer=False),
is_progress_being_made=_MagenticProgressLedgerItem(reason="r", answer=True),
next_speaker=_MagenticProgressLedgerItem(reason="r", answer="agentA"),
instruction_or_question=_MagenticProgressLedgerItem(reason="r", answer="done"),
)
async def prepare_final_answer(self, magentic_context: MagenticContext) -> ChatMessage:
@@ -533,17 +533,10 @@ class StubAssistantsAgent(BaseAgent):
async def _collect_agent_responses_setup(participant_obj: object):
captured: list[ChatMessage] = []
async def sink(event) -> None: # type: ignore[no-untyped-def]
from agent_framework._workflows._magentic import MagenticAgentMessageEvent
if isinstance(event, MagenticAgentMessageEvent) and event.message is not None:
captured.append(event.message)
wf = (
MagenticBuilder()
.participants(agentA=participant_obj) # type: ignore[arg-type]
.with_standard_manager(InvokeOnceManager())
.on_event(sink) # type: ignore
.build()
)
@@ -551,6 +544,10 @@ async def _collect_agent_responses_setup(participant_obj: object):
events: list[WorkflowEvent] = []
async for ev in wf.run_stream("task"): # plan review disabled
events.append(ev)
if isinstance(ev, WorkflowOutputEvent):
break
if isinstance(ev, MagenticAgentMessageEvent) and ev.message is not None:
captured.append(ev.message)
if len(events) > 50:
break
@@ -559,7 +556,7 @@ async def _collect_agent_responses_setup(participant_obj: object):
async def test_agent_executor_invoke_with_thread_chat_client():
captured = await _collect_agent_responses_setup(StubThreadAgent())
# Should have at least one response from agentA via MagenticAgentExecutor path
# Should have at least one response from agentA via _MagenticAgentExecutor path
assert any((m.author_name == "agentA" and "ok" in (m.text or "")) for m in captured)
@@ -685,7 +682,7 @@ async def test_magentic_checkpoint_resume_rejects_participant_renames():
.build()
)
with pytest.raises(RuntimeError, match="participant names do not match"):
with pytest.raises(ValueError, match="Workflow graph has changed"):
async for _ in renamed_workflow.run_stream_from_checkpoint(
target_checkpoint.checkpoint_id, # type: ignore[reportUnknownMemberType]
responses={req_event.request_id: MagenticPlanReviewReply(decision=MagenticPlanReviewDecision.APPROVE)},
@@ -704,13 +701,13 @@ class NotProgressingManager(MagenticManagerBase):
async def replan(self, magentic_context: MagenticContext) -> ChatMessage:
return ChatMessage(role=Role.ASSISTANT, text="re-ledger")
async def create_progress_ledger(self, magentic_context: MagenticContext) -> MagenticProgressLedger:
return MagenticProgressLedger(
is_request_satisfied=MagenticProgressLedgerItem(reason="r", answer=False),
is_in_loop=MagenticProgressLedgerItem(reason="r", answer=True),
is_progress_being_made=MagenticProgressLedgerItem(reason="r", answer=False),
next_speaker=MagenticProgressLedgerItem(reason="r", answer="agentA"),
instruction_or_question=MagenticProgressLedgerItem(reason="r", answer="done"),
async def create_progress_ledger(self, magentic_context: MagenticContext) -> _MagenticProgressLedger:
return _MagenticProgressLedger(
is_request_satisfied=_MagenticProgressLedgerItem(reason="r", answer=False),
is_in_loop=_MagenticProgressLedgerItem(reason="r", answer=True),
is_progress_being_made=_MagenticProgressLedgerItem(reason="r", answer=False),
next_speaker=_MagenticProgressLedgerItem(reason="r", answer="agentA"),
instruction_or_question=_MagenticProgressLedgerItem(reason="r", answer="done"),
)
async def prepare_final_answer(self, magentic_context: MagenticContext) -> ChatMessage:
@@ -6,7 +6,10 @@ import inspect
import json
import logging
from dataclasses import fields, is_dataclass
from typing import Any, get_args, get_origin
from types import UnionType
from typing import Any, Union, get_args, get_origin
from agent_framework import ChatMessage
logger = logging.getLogger(__name__)
@@ -110,10 +113,25 @@ def extract_executor_message_types(executor: Any) -> list[Any]:
return message_types
def _contains_chat_message(type_hint: Any) -> bool:
"""Check whether the provided type hint directly or indirectly references ChatMessage."""
if type_hint is ChatMessage:
return True
origin = get_origin(type_hint)
if origin in (list, tuple):
return any(_contains_chat_message(arg) for arg in get_args(type_hint))
if origin in (Union, UnionType):
return any(_contains_chat_message(arg) for arg in get_args(type_hint))
return False
def select_primary_input_type(message_types: list[Any]) -> Any | None:
"""Choose the most user-friendly input type for workflow inputs.
Prefers str and dict types for better user experience.
Prefers ChatMessage (or containers thereof) and then falls back to primitives.
Args:
message_types: List of possible message types
@@ -124,6 +142,10 @@ def select_primary_input_type(message_types: list[Any]) -> Any | None:
if not message_types:
return None
for message_type in message_types:
if _contains_chat_message(message_type):
return ChatMessage
preferred = (str, dict)
for candidate in preferred: