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
3 changes: 3 additions & 0 deletions src/claude_agent_sdk/_internal/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -145,6 +145,9 @@ async def _process_query_inner(
exclude_dynamic_sections=exclude_dynamic_sections,
skills=configured_options.skills,
forward_subagent_text=configured_options.forward_subagent_text,
is_resuming=bool(
configured_options.resume or configured_options.continue_conversation
),
)

if configured_options.session_store is not None:
Expand Down
21 changes: 16 additions & 5 deletions src/claude_agent_sdk/_internal/query.py
Original file line number Diff line number Diff line change
Expand Up @@ -126,6 +126,7 @@ def __init__(
exclude_dynamic_sections: bool | None = None,
skills: list[str] | Literal["all"] | None = None,
forward_subagent_text: bool = False,
is_resuming: bool = False,
):
"""Initialize Query with transport and callbacks.

Expand All @@ -143,6 +144,7 @@ def __init__(
can filter which skills are loaded into the system prompt
forward_subagent_text: Ask the CLI (via initialize) to forward
subagent text/thinking blocks, not just tool_use/tool_result
is_resuming: Whether the CLI is resuming an existing conversation
"""
self._initialize_timeout = initialize_timeout
self.transport = transport
Expand All @@ -158,6 +160,7 @@ def __init__(
self._exclude_dynamic_sections = exclude_dynamic_sections
self._skills = skills
self._forward_subagent_text = forward_subagent_text
self._is_resuming = is_resuming

# Control protocol state
self.pending_control_responses: dict[str, anyio.Event] = {}
Expand Down Expand Up @@ -865,9 +868,14 @@ async def stream_input(self, stream: AsyncIterable[dict[str, Any]]) -> None:

If SDK MCP servers, hooks, or a ``can_use_tool`` callback are present,
waits for a run-ending result before closing stdin to allow
bidirectional control protocol communication.
bidirectional control protocol communication. A resumed conversation
may issue those control requests without a new prompt message (for
example, when continuing a deferred SDK MCP tool call), so an empty
resumed stream also waits for its result. A prompt stream that fails
before writing anything still closes immediately.
"""
written = 0
failed = False
try:
async for message in stream:
if self._closed:
Expand All @@ -879,14 +887,17 @@ async def stream_input(self, stream: AsyncIterable[dict[str, Any]]) -> None:
# leave stdin open — the CLI would wait for input forever and the
# consumer's `async for` would never finish — fall through and
# close it like a normal end of input.
failed = True
logger.error("Prompt stream failed; closing stdin: %s", e)
try:
if written:
if written or (
not failed and self._is_resuming and self._has_bidirectional_needs()
):
await self.wait_for_result_and_end_input()
else:
# Nothing was sent, so no result will arrive to release the
# hold; close immediately (mirrors the TypeScript SDK's
# messageCount guard).
# A new conversation with no input cannot produce a result to
# release the hold; close immediately (mirrors the TypeScript
# SDK's messageCount guard).
await self.transport.end_input()
except Exception as e:
logger.debug(f"Error closing input stream: {e}")
Expand Down
1 change: 1 addition & 0 deletions src/claude_agent_sdk/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -199,6 +199,7 @@ async def _connect_inner(
exclude_dynamic_sections=exclude_dynamic_sections,
skills=self.options.skills,
forward_subagent_text=self.options.forward_subagent_text,
is_resuming=bool(options.resume or options.continue_conversation),
)

if self.options.session_store is not None:
Expand Down
153 changes: 151 additions & 2 deletions tests/test_query.py
Original file line number Diff line number Diff line change
Expand Up @@ -238,6 +238,45 @@ async def mock_receive():
return mock_transport


def _resumed_mcp_handshake_transport(
writes: list[str], stream_exhausted: anyio.Event
) -> tuple[AsyncMock, dict[str, bool]]:
"""Mock a deferred SDK MCP request arriving after an empty resume stream."""
state = {"ended": False}
responded = anyio.Event()
request = _MCP_CONTROL_REQUESTS[0]
mock_transport = AsyncMock()

async def tracking_write(data):
if state["ended"]:
raise CLIConnectionError("stdin closed")
writes.append(data)
frame = json.loads(data)
if frame.get("type") == "control_response":
responded.set()

async def end_input():
state["ended"] = True

async def mock_receive():
await stream_exhausted.wait()
yield request
with anyio.move_on_after(1):
await responded.wait()
if not responded.is_set():
return
for msg in _ASSISTANT_AND_RESULT:
yield msg

mock_transport.write = tracking_write
mock_transport.read_messages = mock_receive
mock_transport.connect = AsyncMock()
mock_transport.close = AsyncMock()
mock_transport.end_input = end_input
mock_transport.is_ready = Mock(return_value=True)
return mock_transport, state


def _assert_mcp_handshake_succeeded(writes: list[str]) -> None:
control_responses = [
json.loads(w) for w in writes if json.loads(w).get("type") == "control_response"
Expand Down Expand Up @@ -747,6 +786,112 @@ async def mock_receive():
class TestAsyncIterablePromptWithSdkMcpServers:
"""Test that AsyncIterable prompts keep stdin open for SDK MCP servers."""

def test_empty_resumed_stream_waits_for_result(self):
"""A resumed deferred SDK MCP call can emit a control request without
a new prompt message, so stdin must stay open until its result."""

async def _test():
mock_transport = _make_mock_transport(messages=[])
ended = anyio.Event()

async def end_input():
ended.set()

mock_transport.end_input = end_input
q = Query(
transport=mock_transport,
is_streaming_mode=True,
sdk_mcp_servers={"greeter": _make_greet_server()["instance"]},
is_resuming=True,
)

async def empty_stream():
return
yield # pragma: no cover

async with anyio.create_task_group() as tg:
tg.start_soon(q.stream_input, empty_stream())
await anyio.sleep(0.05)
assert not ended.is_set()
q._first_result_event.set()

assert ended.is_set()

anyio.run(_test)

def test_empty_new_stream_closes_immediately(self):
"""A new empty stream has no work that could produce a result."""

async def _test():
mock_transport = _make_mock_transport(messages=[])
q = Query(
transport=mock_transport,
is_streaming_mode=True,
sdk_mcp_servers={"greeter": _make_greet_server()["instance"]},
)

async def empty_stream():
return
yield # pragma: no cover

await q.stream_input(empty_stream())
mock_transport.end_input.assert_called_once()

anyio.run(_test)

def test_empty_resume_handles_deferred_mcp_request(self):
"""The resumed SDK MCP control response must be written before stdin closes."""

async def _test():
server = _make_greet_server()
writes: list[str] = []
stream_exhausted = anyio.Event()
mock_transport, state = _resumed_mcp_handshake_transport(
writes, stream_exhausted
)

async def empty_stream():
stream_exhausted.set()
return
yield # pragma: no cover

with (
patch(
"claude_agent_sdk._internal.client.SubprocessCLITransport"
) as mock_cls,
patch(
"claude_agent_sdk._internal.query.Query.initialize",
new_callable=AsyncMock,
),
):
mock_cls.return_value = mock_transport
with anyio.fail_after(5):
messages = [
msg
async for msg in query(
prompt=empty_stream(),
options=ClaudeAgentOptions(
resume="00000000-0000-4000-8000-000000000000",
mcp_servers={"greeter": server},
),
)
]

assert [type(msg) for msg in messages] == [
AssistantMessage,
ResultMessage,
]
control_responses = [
json.loads(write)
for write in writes
if json.loads(write).get("type") == "control_response"
]
assert len(control_responses) == 1
assert control_responses[0]["response"]["subtype"] == "success"
assert state["ended"] is True

anyio.run(_test)

def test_async_iterable_with_sdk_mcp_servers(self):
"""AsyncIterable prompt path should wait for first result before
closing stdin when SDK MCP servers are present."""
Expand Down Expand Up @@ -1088,7 +1233,8 @@ async def prompt_stream():
assert [type(m) for m in messages] == [AssistantMessage, ResultMessage]
assert state["ended"] is True

def test_prompt_iterable_that_raises_immediately_closes_stdin(self):
@pytest.mark.parametrize("resume", [None, "00000000-0000-4000-8000-000000000000"])
def test_prompt_iterable_that_raises_immediately_closes_stdin(self, resume):
"""Nothing was sent, so no result can release the hold: stdin must be
closed right away or the CLI (and the consumer) would wait forever."""

Expand Down Expand Up @@ -1131,7 +1277,10 @@ async def allow_all(tool_name, tool_input, context):
msg
async for msg in query(
prompt=prompt_stream(),
options=ClaudeAgentOptions(can_use_tool=allow_all),
options=ClaudeAgentOptions(
can_use_tool=allow_all,
resume=resume,
),
)
]
assert messages == []
Expand Down