mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
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:
committed by
GitHub
Unverified
parent
899d8ff775
commit
e3aad8e4e0
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user