diff --git a/src/claude_agent_sdk/_internal/client.py b/src/claude_agent_sdk/_internal/client.py index 4423b4297..7cb7d1246 100644 --- a/src/claude_agent_sdk/_internal/client.py +++ b/src/claude_agent_sdk/_internal/client.py @@ -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: diff --git a/src/claude_agent_sdk/_internal/query.py b/src/claude_agent_sdk/_internal/query.py index 4d5f0070e..b2eb60ab9 100644 --- a/src/claude_agent_sdk/_internal/query.py +++ b/src/claude_agent_sdk/_internal/query.py @@ -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. @@ -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 @@ -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] = {} @@ -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: @@ -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}") diff --git a/src/claude_agent_sdk/client.py b/src/claude_agent_sdk/client.py index bba76b10e..429bcfc98 100644 --- a/src/claude_agent_sdk/client.py +++ b/src/claude_agent_sdk/client.py @@ -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: diff --git a/tests/test_query.py b/tests/test_query.py index b5a254f33..e284f5f9d 100644 --- a/tests/test_query.py +++ b/tests/test_query.py @@ -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" @@ -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.""" @@ -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.""" @@ -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 == []