Files
claude-code-api/tests/test_stream.py
T
hh 621da90623 feat(stream): warm-up turn, because MCP servers aren't up on turn one
The CLI starts connecting its MCP servers when a turn starts and does
not wait for them, so the first turn of a session sees only built-in
tools. Measured on the pi with two HTTP MCP servers: turn 1 reports both
`pending` with 29 tools available, turn 2 reports both `connected` with
90. For an agent whose job is those tools, that first reply is silently
wrong — the model doesn't see them and answers as best it can.

The PTY transport hid this: its multi-second wait for the TUI to settle
happened to cover the connect. Nothing about the headless path does, and
nothing cheap fixes it — an 8s pause before the first prompt changed
nothing, a `/status` slash command returns in 134ms without touching the
MCP client, and a `control_request`/`initialize` handshake answers with
the command list and leaves the servers pending. Only a real turn does
it, so `warmup_turn` spends one deliberately: a few tokens and a couple
of seconds, once per session, against sessions that are pooled for the
whole conversation.

The price is two short messages at the head of the transcript, which is
why the default prompt reads as procedural rather than conversational.

Also log, at WARNING, any MCP server a turn starts without — this is a
failure with no other symptom, and it should be one grep away rather
than a mystery about the agent forgetting a tool it has.
2026-07-28 03:22:30 +02:00

498 lines
16 KiB
Python

"""Unit tests for the headless stream-json transport.
No real `claude` here — `StreamTurnManager` talks to a fake process that
replays a canned record list, which is enough to pin the parts that are
genuinely ours: turn-boundary detection, the content-block index rebase
across multi-request turns, bookkeeping filtering, and error mapping.
The live end-to-end path is covered by the smoke test at the bottom.
"""
from __future__ import annotations
import os
from typing import Any
import pytest
from claude_code_api.errors import AuthError, ProcessError
from claude_code_api.events import (
AssistantMessage,
ResultMessage,
StreamEvent,
UserMessage,
)
from claude_code_api.pty import PtyProcessOptions
from claude_code_api.stream import (
StreamTurnManager,
build_stream_argv,
strip_tui_only_flags,
)
class FakeStreamProcess:
"""Minimal stand-in for `StreamClaudeProcess`.
Records sent prompts, replays a scripted list of stdout records, and
returns ``None`` (EOF) once exhausted.
"""
def __init__(self, records: list[dict[str, Any] | None], *, stderr: str = "") -> None:
self._records = list(records)
self.sent: list[str] = []
self.started = False
self.closed = False
self.session_id = "sess-fake"
self._stderr = stderr
async def start(self) -> None:
self.started = True
async def send_user_message(self, text: str) -> None:
self.sent.append(text)
async def next_record(self) -> dict[str, Any] | None:
if not self._records:
return None
return self._records.pop(0)
def stderr_text(self) -> str:
return self._stderr
def captured_output(self) -> bytes:
return b""
async def aclose(self) -> None:
self.closed = True
def _assistant(*blocks: dict[str, Any], stop: str | None = None) -> dict[str, Any]:
return {
"type": "assistant",
"session_id": "sess-fake",
"uuid": "u1",
"message": {
"role": "assistant",
"model": "claude-test",
"content": list(blocks),
"stop_reason": stop,
"usage": {"input_tokens": 1, "output_tokens": 2},
},
}
def _result(**over: Any) -> dict[str, Any]:
base: dict[str, Any] = {
"type": "result",
"subtype": "success",
"is_error": False,
"stop_reason": "end_turn",
"num_turns": 1,
"duration_ms": 1234,
"session_id": "sess-fake",
"usage": {"input_tokens": 10, "output_tokens": 20},
}
base.update(over)
return base
def _stream(etype: str, index: int | None = None, **extra: Any) -> dict[str, Any]:
event: dict[str, Any] = {"type": etype, **extra}
if index is not None:
event["index"] = index
return {
"type": "stream_event",
"session_id": "sess-fake",
"parent_tool_use_id": None,
"event": event,
}
async def _drain(tm: StreamTurnManager, text: str = "hi") -> list[Any]:
return [ev async for ev in tm.send_user_message(text)]
# --- argv ----------------------------------------------------------------
def test_build_stream_argv_is_headless_and_carries_common_flags() -> None:
opts = PtyProcessOptions(
cwd="/vault",
session_id="SID",
model="claude-opus-5",
system_prompt="SP",
effort="high",
dangerously_skip_permissions=True,
)
argv = build_stream_argv(opts, "SID")
assert argv[:7] == [
"claude",
"-p",
"--input-format",
"stream-json",
"--output-format",
"stream-json",
"--verbose",
]
assert "--include-partial-messages" not in argv
assert argv[argv.index("--session-id") + 1] == "SID"
assert argv[argv.index("--model") + 1] == "claude-opus-5"
assert argv[argv.index("--system-prompt") + 1] == "SP"
assert argv[argv.index("--effort") + 1] == "high"
assert "--dangerously-skip-permissions" in argv
def test_build_stream_argv_resume_wins_over_session_id() -> None:
opts = PtyProcessOptions(cwd="/vault", resume_session_id="RID")
argv = build_stream_argv(opts, "ignored")
assert "--session-id" not in argv
assert argv[argv.index("--resume") + 1] == "RID"
def test_build_stream_argv_include_partial_messages() -> None:
opts = PtyProcessOptions(cwd="/vault", session_id="SID")
argv = build_stream_argv(opts, "SID", include_partial_messages=True)
assert "--include-partial-messages" in argv
def test_strip_tui_only_flags_drops_interactive_escape_hatches() -> None:
kept = strip_tui_only_flags(
("--remote-control", "--keep", "--worktree=wt", "--ide", "--also-keep")
)
assert kept == ("--keep", "--also-keep")
def test_build_stream_argv_strips_tui_only_extra_args() -> None:
opts = PtyProcessOptions(
cwd="/vault", session_id="SID", extra_args=("--remote-control", "--keep")
)
argv = build_stream_argv(opts, "SID")
assert "--remote-control" not in argv
assert "--keep" in argv
# --- turn loop -----------------------------------------------------------
async def test_result_record_closes_the_turn_and_is_not_synthesized() -> None:
proc = FakeStreamProcess(
[_assistant({"type": "text", "text": "hello"}), _result()]
)
tm = StreamTurnManager(proc) # type: ignore[arg-type]
await tm.start()
events = await _drain(tm)
assert [type(e).__name__ for e in events] == ["AssistantMessage", "ResultMessage"]
result = events[-1]
assert isinstance(result, ResultMessage)
# Straight off the wire, not fabricated from the last assistant.
assert result.duration_ms == 1234
assert result.stop_reason == "end_turn"
assert result.usage == {"input_tokens": 10, "output_tokens": 20}
assert result.is_error is False
assert proc.sent == ["hi"]
async def test_assistant_stop_reason_null_does_not_end_the_turn() -> None:
"""The distinguishing property of this transport.
Headless claude leaves `stop_reason` null on every assistant record,
so a terminal-stop_reason heuristic would either end the turn at the
first record or never end it. Only `result` counts.
"""
proc = FakeStreamProcess(
[
_assistant({"type": "tool_use", "id": "t1", "name": "Bash", "input": {}}),
{
"type": "user",
"session_id": "sess-fake",
"message": {
"role": "user",
"content": [
{"type": "tool_result", "tool_use_id": "t1", "content": "ok"}
],
},
},
_assistant({"type": "text", "text": "done"}),
_result(num_turns=2),
]
)
tm = StreamTurnManager(proc) # type: ignore[arg-type]
await tm.start()
events = await _drain(tm)
assert [type(e).__name__ for e in events] == [
"AssistantMessage",
"UserMessage",
"AssistantMessage",
"ResultMessage",
]
assert isinstance(events[1], UserMessage)
assert isinstance(events[2], AssistantMessage)
assert events[2].content[0].text == "done" # type: ignore[union-attr]
async def test_bookkeeping_system_records_are_dropped() -> None:
proc = FakeStreamProcess(
[
{"type": "system", "subtype": "init", "session_id": "s"},
{"type": "system", "subtype": "hook_started", "session_id": "s"},
{"type": "system", "subtype": "status", "session_id": "s"},
{"type": "rate_limit_event", "rate_limit_info": {"status": "allowed"}},
_assistant({"type": "text", "text": "x"}),
_result(),
]
)
tm = StreamTurnManager(proc) # type: ignore[arg-type]
await tm.start()
events = await _drain(tm)
assert [type(e).__name__ for e in events] == ["AssistantMessage", "ResultMessage"]
async def test_partial_messages_are_suppressed_by_default() -> None:
proc = FakeStreamProcess(
[
_stream("content_block_delta", 0, delta={"type": "text_delta", "text": "a"}),
_assistant({"type": "text", "text": "a"}),
_result(),
]
)
tm = StreamTurnManager(proc) # type: ignore[arg-type]
await tm.start()
events = await _drain(tm)
assert not any(isinstance(e, StreamEvent) for e in events)
async def test_partial_message_indices_are_rebased_across_requests() -> None:
"""A turn spans several API requests; the caller sees one envelope.
Each request numbers its content blocks from zero, so without a
rebase the second request's block 0 would collide with the first
request's block 0 in the consumer's accumulator.
"""
proc = FakeStreamProcess(
[
# request 1: two blocks (thinking + tool_use)
_stream("message_start", None),
_stream("content_block_start", 0, content_block={"type": "thinking"}),
_stream("content_block_stop", 0),
_stream("content_block_start", 1, content_block={"type": "tool_use"}),
_stream("content_block_stop", 1),
_stream("message_delta", None),
_stream("message_stop", None),
# request 2: one block, numbered from zero again
_stream("message_start", None),
_stream("content_block_start", 0, content_block={"type": "text"}),
_stream(
"content_block_delta", 0, delta={"type": "text_delta", "text": "hi"}
),
_stream("content_block_stop", 0),
_stream("message_stop", None),
_result(),
],
)
tm = StreamTurnManager(proc, include_partial_messages=True) # type: ignore[arg-type]
await tm.start()
events = await _drain(tm)
partials = [e for e in events if isinstance(e, StreamEvent)]
assert [(e.event["type"], e.event["index"]) for e in partials] == [
("content_block_start", 0),
("content_block_stop", 0),
("content_block_start", 1),
("content_block_stop", 1),
# request 2's block 0 lands at 2, after request 1's two blocks
("content_block_start", 2),
("content_block_delta", 2),
("content_block_stop", 2),
]
# Inner envelopes never leak — the caller owns the outer one.
assert not any(
e.event["type"].startswith("message_") for e in partials
)
async def test_warmup_spends_a_turn_without_counting_it() -> None:
"""The warm-up must not be visible in the conversation's numbering.
It is a handshake with the CLI, not something the caller asked for,
so `ResultMessage.num_turns` on the first real reply should say 1.
"""
proc = FakeStreamProcess(
[
_assistant({"type": "text", "text": "ready"}),
_result(),
_assistant({"type": "text", "text": "real answer"}),
_result(),
]
)
tm = StreamTurnManager(proc) # type: ignore[arg-type]
await tm.start()
await tm.warmup("warm me")
assert proc.sent == ["warm me"]
assert tm.turn_count == 0
await _drain(tm, "the real question")
assert proc.sent == ["warm me", "the real question"]
assert tm.turn_count == 1
async def test_warmup_failure_is_swallowed() -> None:
"""A failed warm-up must not take the session down with it."""
proc = FakeStreamProcess([], stderr="boom")
tm = StreamTurnManager(proc) # type: ignore[arg-type]
await tm.start()
await tm.warmup() # would raise ProcessError if it propagated
assert tm.turn_count == 0
async def test_unconnected_mcp_servers_are_logged(
caplog: Any,
) -> None:
"""A turn starting without its MCP tools has to be greppable.
Nothing else surfaces it: the model simply doesn't see those tools
and answers anyway, which reads as the agent forgetting a capability.
"""
import logging
proc = FakeStreamProcess(
[
{
"type": "system",
"subtype": "init",
"session_id": "s",
"tools": ["Bash"],
"mcp_servers": [
{"name": "telegram", "status": "pending"},
{"name": "firefly", "status": "connected"},
],
},
_assistant({"type": "text", "text": "x"}),
_result(),
]
)
tm = StreamTurnManager(proc) # type: ignore[arg-type]
await tm.start()
with caplog.at_level(logging.WARNING, logger="claude_code_api.stream"):
await _drain(tm)
assert "telegram:pending" in caplog.text
assert "firefly" not in caplog.text
# --- failures ------------------------------------------------------------
async def test_eof_before_result_raises_process_error() -> None:
proc = FakeStreamProcess([_assistant({"type": "text", "text": "partial"})])
tm = StreamTurnManager(proc) # type: ignore[arg-type]
await tm.start()
with pytest.raises(ProcessError, match="exited before completing the turn"):
await _drain(tm)
async def test_error_result_with_no_content_raises() -> None:
proc = FakeStreamProcess(
[_result(subtype="error_during_execution", is_error=True, result="boom")]
)
tm = StreamTurnManager(proc) # type: ignore[arg-type]
await tm.start()
with pytest.raises(ProcessError, match="error_during_execution"):
await _drain(tm)
async def test_error_result_after_content_is_surfaced_not_raised() -> None:
"""Partial output beats an exception that would discard it."""
proc = FakeStreamProcess(
[
_assistant({"type": "text", "text": "got this far"}),
_result(subtype="error_max_turns", is_error=True),
]
)
tm = StreamTurnManager(proc) # type: ignore[arg-type]
await tm.start()
events = await _drain(tm)
assert isinstance(events[-1], ResultMessage)
assert events[-1].is_error is True
assert events[-1].subtype == "error_max_turns"
async def test_auth_failure_is_classified_from_stderr() -> None:
proc = FakeStreamProcess([], stderr="API Error: 403 please run /login")
tm = StreamTurnManager(proc) # type: ignore[arg-type]
await tm.start()
with pytest.raises(AuthError):
await _drain(tm)
async def test_second_turn_reuses_the_same_process() -> None:
proc = FakeStreamProcess(
[
_assistant({"type": "text", "text": "one"}),
_result(),
_assistant({"type": "text", "text": "two"}),
_result(num_turns=1),
]
)
tm = StreamTurnManager(proc) # type: ignore[arg-type]
await tm.start()
await _drain(tm, "first")
await _drain(tm, "second")
assert proc.sent == ["first", "second"]
assert tm.turn_count == 2
async def test_concurrent_turns_are_rejected() -> None:
proc = FakeStreamProcess([_result()])
tm = StreamTurnManager(proc) # type: ignore[arg-type]
await tm.start()
gen = tm.send_user_message("a")
await anext(gen)
with pytest.raises(RuntimeError, match="turn is in progress"):
await _drain(tm, "b")
await gen.aclose()
# --- smoke (real claude) -------------------------------------------------
_SMOKE_ENV = "RUN_CLAUDE_SMOKE"
@pytest.mark.live
@pytest.mark.skipif(
not os.environ.get(_SMOKE_ENV), reason=f"set {_SMOKE_ENV}=1 to run against claude"
)
async def test_live_stream_transport_multi_turn(tmp_path: Any) -> None:
"""Two turns over one live headless claude, second recalling the first."""
from claude_code_api.backend import BackendOptions, ClaudeCodeBackend
opts = BackendOptions(
cwd=str(tmp_path),
model="sonnet",
system_prompt="Answer in one short sentence.",
dangerously_skip_permissions=True,
transport="stream_json",
include_partial_messages=True,
)
async with ClaudeCodeBackend(opts) as backend:
history: list[dict[str, Any]] = [
{"role": "user", "content": "Remember the codeword: OKAPI-13. Acknowledge."}
]
events = [ev async for ev in backend.complete(history)]
assert any(isinstance(e, StreamEvent) for e in events)
assert isinstance(events[-1], ResultMessage)
from claude_code_api.backend import synthesize_turn_messages
history = [*history, *synthesize_turn_messages(events)]
history.append({"role": "user", "content": "What was the codeword?"})
text = ""
async for ev in backend.complete(history):
if isinstance(ev, AssistantMessage):
text += "".join(
b.text for b in ev.content if hasattr(b, "text") # type: ignore[attr-defined]
)
assert "OKAPI" in text.upper()
# One process served both turns.
assert backend.live_session_count == 1