[BREAKING] Python: Refactor workflows kwargs (#5010)

* Refactor workflows kwargs usage

* Update sample

* Add tests

* Update samples

* Fix formatting

* Comments

* Comments 2

* Comments 3

* Fix test and typing
This commit is contained in:
Tao Chen
2026-04-02 02:40:39 -07:00
committed by GitHub
Unverified
parent fd253c0b0e
commit 62595b233f
10 changed files with 1092 additions and 826 deletions
@@ -1,8 +1,7 @@
# Copyright (c) Microsoft. All rights reserved.
import logging
from collections.abc import AsyncIterable, Awaitable
from typing import TYPE_CHECKING, Any, Literal, overload
from typing import Any, Literal, overload
import pytest
@@ -22,9 +21,7 @@ from agent_framework import (
)
from agent_framework._workflows._agent_executor import AgentExecutorResponse
from agent_framework._workflows._checkpoint import InMemoryCheckpointStorage
if TYPE_CHECKING:
from _pytest.logging import LogCaptureFixture
from agent_framework._workflows._const import GLOBAL_KWARGS_KEY
class _CountingAgent(BaseAgent):
@@ -309,87 +306,28 @@ async def test_agent_executor_save_and_restore_state_directly() -> None:
assert restored_session.session_id == session.session_id
async def test_agent_executor_run_with_session_kwarg_does_not_raise() -> None:
"""Passing session= via workflow.run() should not cause a duplicate-keyword TypeError (#4295)."""
agent = _CountingAgent(id="session_kwarg_agent", name="SessionKwargAgent")
executor = AgentExecutor(agent, id="session_kwarg_exec")
workflow = WorkflowBuilder(start_executor=executor).build()
async def test_prepare_agent_run_args_extracts_invocation_kwargs() -> None:
"""_prepare_agent_run_args extracts function_invocation_kwargs and client_kwargs."""
agent = _CountingAgent(id="test_agent", name="TestAgent")
executor = AgentExecutor(agent, id="test_exec")
# This previously raised: TypeError: run() got multiple values for keyword argument 'session'
result = await workflow.run("hello", session="user-supplied-value")
assert result is not None
assert agent.call_count == 1
async def test_agent_executor_run_streaming_with_stream_kwarg_does_not_raise() -> None:
"""Passing stream= via workflow.run() kwargs should not cause a duplicate-keyword TypeError."""
agent = _CountingAgent(id="stream_kwarg_agent", name="StreamKwargAgent")
executor = AgentExecutor(agent, id="stream_kwarg_exec")
workflow = WorkflowBuilder(start_executor=executor).build()
# stream=True at workflow level triggers streaming mode (returns async iterable)
events: list[WorkflowEvent] = []
async for event in workflow.run("hello", stream=True):
events.append(event)
assert len(events) > 0
assert agent.call_count == 1
@pytest.mark.parametrize("reserved_kwarg", ["session", "stream", "messages"])
async def test_prepare_agent_run_args_strips_reserved_kwargs(reserved_kwarg: str, caplog: "LogCaptureFixture") -> None:
"""_prepare_agent_run_args must remove reserved kwargs and log a warning."""
raw: dict[str, Any] = {
reserved_kwarg: "should-be-stripped",
"custom_key": "keep-me",
"function_invocation_kwargs": {"__global__": {"key": "fi_val"}},
"client_kwargs": {"__global__": {"key": "ci_val"}},
}
with caplog.at_level(logging.WARNING):
run_kwargs, options = AgentExecutor._prepare_agent_run_args(raw) # pyright: ignore[reportPrivateUsage]
assert reserved_kwarg not in run_kwargs
assert "custom_key" in run_kwargs
assert options is not None
assert options["additional_function_arguments"]["custom_key"] == "keep-me"
assert any(reserved_kwarg in record.message for record in caplog.records)
fi_kwargs, ci_kwargs = executor._prepare_agent_run_args(raw) # pyright: ignore[reportPrivateUsage]
assert fi_kwargs == {"key": "fi_val"}
assert ci_kwargs == {"key": "ci_val"}
async def test_prepare_agent_run_args_preserves_non_reserved_kwargs() -> None:
"""Non-reserved workflow kwargs should pass through unchanged."""
raw: dict[str, Any] = {"custom_param": "value", "another": 42}
run_kwargs, _options = AgentExecutor._prepare_agent_run_args(raw) # pyright: ignore[reportPrivateUsage]
assert run_kwargs["custom_param"] == "value"
assert run_kwargs["another"] == 42
async def test_prepare_agent_run_args_returns_none_when_no_kwargs() -> None:
"""_prepare_agent_run_args returns None for both when raw dict has no invocation kwargs."""
agent = _CountingAgent(id="test_agent", name="TestAgent")
executor = AgentExecutor(agent, id="test_exec")
async def test_prepare_agent_run_args_strips_all_reserved_kwargs_at_once(
caplog: "LogCaptureFixture",
) -> None:
"""All reserved kwargs should be stripped when supplied together, each emitting a warning."""
raw: dict[str, Any] = {"session": "x", "stream": True, "messages": [], "custom": 1}
with caplog.at_level(logging.WARNING):
run_kwargs, options = AgentExecutor._prepare_agent_run_args(raw) # pyright: ignore[reportPrivateUsage]
assert "session" not in run_kwargs
assert "stream" not in run_kwargs
assert "messages" not in run_kwargs
assert run_kwargs["custom"] == 1
assert options is not None
assert options["additional_function_arguments"]["custom"] == 1
warned_keys = {r.message.split("'")[1] for r in caplog.records if "reserved" in r.message.lower()}
assert warned_keys == {"session", "stream", "messages"}
async def test_agent_executor_run_with_messages_kwarg_does_not_raise() -> None:
"""Passing messages= via workflow.run() kwargs should not cause a duplicate-keyword TypeError."""
agent = _CountingAgent(id="messages_kwarg_agent", name="MessagesKwargAgent")
executor = AgentExecutor(agent, id="messages_kwarg_exec")
workflow = WorkflowBuilder(start_executor=executor).build()
result = await workflow.run("hello", messages=["stale"])
assert result is not None
assert agent.call_count == 1
fi_kwargs, ci_kwargs = executor._prepare_agent_run_args({}) # pyright: ignore[reportPrivateUsage]
assert fi_kwargs is None
assert ci_kwargs is None
class _NonCopyableRaw:
@@ -638,3 +576,126 @@ async def test_checkpoint_restore_works_without_context_mode_in_state() -> None:
assert cache[0].text == "cached msg"
# context_mode should remain as configured in the constructor, not changed by restore
assert executor._context_mode == "last_agent" # pyright: ignore[reportPrivateUsage]
# ---------------------------------------------------------------------------
# Per-executor kwargs resolution tests
# ---------------------------------------------------------------------------
async def test_resolve_executor_kwargs_returns_global_kwargs() -> None:
"""_resolve_executor_kwargs with the global kwargs key returns the global kwargs."""
agent = _CountingAgent(id="a", name="A")
executor = AgentExecutor(agent, id="exec_a")
resolved = {GLOBAL_KWARGS_KEY: {"tool_param": "value"}}
result = executor._resolve_executor_kwargs(resolved) # pyright: ignore[reportPrivateUsage]
assert result == {"tool_param": "value"}
async def test_resolve_executor_kwargs_returns_per_executor_kwargs() -> None:
"""_resolve_executor_kwargs with matching executor ID returns that executor's kwargs."""
agent = _CountingAgent(id="a", name="A")
executor = AgentExecutor(agent, id="exec_a")
resolved = {"exec_a": {"my_param": 42}, "exec_b": {"other_param": 99}}
result = executor._resolve_executor_kwargs(resolved) # pyright: ignore[reportPrivateUsage]
assert result == {"my_param": 42}
async def test_resolve_executor_kwargs_returns_none_for_unmatched_per_executor() -> None:
"""_resolve_executor_kwargs returns None when per-executor dict has no matching ID."""
agent = _CountingAgent(id="a", name="A")
executor = AgentExecutor(agent, id="exec_c")
resolved = {"exec_a": {"my_param": 42}, "exec_b": {"other_param": 99}}
result = executor._resolve_executor_kwargs(resolved) # pyright: ignore[reportPrivateUsage]
assert result is None
async def test_resolve_executor_kwargs_returns_none_for_none_input() -> None:
"""_resolve_executor_kwargs returns None when input is None."""
agent = _CountingAgent(id="a", name="A")
executor = AgentExecutor(agent, id="exec_a")
result = executor._resolve_executor_kwargs(None) # pyright: ignore[reportPrivateUsage]
assert result is None
async def test_resolve_executor_kwargs_prefers_executor_id_over_global() -> None:
"""_resolve_executor_kwargs prefers executor-specific entry over __global__."""
agent = _CountingAgent(id="a", name="A")
executor = AgentExecutor(agent, id="exec_a")
# Dict has both a per-executor entry and a global entry
resolved = {"exec_a": {"specific": True}, GLOBAL_KWARGS_KEY: {"global": True}}
result = executor._resolve_executor_kwargs(resolved) # pyright: ignore[reportPrivateUsage]
assert result == {"specific": True}
async def test_prepare_agent_run_args_extracts_function_invocation_kwargs() -> None:
"""_prepare_agent_run_args extracts function_invocation_kwargs from the state dict."""
agent = _CountingAgent(id="a", name="A")
executor = AgentExecutor(agent, id="exec_a")
raw: dict[str, Any] = {
"function_invocation_kwargs": {GLOBAL_KWARGS_KEY: {"tool_key": "tool_val"}},
}
fi_kwargs, client_kwargs = executor._prepare_agent_run_args(raw) # pyright: ignore[reportPrivateUsage]
assert fi_kwargs == {"tool_key": "tool_val"}
assert client_kwargs is None
async def test_prepare_agent_run_args_extracts_client_kwargs() -> None:
"""_prepare_agent_run_args extracts client_kwargs from the state dict."""
agent = _CountingAgent(id="a", name="A")
executor = AgentExecutor(agent, id="exec_a")
raw: dict[str, Any] = {
"client_kwargs": {GLOBAL_KWARGS_KEY: {"model": "gpt-4"}},
}
fi_kwargs, client_kwargs = executor._prepare_agent_run_args(raw) # pyright: ignore[reportPrivateUsage]
assert fi_kwargs is None
assert client_kwargs == {"model": "gpt-4"}
async def test_prepare_agent_run_args_per_executor_resolution() -> None:
"""_prepare_agent_run_args resolves per-executor function_invocation_kwargs using self.id."""
agent = _CountingAgent(id="a", name="A")
executor = AgentExecutor(agent, id="exec_a")
raw: dict[str, Any] = {
"function_invocation_kwargs": {
"exec_a": {"my_tool_key": "my_val"},
"exec_b": {"other_tool_key": "other_val"},
},
}
fi_kwargs, _ = executor._prepare_agent_run_args(raw) # pyright: ignore[reportPrivateUsage]
assert fi_kwargs == {"my_tool_key": "my_val"}
async def test_prepare_agent_run_args_per_executor_no_match() -> None:
"""_prepare_agent_run_args returns None for function_invocation_kwargs when executor ID not found."""
agent = _CountingAgent(id="a", name="A")
executor = AgentExecutor(agent, id="exec_c")
raw: dict[str, Any] = {
"function_invocation_kwargs": {
"exec_a": {"my_tool_key": "my_val"},
"exec_b": {"other_tool_key": "other_val"},
},
}
fi_kwargs, _ = executor._prepare_agent_run_args(raw) # pyright: ignore[reportPrivateUsage]
assert fi_kwargs is None
async def test_resolve_executor_kwargs_empty_per_executor_does_not_fallback_to_global() -> None:
"""An explicit empty per-executor dict should not fall through to global kwargs."""
agent = _CountingAgent(id="a", name="A")
executor = AgentExecutor(agent, id="exec_a")
# Per-executor entry for exec_a is empty, but global has values.
# The empty dict should be honoured (no fallback to global).
resolved = {"exec_a": {}, GLOBAL_KWARGS_KEY: {"global_key": "global_val"}}
result = executor._resolve_executor_kwargs(resolved) # pyright: ignore[reportPrivateUsage]
assert result == {}
File diff suppressed because it is too large Load Diff