diff --git a/python/packages/claude/agent_framework_claude/_agent.py b/python/packages/claude/agent_framework_claude/_agent.py index 7db8eb7dbfd..89b21964743 100644 --- a/python/packages/claude/agent_framework_claude/_agent.py +++ b/python/packages/claude/agent_framework_claude/_agent.py @@ -385,7 +385,6 @@ def __init__( self._default_options = opts self._started = False - self._current_session_id: str | None = None self._structured_output: Any = None def _normalize_tools( @@ -424,59 +423,78 @@ async def __aexit__(self, exc_type: Any, exc_val: Any, exc_tb: Any) -> None: async def start(self) -> None: """Start the Claude SDK client. - This method initializes the Claude SDK client and establishes a connection - to the Claude Code CLI. It is called automatically when using the agent - as an async context manager. + Owned clients are created per run so that distinct sessions stay isolated; + this only needs to establish a connection for a pre-configured client that + was injected at construction. It is called automatically when using the + agent as an async context manager. Raises: AgentException: If the client fails to start. """ - await self._ensure_session() + if self._client is not None and not self._owns_client and not self._started: + try: + await self._client.connect() + self._started = True + except Exception as ex: + raise AgentException(f"Failed to start Claude SDK client: {ex}") from ex async def stop(self) -> None: """Stop the Claude SDK client and clean up resources. - Stops the client if owned by this agent. Called automatically when - using the agent as an async context manager. + Per-run owned clients are disconnected at the end of each run, so this + only disconnects a long-lived client the agent still owns. A client that + was injected at construction is owned by the caller and left untouched. + Called automatically when using the agent as an async context manager. """ - if self._client and self._owns_client: + if self._client is not None and self._owns_client: with contextlib.suppress(Exception): await self._client.disconnect() self._started = False - self._current_session_id = None - async def _ensure_session(self, session_id: str | None = None) -> None: - """Ensure the client is connected for the specified session. + async def _acquire_client(self, resume_session_id: str | None = None) -> tuple[ClaudeSDKClient, bool]: + """Acquire a Claude SDK client for a single run. - If the requested session differs from the current one, recreates the client. + A ``ClaudeSDKClient`` is stateful and represents exactly one provider + conversation, so obtaining it is an isolation decision, not just a + connection optimization. When a client was injected at construction, that + single client is reused for every run and its lifecycle is owned by the + caller. Otherwise a fresh client scoped to this run is created and + connected, resuming ``resume_session_id`` when the framework session + already carries a provider conversation id. Binding the client to the run + rather than to shared agent state keeps distinct sessions isolated even + when they run concurrently against the same agent instance. Args: - session_id: The session ID to use, or None for a new session. - """ - needs_new_client = ( - not self._started or self._client is None or (session_id and session_id != self._current_session_id) - ) - - if needs_new_client: - # Stop existing client if any - if self._client and self._owns_client: - with contextlib.suppress(Exception): - await self._client.disconnect() - self._started = False + resume_session_id: The provider continuation id to resume, or None for + a fresh conversation. - # Create new client with resume option if needed - opts = self._prepare_client_options(resume_session_id=session_id) - self._client = ClaudeSDKClient(options=opts) - self._owns_client = True + Returns: + A tuple of the client and whether the caller owns it and must + disconnect it when the run completes. An injected client is never + owned by the caller of this method. - try: - await self._client.connect() - self._started = True - self._current_session_id = session_id - except Exception as ex: - self._client = None - raise AgentException(f"Failed to start Claude SDK client: {ex}") from ex + Raises: + AgentException: If the client fails to connect. + """ + if self._client is not None and not self._owns_client: + # Injected client: a single shared conversation with a caller-managed + # lifecycle. Connect it once and reuse it verbatim. + if not self._started: + try: + await self._client.connect() + self._started = True + except Exception as ex: + raise AgentException(f"Failed to start Claude SDK client: {ex}") from ex + return self._client, False + + opts = self._prepare_client_options(resume_session_id=resume_session_id) + client = ClaudeSDKClient(options=opts) + try: + await client.connect() + except Exception as ex: + raise AgentException(f"Failed to start Claude SDK client: {ex}") from ex + return client, True def _prepare_client_options(self, resume_session_id: str | None = None) -> SDKOptions: """Prepare SDK options for client initialization. @@ -635,15 +653,16 @@ async def handler(args: dict[str, Any]) -> dict[str, Any]: handler=handler, ) - async def _apply_runtime_options(self, options: dict[str, Any] | None) -> None: + async def _apply_runtime_options(self, client: ClaudeSDKClient, options: dict[str, Any] | None) -> None: """Apply runtime options that can be changed dynamically. The Claude SDK supports changing model and permission_mode after connection. Args: + client: The per-run client to apply the options to. options: Runtime options to apply. """ - if not options or not self._client: + if not options: return if "on_function_approval" in options: @@ -654,10 +673,10 @@ async def _apply_runtime_options(self, options: dict[str, Any] | None) -> None: ) if "model" in options: - await self._client.set_model(options["model"]) + await client.set_model(options["model"]) if "permission_mode" in options: - await self._client.set_permission_mode(options["permission_mode"]) + await client.set_permission_mode(options["permission_mode"]) def _format_prompt(self, messages: list[Message] | None) -> str: """Format messages into a prompt string. @@ -765,22 +784,46 @@ async def _get_stream( """Internal streaming implementation.""" session = session or self.create_session() - # Ensure we're connected to the right session - await self._ensure_session(self._get_chat_conversation_id(session)) + # A ClaudeSDKClient represents a single provider conversation, so each run + # acquires its own client (resuming the framework session's provider + # conversation when one exists) and releases it when the run completes. + # Binding the client to the run keeps distinct sessions isolated even when + # they run concurrently against the same agent instance. + client, owns_client = await self._acquire_client(self._get_chat_conversation_id(session)) + try: + async for update in self._stream_run(client, session, messages, options): + yield update + finally: + if owns_client: + with contextlib.suppress(Exception): + await client.disconnect() - if not self._client: - raise RuntimeError("Claude SDK client not initialized.") + async def _stream_run( + self, + client: ClaudeSDKClient, + session: AgentSession, + messages: AgentRunInputs | None, + options: OptionsT | None, + ) -> AsyncIterable[AgentResponseUpdate]: + """Run a single query against ``client`` and stream response updates. + Args: + client: The per-run Claude SDK client to query. + session: The active session; its ``service_session_id`` is updated with + the provider conversation id when the run completes. + messages: The input messages for this run. + options: Runtime options (model, permission_mode) for this run. + """ prompt = self._format_prompt(normalize_messages(messages)) # Apply runtime options (model, permission_mode) - await self._apply_runtime_options(dict(options) if options else None) + await self._apply_runtime_options(client, dict(options) if options else None) session_id: str | None = None structured_output: Any = None - await self._client.query(prompt) - async for message in self._client.receive_response(): + await client.query(prompt) + async for message in client.receive_response(): if isinstance(message, StreamEvent): # Handle streaming events - extract text/thinking deltas event = message.event diff --git a/python/packages/claude/tests/test_claude_agent.py b/python/packages/claude/tests/test_claude_agent.py index 116d4114694..87d77ce482f 100644 --- a/python/packages/claude/tests/test_claude_agent.py +++ b/python/packages/claude/tests/test_claude_agent.py @@ -553,59 +553,90 @@ def test_create_session_with_service_session_id(self) -> None: session = agent.create_session(session_id="existing-session-123") assert isinstance(session, AgentSession) - async def test_ensure_session_creates_client(self) -> None: - """Test _ensure_session creates client when not started.""" - with patch("agent_framework_claude._agent.ClaudeSDKClient") as mock_client_class: - mock_client = MagicMock() - mock_client.connect = AsyncMock() - mock_client_class.return_value = mock_client - - agent = ClaudeAgent() - await agent._ensure_session(None) # type: ignore[reportPrivateUsage] - - assert agent._started # type: ignore[reportPrivateUsage] - mock_client.connect.assert_called_once() - - async def test_ensure_session_recreates_for_different_session(self) -> None: - """Test _ensure_session recreates client for different session ID.""" - with patch("agent_framework_claude._agent.ClaudeSDKClient") as mock_client_class: - mock_client1 = MagicMock() - mock_client1.connect = AsyncMock() - mock_client1.disconnect = AsyncMock() - - mock_client2 = MagicMock() - mock_client2.connect = AsyncMock() + @staticmethod + async def _create_async_generator(items: list[Any]) -> Any: + """Yield the given items as an async generator (a mock provider response).""" + for item in items: + yield item - mock_client_class.side_effect = [mock_client1, mock_client2] + def _make_mock_client(self) -> MagicMock: + """Build a mock ClaudeSDKClient exposing the methods a run exercises.""" + mock_client = MagicMock() + mock_client.connect = AsyncMock() + mock_client.disconnect = AsyncMock() + mock_client.query = AsyncMock() + mock_client.set_model = AsyncMock() + mock_client.set_permission_mode = AsyncMock() + mock_client.receive_response = MagicMock(return_value=self._create_async_generator([])) + return mock_client + async def test_acquire_client_creates_fresh_owned_client(self) -> None: + """A fresh run creates and owns a newly connected client.""" + mock_client = self._make_mock_client() + with patch("agent_framework_claude._agent.ClaudeSDKClient", return_value=mock_client): agent = ClaudeAgent() + client, owns_client = await agent._acquire_client(None) # type: ignore[reportPrivateUsage] - # First session - await agent._ensure_session(None) # type: ignore[reportPrivateUsage] - assert agent._started # type: ignore[reportPrivateUsage] - - # Different session should recreate client - await agent._ensure_session("new-session-id") # type: ignore[reportPrivateUsage] - assert agent._current_session_id == "new-session-id" # type: ignore[reportPrivateUsage] - mock_client1.disconnect.assert_called_once() - - async def test_ensure_session_reuses_for_same_session(self) -> None: - """Test _ensure_session reuses client for same session ID.""" - with patch("agent_framework_claude._agent.ClaudeSDKClient") as mock_client_class: - mock_client = MagicMock() - mock_client.connect = AsyncMock() - mock_client_class.return_value = mock_client + assert client is mock_client + assert owns_client is True + mock_client.connect.assert_awaited_once() + async def test_acquire_client_forwards_resume_id(self) -> None: + """An existing provider conversation id is forwarded as the resume id.""" + mock_client = self._make_mock_client() + with patch("agent_framework_claude._agent.ClaudeSDKClient", return_value=mock_client): + agent = ClaudeAgent() + with patch.object( + agent, + "_prepare_client_options", + wraps=agent._prepare_client_options, # type: ignore[reportPrivateUsage] + ) as prepare: + _, owns_client = await agent._acquire_client("provider-session-1") # type: ignore[reportPrivateUsage] + + assert owns_client is True + prepare.assert_called_once_with(resume_session_id="provider-session-1") + + async def test_acquire_client_reuses_injected_client(self) -> None: + """An injected client is reused across runs and never owned by the run.""" + injected = self._make_mock_client() + agent = ClaudeAgent(client=injected) + + client_a, owns_a = await agent._acquire_client(None) # type: ignore[reportPrivateUsage] + client_b, owns_b = await agent._acquire_client("session-123") # type: ignore[reportPrivateUsage] + + assert client_a is injected + assert client_b is injected + assert owns_a is False + assert owns_b is False + # Connected once and reused; never disconnected by the run. + injected.connect.assert_awaited_once() + injected.disconnect.assert_not_called() + + async def test_two_fresh_sessions_use_separate_clients(self) -> None: + """Two distinct fresh sessions on one agent get isolated, separately-owned clients.""" + mock_client1 = self._make_mock_client() + mock_client2 = self._make_mock_client() + with patch( + "agent_framework_claude._agent.ClaudeSDKClient", + side_effect=[mock_client1, mock_client2], + ) as mock_client_class: agent = ClaudeAgent() + await agent.run("hello", session=agent.create_session()) + await agent.run("hello", session=agent.create_session()) - # First call - await agent._ensure_session("session-123") # type: ignore[reportPrivateUsage] + assert mock_client_class.call_count == 2 + # Each per-run client is released when its run completes. + mock_client1.disconnect.assert_awaited_once() + mock_client2.disconnect.assert_awaited_once() - # Same session should not recreate - await agent._ensure_session("session-123") # type: ignore[reportPrivateUsage] + async def test_owned_client_disconnected_after_run(self) -> None: + """A per-run owned client is disconnected when the run finishes.""" + mock_client = self._make_mock_client() + with patch("agent_framework_claude._agent.ClaudeSDKClient", return_value=mock_client): + agent = ClaudeAgent() + await agent.run("hello") - # Only called once - assert mock_client_class.call_count == 1 + mock_client.disconnect.assert_awaited_once() # region Test ClaudeAgent Tool Conversion @@ -992,9 +1023,8 @@ async def test_apply_runtime_model(self) -> None: mock_client.set_permission_mode = AsyncMock() agent = ClaudeAgent() - agent._client = mock_client # type: ignore[reportPrivateUsage] - await agent._apply_runtime_options({"model": "opus"}) # type: ignore[reportPrivateUsage] + await agent._apply_runtime_options(mock_client, {"model": "opus"}) # type: ignore[reportPrivateUsage] mock_client.set_model.assert_called_once_with("opus") async def test_apply_runtime_permission_mode(self) -> None: @@ -1004,9 +1034,8 @@ async def test_apply_runtime_permission_mode(self) -> None: mock_client.set_permission_mode = AsyncMock() agent = ClaudeAgent() - agent._client = mock_client # type: ignore[reportPrivateUsage] - await agent._apply_runtime_options({"permission_mode": "acceptEdits"}) # type: ignore[reportPrivateUsage] + await agent._apply_runtime_options(mock_client, {"permission_mode": "acceptEdits"}) # type: ignore[reportPrivateUsage] mock_client.set_permission_mode.assert_called_once_with("acceptEdits") async def test_apply_runtime_options_none(self) -> None: @@ -1016,9 +1045,8 @@ async def test_apply_runtime_options_none(self) -> None: mock_client.set_permission_mode = AsyncMock() agent = ClaudeAgent() - agent._client = mock_client # type: ignore[reportPrivateUsage] - await agent._apply_runtime_options(None) # type: ignore[reportPrivateUsage] + await agent._apply_runtime_options(mock_client, None) # type: ignore[reportPrivateUsage] mock_client.set_model.assert_not_called() mock_client.set_permission_mode.assert_not_called() @@ -1029,10 +1057,9 @@ async def test_apply_runtime_on_function_approval_rejected(self) -> None: mock_client.set_permission_mode = AsyncMock() agent = ClaudeAgent() - agent._client = mock_client # type: ignore[reportPrivateUsage] with pytest.raises(ValueError, match="on_function_approval"): - await agent._apply_runtime_options({"on_function_approval": lambda _c: True}) # type: ignore[reportPrivateUsage] + await agent._apply_runtime_options(mock_client, {"on_function_approval": lambda _c: True}) # type: ignore[reportPrivateUsage] mock_client.set_model.assert_not_called() mock_client.set_permission_mode.assert_not_called()