Skip to content
Draft
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
4 changes: 4 additions & 0 deletions python/packages/core/AGENTS.md
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,10 @@ agent_framework/
- **`RawAgent.open()` / `close()`** (inherited by `Agent`) - Explicitly enter the client's and configured MCP tools'
async contexts, then release them along with lazily connected MCP tools. `async with agent` delegates to these
methods; a partially failed `open()` closes resources already entered.
- Providers that store service-owned continuation handles in `AgentSession.state` declare their keys through
`service_session_state_keys`. `as_tool(propagate_session=True)` combines declarations from the parent and child
agents and their clients to isolate those handles in both directions while ordinary application-owned state
propagates. Parent declarations travel through private function-invocation metadata, not tool arguments.

### Chat Clients (`_clients.py`)

Expand Down
45 changes: 40 additions & 5 deletions python/packages/core/agent_framework/_agents.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,8 @@
from typing_extensions import Self, TypedDict # pragma: no cover

if TYPE_CHECKING:
from collections.abc import Collection
Comment thread
rogerbarreto marked this conversation as resolved.

from mcp import types
from mcp.server.lowlevel import Server
from pydantic import BaseModel
Expand Down Expand Up @@ -110,6 +112,17 @@ def _tool_approval_source_ids(middleware: Sequence[MiddlewareTypes] | None) -> f
return frozenset(item.source_id for item in middleware or () if isinstance(item, ToolApprovalMiddleware))


def _provider_service_session_state_keys(agent: object) -> frozenset[str]:
"""Return provider-owned session-state keys declared by an agent or its client."""
keys: set[str] = set()
# A generic Agent can use a provider client that persists its own continuation state.
for owner in (agent, getattr(agent, "client", None)):
declared = getattr(owner, "service_session_state_keys", ())
if isinstance(declared, (list, tuple, set, frozenset)):
keys.update(key for key in cast("Collection[Any]", declared) if isinstance(key, str))
return frozenset(keys)


def _merge_delegated_session_state(
parent_state: MutableMapping[str, Any],
initial_child_state: Mapping[str, Any],
Expand Down Expand Up @@ -662,9 +675,10 @@ def as_tool(
propagate_session: If True, the parent agent's session is forwarded
to this sub-agent's ``run()`` call. Application-state changes
propagate back to the parent, while framework approval continuation
state remains isolated. Defaults to False. The sub-agent always
receives an AgentSession so session-backed middleware can run.
When False, that session is private to this invocation.
state and provider-owned service session state remain isolated.
Defaults to False. The sub-agent always receives an AgentSession
so session-backed middleware can run. When False, that session is
private to this invocation.

Returns:
A FunctionTool that can be used as a tool by other agents.
Expand Down Expand Up @@ -730,10 +744,20 @@ async def _agent_wrapper(ctx: FunctionInvocationContext, **kwargs: Any) -> str:
session = AgentSession()
child_approval_source_ids = _tool_approval_source_ids(self.middleware)
parent_approval_source_ids: frozenset[str] = frozenset()
parent_service_session_state_keys: frozenset[str] = frozenset()

if propagate_session and parent_session is not None:
from ._tools import _PARENT_TOOL_APPROVAL_SOURCE_IDS_CONTEXT_KEY # pyright: ignore[reportPrivateUsage]
from ._tools import (
_PARENT_SERVICE_SESSION_STATE_KEYS_CONTEXT_KEY, # pyright: ignore[reportPrivateUsage]
_PARENT_TOOL_APPROVAL_SOURCE_IDS_CONTEXT_KEY, # pyright: ignore[reportPrivateUsage]
)

raw_parent_service_keys = ctx.metadata.get(_PARENT_SERVICE_SESSION_STATE_KEYS_CONTEXT_KEY)
parent_service_session_state_keys = (
cast("frozenset[str]", raw_parent_service_keys)
if isinstance(raw_parent_service_keys, frozenset)
else frozenset()
)
raw_parent_approval_source_ids = ctx.metadata.get(_PARENT_TOOL_APPROVAL_SOURCE_IDS_CONTEXT_KEY)
parent_approval_source_ids = (
cast("frozenset[str]", raw_parent_approval_source_ids)
Expand Down Expand Up @@ -767,6 +791,10 @@ async def _agent_wrapper(ctx: FunctionInvocationContext, **kwargs: Any) -> str:
_TOOL_APPROVAL_STATE_KEY,
_FUNCTION_INVOCATION_BUDGET_STATE_KEY,
_FUNCTION_RESULT_PAYLOAD_BUDGET_STATE_KEY,
# Service handles belong to one agent's remote resources, not shared application state.
# Exclude them in both directions so later children cannot inherit an earlier child's handles.
*_provider_service_session_state_keys(self),
Comment thread
rogerbarreto marked this conversation as resolved.
Comment thread
rogerbarreto marked this conversation as resolved.
*parent_service_session_state_keys,
*child_approval_source_ids,
*parent_approval_source_ids,
})
Expand Down Expand Up @@ -1586,7 +1614,10 @@ async def _prepare_run_context(

agent_name = self._get_agent_name()
from ._mcp import MCPTool
from ._tools import _PARENT_TOOL_APPROVAL_SOURCE_IDS_CONTEXT_KEY # pyright: ignore[reportPrivateUsage]
from ._tools import (
_PARENT_SERVICE_SESSION_STATE_KEYS_CONTEXT_KEY, # pyright: ignore[reportPrivateUsage]
_PARENT_TOOL_APPROVAL_SOURCE_IDS_CONTEXT_KEY, # pyright: ignore[reportPrivateUsage]
)

base_tools = _normalize_tools(chat_options.pop("tools", None))
mcp_duplicate_message = "Tool names must be unique. Consider setting `tool_name_prefix` on the MCPTool."
Expand Down Expand Up @@ -1634,6 +1665,10 @@ async def _prepare_run_context(
additional_function_arguments[_PARENT_TOOL_APPROVAL_SOURCE_IDS_CONTEXT_KEY] = _tool_approval_source_ids(
self.middleware
)
# Recompute ownership for this invoking agent; caller kwargs must not replace its declarations.
additional_function_arguments[_PARENT_SERVICE_SESSION_STATE_KEYS_CONTEXT_KEY] = (
_provider_service_session_state_keys(self)
)

model = opts.pop("model", None)

Expand Down
12 changes: 12 additions & 0 deletions python/packages/core/agent_framework/_tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,6 +107,7 @@ def _generate_function_call_occurrence_id() -> str:
_TOOL_APPROVAL_STATE_KEY: Final[str] = "tool_approval"
_APPROVAL_SESSION_IS_AUTHORITATIVE_KEY: Final[str] = "_approval_session_is_authoritative"
_PARENT_TOOL_APPROVAL_SOURCE_IDS_CONTEXT_KEY: Final[str] = "_parent_tool_approval_source_ids"
_PARENT_SERVICE_SESSION_STATE_KEYS_CONTEXT_KEY: Final[str] = "_parent_service_session_state_keys"


def _has_authoritative_approval_session(invocation_session: AgentSession | None) -> bool:
Expand Down Expand Up @@ -2095,6 +2096,7 @@ async def _auto_invoke_function(
"middleware",
"conversation_id",
_PARENT_TOOL_APPROVAL_SOURCE_IDS_CONTEXT_KEY,
_PARENT_SERVICE_SESSION_STATE_KEYS_CONTEXT_KEY,
}
}
raw_parent_approval_source_ids = (custom_args or {}).get(_PARENT_TOOL_APPROVAL_SOURCE_IDS_CONTEXT_KEY)
Expand All @@ -2104,6 +2106,12 @@ async def _auto_invoke_function(
if isinstance(raw_parent_approval_source_ids, frozenset)
else frozenset()
)
raw_parent_service_keys = (custom_args or {}).get(_PARENT_SERVICE_SESSION_STATE_KEYS_CONTEXT_KEY)
parent_service_session_state_keys: frozenset[str] = (
cast("frozenset[str]", raw_parent_service_keys)
if isinstance(raw_parent_service_keys, frozenset)
else frozenset()
)
if invocation_session is not None:
runtime_kwargs["session"] = invocation_session
args = dict(parsed_args)
Expand All @@ -2123,6 +2131,9 @@ async def _auto_invoke_function(
tools=live_tools,
)
direct_context.metadata[_PARENT_TOOL_APPROVAL_SOURCE_IDS_CONTEXT_KEY] = parent_approval_source_ids
direct_context.metadata[_PARENT_SERVICE_SESSION_STATE_KEYS_CONTEXT_KEY] = (
parent_service_session_state_keys
)
if host_payload_budget is not None:
direct_context.metadata[_FUNCTION_RESULT_PAYLOAD_BUDGET_CONTEXT_KEY] = host_payload_budget
function_result = await tool.invoke(
Expand Down Expand Up @@ -2162,6 +2173,7 @@ async def _auto_invoke_function(
tools=live_tools,
)
middleware_context.metadata[_PARENT_TOOL_APPROVAL_SOURCE_IDS_CONTEXT_KEY] = parent_approval_source_ids
middleware_context.metadata[_PARENT_SERVICE_SESSION_STATE_KEYS_CONTEXT_KEY] = parent_service_session_state_keys
if host_payload_budget is not None:
middleware_context.metadata[_FUNCTION_RESULT_PAYLOAD_BUDGET_CONTEXT_KEY] = host_payload_budget
middleware_context.metadata[_AUTO_ARGUMENT_PREPARATION_CONTEXT_KEY] = True
Expand Down
156 changes: 156 additions & 0 deletions python/packages/core/tests/core/test_agents.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,7 @@
from agent_framework._agents import _get_tool_name, _merge_options, _sanitize_agent_name
from agent_framework._mcp import MCPTool, _build_prefixed_mcp_name, _normalize_mcp_name
from agent_framework._middleware import FunctionInvocationContext
from agent_framework._tools import _PARENT_SERVICE_SESSION_STATE_KEYS_CONTEXT_KEY
from agent_framework.exceptions import (
AgentInvalidRequestException,
ChatClientInvalidResponseException,
Expand Down Expand Up @@ -2585,6 +2586,161 @@ def capturing_run(*args: Any, **kwargs: Any) -> Any:
assert parent_session.state["counter"] == 1


@pytest.mark.parametrize("parent_handle", [None, "parent-provider-session"])
async def test_chat_agent_as_tool_propagate_session_isolates_provider_owned_state(
client: SupportsChatGetResponse, parent_handle: str | None, monkeypatch: pytest.MonkeyPatch
) -> None:
"""Test that provider-owned state stays isolated across sequential delegated calls."""

class ProviderStateAgent(Agent):
service_session_state_keys = frozenset({"provider_session"})

monkeypatch.setattr(client, "service_session_state_keys", frozenset({"client_provider_session"}), raising=False)
parent_session = AgentSession(session_id="shared-session", service_session_id="parent-service-session")
parent_session.state.update({
"counter": 0,
"ordinary": "parent-value",
})
if parent_handle is not None:
parent_session.state["provider_session"] = parent_handle
parent_session.state["client_provider_session"] = parent_handle

observed_child_states: list[dict[str, Any]] = []
agents = [
ProviderStateAgent(client=client, name="FirstChild"),
ProviderStateAgent(client=client, name="SecondChild"),
]

for agent in (agents[0], agents[1], agents[0]):
original_run = agent.run

def capturing_run(*args: Any, _run: Callable[..., Any] = original_run, **kwargs: Any) -> Any:
child_session = cast(AgentSession, kwargs["session"])
observed_child_states.append(dict(child_session.state))
assert child_session.service_session_id is None
child_session.state["counter"] += 1
child_session.state["ordinary"] = f"child-value-{len(observed_child_states)}"
child_session.state["provider_session"] = f"child-provider-session-{len(observed_child_states)}"
child_session.state["client_provider_session"] = f"client-provider-session-{len(observed_child_states)}"
child_session.service_session_id = f"child-service-session-{len(observed_child_states)}"
return _run(*args, **kwargs)

delegated_tool = agent.as_tool(propagate_session=True)
with patch.object(agent, "run", side_effect=capturing_run):
await delegated_tool.invoke(
context=FunctionInvocationContext(
function=delegated_tool,
arguments={"task": "Run child"},
session=parent_session,
)
)

assert [state.get("provider_session") for state in observed_child_states] == [None, None, None]
assert [state.get("client_provider_session") for state in observed_child_states] == [None, None, None]
assert [state["counter"] for state in observed_child_states] == [0, 1, 2]
assert [state["ordinary"] for state in observed_child_states] == ["parent-value", "child-value-1", "child-value-2"]
expected_state: dict[str, Any] = {"counter": 3, "ordinary": "child-value-3"}
if parent_handle is not None:
expected_state["provider_session"] = parent_handle
expected_state["client_provider_session"] = parent_handle
assert parent_session.state == expected_state
assert parent_session.service_session_id == "parent-service-session"


@pytest.mark.parametrize("stream", [False, True])
@pytest.mark.parametrize("with_middleware", [False, True])
@pytest.mark.parametrize("child_declares_keys", [False, True])
async def test_chat_agent_as_tool_isolates_distinct_parent_provider_state(
stream: bool, with_middleware: bool, child_declares_keys: bool
) -> None:
"""Keep parent agent and client state local even when the child uses different declarations."""

class ParentAgent(Agent):
service_session_state_keys = frozenset({"parent_provider", "parent_removed"})

class ParentClient(MockBaseChatClient):
service_session_state_keys = frozenset({"parent_client_provider", "parent_client_removed"})

class ChildAgent(Agent):
service_session_state_keys = frozenset({"child_provider"}) if child_declares_keys else frozenset()

parent_state: dict[str, Any] = {
"parent_provider": {"values": ["parent"]},
"parent_removed": "parent",
"parent_client_provider": "parent-client",
"parent_client_removed": "parent-client",
"counter": 0,
}
parent_session = AgentSession()
parent_session.state.update(parent_state)
observed_states: list[dict[str, Any]] = []

@tool(approval_mode="never_require")
def inspect_state(ctx: FunctionInvocationContext) -> str:
assert ctx.session is not None
observed_states.append(dict(ctx.session.state))
assert _PARENT_SERVICE_SESSION_STATE_KEYS_CONTEXT_KEY not in ctx.kwargs
ctx.session.state.setdefault("parent_provider", {"values": []})["values"].append("child")
ctx.session.state.pop("parent_removed", None)
ctx.session.state["parent_client_provider"] = "child"
ctx.session.state.pop("parent_client_removed", None)
ctx.session.state["counter"] += 1
if child_declares_keys:
ctx.session.state["child_provider"] = "child"
return "State inspected."

child_client = MockBaseChatClient()
child_client.streaming_responses = [
[
ChatResponseUpdate(
role="assistant",
contents=[Content.from_function_call(call_id="inspect", name="inspect_state", arguments={})],
)
],
[ChatResponseUpdate(role="assistant", contents=[Content.from_text("Child completed.")])],
]
child = ChildAgent(client=child_client, tools=[inspect_state])
parent_client = ParentClient()
parent_call = Content.from_function_call(
call_id="delegate", name="delegate", arguments={"task": "Inspect the delegated session"}
)
if stream:
parent_client.streaming_responses = [
[ChatResponseUpdate(role="assistant", contents=[parent_call])],
[ChatResponseUpdate(role="assistant", contents=[Content.from_text("Parent completed.")])],
]
else:
parent_client.run_responses = [
ChatResponse(messages=Message(role="assistant", contents=[parent_call])),
ChatResponse(messages=Message(role="assistant", contents=[Content.from_text("Parent completed.")])),
]
parent = ParentAgent(
client=parent_client,
tools=[child.as_tool(name="delegate", propagate_session=True)],
middleware=[ToolApprovalMiddleware(source_id="parent_approval")] if with_middleware else [],
)
invocation_kwargs: dict[str, frozenset[str]] = {_PARENT_SERVICE_SESSION_STATE_KEYS_CONTEXT_KEY: frozenset()}
if stream:
result = await parent.run(
"Delegate.", session=parent_session, stream=True, function_invocation_kwargs=invocation_kwargs
).get_final_response()
else:
result = await parent.run("Delegate.", session=parent_session, function_invocation_kwargs=invocation_kwargs)

assert result.text == "Parent completed."
assert len(observed_states) == 1
assert all(
key not in observed_states[0]
for key in ParentAgent.service_session_state_keys | ParentClient.service_session_state_keys
)
assert parent_session.state["parent_provider"] == {"values": ["parent"]}
assert parent_session.state["parent_removed"] == "parent"
assert parent_session.state["parent_client_provider"] == "parent-client"
assert parent_session.state["parent_client_removed"] == "parent-client"
assert "child_provider" not in parent_session.state
assert parent_session.state["counter"] == 1


async def test_chat_agent_as_tool_propagate_session_clears_service_session_id(client: SupportsChatGetResponse) -> None:
"""Test that propagate_session=True gives the child a separate session with cleared service_session_id."""
agent = Agent(client=client, name="SubAgent", description="Sub agent")
Expand Down
4 changes: 4 additions & 0 deletions python/packages/foundry/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -82,6 +82,10 @@ entry points for existing traces, response IDs, and registered targets.

## Concurrent reuse

When a Foundry agent is used through `as_tool(propagate_session=True)`, application-owned state propagates
between the parent and child, but service-owned session handles do not. Each delegated invocation has its own
service session. This also applies to an `Agent` configured with a `RawFoundryAgentChatClient`.

A `FoundryChatClient` instance can be shared by concurrent asynchronous calls on the same event loop. Streaming,
non-streaming, and mixed calls are supported. Keep mutable run state isolated by creating a separate `Agent` and
`AgentSession` for each concurrent run and by passing separate messages and options.
Expand Down
3 changes: 3 additions & 0 deletions python/packages/foundry/agent_framework_foundry/_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -181,6 +181,9 @@ class MyClient(FunctionInvocationLayer, RawFoundryAgentChatClient):
OTEL_PROVIDER_NAME: ClassVar[str] = "azure.ai.foundry"
_FEATURE_USAGE_INDEX: ClassVar[int | None] = FeatureIndex.FOUNDRY_AGENT

service_session_state_keys: ClassVar[frozenset[str]] = frozenset({FOUNDRY_HOSTED_AGENT_SESSION_ID_KEY})
"""Service-owned state keys, including when this client is used by a generic Agent."""

def __init__(
self,
*,
Expand Down
Loading
Loading