Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,8 @@

import logging
import os
from typing import TYPE_CHECKING
from contextlib import _AsyncGeneratorContextManager # pyright: ignore[reportPrivateUsage]
from typing import TYPE_CHECKING, Any
from urllib.parse import urlsplit

import httpx
Expand All @@ -18,6 +19,7 @@
SkillsSourceContext,
)
from azure.ai.agentserver.core import get_request_context
from typing_extensions import override

if TYPE_CHECKING:
from collections.abc import Generator
Expand Down Expand Up @@ -174,6 +176,9 @@ def __init__(
auth=_ToolboxAuth(credential, token_scope),
timeout=timeout,
)
self._credential = credential
self._token_scope = token_scope
self._timeout = timeout

super().__init__(
name=tool_name,
Expand All @@ -183,8 +188,22 @@ def __init__(
load_tools=load_tools,
)

@override
def get_mcp_client(self) -> _AsyncGeneratorContextManager[Any, None]:
"""Get an authenticated MCP HTTP client.

Recreates the underlying HTTP client if it was previously closed.
"""
if self._httpx_client is None:
self._httpx_client = httpx.AsyncClient(
auth=_ToolboxAuth(self._credential, self._token_scope),
timeout=self._timeout,
)
return super().get_mcp_client()

@override
async def close(self) -> None:
"""Close the MCP session and the toolbox-owned HTTP client."""
"""Close the MCP session and toolbox HTTP client while preserving credentials and timeout for reconnection."""
try:
await super().close()
finally:
Expand Down
88 changes: 74 additions & 14 deletions python/packages/foundry_hosting/tests/test_toolbox.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
# Copyright (c) Microsoft. All rights reserved.
# pyright: reportPrivateUsage=false

"""Unit tests for FoundryToolbox."""

Expand All @@ -20,7 +21,7 @@
)

from agent_framework_foundry_hosting import FoundryToolbox
from agent_framework_foundry_hosting._toolbox import ( # pyright: ignore[reportPrivateUsage]
from agent_framework_foundry_hosting._toolbox import (
_FoundryToolboxSkillsSource,
_resolve_toolbox_endpoint,
_toolbox_name_from_endpoint,
Expand Down Expand Up @@ -145,9 +146,9 @@ async def test_close_closes_owned_http_client() -> None:
_FakeCredential(), # type: ignore
url="https://h/toolboxes/tb/mcp",
)
client = toolbox._httpx_client # pyright: ignore[reportPrivateUsage]
client = toolbox._httpx_client
assert client is not None
client.aclose = AsyncMock() # type: ignore[method-assign]
client.aclose = AsyncMock() # zuban: ignore

await toolbox.close()

Expand All @@ -174,9 +175,9 @@ def test_as_skills_provider_requires_approval_by_default() -> None:
)
provider = toolbox.as_skills_provider()
# By default every skill tool keeps its approval requirement.
assert provider._disable_load_skill_approval is False # pyright: ignore[reportPrivateUsage]
assert provider._disable_read_skill_resource_approval is False # pyright: ignore[reportPrivateUsage]
assert provider._disable_run_skill_script_approval is False # pyright: ignore[reportPrivateUsage]
assert provider._disable_load_skill_approval is False
assert provider._disable_read_skill_resource_approval is False
assert provider._disable_run_skill_script_approval is False


def test_as_skills_provider_forwards_approval_overrides() -> None:
Expand All @@ -191,9 +192,9 @@ def test_as_skills_provider_forwards_approval_overrides() -> None:
)
# Overrides flow through to the underlying SkillsProvider so an unattended
# host (no AgentSession) can load skills without an approval round-trip.
assert provider._disable_load_skill_approval is True # pyright: ignore[reportPrivateUsage]
assert provider._disable_read_skill_resource_approval is True # pyright: ignore[reportPrivateUsage]
assert provider._disable_run_skill_script_approval is True # pyright: ignore[reportPrivateUsage]
assert provider._disable_load_skill_approval is True
assert provider._disable_read_skill_resource_approval is True
assert provider._disable_run_skill_script_approval is True


async def test_skills_source_requires_connection() -> None:
Expand Down Expand Up @@ -248,9 +249,9 @@ async def test_skills_source_requires_connection_via_provider() -> None:
source = _FoundryToolboxSkillsSource(toolbox)
# Discovery captures the bound provider; a later reconnect gap (session is None)
# surfaces the same clear error when the provider is resolved.
toolbox.session = None # type: ignore
toolbox.session = None
with pytest.raises(RuntimeError, match="not connected"):
source._require_session() # pyright: ignore[reportPrivateUsage]
source._require_session()


class _FakeSkill:
Expand Down Expand Up @@ -291,7 +292,7 @@ async def test_as_skills_provider_caches_by_default(monkeypatch: pytest.MonkeyPa
provider = toolbox.as_skills_provider()
context = _source_context()
for _ in range(3):
await provider._source.get_skills(context) # pyright: ignore[reportPrivateUsage]
await provider._source.get_skills(context)

# By default the toolbox index is read once and reused across agent runs.
assert read_count[0] == 1
Expand All @@ -308,7 +309,7 @@ async def test_as_skills_provider_disable_caching_rereads_every_run(monkeypatch:
provider = toolbox.as_skills_provider(disable_caching=True)
context = _source_context()
for _ in range(3):
await provider._source.get_skills(context) # pyright: ignore[reportPrivateUsage]
await provider._source.get_skills(context)

# With caching disabled the index is re-read on every agent run.
assert read_count[0] == 3
Expand All @@ -329,6 +330,65 @@ async def test_as_skills_provider_cache_refresh_interval_rereads_after_staleness
provider = toolbox.as_skills_provider(cache_refresh_interval=timedelta(0))
context = _source_context()
for _ in range(3):
await provider._source.get_skills(context) # pyright: ignore[reportPrivateUsage]
await provider._source.get_skills(context)

assert read_count[0] == 3


class TestFoundryToolboxReconnection:
async def test_close_preserves_credential_for_reconnection(self) -> None:
"""After close(), get_mcp_client() should recreate an authenticated client."""
cred = _FakeCredential("reconnect-token")
toolbox = FoundryToolbox(
cred, # type: ignore
url="https://h/toolboxes/recon/mcp",
timeout=60.0,
)

assert toolbox._credential is cred
assert toolbox._token_scope == "https://ai.azure.com/.default"
assert toolbox._timeout == 60.0

assert toolbox._httpx_client is not None
assert isinstance(toolbox._httpx_client.auth, _ToolboxAuth)
original_auth = toolbox._httpx_client.auth

client = toolbox._httpx_client
client.aclose = AsyncMock() # zuban: ignore
await toolbox.close()

client.aclose.assert_awaited_once()
assert toolbox._httpx_client is None

assert toolbox._credential is cred
assert toolbox._timeout == 60.0

ctx_manager = toolbox.get_mcp_client()
assert toolbox._httpx_client is not None
assert isinstance(toolbox._httpx_client.auth, _ToolboxAuth)

new_auth = toolbox._httpx_client.auth
assert new_auth is not original_auth
assert new_auth._credential is cred

assert hasattr(ctx_manager, "__aenter__")
assert hasattr(ctx_manager, "__aexit__")

await toolbox.close()

async def test_close_idempotent_with_reconnection(self) -> None:
"""Multiple close() calls don't break reconnection."""
cred = _FakeCredential()
toolbox = FoundryToolbox(
cred, # type: ignore
url="https://h/toolboxes/idem/mcp",
)

await toolbox.close()
await toolbox.close()

toolbox.get_mcp_client()
assert toolbox._httpx_client is not None
assert isinstance(toolbox._httpx_client.auth, _ToolboxAuth)

await toolbox.close()
Loading