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
135 changes: 89 additions & 46 deletions python/packages/claude/agent_framework_claude/_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Comment on lines 386 to 388

def _normalize_tools(
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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:
Expand All @@ -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.
Expand Down Expand Up @@ -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
Expand Down
131 changes: 79 additions & 52 deletions python/packages/claude/tests/test_claude_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand All @@ -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:
Expand All @@ -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()

Expand All @@ -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()

Expand Down
Loading