diff --git a/python/packages/core/AGENTS.md b/python/packages/core/AGENTS.md index 8e9d3a5e393..fa2f2a9df1e 100644 --- a/python/packages/core/AGENTS.md +++ b/python/packages/core/AGENTS.md @@ -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`) diff --git a/python/packages/core/agent_framework/_agents.py b/python/packages/core/agent_framework/_agents.py index db264535aa1..c4b12f844e9 100644 --- a/python/packages/core/agent_framework/_agents.py +++ b/python/packages/core/agent_framework/_agents.py @@ -81,6 +81,8 @@ from typing_extensions import Self, TypedDict # pragma: no cover if TYPE_CHECKING: + from collections.abc import Collection + from mcp import types from mcp.server.lowlevel import Server from pydantic import BaseModel @@ -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], @@ -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. @@ -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) @@ -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), + *parent_service_session_state_keys, *child_approval_source_ids, *parent_approval_source_ids, }) @@ -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." @@ -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) diff --git a/python/packages/core/agent_framework/_tools.py b/python/packages/core/agent_framework/_tools.py index 4ab15a52f79..187ad79155f 100644 --- a/python/packages/core/agent_framework/_tools.py +++ b/python/packages/core/agent_framework/_tools.py @@ -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: @@ -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) @@ -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) @@ -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( @@ -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 diff --git a/python/packages/core/tests/core/test_agents.py b/python/packages/core/tests/core/test_agents.py index e0e79153288..01c4a306405 100644 --- a/python/packages/core/tests/core/test_agents.py +++ b/python/packages/core/tests/core/test_agents.py @@ -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, @@ -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") diff --git a/python/packages/foundry/README.md b/python/packages/foundry/README.md index 56d13cbd872..8b170212cfc 100644 --- a/python/packages/foundry/README.md +++ b/python/packages/foundry/README.md @@ -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. diff --git a/python/packages/foundry/agent_framework_foundry/_agent.py b/python/packages/foundry/agent_framework_foundry/_agent.py index 43374b85d29..fbceea8cff3 100644 --- a/python/packages/foundry/agent_framework_foundry/_agent.py +++ b/python/packages/foundry/agent_framework_foundry/_agent.py @@ -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, *, diff --git a/python/packages/foundry/tests/foundry/test_foundry_agent.py b/python/packages/foundry/tests/foundry/test_foundry_agent.py index 3fe96541499..335cdb5756e 100644 --- a/python/packages/foundry/tests/foundry/test_foundry_agent.py +++ b/python/packages/foundry/tests/foundry/test_foundry_agent.py @@ -1223,6 +1223,8 @@ def test_foundry_agents_declare_hosted_agent_session_id_as_server_owned() -> Non assert FOUNDRY_HOSTED_AGENT_SESSION_ID_KEY in RawFoundryAgent.service_session_state_keys # FoundryAgent is the recommended production class, so it must inherit the same protection. assert FOUNDRY_HOSTED_AGENT_SESSION_ID_KEY in FoundryAgent.service_session_state_keys + assert FOUNDRY_HOSTED_AGENT_SESSION_ID_KEY in RawFoundryAgentChatClient.service_session_state_keys + assert FOUNDRY_HOSTED_AGENT_SESSION_ID_KEY in _FoundryAgentChatClient.service_session_state_keys async def test_raw_foundry_agent_prepare_run_context_injects_agent_session_id_from_state() -> None: @@ -1263,6 +1265,65 @@ async def test_raw_foundry_agent_prepare_run_context_injects_agent_session_id_fr } +@pytest.mark.parametrize("agent_type", [RawFoundryAgent, FoundryAgent, None]) +@pytest.mark.parametrize("parent_handle", [None, "parent-agent-session"]) +async def test_foundry_agent_tools_isolate_service_state_between_children( + agent_type: type[RawFoundryAgent] | None, parent_handle: str | None +) -> None: + """Delegated Foundry calls retain application state without sharing service handles.""" + project = MagicMock() + project.get_openai_client.return_value = MagicMock() + children = [ + agent_type(project_client=project, agent_name=name) + if agent_type is not None + else Agent(client=_FoundryAgentChatClient(project_client=project, agent_name=name), name=name) + for name in ("first-child", "second-child") + ] + parent = AgentSession(service_session_id="parent-response") + parent.state["ordinary"] = "parent-value" + if parent_handle is not None: + parent.state[FOUNDRY_HOSTED_AGENT_SESSION_ID_KEY] = parent_handle + observed_options: list[dict[str, Any]] = [] + updates: list[AgentResponseUpdate] = [] + + def respond(**kwargs: Any) -> ResponseStream[ChatResponseUpdate, ChatResponse]: + assert kwargs["stream"] is True + observed_options.append(dict(kwargs["options"])) + call_number = len(observed_options) + + async def stream() -> AsyncIterator[ChatResponseUpdate]: + yield ChatResponseUpdate( + role="assistant", + contents=[Content.from_text(text="done")], + conversation_id=f"child-response-{call_number}", + additional_properties={"agent_session_id": f"child-agent-session-{call_number}"}, + ) + + return ResponseStream(stream(), finalizer=ChatResponse.from_updates) + + with patch.object(RawFoundryAgentChatClient, "_inner_get_response", side_effect=respond): + for child in (children[0], children[1], children[0]): + delegated_tool = child.as_tool(propagate_session=True, stream_callback=updates.append) + result = await delegated_tool.invoke( + context=FunctionInvocationContext( + function=delegated_tool, + arguments={"task": "Run child"}, + session=parent, + ) + ) + assert result[0].text == "done" + + assert len(observed_options) == 3 + assert all("agent_session_id" not in options.get("extra_body", {}) for options in observed_options) + assert all(options.get("conversation_id") is None for options in observed_options) + assert len(updates) == 3 + expected_state = {"ordinary": "parent-value"} + if parent_handle is not None: + expected_state[FOUNDRY_HOSTED_AGENT_SESSION_ID_KEY] = parent_handle + assert parent.state == expected_state + assert parent.service_session_id == "parent-response" + + def test_foundry_agent_updates_session_from_response_ids() -> None: """Test that response and hosted-agent session IDs persist independently."""