Merge branch 'main' into local-branch-python-add-reset-to-workflow

This commit is contained in:
Tao Chen
2026-06-11 15:28:27 -07:00
Unverified
304 changed files with 12129 additions and 1067 deletions
@@ -194,6 +194,63 @@ def test_create_harness_agent_returns_full_agent() -> None:
assert isinstance(agent, FullAgent)
def test_create_harness_agent_no_token_params_disables_compaction() -> None:
"""When token params are omitted, compaction is automatically disabled."""
agent = create_harness_agent(
client=_FakeChatClient(), # type: ignore[arg-type]
)
provider_types = [type(p) for p in agent.context_providers]
assert CompactionProvider not in provider_types
def test_create_harness_agent_no_token_params_skips_max_tokens_option() -> None:
"""When max_output_tokens is omitted, max_tokens should not be set in default options."""
agent = create_harness_agent(
client=_FakeChatClient(), # type: ignore[arg-type]
)
assert agent.default_options.get("max_tokens") is None
def test_create_harness_agent_custom_before_strategy_enables_compaction_without_tokens() -> None:
"""A custom before_compaction_strategy enables compaction even when token params are omitted."""
from agent_framework import ToolResultCompactionStrategy
agent = create_harness_agent(
client=_FakeChatClient(), # type: ignore[arg-type]
before_compaction_strategy=ToolResultCompactionStrategy(),
)
provider_types = [type(p) for p in agent.context_providers]
assert CompactionProvider in provider_types
def test_create_harness_agent_disable_compaction_overrides_custom_before_strategy() -> None:
"""disable_compaction=True wins even when a custom before strategy is provided."""
from agent_framework import ToolResultCompactionStrategy
agent = create_harness_agent(
client=_FakeChatClient(), # type: ignore[arg-type]
before_compaction_strategy=ToolResultCompactionStrategy(),
disable_compaction=True,
)
provider_types = [type(p) for p in agent.context_providers]
assert CompactionProvider not in provider_types
def test_create_harness_agent_custom_after_strategy_enables_compaction_without_tokens() -> None:
"""A custom after_compaction_strategy enables compaction even when token params are omitted."""
from agent_framework import ToolResultCompactionStrategy
agent = create_harness_agent(
client=_FakeChatClient(), # type: ignore[arg-type]
after_compaction_strategy=ToolResultCompactionStrategy(),
)
compaction_providers = [p for p in agent.context_providers if isinstance(p, CompactionProvider)]
assert len(compaction_providers) == 1
# Before phase is skipped (no token budget, no custom before strategy), after phase is set.
assert compaction_providers[0].before_strategy is None
assert compaction_providers[0].after_strategy is not None
# --- Validation Tests ---
@@ -207,14 +264,15 @@ def test_create_harness_agent_rejects_invalid_context_tokens() -> None:
)
def test_create_harness_agent_rejects_negative_output_tokens() -> None:
"""max_output_tokens must be non-negative."""
with pytest.raises(ValueError, match="max_output_tokens must be non-negative"):
create_harness_agent(
client=_FakeChatClient(), # type: ignore[arg-type]
max_context_window_tokens=1000,
max_output_tokens=-1,
)
def test_create_harness_agent_rejects_non_positive_output_tokens() -> None:
"""max_output_tokens must be positive when provided."""
for invalid_value in (0, -1):
with pytest.raises(ValueError, match="max_output_tokens must be positive"):
create_harness_agent(
client=_FakeChatClient(), # type: ignore[arg-type]
max_context_window_tokens=1000,
max_output_tokens=invalid_value,
)
def test_create_harness_agent_rejects_output_gte_context() -> None:
@@ -485,3 +543,127 @@ def test_create_harness_agent_empty_background_agents_list() -> None:
)
providers = agent.context_providers or []
assert not any(isinstance(p, BackgroundAgentsProvider) for p in providers)
# --- Shell Tool Tests ---
class _FakeShellTool:
"""Fake shell executor/tool exposing as_function()."""
def as_function(self) -> str:
return "shell_fn"
class _FakeShellClient(_FakeChatClient):
"""Fake client that supports the shell tool."""
def __init__(self) -> None:
self.shell_func: Any = None
def get_shell_tool(self, *, func: Any = None, **kwargs: Any) -> str:
self.shell_func = func
return "shell_tool_instance"
def test_create_harness_agent_adds_shell_tool_and_provider() -> None:
"""Shell tool and ShellEnvironmentProvider should be added when a shell executor is supplied."""
from agent_framework_tools.shell import ShellEnvironmentProvider
client = _FakeShellClient()
agent = create_harness_agent(
client=client, # type: ignore[arg-type]
max_context_window_tokens=128_000,
max_output_tokens=16_384,
disable_web_search=True,
shell_executor=_FakeShellTool(),
)
tools = agent.default_options.get("tools", [])
assert "shell_tool_instance" in tools
assert client.shell_func == "shell_fn"
providers = agent.context_providers or []
assert any(isinstance(p, ShellEnvironmentProvider) for p in providers)
def test_create_harness_agent_shell_passes_custom_options() -> None:
"""Custom ShellEnvironmentProviderOptions should be forwarded to the provider."""
from agent_framework_tools.shell import ShellEnvironmentProvider, ShellEnvironmentProviderOptions
options = ShellEnvironmentProviderOptions(probe_tools=("git",))
agent = create_harness_agent(
client=_FakeShellClient(), # type: ignore[arg-type]
max_context_window_tokens=128_000,
max_output_tokens=16_384,
disable_web_search=True,
shell_executor=_FakeShellTool(),
shell_environment_provider_options=options,
)
providers = agent.context_providers or []
provider = next(p for p in providers if isinstance(p, ShellEnvironmentProvider))
assert provider._options is options
def test_create_harness_agent_shell_skipped_when_unsupported(caplog: pytest.LogCaptureFixture) -> None:
"""When the client lacks get_shell_tool, both the tool and provider are skipped with a warning."""
import logging
from agent_framework_tools.shell import ShellEnvironmentProvider
with caplog.at_level(logging.WARNING, logger="agent_framework._harness._agent"):
agent = create_harness_agent(
client=_FakeChatClient(), # type: ignore[arg-type]
max_context_window_tokens=128_000,
max_output_tokens=16_384,
disable_web_search=True,
shell_executor=_FakeShellTool(),
)
assert any("SupportsShellTool" in msg for msg in caplog.messages)
providers = agent.context_providers or []
assert not any(isinstance(p, ShellEnvironmentProvider) for p in providers)
assert "tools" not in agent.default_options or not agent.default_options.get("tools")
def test_create_harness_agent_no_shell_by_default() -> None:
"""No shell tool or provider should be added when shell_executor is not provided."""
from agent_framework_tools.shell import ShellEnvironmentProvider
agent = create_harness_agent(
client=_FakeShellClient(), # type: ignore[arg-type]
max_context_window_tokens=128_000,
max_output_tokens=16_384,
disable_web_search=True,
)
providers = agent.context_providers or []
assert not any(isinstance(p, ShellEnvironmentProvider) for p in providers)
def test_create_harness_agent_shell_executor_without_as_function_raises() -> None:
"""A shell_executor lacking a callable as_function() should raise a clear TypeError."""
class _BadExecutor:
pass
with pytest.raises(TypeError, match="as_function"):
create_harness_agent(
client=_FakeShellClient(), # type: ignore[arg-type]
max_context_window_tokens=128_000,
max_output_tokens=16_384,
disable_web_search=True,
shell_executor=_BadExecutor(),
)
def test_create_harness_agent_shell_executor_validated_before_client_check() -> None:
"""The as_function() contract is validated upfront, even when the client lacks shell support."""
class _BadExecutor:
pass
with pytest.raises(TypeError, match="as_function"):
create_harness_agent(
client=_FakeChatClient(), # type: ignore[arg-type]
max_context_window_tokens=128_000,
max_output_tokens=16_384,
disable_web_search=True,
shell_executor=_BadExecutor(),
)
@@ -0,0 +1,817 @@
# Copyright (c) Microsoft. All rights reserved.
from __future__ import annotations
from agent_framework import (
DEFAULT_TOOL_APPROVAL_SOURCE_ID,
Agent,
AgentSession,
ChatResponse,
ChatResponseUpdate,
Content,
Message,
SupportsChatGetResponse,
ToolApprovalMiddleware,
ToolApprovalState,
create_always_approve_tool_response,
create_always_approve_tool_with_arguments_response,
tool,
)
def _approval_requests(messages: list[Message]) -> list[Content]:
return [
content for message in messages for content in message.contents if content.type == "function_approval_request"
]
async def test_mixed_batch_hides_already_approved_request_until_approval_replay(
chat_client_base: SupportsChatGetResponse,
) -> None:
"""Mixed batches should only show real approval requests when a session can store hidden requests."""
no_approval_calls = 0
approval_calls = 0
@tool(name="lookup_work_items", approval_mode="never_require")
def lookup_work_items(query: str) -> str:
nonlocal no_approval_calls
no_approval_calls += 1
return f"found {query}"
@tool(name="add_comment", approval_mode="always_require")
def add_comment(comment: str) -> str:
nonlocal approval_calls
approval_calls += 1
return f"added {comment}"
agent = Agent(client=chat_client_base, tools=[lookup_work_items, add_comment])
session = AgentSession(session_id="approval-session")
chat_client_base.run_responses = [
ChatResponse(
messages=Message(
role="assistant",
contents=[
Content.from_function_call(
call_id="call_lookup",
name="lookup_work_items",
arguments='{"query": "mine"}',
),
Content.from_function_call(
call_id="call_comment",
name="add_comment",
arguments='{"comment": "done"}',
),
],
)
)
]
first_response = await agent.run("update work item", session=session)
requests = _approval_requests(first_response.messages)
assert [request.function_call.name for request in requests] == ["add_comment"]
assert no_approval_calls == 0
assert approval_calls == 0
chat_client_base.run_responses = [ChatResponse(messages=Message(role="assistant", contents=["complete"]))]
second_response = await agent.run(requests[0].to_function_approval_response(approved=True), session=session)
assert second_response.text == "complete"
assert no_approval_calls == 1
assert approval_calls == 1
async def test_mixed_batch_accepts_restored_tool_approval_state(
chat_client_base: SupportsChatGetResponse,
) -> None:
"""Mixed-batch bypass should work when session state contains ToolApprovalState."""
safe_calls = 0
risky_calls = 0
@tool(name="safe_read", approval_mode="never_require")
def safe_read() -> str:
nonlocal safe_calls
safe_calls += 1
return "safe"
@tool(name="risky_write", approval_mode="always_require")
def risky_write() -> str:
nonlocal risky_calls
risky_calls += 1
return "risky"
agent = Agent(client=chat_client_base, tools=[safe_read, risky_write])
session = AgentSession(session_id="restored-state-session")
session.state[DEFAULT_TOOL_APPROVAL_SOURCE_ID] = ToolApprovalState()
chat_client_base.run_responses = [
ChatResponse(
messages=Message(
role="assistant",
contents=[
Content.from_function_call(call_id="call_safe", name="safe_read", arguments="{}"),
Content.from_function_call(call_id="call_risky", name="risky_write", arguments="{}"),
],
)
)
]
first_response = await agent.run("read and write", session=session)
requests = _approval_requests(first_response.messages)
assert [request.function_call.name for request in requests] == ["risky_write"]
assert safe_calls == 0
assert risky_calls == 0
chat_client_base.run_responses = [ChatResponse(messages=Message(role="assistant", contents=["done"]))]
final_response = await agent.run(requests[0].to_function_approval_response(approved=True), session=session)
assert final_response.text == "done"
assert safe_calls == 1
assert risky_calls == 1
async def test_hidden_mixed_batch_requests_do_not_replay_on_unrelated_turn(
chat_client_base: SupportsChatGetResponse,
) -> None:
"""Stored hidden approvals should only replay when an approval response resumes the flow."""
safe_calls = 0
risky_calls = 0
@tool(name="safe_lookup", approval_mode="never_require")
def safe_lookup() -> str:
nonlocal safe_calls
safe_calls += 1
return "safe"
@tool(name="risky_update", approval_mode="always_require")
def risky_update() -> str:
nonlocal risky_calls
risky_calls += 1
return "risky"
agent = Agent(client=chat_client_base, tools=[safe_lookup, risky_update])
session = AgentSession(session_id="stale-hidden-session")
chat_client_base.run_responses = [
ChatResponse(
messages=Message(
role="assistant",
contents=[
Content.from_function_call(call_id="call_safe", name="safe_lookup", arguments="{}"),
Content.from_function_call(call_id="call_risky", name="risky_update", arguments="{}"),
],
)
)
]
first_response = await agent.run("lookup and update", session=session)
request = _approval_requests(first_response.messages)[0]
chat_client_base.run_responses = [ChatResponse(messages=Message(role="assistant", contents=["unrelated"]))]
unrelated_response = await agent.run("never mind, answer something else", session=session)
assert unrelated_response.text == "unrelated"
assert safe_calls == 0
assert risky_calls == 0
chat_client_base.run_responses = [ChatResponse(messages=Message(role="assistant", contents=["done"]))]
final_response = await agent.run(request.to_function_approval_response(approved=True), session=session)
assert final_response.text == "done"
assert safe_calls == 1
assert risky_calls == 1
async def test_hidden_mixed_batch_requests_replay_only_for_matching_visible_approval(
chat_client_base: SupportsChatGetResponse,
) -> None:
"""Approving one mixed batch must not replay hidden calls from another abandoned batch."""
safe_a_calls = 0
safe_b_calls = 0
risky_a_calls = 0
risky_b_calls = 0
@tool(name="safe_a", approval_mode="never_require")
def safe_a() -> str:
nonlocal safe_a_calls
safe_a_calls += 1
return "safe-a"
@tool(name="safe_b", approval_mode="never_require")
def safe_b() -> str:
nonlocal safe_b_calls
safe_b_calls += 1
return "safe-b"
@tool(name="risky_a", approval_mode="always_require")
def risky_a() -> str:
nonlocal risky_a_calls
risky_a_calls += 1
return "risky-a"
@tool(name="risky_b", approval_mode="always_require")
def risky_b() -> str:
nonlocal risky_b_calls
risky_b_calls += 1
return "risky-b"
agent = Agent(client=chat_client_base, tools=[safe_a, safe_b, risky_a, risky_b])
session = AgentSession(session_id="grouped-hidden-session")
chat_client_base.run_responses = [
ChatResponse(
messages=Message(
role="assistant",
contents=[
Content.from_function_call(call_id="call_safe_a", name="safe_a", arguments="{}"),
Content.from_function_call(call_id="call_risky_a", name="risky_a", arguments="{}"),
],
)
)
]
first_response = await agent.run("batch a", session=session)
assert [request.function_call.name for request in _approval_requests(first_response.messages)] == ["risky_a"]
chat_client_base.run_responses = [
ChatResponse(
messages=Message(
role="assistant",
contents=[
Content.from_function_call(call_id="call_safe_b", name="safe_b", arguments="{}"),
Content.from_function_call(call_id="call_risky_b", name="risky_b", arguments="{}"),
],
)
)
]
second_response = await agent.run("batch b", session=session)
second_request = _approval_requests(second_response.messages)[0]
chat_client_base.run_responses = [ChatResponse(messages=Message(role="assistant", contents=["done"]))]
final_response = await agent.run(second_request.to_function_approval_response(approved=True), session=session)
assert final_response.text == "done"
assert safe_a_calls == 0
assert risky_a_calls == 0
assert safe_b_calls == 1
assert risky_b_calls == 1
async def test_tool_approval_middleware_queues_multiple_approval_requests(
chat_client_base: SupportsChatGetResponse,
) -> None:
"""The opt-in middleware should present multiple unresolved approvals one at a time."""
first_calls = 0
second_calls = 0
@tool(name="first_tool", approval_mode="always_require")
def first_tool() -> str:
nonlocal first_calls
first_calls += 1
return "first"
@tool(name="second_tool", approval_mode="always_require")
def second_tool() -> str:
nonlocal second_calls
second_calls += 1
return "second"
agent = Agent(
client=chat_client_base,
tools=[first_tool, second_tool],
middleware=[ToolApprovalMiddleware()],
)
session = AgentSession(session_id="queue-session")
chat_client_base.run_responses = [
ChatResponse(
messages=Message(
role="assistant",
contents=[
Content.from_function_call(call_id="call_first", name="first_tool", arguments="{}"),
Content.from_function_call(call_id="call_second", name="second_tool", arguments="{}"),
],
)
)
]
first_response = await agent.run("call both", session=session)
first_requests = _approval_requests(first_response.messages)
assert [request.function_call.name for request in first_requests] == ["first_tool"]
assert first_calls == 0
assert second_calls == 0
second_response = await agent.run(first_requests[0].to_function_approval_response(approved=True), session=session)
second_requests = _approval_requests(second_response.messages)
assert [request.function_call.name for request in second_requests] == ["second_tool"]
assert first_calls == 0
assert second_calls == 0
chat_client_base.run_responses = [ChatResponse(messages=Message(role="assistant", contents=["done"]))]
final_response = await agent.run(second_requests[0].to_function_approval_response(approved=True), session=session)
assert final_response.text == "done"
assert first_calls == 1
assert second_calls == 1
async def test_tool_approval_middleware_preserves_hidden_mixed_batch_requests(
chat_client_base: SupportsChatGetResponse,
) -> None:
"""Middleware state saves should not discard core hidden already-approved requests."""
lookup_calls = 0
write_calls = 0
@tool(name="lookup_records", approval_mode="never_require")
def lookup_records() -> str:
nonlocal lookup_calls
lookup_calls += 1
return "records"
@tool(name="write_record", approval_mode="always_require")
def write_record() -> str:
nonlocal write_calls
write_calls += 1
return "written"
agent = Agent(
client=chat_client_base,
tools=[lookup_records, write_record],
middleware=[ToolApprovalMiddleware()],
)
session = AgentSession(session_id="mixed-middleware-session")
chat_client_base.run_responses = [
ChatResponse(
messages=Message(
role="assistant",
contents=[
Content.from_function_call(call_id="call_lookup", name="lookup_records", arguments="{}"),
Content.from_function_call(call_id="call_write", name="write_record", arguments="{}"),
],
)
)
]
first_response = await agent.run("lookup and write", session=session)
request = _approval_requests(first_response.messages)[0]
chat_client_base.run_responses = [ChatResponse(messages=Message(role="assistant", contents=["done"]))]
second_response = await agent.run(request.to_function_approval_response(approved=True), session=session)
assert second_response.text == "done"
assert lookup_calls == 1
assert write_calls == 1
async def test_tool_approval_middleware_auto_approval_rule_receives_function_call(
chat_client_base: SupportsChatGetResponse,
) -> None:
"""Heuristic auto-approval callbacks should receive function-call content and approve matching calls."""
auto_calls = 0
manual_calls = 0
seen_calls: list[tuple[str, str | None]] = []
@tool(name="auto_write", approval_mode="always_require")
def auto_write() -> str:
nonlocal auto_calls
auto_calls += 1
return "auto"
@tool(name="manual_write", approval_mode="always_require")
def manual_write() -> str:
nonlocal manual_calls
manual_calls += 1
return "manual"
async def auto_approve_auto_write(function_call: Content) -> bool:
seen_calls.append((function_call.type, function_call.name))
return function_call.name == "auto_write"
agent = Agent(
client=chat_client_base,
tools=[auto_write, manual_write],
middleware=[ToolApprovalMiddleware(auto_approval_rules=[auto_approve_auto_write])],
)
session = AgentSession(session_id="heuristic-session")
chat_client_base.run_responses = [
ChatResponse(
messages=Message(
role="assistant",
contents=[
Content.from_function_call(call_id="call_auto", name="auto_write", arguments="{}"),
Content.from_function_call(call_id="call_manual", name="manual_write", arguments="{}"),
],
)
)
]
first_response = await agent.run("write both", session=session)
requests = _approval_requests(first_response.messages)
assert [request.function_call.name for request in requests] == ["manual_write"]
assert seen_calls == [("function_call", "auto_write"), ("function_call", "manual_write")]
assert auto_calls == 0
assert manual_calls == 0
chat_client_base.run_responses = [ChatResponse(messages=Message(role="assistant", contents=["done"]))]
final_response = await agent.run(requests[0].to_function_approval_response(approved=True), session=session)
assert final_response.text == "done"
assert auto_calls == 1
assert manual_calls == 1
async def test_tool_approval_middleware_auto_approved_loops_share_function_call_budget(
chat_client_base: SupportsChatGetResponse,
) -> None:
"""Auto-approved re-entry should not reset max_function_calls."""
calls = 0
@tool(name="budgeted_tool", approval_mode="always_require")
def budgeted_tool(value: str) -> str:
nonlocal calls
calls += 1
return value
def auto_approve_budgeted_tool(function_call: Content) -> bool:
return function_call.name == "budgeted_tool"
chat_client_base.function_invocation_configuration["max_function_calls"] = 1 # type: ignore[attr-defined]
agent = Agent(
client=chat_client_base,
tools=[budgeted_tool],
middleware=[ToolApprovalMiddleware(auto_approval_rules=[auto_approve_budgeted_tool])],
)
session = AgentSession(session_id="shared-budget-session")
chat_client_base.run_responses = [
ChatResponse(
messages=Message(
role="assistant",
contents=[
Content.from_function_call(
call_id="call_first",
name="budgeted_tool",
arguments='{"value": "first"}',
)
],
)
),
ChatResponse(
messages=Message(
role="assistant",
contents=[
Content.from_function_call(
call_id="call_second",
name="budgeted_tool",
arguments='{"value": "second"}',
)
],
)
),
]
response = await agent.run("call repeatedly", session=session)
assert response.text == "I broke out of the function invocation loop..."
assert calls == 1
async def test_tool_approval_middleware_queues_streamed_approval_requests(
chat_client_base: SupportsChatGetResponse,
) -> None:
"""Streaming approval requests should also be queued one at a time."""
calls = 0
@tool(name="first_streamed_tool", approval_mode="always_require")
def first_streamed_tool() -> str:
nonlocal calls
calls += 1
return "first"
@tool(name="second_streamed_tool", approval_mode="always_require")
def second_streamed_tool() -> str:
nonlocal calls
calls += 1
return "second"
agent = Agent(
client=chat_client_base,
tools=[first_streamed_tool, second_streamed_tool],
middleware=[ToolApprovalMiddleware()],
)
session = AgentSession(session_id="stream-queue-session")
chat_client_base.streaming_responses = [
[
ChatResponseUpdate(
contents=[Content.from_function_call(call_id="call_first", name="first_streamed_tool", arguments="{}")],
role="assistant",
),
ChatResponseUpdate(
contents=[
Content.from_function_call(call_id="call_second", name="second_streamed_tool", arguments="{}")
],
role="assistant",
),
]
]
first_stream = agent.run("call both", stream=True, session=session)
first_updates = [update async for update in first_stream]
first_requests = [content for update in first_updates for content in update.user_input_requests]
assert [request.function_call.name for request in first_requests] == ["first_streamed_tool"]
assert calls == 0
second_stream = agent.run(
first_requests[0].to_function_approval_response(approved=True),
stream=True,
session=session,
)
second_updates = [update async for update in second_stream]
second_requests = [content for update in second_updates for content in update.user_input_requests]
assert [request.function_call.name for request in second_requests] == ["second_streamed_tool"]
assert calls == 0
chat_client_base.streaming_responses = [
[ChatResponseUpdate(contents=[Content.from_text("done")], role="assistant")]
]
final_stream = agent.run(
second_requests[0].to_function_approval_response(approved=True),
stream=True,
session=session,
)
final_updates = [update async for update in final_stream]
final_response = await final_stream.get_final_response()
assert final_updates[-1].text == "done"
assert final_response.text == "done"
assert calls == 2
async def test_tool_approval_middleware_always_approve_tool_rule(
chat_client_base: SupportsChatGetResponse,
) -> None:
"""An always-approve response should add a standing tool-level approval rule."""
calls = 0
@tool(name="dangerous_tool", approval_mode="always_require")
def dangerous_tool(value: str) -> str:
nonlocal calls
calls += 1
return value
agent = Agent(
client=chat_client_base,
tools=[dangerous_tool],
middleware=[ToolApprovalMiddleware()],
)
session = AgentSession(session_id="standing-rule-session")
chat_client_base.run_responses = [
ChatResponse(
messages=Message(
role="assistant",
contents=[
Content.from_function_call(
call_id="call_initial",
name="dangerous_tool",
arguments='{"value": "one"}',
)
],
)
)
]
first_response = await agent.run("call once", session=session)
first_request = _approval_requests(first_response.messages)[0]
chat_client_base.run_responses = [ChatResponse(messages=Message(role="assistant", contents=["first done"]))]
await agent.run(create_always_approve_tool_response(first_request), session=session)
assert calls == 1
chat_client_base.run_responses = [
ChatResponse(
messages=Message(
role="assistant",
contents=[
Content.from_function_call(
call_id="call_auto",
name="dangerous_tool",
arguments='{"value": "two"}',
)
],
)
),
ChatResponse(messages=Message(role="assistant", contents=["second done"])),
]
second_response = await agent.run("call again", session=session)
assert second_response.text == "second done"
assert calls == 2
async def test_tool_approval_middleware_standing_rules_include_hosted_server_boundary(
chat_client_base: SupportsChatGetResponse,
) -> None:
"""A standing hosted-tool rule should only match the same server_label."""
calls = 0
@tool(name="hosted_tool", approval_mode="always_require")
def hosted_tool() -> str:
nonlocal calls
calls += 1
return "hosted"
def hosted_call(call_id: str, server_label: str) -> Content:
return Content.from_function_call(
call_id=call_id,
name="hosted_tool",
arguments="{}",
additional_properties={"server_label": server_label},
)
agent = Agent(
client=chat_client_base,
tools=[hosted_tool],
middleware=[ToolApprovalMiddleware()],
)
session = AgentSession(session_id="hosted-boundary-session")
chat_client_base.run_responses = [
ChatResponse(messages=Message(role="assistant", contents=[hosted_call("call_initial", "server-a")]))
]
first_response = await agent.run("call hosted a", session=session)
first_request = _approval_requests(first_response.messages)[0]
chat_client_base.run_responses = [ChatResponse(messages=Message(role="assistant", contents=["server a done"]))]
await agent.run(create_always_approve_tool_response(first_request), session=session)
assert calls == 0
chat_client_base.run_responses = [
ChatResponse(messages=Message(role="assistant", contents=[hosted_call("call_same_server", "server-a")])),
ChatResponse(messages=Message(role="assistant", contents=["same server done"])),
]
same_server_response = await agent.run("call hosted a again", session=session)
assert same_server_response.text == "same server done"
assert _approval_requests(same_server_response.messages) == []
assert calls == 0
chat_client_base.run_responses = [
ChatResponse(messages=Message(role="assistant", contents=[hosted_call("call_other_server", "server-b")]))
]
other_server_response = await agent.run("call hosted b", session=session)
requests = _approval_requests(other_server_response.messages)
assert [request.function_call.additional_properties["server_label"] for request in requests] == ["server-b"]
assert calls == 0
async def test_tool_approval_middleware_always_approve_tool_with_arguments_rule(
chat_client_base: SupportsChatGetResponse,
) -> None:
"""Argument-scoped always-approve rules should require exact argument matches."""
calls = 0
@tool(name="argument_scoped_tool", approval_mode="always_require")
def argument_scoped_tool(value: str) -> str:
nonlocal calls
calls += 1
return value
agent = Agent(
client=chat_client_base,
tools=[argument_scoped_tool],
middleware=[ToolApprovalMiddleware()],
)
session = AgentSession(session_id="argument-rule-session")
chat_client_base.run_responses = [
ChatResponse(
messages=Message(
role="assistant",
contents=[
Content.from_function_call(
call_id="call_initial",
name="argument_scoped_tool",
arguments='{"value": "same"}',
)
],
)
)
]
first_response = await agent.run("call with same", session=session)
first_request = _approval_requests(first_response.messages)[0]
chat_client_base.run_responses = [ChatResponse(messages=Message(role="assistant", contents=["first done"]))]
await agent.run(create_always_approve_tool_with_arguments_response(first_request), session=session)
assert calls == 1
chat_client_base.run_responses = [
ChatResponse(
messages=Message(
role="assistant",
contents=[
Content.from_function_call(
call_id="call_same",
name="argument_scoped_tool",
arguments='{"value": "same"}',
)
],
)
),
ChatResponse(messages=Message(role="assistant", contents=["same done"])),
]
second_response = await agent.run("call with same again", session=session)
assert second_response.text == "same done"
assert calls == 2
chat_client_base.run_responses = [
ChatResponse(
messages=Message(
role="assistant",
contents=[
Content.from_function_call(
call_id="call_different",
name="argument_scoped_tool",
arguments='{"value": "different"}',
)
],
)
)
]
third_response = await agent.run("call with different args", session=session)
requests = _approval_requests(third_response.messages)
assert [request.function_call.arguments for request in requests] == ['{"value": "different"}']
assert calls == 2
async def test_tool_approval_middleware_empty_arguments_rule_is_not_tool_wide(
chat_client_base: SupportsChatGetResponse,
) -> None:
"""An argument-scoped no-argument approval should not become a wildcard."""
calls = 0
@tool(name="optional_args_tool", approval_mode="always_require")
def optional_args_tool(value: str = "default") -> str:
nonlocal calls
calls += 1
return value
agent = Agent(
client=chat_client_base,
tools=[optional_args_tool],
middleware=[ToolApprovalMiddleware()],
)
session = AgentSession(session_id="empty-arguments-rule-session")
chat_client_base.run_responses = [
ChatResponse(
messages=Message(
role="assistant",
contents=[
Content.from_function_call(
call_id="call_empty",
name="optional_args_tool",
arguments="{}",
)
],
)
)
]
first_response = await agent.run("call without args", session=session)
first_request = _approval_requests(first_response.messages)[0]
chat_client_base.run_responses = [ChatResponse(messages=Message(role="assistant", contents=["empty done"]))]
await agent.run(create_always_approve_tool_with_arguments_response(first_request), session=session)
assert calls == 1
chat_client_base.run_responses = [
ChatResponse(
messages=Message(
role="assistant",
contents=[
Content.from_function_call(
call_id="call_non_empty",
name="optional_args_tool",
arguments='{"value": "custom"}',
)
],
)
)
]
second_response = await agent.run("call with args", session=session)
requests = _approval_requests(second_response.messages)
assert [request.function_call.arguments for request in requests] == ['{"value": "custom"}']
assert calls == 1
+277 -20
View File
@@ -342,6 +342,69 @@ def test_parse_tool_result_from_mcp_resource_link_text_resource_and_unknown():
assert result[1].text == "Embedded result"
def test_parse_tool_result_from_mcp_structured_content_only():
"""Test that structuredContent is parsed when content list is empty."""
mcp_result = types.CallToolResult(
content=[],
structuredContent={"Tables": [{"Name": "Sales", "Columns": ["Amount", "Date"]}]},
)
result = _HELPER_MCP_TOOL._parse_tool_result_from_mcp(mcp_result)
assert isinstance(result, list)
assert len(result) == 1
assert result[0].type == "text"
parsed = json.loads(result[0].text)
assert parsed == {"Tables": [{"Name": "Sales", "Columns": ["Amount", "Date"]}]}
def test_parse_tool_result_from_mcp_structured_content_with_text():
"""Test that structuredContent is appended alongside regular content items."""
mcp_result = types.CallToolResult(
content=[types.TextContent(type="text", text="Summary")],
structuredContent={"data": [1, 2, 3]},
)
result = _HELPER_MCP_TOOL._parse_tool_result_from_mcp(mcp_result)
assert isinstance(result, list)
assert len(result) == 2
assert result[0].type == "text"
assert result[0].text == "Summary"
assert result[1].type == "text"
parsed = json.loads(result[1].text)
assert parsed == {"data": [1, 2, 3]}
def test_parse_tool_result_from_mcp_structured_content_none():
"""Test that None structuredContent does not affect results."""
mcp_result = types.CallToolResult(
content=[types.TextContent(type="text", text="Hello")],
structuredContent=None,
)
result = _HELPER_MCP_TOOL._parse_tool_result_from_mcp(mcp_result)
assert isinstance(result, list)
assert len(result) == 1
assert result[0].type == "text"
assert result[0].text == "Hello"
def test_parse_tool_result_from_mcp_structured_content_non_serializable():
"""Test that non-JSON-serializable values in structuredContent degrade gracefully."""
mcp_result = types.CallToolResult(
content=[],
structuredContent={"data": b"raw bytes", "count": 42},
)
result = _HELPER_MCP_TOOL._parse_tool_result_from_mcp(mcp_result)
assert isinstance(result, list)
assert len(result) == 1
assert result[0].type == "text"
parsed = json.loads(result[0].text)
assert parsed["count"] == 42
# bytes should be converted to string representation via default=str
assert "raw bytes" in parsed["data"]
def test_mcp_content_types_to_ai_content_text():
"""Test conversion of MCP text content to AI content."""
mcp_content = types.TextContent(type="text", text="Sample text")
@@ -1467,6 +1530,7 @@ def test_mcp_tool_approval_mode_returns_none_for_unmatched_names() -> None:
3,
["tool_one", "tool_two", "tool_three"],
), # None means all tools are allowed
([], 0, []), # Empty list means no tools are allowed
(["tool_one"], 1, ["tool_one"]), # Only tool_one is allowed
(
["tool_one", "tool_three"],
@@ -1813,6 +1877,18 @@ async def test_mcp_tool_message_handler_cancel_and_replace():
assert len(tool._pending_reload_tasks) == 0
def _approve(_params: object) -> bool:
"""Approving sampling gate used by tests that exercise forwarding behavior."""
return True
def _make_sampling_response(text: str = "response", model: str = "test-model") -> Mock:
mock_response = Mock()
mock_response.messages = [Message(role="assistant", contents=[Content.from_text(text)])]
mock_response.model = model
return mock_response
async def test_mcp_tool_sampling_callback_no_client():
"""Test sampling callback error path when no chat client is available."""
tool = MCPStdioTool(name="test_tool", command="python")
@@ -1828,9 +1904,190 @@ async def test_mcp_tool_sampling_callback_no_client():
assert "No chat client available" in result.message
async def test_mcp_tool_sampling_callback_denies_by_default():
"""Sampling is denied when no approval callback is configured (safe default)."""
tool = MCPStdioTool(name="test_tool", command="python")
mock_chat_client = AsyncMock()
tool.client = mock_chat_client
params = Mock()
params.messages = []
params.maxTokens = 128
result = await tool.sampling_callback(Mock(), params)
assert isinstance(result, types.ErrorData)
assert result.code == types.INVALID_REQUEST
assert "denied" in result.message
assert "sampling_approval_callback" in result.message
mock_chat_client.get_response.assert_not_called()
async def test_mcp_tool_sampling_callback_denied_by_callback():
"""Sampling is denied when the approval callback returns a falsy value."""
tool = MCPStdioTool(name="test_tool", command="python", sampling_approval_callback=lambda params: False)
mock_chat_client = AsyncMock()
tool.client = mock_chat_client
params = Mock()
params.messages = []
params.maxTokens = 128
result = await tool.sampling_callback(Mock(), params)
assert isinstance(result, types.ErrorData)
assert result.code == types.INVALID_REQUEST
assert "denied by the 'sampling_approval_callback'" in result.message
mock_chat_client.get_response.assert_not_called()
async def test_mcp_tool_sampling_callback_callback_exception_denies():
"""An approval callback that raises results in denial, not an LLM call."""
def boom(_params: object) -> bool:
raise RuntimeError("approval error")
tool = MCPStdioTool(name="test_tool", command="python", sampling_approval_callback=boom)
mock_chat_client = AsyncMock()
tool.client = mock_chat_client
params = Mock()
params.messages = []
params.maxTokens = 128
result = await tool.sampling_callback(Mock(), params)
assert isinstance(result, types.ErrorData)
assert result.code == types.INVALID_REQUEST
mock_chat_client.get_response.assert_not_called()
async def test_mcp_tool_sampling_callback_async_approval():
"""An async approval callback that approves allows the request through."""
async def approve(_params: object) -> bool:
return True
tool = MCPStdioTool(name="test_tool", command="python", sampling_approval_callback=approve)
mock_chat_client = AsyncMock()
mock_chat_client.get_response.return_value = _make_sampling_response("ok")
tool.client = mock_chat_client
params = Mock()
params.messages = [types.PromptMessage(role="user", content=types.TextContent(type="text", text="Hi"))]
params.temperature = None
params.maxTokens = 100
params.stopSequences = None
params.systemPrompt = None
params.tools = None
params.toolChoice = None
result = await tool.sampling_callback(Mock(), params)
assert isinstance(result, types.CreateMessageResult)
assert result.content.text == "ok"
mock_chat_client.get_response.assert_awaited_once()
async def test_mcp_tool_sampling_callback_clamps_max_tokens():
"""An approved request's maxTokens is clamped to sampling_max_tokens."""
tool = MCPStdioTool(
name="test_tool",
command="python",
sampling_approval_callback=_approve,
sampling_max_tokens=512,
)
mock_chat_client = AsyncMock()
mock_chat_client.get_response.return_value = _make_sampling_response()
tool.client = mock_chat_client
params = Mock()
params.messages = [types.PromptMessage(role="user", content=types.TextContent(type="text", text="Hi"))]
params.temperature = None
params.maxTokens = 1_000_000
params.stopSequences = None
params.systemPrompt = None
params.tools = None
params.toolChoice = None
result = await tool.sampling_callback(Mock(), params)
assert isinstance(result, types.CreateMessageResult)
options = mock_chat_client.get_response.call_args.kwargs.get("options") or {}
assert options["max_tokens"] == 512
async def test_mcp_tool_sampling_callback_does_not_clamp_under_cap():
"""A request below the cap keeps its requested maxTokens."""
tool = MCPStdioTool(
name="test_tool",
command="python",
sampling_approval_callback=_approve,
sampling_max_tokens=512,
)
mock_chat_client = AsyncMock()
mock_chat_client.get_response.return_value = _make_sampling_response()
tool.client = mock_chat_client
params = Mock()
params.messages = [types.PromptMessage(role="user", content=types.TextContent(type="text", text="Hi"))]
params.temperature = None
params.maxTokens = 100
params.stopSequences = None
params.systemPrompt = None
params.tools = None
params.toolChoice = None
result = await tool.sampling_callback(Mock(), params)
assert isinstance(result, types.CreateMessageResult)
options = mock_chat_client.get_response.call_args.kwargs.get("options") or {}
assert options["max_tokens"] == 100
async def test_mcp_tool_sampling_callback_rate_limited():
"""Sampling requests beyond sampling_max_requests are rejected per session."""
tool = MCPStdioTool(
name="test_tool",
command="python",
sampling_approval_callback=_approve,
sampling_max_requests=2,
)
mock_chat_client = AsyncMock()
mock_chat_client.get_response.return_value = _make_sampling_response()
tool.client = mock_chat_client
def make_params() -> Mock:
params = Mock()
params.messages = [types.PromptMessage(role="user", content=types.TextContent(type="text", text="Hi"))]
params.temperature = None
params.maxTokens = 100
params.stopSequences = None
params.systemPrompt = None
params.tools = None
params.toolChoice = None
return params
first = await tool.sampling_callback(Mock(), make_params())
second = await tool.sampling_callback(Mock(), make_params())
third = await tool.sampling_callback(Mock(), make_params())
assert isinstance(first, types.CreateMessageResult)
assert isinstance(second, types.CreateMessageResult)
assert isinstance(third, types.ErrorData)
assert third.code == types.INVALID_REQUEST
assert "rate limit" in third.message.lower()
assert mock_chat_client.get_response.await_count == 2
# The counter resets on a session reset.
tool._reset_session_state()
fourth = await tool.sampling_callback(Mock(), make_params())
assert isinstance(fourth, types.CreateMessageResult)
async def test_mcp_tool_sampling_callback_chat_client_exception():
"""Test sampling callback when chat client raises exception."""
tool = MCPStdioTool(name="test_tool", command="python")
tool = MCPStdioTool(name="test_tool", command="python", sampling_approval_callback=_approve)
# Mock chat client that raises exception
mock_chat_client = AsyncMock()
@@ -1846,7 +2103,7 @@ async def test_mcp_tool_sampling_callback_chat_client_exception():
mock_message.content.text = "Test question"
params.messages = [mock_message]
params.temperature = None
params.maxTokens = None
params.maxTokens = 100
params.stopSequences = None
params.systemPrompt = None
params.tools = None
@@ -1863,7 +2120,7 @@ async def test_mcp_tool_sampling_callback_no_valid_content():
"""Test sampling callback when response has no valid content types."""
from agent_framework import Message
tool = MCPStdioTool(name="test_tool", command="python")
tool = MCPStdioTool(name="test_tool", command="python", sampling_approval_callback=_approve)
# Mock chat client with response containing only invalid content types
mock_chat_client = AsyncMock()
@@ -1892,7 +2149,7 @@ async def test_mcp_tool_sampling_callback_no_valid_content():
mock_message.content.text = "Test question"
params.messages = [mock_message]
params.temperature = None
params.maxTokens = None
params.maxTokens = 100
params.stopSequences = None
params.systemPrompt = None
params.tools = None
@@ -1905,18 +2162,18 @@ async def test_mcp_tool_sampling_callback_no_valid_content():
assert "Failed to get right content types from the response." in result.message
mock_chat_client.get_response.assert_awaited_once()
_, kwargs = mock_chat_client.get_response.await_args
assert kwargs["options"] == {"max_tokens": None}
assert kwargs["options"] == {"max_tokens": 100}
async def test_mcp_tool_sampling_callback_no_response_and_successful_message_creation():
"""Test sampling callback when the chat client returns no response and then valid content."""
tool = MCPStdioTool(name="test_tool", command="python")
tool = MCPStdioTool(name="test_tool", command="python", sampling_approval_callback=_approve)
tool.client = AsyncMock()
params = Mock()
params.messages = [types.PromptMessage(role="user", content=types.TextContent(type="text", text="Hi"))]
params.temperature = None
params.maxTokens = None
params.maxTokens = 100
params.stopSequences = None
params.systemPrompt = None
params.tools = None
@@ -1955,7 +2212,7 @@ async def test_mcp_tool_sampling_callback_forwards_system_prompt():
"""Test sampling callback passes systemPrompt as instructions in options."""
from agent_framework import Message
tool = MCPStdioTool(name="test_tool", command="python")
tool = MCPStdioTool(name="test_tool", command="python", sampling_approval_callback=_approve)
mock_chat_client = AsyncMock()
mock_response = Mock()
@@ -1972,7 +2229,7 @@ async def test_mcp_tool_sampling_callback_forwards_system_prompt():
mock_message.content.text = "Test question"
params.messages = [mock_message]
params.temperature = None
params.maxTokens = None
params.maxTokens = 100
params.stopSequences = None
params.systemPrompt = "You are a helpful assistant"
params.tools = None
@@ -1990,7 +2247,7 @@ async def test_mcp_tool_sampling_callback_forwards_tools():
"""Test sampling callback converts MCP tools to FunctionTools and passes them in options."""
from agent_framework import FunctionTool, Message
tool = MCPStdioTool(name="test_tool", command="python")
tool = MCPStdioTool(name="test_tool", command="python", sampling_approval_callback=_approve)
mock_chat_client = AsyncMock()
mock_response = Mock()
@@ -2013,7 +2270,7 @@ async def test_mcp_tool_sampling_callback_forwards_tools():
mock_message.content.text = "Test question"
params.messages = [mock_message]
params.temperature = None
params.maxTokens = None
params.maxTokens = 100
params.stopSequences = None
params.systemPrompt = None
params.tools = [mcp_tool]
@@ -2036,7 +2293,7 @@ async def test_mcp_tool_sampling_callback_forwards_tool_choice():
"""Test sampling callback passes toolChoice mode in options."""
from agent_framework import Message
tool = MCPStdioTool(name="test_tool", command="python")
tool = MCPStdioTool(name="test_tool", command="python", sampling_approval_callback=_approve)
mock_chat_client = AsyncMock()
mock_response = Mock()
@@ -2053,7 +2310,7 @@ async def test_mcp_tool_sampling_callback_forwards_tool_choice():
mock_message.content.text = "Test question"
params.messages = [mock_message]
params.temperature = None
params.maxTokens = None
params.maxTokens = 100
params.stopSequences = None
params.systemPrompt = None
params.tools = None
@@ -2071,7 +2328,7 @@ async def test_mcp_tool_sampling_callback_forwards_empty_system_prompt():
"""Test sampling callback forwards empty string systemPrompt as instructions."""
from agent_framework import Message
tool = MCPStdioTool(name="test_tool", command="python")
tool = MCPStdioTool(name="test_tool", command="python", sampling_approval_callback=_approve)
mock_chat_client = AsyncMock()
mock_response = Mock()
@@ -2088,7 +2345,7 @@ async def test_mcp_tool_sampling_callback_forwards_empty_system_prompt():
mock_message.content.text = "Test question"
params.messages = [mock_message]
params.temperature = None
params.maxTokens = None
params.maxTokens = 100
params.stopSequences = None
params.systemPrompt = ""
params.tools = None
@@ -2106,7 +2363,7 @@ async def test_mcp_tool_sampling_callback_forwards_empty_tools_list():
"""Test sampling callback forwards empty tools list in options."""
from agent_framework import Message
tool = MCPStdioTool(name="test_tool", command="python")
tool = MCPStdioTool(name="test_tool", command="python", sampling_approval_callback=_approve)
mock_chat_client = AsyncMock()
mock_response = Mock()
@@ -2123,7 +2380,7 @@ async def test_mcp_tool_sampling_callback_forwards_empty_tools_list():
mock_message.content.text = "Test question"
params.messages = [mock_message]
params.temperature = None
params.maxTokens = None
params.maxTokens = 100
params.stopSequences = None
params.systemPrompt = None
params.tools = []
@@ -2141,7 +2398,7 @@ async def test_mcp_tool_sampling_callback_forwards_generation_params_in_options(
"""Test sampling callback passes temperature, max_tokens, and stop in options."""
from agent_framework import Message
tool = MCPStdioTool(name="test_tool", command="python")
tool = MCPStdioTool(name="test_tool", command="python", sampling_approval_callback=_approve)
mock_chat_client = AsyncMock()
mock_response = Mock()
@@ -2182,7 +2439,7 @@ async def test_mcp_tool_sampling_callback_omits_temperature_when_none():
"""Test sampling callback does not set temperature in options when it is None."""
from agent_framework import Message
tool = MCPStdioTool(name="test_tool", command="python")
tool = MCPStdioTool(name="test_tool", command="python", sampling_approval_callback=_approve)
mock_chat_client = AsyncMock()
mock_response = Mock()
@@ -2219,7 +2476,7 @@ async def test_mcp_tool_sampling_callback_always_passes_max_tokens():
"""Test sampling callback always sets max_tokens in options since maxTokens is a required int field."""
from agent_framework import Message
tool = MCPStdioTool(name="test_tool", command="python")
tool = MCPStdioTool(name="test_tool", command="python", sampling_approval_callback=_approve)
mock_chat_client = AsyncMock()
mock_response = Mock()
@@ -76,6 +76,7 @@ def _make_call_tool_result(text: str = "result", is_error: bool = False) -> Mock
result = Mock()
result.isError = is_error
result.content = [types.TextContent(type="text", text=text)]
result.structuredContent = None
return result