mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: [BREAKING] Moved to a single get_response and run API (#3379)
* WIP * big update to new ResponseStream model * fixed tests and typing * fixed tests and typing * fixed tools typevar import * fix * mypy fix * mypy fixes and some cleanup * fix missing quoted names * and client * fix imports agui * fix anthropic override * fix agui * fix ag ui * fix import * fix anthropic types * fix mypy * refactoring * updated typing * fix 3.11 * fixes * redid layering of chat clients and agents * redid layering of chat clients and agents * Fix lint, type, and test issues after rebase - Add @overload decorators to AgentProtocol.run() for type compatibility - Add missing docstring params (middleware, function_invocation_configuration) - Fix TODO format (TD002) by adding author tags - Fix broken observability tests from upstream: - Replace non-existent use_instrumentation with direct instantiation - Replace non-existent use_agent_instrumentation with AgentTelemetryLayer mixin - Fix get_streaming_response to use get_response(stream=True) - Add AgentInitializationError import - Update streaming exception tests to match actual behavior * Fix AgentExecutionException import error in test_agents.py - Replace non-existent AgentExecutionException with AgentRunException * Fix test import and asyncio deprecation issues - Add 'tests' to pythonpath in ag-ui pyproject.toml for utils_test_ag_ui import - Replace deprecated asyncio.get_event_loop().run_until_complete with asyncio.run * Fix azure-ai test failures - Update _prepare_options patching to use correct class path - Fix test_to_azure_ai_agent_tools_web_search_missing_connection to clear env vars * Convert ag-ui utils_test_ag_ui.py to conftest.py - Move test utilities to conftest.py for proper pytest discovery - Update all test imports to use conftest instead of utils_test_ag_ui - Remove old utils_test_ag_ui.py file - Revert pythonpath change in pyproject.toml * fix: use relative imports for ag-ui test utilities * fix agui * Rename Bare*Client to Raw*Client and BaseChatClient - Renamed BareChatClient to BaseChatClient (abstract base class) - Renamed BareOpenAIChatClient to RawOpenAIChatClient - Renamed BareOpenAIResponsesClient to RawOpenAIResponsesClient - Renamed BareAzureAIClient to RawAzureAIClient - Added warning docstrings to Raw* classes about layer ordering - Updated README in samples/getting_started/agents/custom with layer docs - Added test for span ordering with function calling * Fix layer ordering: FunctionInvocationLayer before ChatTelemetryLayer This ensures each inner LLM call gets its own telemetry span, resulting in the correct span sequence: chat -> execute_tool -> chat Updated all production clients and test mocks to use correct ordering: - ChatMiddlewareLayer (first) - FunctionInvocationLayer (second) - ChatTelemetryLayer (third) - BaseChatClient/Raw...Client (fourth) * Remove run_stream usage * Fix conversation_id propagation * Python: Add BaseAgent implementation for Claude Agent SDK (#3509) * Added ClaudeAgent implementation * Updated streaming logic * Small updates * Small update * Fixes * Small fix * Naming improvements * Updated imports * Addressed comments * Updated package versions * Update Claude agent connector layering * fix test and plugin * Store function middleware in invocation layer * Fix telemetry streaming and ag-ui tests * Remove legacy ag-ui tests folder * updates * Remove terminate flag from FunctionInvocationContext, use MiddlewareTermination instead - Remove terminate attribute from FunctionInvocationContext - Add result attribute to MiddlewareTermination to carry function results - FunctionMiddlewarePipeline.execute() now lets MiddlewareTermination propagate - _auto_invoke_function captures context.result in exception before re-raising - _try_execute_function_calls catches MiddlewareTermination and sets should_terminate - Fix handoff middleware to append to chat_client.function_middleware directly - Update tests to use raise MiddlewareTermination instead of context.terminate - Add middleware flow documentation in samples/concepts/tools/README.md - Fix ag-ui to use FunctionMiddlewarePipeline instead of removed create_function_middleware_pipeline * fix: remove references to removed terminate flag in purview tests, add type ignore * fix: move _test_utils.py from package to test folder * fix: call get_final_response() to trigger context provider notification in streaming test * fix: correct broken links in tools README * docs: clarify default middleware behavior in summary table * fix: ensure inner stream result hooks are called when using map()/from_awaitable() * Fix mypy type errors * Address PR review comments on observability.py - Remove TODO comment about unconsumed streams, add explanatory note instead - Remove redundant _close_span cleanup hook (already called in _finalize_stream) - Clarify behavior: cleanup hooks run after stream iteration, if stream is not consumed the span remains open until garbage collected * Remove gen_ai.client.operation.duration from span attributes Duration is a metrics-only attribute per OpenTelemetry semantic conventions. It should be recorded to the histogram but not set as a span attribute. * Remove duration from _get_response_attributes, pass directly to _capture_response Duration is a metrics-only attribute. It's now passed directly to _capture_response instead of being included in the attributes dict that gets set on the span. * Remove redundant _close_span cleanup hook in AgentTelemetryLayer _finalize_stream already calls _close_span() in its finally block, so adding it as a separate cleanup hook is redundant. * Use weakref.finalize to close span when stream is garbage collected If a user creates a streaming response but never consumes it, the cleanup hooks won't run. Now we register a weak reference finalizer that will close the span when the stream object is garbage collected, ensuring spans don't leak in this scenario. * Fix _get_finalizers_from_stream to use _result_hooks attribute Renamed function to _get_result_hooks_from_stream and fixed it to look for the _result_hooks attribute which is the correct name in ResponseStream class. * Add missing asyncio import in test_request_info_mixin.py * Fix leftover merge conflict marker in image_generation sample * Update integration tests * Fix integration tests: increase max_iterations from 1 to 2 Tests with tool_choice options require at least 2 iterations: 1. First iteration to get function call and execute the tool 2. Second iteration to get the final text response With max_iterations=1, streaming tests would return early with only the function call/result but no final text content. * Fix duplicate function call error in conversation-based APIs When using conversation_id (for Responses/Assistants APIs), the server already has the function call message from the previous response. We should only send the new function result message, not all messages including the function call which would cause a duplicate ID error. Fix: When conversation_id is set, only send the last message (the tool result) instead of all response.messages. * Add regression test for conversation_id propagation between tool iterations Port test from PR #3664 with updates for new streaming API pattern. Tests that conversation_id is properly updated in options dict during function invocation loop iterations. * Fix tool_choice=required to return after tool execution When tool_choice is 'required', the user's intent is to force exactly one tool call. After the tool executes, return immediately with the function call and result - don't continue to call the model again. This fixes integration tests that were failing with empty text responses because with tool_choice=required, the model would keep returning function calls instead of text. Also adds regression tests for: - conversation_id propagation between tool iterations (from PR #3664) - tool_choice=required returns after tool execution * Document tool_choice behavior in tools README - Add table explaining tool_choice values (auto, none, required) - Explain why tool_choice=required returns immediately after tool execution - Add code example showing the difference between required and auto - Update flow diagram to show the early return path for tool_choice=required * Fix tool_choice=None behavior - don't default to 'auto' Remove the hardcoded default of 'auto' for tool_choice in ChatAgent init. When tool_choice is not specified (None), it will now not be sent to the API, allowing the API's default behavior to be used. Users who want tool_choice='auto' can still explicitly set it either in default_options or at runtime. Fixes #3585 * Fix tool_choice=none should not remove tools In OpenAI Assistants client, tools were not being sent when tool_choice='none'. This was incorrect - tool_choice='none' means the model won't call tools, but tools should still be available in the request (they may be used later in the conversation). Fixes #3585 * Add test for tool_choice=none preserving tools Adds a regression test to ensure that when tool_choice='none' is set but tools are provided, the tools are still sent to the API. This verifies the fix for #3585. * Fix tool_choice=none should not remove tools in all clients Apply the same fix to OpenAI Responses client and Azure AI client: - OpenAI Responses: Remove else block that popped tool_choice/parallel_tool_calls - Azure AI: Remove tool_choice != 'none' check when adding tools When tool_choice='none', the model won't call tools, but tools should still be sent to the API so they're available for future turns. Also update README to clarify tool_choice=required supports multiple tools. Fixes #3585 * Keep tool_choice even when tools is None Move tool_choice processing outside of the 'if tools' block in OpenAI Responses client so tool_choice is sent to the API even when no tools are provided. * Update test to match new parallel_tool_calls behavior Changed test_prepare_options_removes_parallel_tool_calls_when_no_tools to test_prepare_options_preserves_parallel_tool_calls_when_no_tools to reflect that parallel_tool_calls is now preserved even when no tools are present, consistent with the tool_choice behavior. * Fix ChatMessage API and Role enum usage after rebase - Update ChatMessage instantiation to use keyword args (role=, text=, contents=) - Fix Role enum comparisons to use .value for string comparison - Add created_at to AgentResponse in error handling - Fix AgentResponse.from_updates -> from_agent_run_response_updates - Fix DurableAgentStateMessage.from_chat_message to convert Role enum to string - Add Role import where needed * Fix additional ChatMessage API and method name changes - Fix ChatMessage usage in workflow files (use text= instead of contents= for strings) - Fix AgentResponse.from_updates -> from_agent_run_response_updates in workflow files - Fix test files for ChatMessage and Role enum usage * Fix remaining ChatMessage API usage in test files * Fix more ChatMessage and Role API changes in source and test files - Fix ChatMessage in _magentic.py replan method - Fix Role enum comparison in test assertions - Fix remaining test files with old ChatMessage syntax * Fix ChatMessage and Role API changes across packages - Add Role import where missing - Fix ChatMessage signature: positional args to keyword args (role=, text=, contents=) - Fix Role enum comparisons: .role.value instead of .role string - Fix FinishReason enum usage in ag-ui event converters - Rename AgentResponse.from_updates to from_agent_run_response_updates in ag-ui Fixes API compatibility after Types API Review improvements merge * Fix ChatMessage and Role API changes in github_copilot tests * Fix ChatMessage and Role API changes in redis and github_copilot packages - Fix redis provider: Role enum comparison using .value - Fix redis tests: ChatMessage signature and Role comparisons - Fix github_copilot tests: ChatMessage signature and Role comparisons - Update docstring examples in redis chat message store * Fix ChatMessage and Role API changes in devui package - Fix executor: ChatMessage signature change - Fix conversations: Role enum to string conversion in two places - Fix tests: ChatMessage signatures and Role comparisons * Fix ChatMessage and Role API changes in a2a and lab packages - Fix a2a tests: Role comparisons and ChatMessage signatures - Fix lab tau2 source: Role enum comparison in flip_messages, log_messages, sliding_window - Fix lab tau2 tests: ChatMessage signatures and Role comparisons * Remove duplicate test files from ag-ui/tests (tests are in ag_ui_tests) * Fix ChatMessage and Role API changes across packages After rebasing on upstream/main which merged PR #3647 (Types API Review improvements), fix all packages to use the new API: - ChatMessage: Use keyword args (role=, text=, contents=) instead of positional args - Role: Compare using .value attribute since it's now an enum Packages fixed: - ag-ui: Fixed Role value extraction bugs in _message_adapters.py - anthropic: Fixed ChatMessage and Role comparisons in tests - azure-ai: Fixed Role comparison in _client.py - azure-ai-search: Fixed ChatMessage and Role in source/tests - bedrock: Fixed ChatMessage signatures in tests - chatkit: Fixed ChatMessage and Role in source/tests - copilotstudio: Fixed ChatMessage and Role in tests - declarative: Fixed ChatMessage in _executors_agents.py - mem0: Fixed ChatMessage and Role in source/tests - purview: Fixed ChatMessage in source/tests * Fix mypy errors for ChatMessage and Role API changes - durabletask: Use str() fallback in role value extraction - core: Fix ChatMessage in _orchestrator_helpers.py to use keyword args - core: Add type ignore for _conversation_state.py contents deserialization - ag-ui: Fix type ignore comments (call-overload instead of arg-type) - azure-ai-search: Fix get_role_value type hint to accept Any - lab: Move get_role_value to module level with Any type hint * Improve CI test timeout configuration - Increase job timeout from 10 to 15 minutes - Reduce per-test timeout to 60s (was 900s/300s) - Add --timeout_method thread for better timeout handling - Add --timeout-verbose to see which tests are slow - Reduce retries from 3 to 2 and delay from 10s to 5s This ensures individual test timeouts are shorter than the job timeout, providing better visibility when tests hang. With 60s timeout and 2 retries, worst case per test is ~180s. * Fix ChatMessage API usage in docstrings and source - Fix ChatMessage positional args in docstrings: _serialization.py, _threads.py, _middleware.py - Fix ChatMessage in tau2 runner.py - Fix role comparison in _orchestrator_helpers.py to use .value - Fix role comparison in _group_chat.py docstring example - Fix role assertions in test_durable_entities.py to use .value * Revert tool_choice/parallel_tool_calls changes - must be removed when no tools OpenAI API requires tool_choice and parallel_tool_calls to only be present when tools are specified. Restored the logic that removes these options when there are no tools. - Restored check in _chat_client.py to remove tool_choice and parallel_tool_calls when no tools present - Restored same logic in _responses_client.py - Reverted test to expect the correct behavior * fixed issue in tests * fix: resolve merge conflict markers in ag-ui tests * fix: restructure ag-ui tests and fix Role/FinishReason to use string types * fix: streaming function invocation and middleware termination - Refactor streaming function invocation to use get_final_response() on inner streams - Fix MiddlewareTermination to accept result parameter for passing results - Fix _AutoHandoffMiddleware to use MiddlewareTermination instead of context.terminate - Fix AgentMiddlewareLayer.run() to properly forward function/chat middleware - Remove duplicate middleware registration in AgentMiddlewareLayer.__init__ - Fix exception handling in _auto_invoke_function to properly capture termination - Fix mypy errors in core package - Update tests to use stream=True parameter for unified run API * fix all tests command * Refactor integration tests to use pytest fixtures - Merge testutils.py into conftest.py for azurefunctions integration tests - Merge dt_testutils.py into conftest.py for durabletask integration tests - Convert all integration tests to use fixtures instead of direct imports (fixes ModuleNotFoundError with --import-mode=importlib) - Add sample_helper fixture for azurefunctions tests - Add agent_client_factory and orchestration_helper fixtures for durabletask - Integration tests now skip with descriptive messages when services unavailable - Restructure devui tests into tests/devui/ with proper conftest.py - Add test organization guidelines to CODING_STANDARD.md - Remove __init__.py from test directories per pytest best practices * Fix pytest_collection_modifyitems to only skip integration tests The hook was skipping all tests in the test session, not just integration tests. Now it only skips items in the integration_tests directory. * Fix mem0 tests failing on Python 3.13 Use patch.object on the imported module instead of @patch with string path to ensure the mock takes effect regardless of import timing. * fix mem0 * another attempt for mem0 * fix for mem0 * fix mem0 * Increase worker initialization wait time in durabletask tests Increase from 2 to 8 seconds to allow time for: - Python startup and module imports - Azure OpenAI client creation - Agent registration with DTS worker - Worker connection to DTS This helps prevent test failures in CI where the first tests may run before the worker is fully ready to process requests. * Fix streaming test to use ResponseStream with finalizer The _consume_stream method now expects a ResponseStream that can provide a final AgentResponse via get_final_response(). Update the test to use ResponseStream with AgentResponse.from_updates as the finalizer. * Fix MockToolCallingAgent to use new ResponseStream API and update samples * small updates to run_stream to run * fix sub workflow * temp fix for az func test --------- Co-authored-by: Dmytro Struk <13853051+dmytrostruk@users.noreply.github.com>
This commit is contained in:
committed by
GitHub
Unverified
parent
d1205896a1
commit
3dc59c83b5
@@ -423,7 +423,7 @@ class AgentBasedGroupChatOrchestrator(BaseGroupChatOrchestrator):
|
||||
])
|
||||
)
|
||||
# Prepend instruction as system message
|
||||
current_conversation.append(ChatMessage("user", [instruction]))
|
||||
current_conversation.append(ChatMessage(role="user", text=instruction))
|
||||
|
||||
retry_attempts = self._retry_attempts
|
||||
while True:
|
||||
|
||||
@@ -141,9 +141,11 @@ class _AutoHandoffMiddleware(FunctionMiddleware):
|
||||
await next(context)
|
||||
return
|
||||
|
||||
from agent_framework._middleware import MiddlewareTermination
|
||||
|
||||
# Short-circuit execution and provide deterministic response payload for the tool call.
|
||||
context.result = {HANDOFF_FUNCTION_RESULT_KEY: self._handoff_functions[context.function.name]}
|
||||
context.terminate = True
|
||||
raise MiddlewareTermination(result=context.result)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -161,7 +163,7 @@ class HandoffAgentUserRequest:
|
||||
"""Create a HandoffAgentUserRequest from a simple text response."""
|
||||
messages: list[ChatMessage] = []
|
||||
if isinstance(response, str):
|
||||
messages.append(ChatMessage("user", [response]))
|
||||
messages.append(ChatMessage(role="user", text=response))
|
||||
elif isinstance(response, ChatMessage):
|
||||
messages.append(response)
|
||||
elif isinstance(response, list):
|
||||
@@ -169,7 +171,7 @@ class HandoffAgentUserRequest:
|
||||
if isinstance(item, ChatMessage):
|
||||
messages.append(item)
|
||||
elif isinstance(item, str):
|
||||
messages.append(ChatMessage("user", [item]))
|
||||
messages.append(ChatMessage(role="user", text=item))
|
||||
else:
|
||||
raise TypeError("List items must be either str or ChatMessage instances")
|
||||
else:
|
||||
@@ -428,7 +430,7 @@ class HandoffAgentExecutor(AgentExecutor):
|
||||
# or a termination condition is met.
|
||||
# This allows the agent to perform long-running tasks without returning control
|
||||
# to the coordinator or user prematurely.
|
||||
self._cache.extend([ChatMessage("user", [self._autonomous_mode_prompt])])
|
||||
self._cache.extend([ChatMessage(role="user", text=self._autonomous_mode_prompt)])
|
||||
self._autonomous_mode_turns += 1
|
||||
await self._run_agent_and_emit(ctx)
|
||||
else:
|
||||
@@ -975,12 +977,12 @@ class HandoffBuilder:
|
||||
workflow = HandoffBuilder(participants=[triage, refund, billing]).with_checkpointing(storage).build()
|
||||
|
||||
# Run workflow with a session ID for resumption
|
||||
async for event in workflow.run_stream("Help me", session_id="user_123"):
|
||||
async for event in workflow.run("Help me", session_id="user_123", stream=True):
|
||||
# Process events...
|
||||
pass
|
||||
|
||||
# Later, resume the same conversation
|
||||
async for event in workflow.run_stream("I need a refund", session_id="user_123"):
|
||||
async for event in workflow.run("I need a refund", session_id="user_123", stream=True):
|
||||
# Conversation continues from where it left off
|
||||
pass
|
||||
|
||||
@@ -1039,7 +1041,7 @@ class HandoffBuilder:
|
||||
- Request/response handling
|
||||
|
||||
Returns:
|
||||
A fully configured Workflow ready to execute via `.run()` or `.run_stream()`.
|
||||
A fully configured Workflow ready to execute via `.run()` with optional `stream=True` parameter.
|
||||
|
||||
Raises:
|
||||
ValueError: If participants or coordinator were not configured, or if
|
||||
|
||||
@@ -629,7 +629,7 @@ class StandardMagenticManager(MagenticManagerBase):
|
||||
facts=facts_msg.text,
|
||||
plan=plan_msg.text,
|
||||
)
|
||||
return ChatMessage("assistant", [combined], author_name=MAGENTIC_MANAGER_NAME)
|
||||
return ChatMessage(role="assistant", text=combined, author_name=MAGENTIC_MANAGER_NAME)
|
||||
|
||||
async def replan(self, magentic_context: MagenticContext) -> ChatMessage:
|
||||
"""Update facts and plan when stalling or looping has been detected."""
|
||||
@@ -640,19 +640,17 @@ class StandardMagenticManager(MagenticManagerBase):
|
||||
|
||||
# Update facts
|
||||
facts_update_user = ChatMessage(
|
||||
"user",
|
||||
[
|
||||
self.task_ledger_facts_update_prompt.format(
|
||||
task=magentic_context.task, old_facts=self.task_ledger.facts.text
|
||||
)
|
||||
],
|
||||
role="user",
|
||||
text=self.task_ledger_facts_update_prompt.format(
|
||||
task=magentic_context.task, old_facts=self.task_ledger.facts.text
|
||||
),
|
||||
)
|
||||
updated_facts = await self._complete([*magentic_context.chat_history, facts_update_user])
|
||||
|
||||
# Update plan
|
||||
plan_update_user = ChatMessage(
|
||||
"user",
|
||||
[self.task_ledger_plan_update_prompt.format(team=team_text)],
|
||||
role="user",
|
||||
text=self.task_ledger_plan_update_prompt.format(team=team_text),
|
||||
)
|
||||
updated_plan = await self._complete([
|
||||
*magentic_context.chat_history,
|
||||
@@ -674,7 +672,7 @@ class StandardMagenticManager(MagenticManagerBase):
|
||||
facts=updated_facts.text,
|
||||
plan=updated_plan.text,
|
||||
)
|
||||
return ChatMessage("assistant", [combined], author_name=MAGENTIC_MANAGER_NAME)
|
||||
return ChatMessage(role="assistant", text=combined, author_name=MAGENTIC_MANAGER_NAME)
|
||||
|
||||
async def create_progress_ledger(self, magentic_context: MagenticContext) -> MagenticProgressLedger:
|
||||
"""Use the model to produce a JSON progress ledger based on the conversation so far.
|
||||
@@ -694,7 +692,7 @@ class StandardMagenticManager(MagenticManagerBase):
|
||||
team=team_text,
|
||||
names=names_csv,
|
||||
)
|
||||
user_message = ChatMessage("user", [prompt])
|
||||
user_message = ChatMessage(role="user", text=prompt)
|
||||
|
||||
# Include full context to help the model decide current stage, with small retry loop
|
||||
attempts = 0
|
||||
@@ -721,7 +719,7 @@ class StandardMagenticManager(MagenticManagerBase):
|
||||
async def prepare_final_answer(self, magentic_context: MagenticContext) -> ChatMessage:
|
||||
"""Ask the model to produce the final answer addressed to the user."""
|
||||
prompt = self.final_answer_prompt.format(task=magentic_context.task)
|
||||
user_message = ChatMessage("user", [prompt])
|
||||
user_message = ChatMessage(role="user", text=prompt)
|
||||
response = await self._complete([*magentic_context.chat_history, user_message])
|
||||
# Ensure role is assistant
|
||||
return ChatMessage(
|
||||
@@ -811,11 +809,11 @@ class MagenticPlanReviewResponse:
|
||||
def revise(feedback: str | list[str] | ChatMessage | list[ChatMessage]) -> "MagenticPlanReviewResponse":
|
||||
"""Create a revision response with feedback."""
|
||||
if isinstance(feedback, str):
|
||||
feedback = [ChatMessage("user", [feedback])]
|
||||
feedback = [ChatMessage(role="user", text=feedback)]
|
||||
elif isinstance(feedback, ChatMessage):
|
||||
feedback = [feedback]
|
||||
elif isinstance(feedback, list):
|
||||
feedback = [ChatMessage("user", [item]) if isinstance(item, str) else item for item in feedback]
|
||||
feedback = [ChatMessage(role="user", text=item) if isinstance(item, str) else item for item in feedback]
|
||||
|
||||
return MagenticPlanReviewResponse(review=feedback)
|
||||
|
||||
@@ -1515,7 +1513,7 @@ class MagenticBuilder:
|
||||
)
|
||||
|
||||
# During execution, handle plan review
|
||||
async for event in workflow.run_stream("task"):
|
||||
async for event in workflow.run("task", stream=True):
|
||||
if isinstance(event, RequestInfoEvent):
|
||||
request = event.data
|
||||
if isinstance(request, MagenticHumanInterventionRequest):
|
||||
@@ -1563,11 +1561,11 @@ class MagenticBuilder:
|
||||
|
||||
# First run
|
||||
thread_id = "task-123"
|
||||
async for msg in workflow.run("task", thread_id=thread_id):
|
||||
async for msg in workflow.run("task", thread_id=thread_id, stream=True):
|
||||
print(msg.text)
|
||||
|
||||
# Resume from checkpoint
|
||||
async for msg in workflow.run("continue", thread_id=thread_id):
|
||||
async for msg in workflow.run("continue", thread_id=thread_id, stream=True):
|
||||
print(msg.text)
|
||||
|
||||
Notes:
|
||||
@@ -1812,7 +1810,7 @@ class MagenticBuilder:
|
||||
class MyManager(MagenticManagerBase):
|
||||
async def plan(self, context: MagenticContext) -> ChatMessage:
|
||||
# Custom planning logic
|
||||
return ChatMessage("assistant", ["..."])
|
||||
return ChatMessage(role="assistant", text="...")
|
||||
|
||||
|
||||
manager = MyManager()
|
||||
|
||||
@@ -34,7 +34,7 @@ class _FakeAgentExec(Executor):
|
||||
|
||||
@handler
|
||||
async def run(self, request: AgentExecutorRequest, ctx: WorkflowContext[AgentExecutorResponse]) -> None:
|
||||
response = AgentResponse(messages=ChatMessage("assistant", text=self._reply_text))
|
||||
response = AgentResponse(messages=ChatMessage(role="assistant", text=self._reply_text))
|
||||
full_conversation = list(request.messages) + list(response.messages)
|
||||
await ctx.send_message(AgentExecutorResponse(self.id, response, full_conversation=full_conversation))
|
||||
|
||||
@@ -110,7 +110,7 @@ async def test_concurrent_default_aggregator_emits_single_user_and_assistants()
|
||||
|
||||
completed = False
|
||||
output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("prompt: hello world"):
|
||||
async for ev in wf.run("prompt: hello world", stream=True):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
@@ -148,7 +148,7 @@ async def test_concurrent_custom_aggregator_callback_is_used() -> None:
|
||||
|
||||
completed = False
|
||||
output: str | None = None
|
||||
async for ev in wf.run_stream("prompt: custom"):
|
||||
async for ev in wf.run("prompt: custom", stream=True):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
@@ -179,7 +179,7 @@ async def test_concurrent_custom_aggregator_sync_callback_is_used() -> None:
|
||||
|
||||
completed = False
|
||||
output: str | None = None
|
||||
async for ev in wf.run_stream("prompt: custom sync"):
|
||||
async for ev in wf.run("prompt: custom sync", stream=True):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
@@ -227,7 +227,7 @@ async def test_concurrent_with_aggregator_executor_instance() -> None:
|
||||
|
||||
completed = False
|
||||
output: str | None = None
|
||||
async for ev in wf.run_stream("prompt: instance test"):
|
||||
async for ev in wf.run("prompt: instance test", stream=True):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
@@ -265,7 +265,7 @@ async def test_concurrent_with_aggregator_executor_factory() -> None:
|
||||
|
||||
completed = False
|
||||
output: str | None = None
|
||||
async for ev in wf.run_stream("prompt: factory test"):
|
||||
async for ev in wf.run("prompt: factory test", stream=True):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
@@ -301,7 +301,7 @@ async def test_concurrent_with_aggregator_executor_factory_with_default_id() ->
|
||||
|
||||
completed = False
|
||||
output: str | None = None
|
||||
async for ev in wf.run_stream("prompt: factory test"):
|
||||
async for ev in wf.run("prompt: factory test", stream=True):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
@@ -351,7 +351,7 @@ async def test_concurrent_checkpoint_resume_round_trip() -> None:
|
||||
wf = ConcurrentBuilder().participants(list(participants)).with_checkpointing(storage).build()
|
||||
|
||||
baseline_output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("checkpoint concurrent"):
|
||||
async for ev in wf.run("checkpoint concurrent", stream=True):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
baseline_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
@@ -375,7 +375,7 @@ async def test_concurrent_checkpoint_resume_round_trip() -> None:
|
||||
wf_resume = ConcurrentBuilder().participants(list(resumed_participants)).with_checkpointing(storage).build()
|
||||
|
||||
resumed_output: list[ChatMessage] | None = None
|
||||
async for ev in wf_resume.run_stream(checkpoint_id=resume_checkpoint.checkpoint_id):
|
||||
async for ev in wf_resume.run(checkpoint_id=resume_checkpoint.checkpoint_id, stream=True):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
resumed_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state in (
|
||||
@@ -397,7 +397,7 @@ async def test_concurrent_checkpoint_runtime_only() -> None:
|
||||
wf = ConcurrentBuilder().participants(agents).build()
|
||||
|
||||
baseline_output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("runtime checkpoint test", checkpoint_storage=storage):
|
||||
async for ev in wf.run("runtime checkpoint test", checkpoint_storage=storage, stream=True):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
baseline_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
@@ -418,7 +418,9 @@ async def test_concurrent_checkpoint_runtime_only() -> None:
|
||||
wf_resume = ConcurrentBuilder().participants(resumed_agents).build()
|
||||
|
||||
resumed_output: list[ChatMessage] | None = None
|
||||
async for ev in wf_resume.run_stream(checkpoint_id=resume_checkpoint.checkpoint_id, checkpoint_storage=storage):
|
||||
async for ev in wf_resume.run(
|
||||
checkpoint_id=resume_checkpoint.checkpoint_id, checkpoint_storage=storage, stream=True
|
||||
):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
resumed_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state in (
|
||||
@@ -445,7 +447,7 @@ async def test_concurrent_checkpoint_runtime_overrides_buildtime() -> None:
|
||||
wf = ConcurrentBuilder().participants(agents).with_checkpointing(buildtime_storage).build()
|
||||
|
||||
baseline_output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("override test", checkpoint_storage=runtime_storage):
|
||||
async for ev in wf.run("override test", checkpoint_storage=runtime_storage, stream=True):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
baseline_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
@@ -527,7 +529,7 @@ async def test_concurrent_with_register_participants() -> None:
|
||||
|
||||
completed = False
|
||||
output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("test prompt"):
|
||||
async for ev in wf.run("test prompt", stream=True):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from collections.abc import AsyncIterable, Callable, Sequence
|
||||
from collections.abc import AsyncIterable, Awaitable, Callable, Sequence
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
@@ -38,29 +38,26 @@ class StubAgent(BaseAgent):
|
||||
super().__init__(name=agent_name, description=f"Stub agent {agent_name}", **kwargs)
|
||||
self._reply_text = reply_text
|
||||
|
||||
async def run( # type: ignore[override]
|
||||
def run( # type: ignore[override]
|
||||
self,
|
||||
messages: str | ChatMessage | Sequence[str | ChatMessage] | None = None,
|
||||
*,
|
||||
stream: bool = False,
|
||||
thread: AgentThread | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AgentResponse:
|
||||
response = ChatMessage("assistant", [self._reply_text], author_name=self.name)
|
||||
) -> Awaitable[AgentResponse] | AsyncIterable[AgentResponseUpdate]:
|
||||
if stream:
|
||||
return self._run_stream_impl()
|
||||
return self._run_impl()
|
||||
|
||||
async def _run_impl(self) -> AgentResponse:
|
||||
response = ChatMessage(role="assistant", text=self._reply_text, author_name=self.name)
|
||||
return AgentResponse(messages=[response])
|
||||
|
||||
def run_stream( # type: ignore[override]
|
||||
self,
|
||||
messages: str | ChatMessage | Sequence[str | ChatMessage] | None = None,
|
||||
*,
|
||||
thread: AgentThread | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterable[AgentResponseUpdate]:
|
||||
async def _stream() -> AsyncIterable[AgentResponseUpdate]:
|
||||
yield AgentResponseUpdate(
|
||||
contents=[Content.from_text(text=self._reply_text)], role="assistant", author_name=self.name
|
||||
)
|
||||
|
||||
return _stream()
|
||||
async def _run_stream_impl(self) -> AsyncIterable[AgentResponseUpdate]:
|
||||
yield AgentResponseUpdate(
|
||||
contents=[Content.from_text(text=self._reply_text)], role="assistant", author_name=self.name
|
||||
)
|
||||
|
||||
|
||||
class MockChatClient:
|
||||
@@ -68,10 +65,9 @@ class MockChatClient:
|
||||
|
||||
additional_properties: dict[str, Any]
|
||||
|
||||
async def get_response(self, messages: Any, **kwargs: Any) -> ChatResponse:
|
||||
raise NotImplementedError
|
||||
|
||||
def get_streaming_response(self, messages: Any, **kwargs: Any) -> AsyncIterable[ChatResponseUpdate]:
|
||||
async def get_response(
|
||||
self, messages: Any, stream: bool = False, **kwargs: Any
|
||||
) -> ChatResponse | AsyncIterable[ChatResponseUpdate]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@@ -126,48 +122,6 @@ class StubManagerAgent(ChatAgent):
|
||||
value=payload,
|
||||
)
|
||||
|
||||
def run_stream(
|
||||
self,
|
||||
messages: str | ChatMessage | Sequence[str | ChatMessage] | None = None,
|
||||
*,
|
||||
thread: AgentThread | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterable[AgentResponseUpdate]:
|
||||
if self._call_count == 0:
|
||||
self._call_count += 1
|
||||
|
||||
async def _stream_initial() -> AsyncIterable[AgentResponseUpdate]:
|
||||
yield AgentResponseUpdate(
|
||||
contents=[
|
||||
Content.from_text(
|
||||
text=(
|
||||
'{"terminate": false, "reason": "Selecting agent", '
|
||||
'"next_speaker": "agent", "final_message": null}'
|
||||
)
|
||||
)
|
||||
],
|
||||
role="assistant",
|
||||
author_name=self.name,
|
||||
)
|
||||
|
||||
return _stream_initial()
|
||||
|
||||
async def _stream_final() -> AsyncIterable[AgentResponseUpdate]:
|
||||
yield AgentResponseUpdate(
|
||||
contents=[
|
||||
Content.from_text(
|
||||
text=(
|
||||
'{"terminate": true, "reason": "Task complete", '
|
||||
'"next_speaker": null, "final_message": "agent manager final"}'
|
||||
)
|
||||
)
|
||||
],
|
||||
role="assistant",
|
||||
author_name=self.name,
|
||||
)
|
||||
|
||||
return _stream_final()
|
||||
|
||||
|
||||
def make_sequence_selector() -> Callable[[GroupChatState], str]:
|
||||
state_counter = {"value": 0}
|
||||
@@ -192,7 +146,7 @@ class StubMagenticManager(MagenticManagerBase):
|
||||
self._round = 0
|
||||
|
||||
async def plan(self, magentic_context: MagenticContext) -> ChatMessage:
|
||||
return ChatMessage("assistant", ["plan"], author_name="magentic_manager")
|
||||
return ChatMessage(role="assistant", text="plan", author_name="magentic_manager")
|
||||
|
||||
async def replan(self, magentic_context: MagenticContext) -> ChatMessage:
|
||||
return await self.plan(magentic_context)
|
||||
@@ -218,7 +172,7 @@ class StubMagenticManager(MagenticManagerBase):
|
||||
)
|
||||
|
||||
async def prepare_final_answer(self, magentic_context: MagenticContext) -> ChatMessage:
|
||||
return ChatMessage("assistant", ["final"], author_name="magentic_manager")
|
||||
return ChatMessage(role="assistant", text="final", author_name="magentic_manager")
|
||||
|
||||
|
||||
async def test_group_chat_builder_basic_flow() -> None:
|
||||
@@ -235,7 +189,7 @@ async def test_group_chat_builder_basic_flow() -> None:
|
||||
)
|
||||
|
||||
outputs: list[list[ChatMessage]] = []
|
||||
async for event in workflow.run_stream("coordinate task"):
|
||||
async for event in workflow.run("coordinate task", stream=True):
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
data = event.data
|
||||
if isinstance(data, list):
|
||||
@@ -263,8 +217,8 @@ async def test_group_chat_as_agent_accepts_conversation() -> None:
|
||||
|
||||
agent = workflow.as_agent(name="group-chat-agent")
|
||||
conversation = [
|
||||
ChatMessage("user", ["kickoff"], author_name="user"),
|
||||
ChatMessage("assistant", ["noted"], author_name="alpha"),
|
||||
ChatMessage(role="user", text="kickoff", author_name="user"),
|
||||
ChatMessage(role="assistant", text="noted", author_name="alpha"),
|
||||
]
|
||||
response = await agent.run(conversation)
|
||||
|
||||
@@ -347,17 +301,20 @@ class TestGroupChatBuilder:
|
||||
def __init__(self) -> None:
|
||||
super().__init__(name="", description="test")
|
||||
|
||||
async def run(self, messages: Any = None, *, thread: Any = None, **kwargs: Any) -> AgentResponse:
|
||||
def run(
|
||||
self, messages: Any = None, *, stream: bool = False, thread: Any = None, **kwargs: Any
|
||||
) -> AgentResponse | AsyncIterable[AgentResponseUpdate]:
|
||||
if stream:
|
||||
|
||||
async def _stream() -> AsyncIterable[AgentResponseUpdate]:
|
||||
yield AgentResponseUpdate(contents=[])
|
||||
|
||||
return _stream()
|
||||
return self._run_impl()
|
||||
|
||||
async def _run_impl(self) -> AgentResponse:
|
||||
return AgentResponse(messages=[])
|
||||
|
||||
def run_stream(
|
||||
self, messages: Any = None, *, thread: Any = None, **kwargs: Any
|
||||
) -> AsyncIterable[AgentResponseUpdate]:
|
||||
async def _stream() -> AsyncIterable[AgentResponseUpdate]:
|
||||
yield AgentResponseUpdate(contents=[])
|
||||
|
||||
return _stream()
|
||||
|
||||
agent = AgentWithoutName()
|
||||
|
||||
def selector(state: GroupChatState) -> str:
|
||||
@@ -404,7 +361,7 @@ class TestGroupChatWorkflow:
|
||||
)
|
||||
|
||||
outputs: list[list[ChatMessage]] = []
|
||||
async for event in workflow.run_stream("test task"):
|
||||
async for event in workflow.run("test task", stream=True):
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
data = event.data
|
||||
if isinstance(data, list):
|
||||
@@ -439,7 +396,7 @@ class TestGroupChatWorkflow:
|
||||
)
|
||||
|
||||
outputs: list[list[ChatMessage]] = []
|
||||
async for event in workflow.run_stream("test task"):
|
||||
async for event in workflow.run("test task", stream=True):
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
data = event.data
|
||||
if isinstance(data, list):
|
||||
@@ -467,7 +424,7 @@ class TestGroupChatWorkflow:
|
||||
)
|
||||
|
||||
outputs: list[list[ChatMessage]] = []
|
||||
async for event in workflow.run_stream("test task"):
|
||||
async for event in workflow.run("test task", stream=True):
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
data = event.data
|
||||
if isinstance(data, list):
|
||||
@@ -489,7 +446,7 @@ class TestGroupChatWorkflow:
|
||||
workflow = GroupChatBuilder().with_orchestrator(selection_func=selector).participants([agent]).build()
|
||||
|
||||
with pytest.raises(RuntimeError, match="Selection function returned unknown participant 'unknown_agent'"):
|
||||
async for _ in workflow.run_stream("test task"):
|
||||
async for _ in workflow.run("test task", stream=True):
|
||||
pass
|
||||
|
||||
|
||||
@@ -515,7 +472,7 @@ class TestCheckpointing:
|
||||
)
|
||||
|
||||
outputs: list[list[ChatMessage]] = []
|
||||
async for event in workflow.run_stream("test task"):
|
||||
async for event in workflow.run("test task", stream=True):
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
data = event.data
|
||||
if isinstance(data, list):
|
||||
@@ -544,7 +501,7 @@ class TestConversationHandling:
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="At least one ChatMessage is required to start the group chat workflow."):
|
||||
async for _ in workflow.run_stream([]):
|
||||
async for _ in workflow.run([], stream=True):
|
||||
pass
|
||||
|
||||
async def test_handle_string_input(self) -> None:
|
||||
@@ -568,7 +525,7 @@ class TestConversationHandling:
|
||||
)
|
||||
|
||||
outputs: list[list[ChatMessage]] = []
|
||||
async for event in workflow.run_stream("test string"):
|
||||
async for event in workflow.run("test string", stream=True):
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
data = event.data
|
||||
if isinstance(data, list):
|
||||
@@ -578,7 +535,7 @@ class TestConversationHandling:
|
||||
|
||||
async def test_handle_chat_message_input(self) -> None:
|
||||
"""Test handling ChatMessage input directly."""
|
||||
task_message = ChatMessage("user", ["test message"])
|
||||
task_message = ChatMessage(role="user", text="test message")
|
||||
|
||||
def selector(state: GroupChatState) -> str:
|
||||
# Verify the task message was preserved in conversation
|
||||
@@ -597,7 +554,7 @@ class TestConversationHandling:
|
||||
)
|
||||
|
||||
outputs: list[list[ChatMessage]] = []
|
||||
async for event in workflow.run_stream(task_message):
|
||||
async for event in workflow.run(task_message, stream=True):
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
data = event.data
|
||||
if isinstance(data, list):
|
||||
@@ -608,8 +565,8 @@ class TestConversationHandling:
|
||||
async def test_handle_conversation_list_input(self) -> None:
|
||||
"""Test handling conversation list preserves context."""
|
||||
conversation = [
|
||||
ChatMessage("system", ["system message"]),
|
||||
ChatMessage("user", ["user message"]),
|
||||
ChatMessage(role="system", text="system message"),
|
||||
ChatMessage(role="user", text="user message"),
|
||||
]
|
||||
|
||||
def selector(state: GroupChatState) -> str:
|
||||
@@ -629,7 +586,7 @@ class TestConversationHandling:
|
||||
)
|
||||
|
||||
outputs: list[list[ChatMessage]] = []
|
||||
async for event in workflow.run_stream(conversation):
|
||||
async for event in workflow.run(conversation, stream=True):
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
data = event.data
|
||||
if isinstance(data, list):
|
||||
@@ -661,7 +618,7 @@ class TestRoundLimitEnforcement:
|
||||
)
|
||||
|
||||
outputs: list[list[ChatMessage]] = []
|
||||
async for event in workflow.run_stream("test"):
|
||||
async for event in workflow.run("test", stream=True):
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
data = event.data
|
||||
if isinstance(data, list):
|
||||
@@ -696,7 +653,7 @@ class TestRoundLimitEnforcement:
|
||||
)
|
||||
|
||||
outputs: list[list[ChatMessage]] = []
|
||||
async for event in workflow.run_stream("test"):
|
||||
async for event in workflow.run("test", stream=True):
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
data = event.data
|
||||
if isinstance(data, list):
|
||||
@@ -728,7 +685,7 @@ async def test_group_chat_checkpoint_runtime_only() -> None:
|
||||
)
|
||||
|
||||
baseline_output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("runtime checkpoint test", checkpoint_storage=storage):
|
||||
async for ev in wf.run("runtime checkpoint test", checkpoint_storage=storage, stream=True):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
baseline_output = cast(list[ChatMessage], ev.data) if isinstance(ev.data, list) else None # type: ignore
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state in (
|
||||
@@ -766,7 +723,7 @@ async def test_group_chat_checkpoint_runtime_overrides_buildtime() -> None:
|
||||
.build()
|
||||
)
|
||||
baseline_output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("override test", checkpoint_storage=runtime_storage):
|
||||
async for ev in wf.run("override test", checkpoint_storage=runtime_storage, stream=True):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
baseline_output = cast(list[ChatMessage], ev.data) if isinstance(ev.data, list) else None # type: ignore
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state in (
|
||||
@@ -814,7 +771,7 @@ async def test_group_chat_with_request_info_filtering():
|
||||
|
||||
# Run until we get a request info event (should be before beta, not alpha)
|
||||
request_events: list[RequestInfoEvent] = []
|
||||
async for event in workflow.run_stream("test task"):
|
||||
async for event in workflow.run("test task", stream=True):
|
||||
if isinstance(event, RequestInfoEvent) and isinstance(event.data, AgentExecutorResponse):
|
||||
request_events.append(event)
|
||||
# Don't break - let stream complete naturally when paused
|
||||
@@ -866,7 +823,7 @@ async def test_group_chat_with_request_info_no_filter_pauses_all():
|
||||
|
||||
# Run until we get a request info event
|
||||
request_events: list[RequestInfoEvent] = []
|
||||
async for event in workflow.run_stream("test task"):
|
||||
async for event in workflow.run("test task", stream=True):
|
||||
if isinstance(event, RequestInfoEvent) and isinstance(event.data, AgentExecutorResponse):
|
||||
request_events.append(event)
|
||||
break
|
||||
@@ -970,7 +927,7 @@ async def test_group_chat_with_participant_factories():
|
||||
assert call_count == 2
|
||||
|
||||
outputs: list[WorkflowOutputEvent] = []
|
||||
async for event in workflow.run_stream("coordinate task"):
|
||||
async for event in workflow.run("coordinate task", stream=True):
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
outputs.append(event)
|
||||
|
||||
@@ -1035,7 +992,7 @@ async def test_group_chat_participant_factories_with_checkpointing():
|
||||
)
|
||||
|
||||
outputs: list[WorkflowOutputEvent] = []
|
||||
async for event in workflow.run_stream("checkpoint test"):
|
||||
async for event in workflow.run("checkpoint test", stream=True):
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
outputs.append(event)
|
||||
|
||||
@@ -1163,7 +1120,7 @@ async def test_group_chat_with_orchestrator_factory_returning_chat_agent():
|
||||
assert factory_call_count == 1
|
||||
|
||||
outputs: list[WorkflowOutputEvent] = []
|
||||
async for event in workflow.run_stream("coordinate task"):
|
||||
async for event in workflow.run("coordinate task", stream=True):
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
outputs.append(event)
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from collections.abc import AsyncIterable
|
||||
from collections.abc import AsyncIterable, Awaitable, Mapping, Sequence
|
||||
from typing import Any, cast
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
@@ -12,25 +12,26 @@ from agent_framework import (
|
||||
ChatResponseUpdate,
|
||||
Content,
|
||||
RequestInfoEvent,
|
||||
ResponseStream,
|
||||
WorkflowEvent,
|
||||
WorkflowOutputEvent,
|
||||
resolve_agent_id,
|
||||
use_function_invocation,
|
||||
)
|
||||
from agent_framework._clients import BaseChatClient
|
||||
from agent_framework._middleware import ChatMiddlewareLayer
|
||||
from agent_framework._tools import FunctionInvocationLayer
|
||||
from agent_framework.orchestrations import HandoffAgentUserRequest, HandoffBuilder
|
||||
|
||||
|
||||
@use_function_invocation
|
||||
class MockChatClient:
|
||||
class MockChatClient(ChatMiddlewareLayer[Any], FunctionInvocationLayer[Any], BaseChatClient[Any]):
|
||||
"""Mock chat client for testing handoff workflows."""
|
||||
|
||||
additional_properties: dict[str, Any]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
name: str,
|
||||
*,
|
||||
name: str = "",
|
||||
handoff_to: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize the mock chat client.
|
||||
|
||||
@@ -39,24 +40,45 @@ class MockChatClient:
|
||||
handoff_to: The name of the agent to hand off to, or None for no handoff.
|
||||
This is hardcoded for testing purposes so that the agent always attempts to hand off.
|
||||
"""
|
||||
ChatMiddlewareLayer.__init__(self)
|
||||
FunctionInvocationLayer.__init__(self)
|
||||
BaseChatClient.__init__(self)
|
||||
self._name = name
|
||||
self._handoff_to = handoff_to
|
||||
self._call_index = 0
|
||||
|
||||
async def get_response(self, messages: Any, **kwargs: Any) -> ChatResponse:
|
||||
contents = _build_reply_contents(self._name, self._handoff_to, self._next_call_id())
|
||||
reply = ChatMessage(
|
||||
role="assistant",
|
||||
contents=contents,
|
||||
)
|
||||
return ChatResponse(messages=reply, response_id="mock_response")
|
||||
def _inner_get_response(
|
||||
self,
|
||||
*,
|
||||
messages: Sequence[ChatMessage],
|
||||
stream: bool,
|
||||
options: Mapping[str, Any],
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse] | ResponseStream[ChatResponseUpdate, ChatResponse]:
|
||||
if stream:
|
||||
return self._build_streaming_response(options=dict(options))
|
||||
|
||||
def get_streaming_response(self, messages: Any, **kwargs: Any) -> AsyncIterable[ChatResponseUpdate]:
|
||||
async def _get() -> ChatResponse:
|
||||
contents = _build_reply_contents(self._name, self._handoff_to, self._next_call_id())
|
||||
reply = ChatMessage(
|
||||
role="assistant",
|
||||
contents=contents,
|
||||
)
|
||||
return ChatResponse(messages=reply, response_id="mock_response")
|
||||
|
||||
return _get()
|
||||
|
||||
def _build_streaming_response(self, *, options: dict[str, Any]) -> ResponseStream[ChatResponseUpdate, ChatResponse]:
|
||||
async def _stream() -> AsyncIterable[ChatResponseUpdate]:
|
||||
contents = _build_reply_contents(self._name, self._handoff_to, self._next_call_id())
|
||||
yield ChatResponseUpdate(contents=contents, role="assistant")
|
||||
yield ChatResponseUpdate(contents=contents, role="assistant", finish_reason="stop")
|
||||
|
||||
return _stream()
|
||||
def _finalize(updates: Sequence[ChatResponseUpdate]) -> ChatResponse:
|
||||
response_format = options.get("response_format")
|
||||
output_format_type = response_format if isinstance(response_format, type) else None
|
||||
return ChatResponse.from_updates(updates, output_format_type=output_format_type)
|
||||
|
||||
return ResponseStream(_stream(), finalizer=_finalize)
|
||||
|
||||
def _next_call_id(self) -> str | None:
|
||||
if not self._handoff_to:
|
||||
@@ -99,7 +121,7 @@ class MockHandoffAgent(ChatAgent):
|
||||
handoff_to: The name of the agent to hand off to, or None for no handoff.
|
||||
This is hardcoded for testing purposes so that the agent always attempts to hand off.
|
||||
"""
|
||||
super().__init__(chat_client=MockChatClient(name, handoff_to=handoff_to), name=name, id=name)
|
||||
super().__init__(chat_client=MockChatClient(name=name, handoff_to=handoff_to), name=name, id=name)
|
||||
|
||||
|
||||
async def _drain(stream: AsyncIterable[WorkflowEvent]) -> list[WorkflowEvent]:
|
||||
@@ -127,7 +149,7 @@ async def test_handoff():
|
||||
# Start conversation - triage hands off to specialist then escalation
|
||||
# escalation won't trigger a handoff, so the response from it will become
|
||||
# a request for user input because autonomous mode is not enabled by default.
|
||||
events = await _drain(workflow.run_stream("Need technical support"))
|
||||
events = await _drain(workflow.run("Need technical support", stream=True))
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
|
||||
assert requests
|
||||
@@ -161,7 +183,7 @@ async def test_autonomous_mode_yields_output_without_user_request():
|
||||
.build()
|
||||
)
|
||||
|
||||
events = await _drain(workflow.run_stream("Package arrived broken"))
|
||||
events = await _drain(workflow.run("Package arrived broken", stream=True))
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert not requests, "Autonomous mode should not request additional user input"
|
||||
|
||||
@@ -187,7 +209,7 @@ async def test_autonomous_mode_resumes_user_input_on_turn_limit():
|
||||
.build()
|
||||
)
|
||||
|
||||
events = await _drain(workflow.run_stream("Start"))
|
||||
events = await _drain(workflow.run("Start", stream=True))
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert requests and len(requests) == 1, "Turn limit should force a user input request"
|
||||
assert requests[0].source_executor_id == worker.name
|
||||
@@ -230,12 +252,14 @@ async def test_handoff_async_termination_condition() -> None:
|
||||
.build()
|
||||
)
|
||||
|
||||
events = await _drain(workflow.run_stream("First user message"))
|
||||
events = await _drain(workflow.run("First user message", stream=True))
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert requests
|
||||
|
||||
events = await _drain(
|
||||
workflow.send_responses_streaming({requests[-1].request_id: [ChatMessage("user", ["Second user message"])]})
|
||||
workflow.send_responses_streaming({
|
||||
requests[-1].request_id: [ChatMessage(role="user", text="Second user message")]
|
||||
})
|
||||
)
|
||||
outputs = [ev for ev in events if isinstance(ev, WorkflowOutputEvent)]
|
||||
assert len(outputs) == 1
|
||||
@@ -257,7 +281,7 @@ async def test_tool_choice_preserved_from_agent_config():
|
||||
if options:
|
||||
recorded_tool_choices.append(options.get("tool_choice"))
|
||||
return ChatResponse(
|
||||
messages=[ChatMessage("assistant", ["Response"])],
|
||||
messages=[ChatMessage(role="assistant", text="Response")],
|
||||
response_id="test_response",
|
||||
)
|
||||
|
||||
@@ -480,13 +504,13 @@ async def test_handoff_with_participant_factories():
|
||||
# Factories should be called during build
|
||||
assert call_count == 2
|
||||
|
||||
events = await _drain(workflow.run_stream("Need help"))
|
||||
events = await _drain(workflow.run("Need help", stream=True))
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert requests
|
||||
|
||||
# Follow-up message
|
||||
events = await _drain(
|
||||
workflow.send_responses_streaming({requests[-1].request_id: [ChatMessage("user", ["More details"])]})
|
||||
workflow.send_responses_streaming({requests[-1].request_id: [ChatMessage(role="user", text="More details")]})
|
||||
)
|
||||
outputs = [ev for ev in events if isinstance(ev, WorkflowOutputEvent)]
|
||||
assert outputs
|
||||
@@ -551,7 +575,7 @@ async def test_handoff_with_participant_factories_and_add_handoff():
|
||||
)
|
||||
|
||||
# Start conversation - triage hands off to specialist_a
|
||||
events = await _drain(workflow.run_stream("Initial request"))
|
||||
events = await _drain(workflow.run("Initial request", stream=True))
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert requests
|
||||
|
||||
@@ -560,7 +584,7 @@ async def test_handoff_with_participant_factories_and_add_handoff():
|
||||
|
||||
# Second user message - specialist_a hands off to specialist_b
|
||||
events = await _drain(
|
||||
workflow.send_responses_streaming({requests[-1].request_id: [ChatMessage("user", ["Need escalation"])]})
|
||||
workflow.send_responses_streaming({requests[-1].request_id: [ChatMessage(role="user", text="Need escalation")]})
|
||||
)
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert requests
|
||||
@@ -590,12 +614,12 @@ async def test_handoff_participant_factories_with_checkpointing():
|
||||
)
|
||||
|
||||
# Run workflow and capture output
|
||||
events = await _drain(workflow.run_stream("checkpoint test"))
|
||||
events = await _drain(workflow.run("checkpoint test", stream=True))
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert requests
|
||||
|
||||
events = await _drain(
|
||||
workflow.send_responses_streaming({requests[-1].request_id: [ChatMessage("user", ["follow up"])]})
|
||||
workflow.send_responses_streaming({requests[-1].request_id: [ChatMessage(role="user", text="follow up")]})
|
||||
)
|
||||
outputs = [ev for ev in events if isinstance(ev, WorkflowOutputEvent)]
|
||||
assert outputs, "Should have workflow output after termination condition is met"
|
||||
@@ -668,7 +692,7 @@ async def test_handoff_participant_factories_autonomous_mode():
|
||||
.build()
|
||||
)
|
||||
|
||||
events = await _drain(workflow.run_stream("Issue"))
|
||||
events = await _drain(workflow.run("Issue", stream=True))
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert requests and len(requests) == 1
|
||||
assert requests[0].source_executor_id == "specialist"
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import sys
|
||||
from collections.abc import AsyncIterable, Sequence
|
||||
from collections.abc import AsyncIterable, Awaitable, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, ClassVar, cast
|
||||
|
||||
@@ -152,29 +152,27 @@ class StubAgent(BaseAgent):
|
||||
super().__init__(name=agent_name, description=f"Stub agent {agent_name}", **kwargs)
|
||||
self._reply_text = reply_text
|
||||
|
||||
async def run( # type: ignore[override]
|
||||
def run( # type: ignore[override]
|
||||
self,
|
||||
messages: str | ChatMessage | Sequence[str | ChatMessage] | None = None,
|
||||
*,
|
||||
stream: bool = False,
|
||||
thread: AgentThread | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AgentResponse:
|
||||
response = ChatMessage("assistant", [self._reply_text], author_name=self.name)
|
||||
return AgentResponse(messages=[response])
|
||||
) -> Awaitable[AgentResponse] | AsyncIterable[AgentResponseUpdate]:
|
||||
if stream:
|
||||
return self._run_stream()
|
||||
|
||||
def run_stream( # type: ignore[override]
|
||||
self,
|
||||
messages: str | ChatMessage | Sequence[str | ChatMessage] | None = None,
|
||||
*,
|
||||
thread: AgentThread | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterable[AgentResponseUpdate]:
|
||||
async def _stream() -> AsyncIterable[AgentResponseUpdate]:
|
||||
yield AgentResponseUpdate(
|
||||
contents=[Content.from_text(text=self._reply_text)], role="assistant", author_name=self.name
|
||||
)
|
||||
async def _run() -> AgentResponse:
|
||||
response = ChatMessage("assistant", [self._reply_text], author_name=self.name)
|
||||
return AgentResponse(messages=[response])
|
||||
|
||||
return _stream()
|
||||
return _run()
|
||||
|
||||
async def _run_stream(self) -> AsyncIterable[AgentResponseUpdate]:
|
||||
yield AgentResponseUpdate(
|
||||
contents=[Content.from_text(text=self._reply_text)], role="assistant", author_name=self.name
|
||||
)
|
||||
|
||||
|
||||
class DummyExec(Executor):
|
||||
@@ -198,7 +196,7 @@ async def test_magentic_builder_returns_workflow_and_runs() -> None:
|
||||
|
||||
outputs: list[ChatMessage] = []
|
||||
orchestrator_event_count = 0
|
||||
async for event in workflow.run_stream("compose summary"):
|
||||
async for event in workflow.run("compose summary", stream=True):
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
msg = event.data
|
||||
if isinstance(msg, list):
|
||||
@@ -249,7 +247,7 @@ async def test_magentic_workflow_plan_review_approval_to_completion():
|
||||
wf = MagenticBuilder().participants([DummyExec("agentA")]).with_manager(manager=manager).with_plan_review().build()
|
||||
|
||||
req_event: RequestInfoEvent | None = None
|
||||
async for ev in wf.run_stream("do work"):
|
||||
async for ev in wf.run("do work", stream=True):
|
||||
if isinstance(ev, RequestInfoEvent) and ev.request_type is MagenticPlanReviewRequest:
|
||||
req_event = ev
|
||||
assert req_event is not None
|
||||
@@ -294,7 +292,7 @@ async def test_magentic_plan_review_with_revise():
|
||||
|
||||
# Wait for the initial plan review request
|
||||
req_event: RequestInfoEvent | None = None
|
||||
async for ev in wf.run_stream("do work"):
|
||||
async for ev in wf.run("do work", stream=True):
|
||||
if isinstance(ev, RequestInfoEvent) and ev.request_type is MagenticPlanReviewRequest:
|
||||
req_event = ev
|
||||
assert req_event is not None
|
||||
@@ -337,7 +335,7 @@ async def test_magentic_orchestrator_round_limit_produces_partial_result():
|
||||
)
|
||||
|
||||
events: list[WorkflowEvent] = []
|
||||
async for ev in wf.run_stream("round limit test"):
|
||||
async for ev in wf.run("round limit test", stream=True):
|
||||
events.append(ev)
|
||||
|
||||
idle_status = next(
|
||||
@@ -370,7 +368,7 @@ async def test_magentic_checkpoint_resume_round_trip():
|
||||
|
||||
task_text = "checkpoint task"
|
||||
req_event: RequestInfoEvent | None = None
|
||||
async for ev in wf.run_stream(task_text):
|
||||
async for ev in wf.run(task_text, stream=True):
|
||||
if isinstance(ev, RequestInfoEvent) and ev.request_type is MagenticPlanReviewRequest:
|
||||
req_event = ev
|
||||
assert req_event is not None
|
||||
@@ -393,8 +391,9 @@ async def test_magentic_checkpoint_resume_round_trip():
|
||||
|
||||
completed: WorkflowOutputEvent | None = None
|
||||
req_event = None
|
||||
async for event in wf_resume.run_stream(
|
||||
async for event in wf_resume.run(
|
||||
resume_checkpoint.checkpoint_id,
|
||||
stream=True,
|
||||
):
|
||||
if isinstance(event, RequestInfoEvent) and event.request_type is MagenticPlanReviewRequest:
|
||||
req_event = event
|
||||
@@ -419,26 +418,24 @@ async def test_magentic_checkpoint_resume_round_trip():
|
||||
class StubManagerAgent(BaseAgent):
|
||||
"""Stub agent for testing StandardMagenticManager."""
|
||||
|
||||
async def run(
|
||||
def run(
|
||||
self,
|
||||
messages: str | ChatMessage | Sequence[str | ChatMessage] | None = None,
|
||||
*,
|
||||
stream: bool = False,
|
||||
thread: Any = None,
|
||||
**kwargs: Any,
|
||||
) -> AgentResponse:
|
||||
return AgentResponse(messages=[ChatMessage("assistant", ["ok"])])
|
||||
) -> Awaitable[AgentResponse] | AsyncIterable[AgentResponseUpdate]:
|
||||
if stream:
|
||||
return self._run_stream()
|
||||
|
||||
def run_stream(
|
||||
self,
|
||||
messages: str | ChatMessage | Sequence[str | ChatMessage] | None = None,
|
||||
*,
|
||||
thread: Any = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterable[AgentResponseUpdate]:
|
||||
async def _gen() -> AsyncIterable[AgentResponseUpdate]:
|
||||
yield AgentResponseUpdate(message_deltas=[ChatMessage("assistant", ["ok"])])
|
||||
async def _run() -> AgentResponse:
|
||||
return AgentResponse(messages=[ChatMessage("assistant", ["ok"])])
|
||||
|
||||
return _gen()
|
||||
return _run()
|
||||
|
||||
async def _run_stream(self) -> AsyncIterable[AgentResponseUpdate]:
|
||||
yield AgentResponseUpdate(message_deltas=[ChatMessage("assistant", ["ok"])])
|
||||
|
||||
|
||||
async def test_standard_manager_plan_and_replan_via_complete_monkeypatch():
|
||||
@@ -538,16 +535,22 @@ class StubThreadAgent(BaseAgent):
|
||||
def __init__(self, name: str | None = None) -> None:
|
||||
super().__init__(name=name or "agentA")
|
||||
|
||||
async def run_stream(self, messages=None, *, thread=None, **kwargs): # type: ignore[override]
|
||||
def run(self, messages=None, *, stream: bool = False, thread=None, **kwargs): # type: ignore[override]
|
||||
if stream:
|
||||
return self._run_stream()
|
||||
|
||||
async def _run():
|
||||
return AgentResponse(messages=[ChatMessage("assistant", ["thread-ok"], author_name=self.name)])
|
||||
|
||||
return _run()
|
||||
|
||||
async def _run_stream(self):
|
||||
yield AgentResponseUpdate(
|
||||
contents=[Content.from_text(text="thread-ok")],
|
||||
author_name=self.name,
|
||||
role="assistant",
|
||||
)
|
||||
|
||||
async def run(self, messages=None, *, thread=None, **kwargs): # type: ignore[override]
|
||||
return AgentResponse(messages=[ChatMessage("assistant", ["thread-ok"], author_name=self.name)])
|
||||
|
||||
|
||||
class StubAssistantsClient:
|
||||
pass # class name used for branch detection
|
||||
@@ -560,16 +563,22 @@ class StubAssistantsAgent(BaseAgent):
|
||||
super().__init__(name="agentA")
|
||||
self.chat_client = StubAssistantsClient() # type name contains 'AssistantsClient'
|
||||
|
||||
async def run_stream(self, messages=None, *, thread=None, **kwargs): # type: ignore[override]
|
||||
def run(self, messages=None, *, stream: bool = False, thread=None, **kwargs): # type: ignore[override]
|
||||
if stream:
|
||||
return self._run_stream()
|
||||
|
||||
async def _run():
|
||||
return AgentResponse(messages=[ChatMessage("assistant", ["assistants-ok"], author_name=self.name)])
|
||||
|
||||
return _run()
|
||||
|
||||
async def _run_stream(self):
|
||||
yield AgentResponseUpdate(
|
||||
contents=[Content.from_text(text="assistants-ok")],
|
||||
author_name=self.name,
|
||||
role="assistant",
|
||||
)
|
||||
|
||||
async def run(self, messages=None, *, thread=None, **kwargs): # type: ignore[override]
|
||||
return AgentResponse(messages=[ChatMessage("assistant", ["assistants-ok"], author_name=self.name)])
|
||||
|
||||
|
||||
async def _collect_agent_responses_setup(participant: AgentProtocol) -> list[ChatMessage]:
|
||||
captured: list[ChatMessage] = []
|
||||
@@ -584,7 +593,7 @@ async def _collect_agent_responses_setup(participant: AgentProtocol) -> list[Cha
|
||||
|
||||
# Run a bounded stream to allow one invoke and then completion
|
||||
events: list[WorkflowEvent] = []
|
||||
async for ev in wf.run_stream("task"): # plan review disabled
|
||||
async for ev in wf.run("task", stream=True): # plan review disabled
|
||||
events.append(ev)
|
||||
if isinstance(ev, WorkflowOutputEvent) and isinstance(ev.data, AgentResponseUpdate):
|
||||
captured.append(
|
||||
@@ -630,7 +639,7 @@ async def test_magentic_checkpoint_resume_inner_loop_superstep():
|
||||
.build()
|
||||
)
|
||||
|
||||
async for event in workflow.run_stream("inner-loop task"):
|
||||
async for event in workflow.run("inner-loop task", stream=True):
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
break
|
||||
|
||||
@@ -646,7 +655,7 @@ async def test_magentic_checkpoint_resume_inner_loop_superstep():
|
||||
)
|
||||
|
||||
completed: WorkflowOutputEvent | None = None
|
||||
async for event in resumed.run_stream(checkpoint_id=inner_loop_checkpoint.checkpoint_id): # type: ignore[reportUnknownMemberType]
|
||||
async for event in resumed.run(checkpoint_id=inner_loop_checkpoint.checkpoint_id, stream=True): # type: ignore[reportUnknownMemberType]
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
completed = event
|
||||
|
||||
@@ -668,7 +677,7 @@ async def test_magentic_checkpoint_resume_from_saved_state():
|
||||
.build()
|
||||
)
|
||||
|
||||
async for event in workflow.run_stream("checkpoint resume task"):
|
||||
async for event in workflow.run("checkpoint resume task", stream=True):
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
break
|
||||
|
||||
@@ -686,7 +695,7 @@ async def test_magentic_checkpoint_resume_from_saved_state():
|
||||
)
|
||||
|
||||
completed: WorkflowOutputEvent | None = None
|
||||
async for event in resumed_workflow.run_stream(checkpoint_id=resumed_state.checkpoint_id):
|
||||
async for event in resumed_workflow.run(checkpoint_id=resumed_state.checkpoint_id, stream=True):
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
completed = event
|
||||
|
||||
@@ -708,7 +717,7 @@ async def test_magentic_checkpoint_resume_rejects_participant_renames():
|
||||
)
|
||||
|
||||
req_event: RequestInfoEvent | None = None
|
||||
async for event in workflow.run_stream("task"):
|
||||
async for event in workflow.run("task", stream=True):
|
||||
if isinstance(event, RequestInfoEvent) and event.request_type is MagenticPlanReviewRequest:
|
||||
req_event = event
|
||||
|
||||
@@ -728,7 +737,8 @@ async def test_magentic_checkpoint_resume_rejects_participant_renames():
|
||||
)
|
||||
|
||||
with pytest.raises(WorkflowCheckpointException, match="Workflow graph has changed"):
|
||||
async for _ in renamed_workflow.run_stream(
|
||||
async for _ in renamed_workflow.run(
|
||||
stream=True,
|
||||
checkpoint_id=target_checkpoint.checkpoint_id, # type: ignore[reportUnknownMemberType]
|
||||
):
|
||||
pass
|
||||
@@ -764,7 +774,7 @@ async def test_magentic_stall_and_reset_reach_limits():
|
||||
wf = MagenticBuilder().participants([DummyExec("agentA")]).with_manager(manager=manager).build()
|
||||
|
||||
events: list[WorkflowEvent] = []
|
||||
async for ev in wf.run_stream("test limits"):
|
||||
async for ev in wf.run("test limits", stream=True):
|
||||
events.append(ev)
|
||||
|
||||
idle_status = next(
|
||||
@@ -789,7 +799,7 @@ async def test_magentic_checkpoint_runtime_only() -> None:
|
||||
wf = MagenticBuilder().participants([DummyExec("agentA")]).with_manager(manager=manager).build()
|
||||
|
||||
baseline_output: ChatMessage | None = None
|
||||
async for ev in wf.run_stream("runtime checkpoint test", checkpoint_storage=storage):
|
||||
async for ev in wf.run("runtime checkpoint test", checkpoint_storage=storage, stream=True):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
baseline_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state in (
|
||||
@@ -827,7 +837,7 @@ async def test_magentic_checkpoint_runtime_overrides_buildtime() -> None:
|
||||
)
|
||||
|
||||
baseline_output: ChatMessage | None = None
|
||||
async for ev in wf.run_stream("override test", checkpoint_storage=runtime_storage):
|
||||
async for ev in wf.run("override test", checkpoint_storage=runtime_storage, stream=True):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
baseline_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state in (
|
||||
@@ -886,7 +896,7 @@ async def test_magentic_checkpoint_restore_no_duplicate_history():
|
||||
ChatMessage("user", ["task_msg"]),
|
||||
]
|
||||
|
||||
async for event in wf.run_stream(conversation):
|
||||
async for event in wf.run(conversation, stream=True):
|
||||
if isinstance(event, WorkflowStatusEvent) and event.state in (
|
||||
WorkflowRunState.IDLE,
|
||||
WorkflowRunState.IDLE_WITH_PENDING_REQUESTS,
|
||||
@@ -996,7 +1006,7 @@ async def test_magentic_with_participant_factories():
|
||||
assert call_count == 1
|
||||
|
||||
outputs: list[WorkflowOutputEvent] = []
|
||||
async for event in workflow.run_stream("test task"):
|
||||
async for event in workflow.run("test task", stream=True):
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
outputs.append(event)
|
||||
|
||||
@@ -1043,7 +1053,7 @@ async def test_magentic_participant_factories_with_checkpointing():
|
||||
)
|
||||
|
||||
outputs: list[WorkflowOutputEvent] = []
|
||||
async for event in workflow.run_stream("checkpoint test"):
|
||||
async for event in workflow.run("checkpoint test", stream=True):
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
outputs.append(event)
|
||||
|
||||
@@ -1100,7 +1110,7 @@ async def test_magentic_with_manager_factory():
|
||||
assert factory_call_count == 1
|
||||
|
||||
outputs: list[WorkflowOutputEvent] = []
|
||||
async for event in workflow.run_stream("test task"):
|
||||
async for event in workflow.run("test task", stream=True):
|
||||
if isinstance(event, WorkflowOutputEvent):
|
||||
outputs.append(event)
|
||||
|
||||
@@ -1129,7 +1139,7 @@ async def test_magentic_with_agent_factory():
|
||||
|
||||
# Verify workflow can be started (may not complete successfully due to stub behavior)
|
||||
event_count = 0
|
||||
async for _ in workflow.run_stream("test task"):
|
||||
async for _ in workflow.run("test task", stream=True):
|
||||
event_count += 1
|
||||
if event_count > 10:
|
||||
break
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from collections.abc import AsyncIterable
|
||||
from collections.abc import AsyncIterable, Awaitable
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
@@ -27,22 +27,23 @@ from agent_framework.orchestrations import SequentialBuilder
|
||||
class _EchoAgent(BaseAgent):
|
||||
"""Simple agent that appends a single assistant message with its name."""
|
||||
|
||||
async def run( # type: ignore[override]
|
||||
def run( # type: ignore[override]
|
||||
self,
|
||||
messages: str | ChatMessage | list[str] | list[ChatMessage] | None = None,
|
||||
*,
|
||||
stream: bool = False,
|
||||
thread: AgentThread | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AgentResponse:
|
||||
return AgentResponse(messages=[ChatMessage("assistant", [f"{self.name} reply"])])
|
||||
) -> Awaitable[AgentResponse] | AsyncIterable[AgentResponseUpdate]:
|
||||
if stream:
|
||||
return self._run_stream()
|
||||
|
||||
async def run_stream( # type: ignore[override]
|
||||
self,
|
||||
messages: str | ChatMessage | list[str] | list[ChatMessage] | None = None,
|
||||
*,
|
||||
thread: AgentThread | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterable[AgentResponseUpdate]:
|
||||
async def _run() -> AgentResponse:
|
||||
return AgentResponse(messages=[ChatMessage("assistant", [f"{self.name} reply"])])
|
||||
|
||||
return _run()
|
||||
|
||||
async def _run_stream(self) -> AsyncIterable[AgentResponseUpdate]:
|
||||
# Minimal async generator with one assistant update
|
||||
yield AgentResponseUpdate(contents=[Content.from_text(text=f"{self.name} reply")])
|
||||
|
||||
@@ -104,7 +105,7 @@ async def test_sequential_agents_append_to_context() -> None:
|
||||
|
||||
completed = False
|
||||
output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("hello sequential"):
|
||||
async for ev in wf.run("hello sequential", stream=True):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
@@ -137,7 +138,7 @@ async def test_sequential_register_participants_with_agent_factories() -> None:
|
||||
|
||||
completed = False
|
||||
output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("hello factories"):
|
||||
async for ev in wf.run("hello factories", stream=True):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
@@ -163,7 +164,7 @@ async def test_sequential_with_custom_executor_summary() -> None:
|
||||
|
||||
completed = False
|
||||
output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("topic X"):
|
||||
async for ev in wf.run("topic X", stream=True):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
@@ -194,7 +195,7 @@ async def test_sequential_register_participants_mixed_agents_and_executors() ->
|
||||
|
||||
completed = False
|
||||
output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("topic Y"):
|
||||
async for ev in wf.run("topic Y", stream=True):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
@@ -219,7 +220,7 @@ async def test_sequential_checkpoint_resume_round_trip() -> None:
|
||||
wf = SequentialBuilder().participants(list(initial_agents)).with_checkpointing(storage).build()
|
||||
|
||||
baseline_output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("checkpoint sequential"):
|
||||
async for ev in wf.run("checkpoint sequential", stream=True):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
baseline_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
@@ -240,7 +241,7 @@ async def test_sequential_checkpoint_resume_round_trip() -> None:
|
||||
wf_resume = SequentialBuilder().participants(list(resumed_agents)).with_checkpointing(storage).build()
|
||||
|
||||
resumed_output: list[ChatMessage] | None = None
|
||||
async for ev in wf_resume.run_stream(checkpoint_id=resume_checkpoint.checkpoint_id):
|
||||
async for ev in wf_resume.run(checkpoint_id=resume_checkpoint.checkpoint_id, stream=True):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
resumed_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state in (
|
||||
@@ -262,7 +263,7 @@ async def test_sequential_checkpoint_runtime_only() -> None:
|
||||
wf = SequentialBuilder().participants(list(agents)).build()
|
||||
|
||||
baseline_output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("runtime checkpoint test", checkpoint_storage=storage):
|
||||
async for ev in wf.run("runtime checkpoint test", checkpoint_storage=storage, stream=True):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
baseline_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
@@ -283,7 +284,9 @@ async def test_sequential_checkpoint_runtime_only() -> None:
|
||||
wf_resume = SequentialBuilder().participants(list(resumed_agents)).build()
|
||||
|
||||
resumed_output: list[ChatMessage] | None = None
|
||||
async for ev in wf_resume.run_stream(checkpoint_id=resume_checkpoint.checkpoint_id, checkpoint_storage=storage):
|
||||
async for ev in wf_resume.run(
|
||||
checkpoint_id=resume_checkpoint.checkpoint_id, checkpoint_storage=storage, stream=True
|
||||
):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
resumed_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state in (
|
||||
@@ -311,7 +314,7 @@ async def test_sequential_checkpoint_runtime_overrides_buildtime() -> None:
|
||||
wf = SequentialBuilder().participants(list(agents)).with_checkpointing(buildtime_storage).build()
|
||||
|
||||
baseline_output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("override test", checkpoint_storage=runtime_storage):
|
||||
async for ev in wf.run("override test", checkpoint_storage=runtime_storage, stream=True):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
baseline_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
@@ -339,7 +342,7 @@ async def test_sequential_register_participants_with_checkpointing() -> None:
|
||||
wf = SequentialBuilder().register_participants([create_agent1, create_agent2]).with_checkpointing(storage).build()
|
||||
|
||||
baseline_output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("checkpoint with factories"):
|
||||
async for ev in wf.run("checkpoint with factories", stream=True):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
baseline_output = ev.data
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
@@ -361,7 +364,7 @@ async def test_sequential_register_participants_with_checkpointing() -> None:
|
||||
)
|
||||
|
||||
resumed_output: list[ChatMessage] | None = None
|
||||
async for ev in wf_resume.run_stream(checkpoint_id=resume_checkpoint.checkpoint_id):
|
||||
async for ev in wf_resume.run(checkpoint_id=resume_checkpoint.checkpoint_id, stream=True):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
resumed_output = ev.data
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state in (
|
||||
@@ -397,7 +400,7 @@ async def test_sequential_register_participants_factories_called_on_build() -> N
|
||||
# Run the workflow to ensure it works
|
||||
completed = False
|
||||
output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("test factories timing"):
|
||||
async for ev in wf.run("test factories timing", stream=True):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
|
||||
Reference in New Issue
Block a user