mirror of
https://github.com/turnstonelabs/turnstone.git
synced 2026-08-12 23:12:23 -06:00
test(832): port test_cancel to the folded seam; add hook + orphan + pre-dispatch pins
Provider-level armed fakes drive the REAL wrapper/consumer/drain path (the seam these tests exist to pin), with title generation quieted — the best-effort title lane consumed one-shot scripts once fakes moved to the provider level. The shared-ref architecture pins become their new-world equivalents (no shared _cancel_ref attribute; _cancel_stream lifecycle via the eager append), and three new pin classes land: on_first_append fires once and never for a superseded arrival; a force-cancelled generation's mid-stream death is never re-issued and touches no UI finalize; a pre-set Stop issues no request and mints no credential on a dynamically authenticated alias.
This commit is contained in:
+214
-133
@@ -25,6 +25,51 @@ from turnstone.core.trajectory import (
|
||||
)
|
||||
|
||||
|
||||
def _armed_provider(session, *streams, name="openai-compatible"):
|
||||
"""Install a provider fake for the post-#832 seam on *session*.
|
||||
|
||||
``create_streaming`` arms the
|
||||
per-attempt ``cancel_ref`` with a closeable handle (mirroring every
|
||||
real adapter's eager append — the creation-vs-midstream classifier and
|
||||
``cancel()``'s stream-close path both depend on it) and serves the
|
||||
given *streams* SEQUENTIALLY, one per create call, each consumed once —
|
||||
so side-effectful generator scripts (yield → ``session.cancel()`` →
|
||||
yield) drive the REAL wrapper+consumer+drain path these tests assert
|
||||
against, and a multi-turn send (tool round-trips) scripts one stream
|
||||
per turn. A create call beyond the script list fails LOUDLY — the
|
||||
old silent-empty-commit that used to absorb it is gone (the strict
|
||||
finish gate, ruled on #832).
|
||||
|
||||
Title generation is quieted on the session: with a provider-LEVEL
|
||||
fake, the best-effort title lane would otherwise consume the one-shot
|
||||
script before the main loop ran (pre-fold, the session-level seam
|
||||
patches hid that lane and its attempt failed silently against the
|
||||
MagicMock client).
|
||||
"""
|
||||
from turnstone.core.providers._protocol import ModelCapabilities
|
||||
|
||||
session._generate_title = lambda *a, **kw: None
|
||||
provider = MagicMock()
|
||||
provider.provider_name = name
|
||||
provider.get_capabilities.return_value = ModelCapabilities()
|
||||
provider.retryable_error_names = frozenset({"IncompleteStreamError"})
|
||||
handle = MagicMock()
|
||||
remaining = list(streams)
|
||||
|
||||
def _create(**kwargs):
|
||||
ref = kwargs.get("cancel_ref")
|
||||
if ref is not None:
|
||||
ref.append(handle)
|
||||
assert remaining, "scripted provider exhausted: send looped for more turns than scripted"
|
||||
stream = remaining.pop(0)
|
||||
return iter(stream) if not hasattr(stream, "__next__") else stream
|
||||
|
||||
provider.create_streaming = MagicMock(side_effect=_create)
|
||||
provider._armed_handle = handle
|
||||
session._provider = provider
|
||||
return provider
|
||||
|
||||
|
||||
class NullUI:
|
||||
"""UI adapter that records state changes and discards other output."""
|
||||
|
||||
@@ -156,18 +201,15 @@ class TestCancelEvent:
|
||||
|
||||
fake_stream = iter([FakeChunk(content_delta="Hello", finish_reason="stop")])
|
||||
|
||||
with (
|
||||
patch.object(session, "_create_stream_with_retry", return_value=fake_stream),
|
||||
patch.object(session, "_full_messages", return_value=[]),
|
||||
):
|
||||
session.send("test")
|
||||
_armed_provider(session, fake_stream)
|
||||
session.send("test")
|
||||
|
||||
# Should complete normally — cancel flag was cleared
|
||||
assert "idle" in ui.states
|
||||
|
||||
|
||||
class TestCancelDuringStreaming:
|
||||
"""Cancel while _stream_attempt is iterating chunks."""
|
||||
"""Cancel while the streaming consumer is observing chunks."""
|
||||
|
||||
def test_preserves_partial_content(self, tmp_db):
|
||||
"""Partial content already streamed should be preserved in messages."""
|
||||
@@ -191,11 +233,8 @@ class TestCancelDuringStreaming:
|
||||
session.cancel()
|
||||
yield FakeChunk(content_delta=" — this should not appear")
|
||||
|
||||
with (
|
||||
patch.object(session, "_create_stream_with_retry", return_value=cancelling_stream()),
|
||||
patch.object(session, "_full_messages", return_value=[]),
|
||||
):
|
||||
session.send("test")
|
||||
_armed_provider(session, cancelling_stream())
|
||||
session.send("test")
|
||||
|
||||
# Session should be idle (not error)
|
||||
assert ui.states[-1] == "idle"
|
||||
@@ -253,28 +292,18 @@ class TestCancelDuringToolExecution:
|
||||
finish_reason="tool_calls",
|
||||
)
|
||||
|
||||
call_count = 0
|
||||
|
||||
def fake_create_stream(msgs):
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
if call_count == 1:
|
||||
return stream_with_tool()
|
||||
# Should not be called a second time since cancel happens before phase 3
|
||||
raise AssertionError("Should not stream again after cancel")
|
||||
|
||||
def cancel_before_execute(tool_calls):
|
||||
"""Simulate cancel happening before tool execution."""
|
||||
session.cancel()
|
||||
raise GenerationCancelled()
|
||||
|
||||
with (
|
||||
patch.object(session, "_create_stream_with_retry", side_effect=fake_create_stream),
|
||||
patch.object(session, "_full_messages", return_value=[]),
|
||||
patch.object(session, "_execute_tools", side_effect=cancel_before_execute),
|
||||
):
|
||||
_armed_provider(session, stream_with_tool())
|
||||
with patch.object(session, "_execute_tools", side_effect=cancel_before_execute):
|
||||
session.send("run something")
|
||||
|
||||
# The stream must not be re-created after the cancel landed.
|
||||
assert session._provider.create_streaming.call_count == 1
|
||||
|
||||
# Session should be idle
|
||||
assert ui.states[-1] == "idle"
|
||||
# Cancelled tool calls should have synthesized results
|
||||
@@ -308,11 +337,8 @@ class TestCancelWhenIdle:
|
||||
provider_blocks: list = field(default_factory=list)
|
||||
|
||||
fake_stream = iter([FakeChunk(content_delta="ok", finish_reason="stop")])
|
||||
with (
|
||||
patch.object(session, "_create_stream_with_retry", return_value=fake_stream),
|
||||
patch.object(session, "_full_messages", return_value=[]),
|
||||
):
|
||||
session.send("hello")
|
||||
_armed_provider(session, fake_stream)
|
||||
session.send("hello")
|
||||
|
||||
# Should complete normally
|
||||
assistant_msgs = [m for m in dicts_from_turns(session.messages) if m["role"] == "assistant"]
|
||||
@@ -345,10 +371,8 @@ class TestCancelThreadSafety:
|
||||
time.sleep(2) # Simulate slow streaming
|
||||
yield FakeChunk(content_delta=" end", finish_reason="stop")
|
||||
|
||||
with (
|
||||
patch.object(session, "_create_stream_with_retry", return_value=slow_stream()),
|
||||
patch.object(session, "_full_messages", return_value=[]),
|
||||
):
|
||||
_armed_provider(session, slow_stream())
|
||||
with contextlib.nullcontext():
|
||||
# Run send() in a thread
|
||||
error = []
|
||||
|
||||
@@ -433,13 +457,17 @@ class TestStreamFlushBeforeToolCalls:
|
||||
finish_reason="tool_calls",
|
||||
)
|
||||
|
||||
# Two scripted turns: the tool round-trip makes send() loop, and
|
||||
# the follow-up turn must carry a real finish reason — the old
|
||||
# exhausted-iterator path was silently committed as an empty turn
|
||||
# by the pre-fold lax consumer; the strict finish gate (ruled,
|
||||
# #832) rejects it now.
|
||||
_armed_provider(
|
||||
session,
|
||||
stream_content_then_tool(),
|
||||
iter([FakeChunk(finish_reason="stop")]),
|
||||
)
|
||||
with (
|
||||
patch.object(
|
||||
session,
|
||||
"_create_stream_with_retry",
|
||||
return_value=stream_content_then_tool(),
|
||||
),
|
||||
patch.object(session, "_full_messages", return_value=[]),
|
||||
# Prevent real tool execution (e.g., bash) during this test.
|
||||
patch.object(session, "_execute_tools", return_value=([], None)),
|
||||
):
|
||||
@@ -483,9 +511,12 @@ class TestStreamAbort:
|
||||
session.cancel() # Should not raise
|
||||
assert session._cancel_event.is_set()
|
||||
|
||||
def test_cancel_ref_populated_after_first_chunk(self, tmp_db):
|
||||
"""_cancel_ref is populated by the provider after the first chunk
|
||||
arrives (lazy generator evaluation)."""
|
||||
def test_stream_handle_registered_during_stream_cleared_after(self, tmp_db):
|
||||
"""The per-attempt ref's eager append registers the SDK handle in
|
||||
``_cancel_stream`` before the first chunk is consumed — the handle
|
||||
``cancel()`` closes to unblock a stuck read — and send()'s finally
|
||||
clears it (#832: the shared ref list is gone; the handle slot is
|
||||
the surviving cancel surface)."""
|
||||
ui = NullUI()
|
||||
session = _make_session(ui=ui)
|
||||
|
||||
@@ -499,26 +530,21 @@ class TestStreamAbort:
|
||||
info_delta: str = ""
|
||||
provider_blocks: list = field(default_factory=list)
|
||||
|
||||
sdk_stream = MagicMock()
|
||||
seen: dict = {}
|
||||
|
||||
def fake_provider_stream():
|
||||
# Simulate provider appending to cancel_ref before first yield
|
||||
session._cancel_ref.append(sdk_stream)
|
||||
def observing_stream():
|
||||
# Runs at first next(): the armed handle must ALREADY be
|
||||
# registered (append happens inside create_streaming, before
|
||||
# the iterator is handed back).
|
||||
seen["handle_at_first_chunk"] = session._cancel_stream
|
||||
yield FakeChunk(content_delta="hi", finish_reason="stop")
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
session,
|
||||
"_create_stream_with_retry",
|
||||
return_value=fake_provider_stream(),
|
||||
),
|
||||
patch.object(session, "_full_messages", return_value=[]),
|
||||
):
|
||||
session.send("test")
|
||||
provider = _armed_provider(session, observing_stream())
|
||||
session.send("test")
|
||||
|
||||
assert seen["handle_at_first_chunk"] is provider._armed_handle
|
||||
# After stream completes, cancel_stream should be cleared
|
||||
assert session._cancel_stream is None
|
||||
assert len(session._cancel_ref) == 0
|
||||
|
||||
def test_transport_error_during_cancel_becomes_generation_cancelled(self, tmp_db):
|
||||
"""When cancel() closes the stream, the resulting transport error
|
||||
@@ -541,15 +567,8 @@ class TestStreamAbort:
|
||||
session._cancel_event.set()
|
||||
raise ConnectionError("stream closed")
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
session,
|
||||
"_create_stream_with_retry",
|
||||
return_value=stream_that_errors(),
|
||||
),
|
||||
patch.object(session, "_full_messages", return_value=[]),
|
||||
):
|
||||
session.send("test")
|
||||
_armed_provider(session, stream_that_errors())
|
||||
session.send("test")
|
||||
|
||||
# Should complete as cancelled, not error
|
||||
assert "idle" in ui.states
|
||||
@@ -582,28 +601,33 @@ class TestStreamAbort:
|
||||
yield FakeChunk(content_delta="Hello")
|
||||
raise ValueError("unexpected error")
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
session,
|
||||
"_create_stream_with_retry",
|
||||
return_value=stream_that_errors(),
|
||||
),
|
||||
patch.object(session, "_full_messages", return_value=[]),
|
||||
pytest.raises(ValueError, match="unexpected error"),
|
||||
):
|
||||
_armed_provider(session, stream_that_errors())
|
||||
with pytest.raises(ValueError, match="unexpected error"):
|
||||
session.send("test")
|
||||
|
||||
def test_check_cancelled_between_retries(self, tmp_db):
|
||||
"""_try_stream checks for cancellation between retry attempts."""
|
||||
def test_pre_set_cancel_no_mint_no_dispatch(self, tmp_db):
|
||||
"""A Stop already set when the streaming phase is entered issues NO
|
||||
request and — on a dynamically authenticated alias — mints NO
|
||||
credential (#832; the #972 pre-dispatch rule extended to the
|
||||
interactive path). Pre-fold, ``_try_stream`` resolved the backend
|
||||
token BEFORE its first cancel check, so an abandoned turn still
|
||||
paid a mint (on a cache miss: a network round trip under the
|
||||
cluster-wide advisory lock). Post-fold the resolver runs inside
|
||||
``model_turn`` after its entry abort read, and the ladder's own
|
||||
loop-top check precedes even the lane build."""
|
||||
session = _make_session()
|
||||
session.cancel()
|
||||
session._cancel_event.set()
|
||||
resolver = MagicMock(return_value=None)
|
||||
provider = _armed_provider(session, iter(()))
|
||||
|
||||
with pytest.raises(GenerationCancelled):
|
||||
session._try_stream(
|
||||
client=MagicMock(),
|
||||
model="test",
|
||||
msgs=[],
|
||||
)
|
||||
with (
|
||||
patch.object(session, "_model_backend_auth_token", resolver),
|
||||
pytest.raises(GenerationCancelled),
|
||||
):
|
||||
session._stream_response(0)
|
||||
|
||||
resolver.assert_not_called()
|
||||
provider.create_streaming.assert_not_called()
|
||||
|
||||
|
||||
class TestCancelRef:
|
||||
@@ -615,7 +639,7 @@ class TestCancelRef:
|
||||
mock_stream = MagicMock()
|
||||
assert session._cancel_stream is None
|
||||
|
||||
session._cancel_ref.append(mock_stream)
|
||||
_CancelRef(session).append(mock_stream)
|
||||
|
||||
assert session._cancel_stream is mock_stream
|
||||
|
||||
@@ -625,7 +649,7 @@ class TestCancelRef:
|
||||
session.cancel() # Set cancel event before stream is created
|
||||
|
||||
mock_stream = MagicMock()
|
||||
session._cancel_ref.append(mock_stream)
|
||||
_CancelRef(session).append(mock_stream)
|
||||
|
||||
mock_stream.close.assert_called_once()
|
||||
|
||||
@@ -634,7 +658,7 @@ class TestCancelRef:
|
||||
session = _make_session()
|
||||
mock_stream = MagicMock()
|
||||
|
||||
session._cancel_ref.append(mock_stream)
|
||||
_CancelRef(session).append(mock_stream)
|
||||
|
||||
mock_stream.close.assert_not_called()
|
||||
assert session._cancel_stream is mock_stream
|
||||
@@ -647,20 +671,24 @@ class TestCancelRef:
|
||||
mock_stream = MagicMock()
|
||||
mock_stream.close.side_effect = RuntimeError("already closed")
|
||||
|
||||
session._cancel_ref.append(mock_stream) # Should not raise
|
||||
_CancelRef(session).append(mock_stream) # Should not raise
|
||||
|
||||
def test_cancel_ref_is_cancel_ref_instance(self, tmp_db):
|
||||
"""ChatSession._cancel_ref is a _CancelRef instance."""
|
||||
def test_no_shared_cancel_ref_attribute(self, tmp_db):
|
||||
"""The long-lived shared ref is GONE (#832): every model-call site
|
||||
builds a fresh per-attempt, generation-scoped _CancelRef, so a
|
||||
force-cancelled generation's ref reads aborted via supersession.
|
||||
This pin holds the line against the gen-0 shared instance coming
|
||||
back through a fixture or a convenience refactor."""
|
||||
session = _make_session()
|
||||
assert isinstance(session._cancel_ref, _CancelRef)
|
||||
assert not hasattr(session, "_cancel_ref")
|
||||
|
||||
def test_cancel_ref_cleared_after_stream_ends(self, tmp_db):
|
||||
"""_cancel_ref is cleared in the send() finally block after streaming."""
|
||||
"""send()'s finally clears _cancel_stream after streaming (the
|
||||
handle registered by the per-attempt ref's eager append must not
|
||||
linger into tool execution, where cancel() would close a dead
|
||||
handle instead of nothing)."""
|
||||
ui = NullUI()
|
||||
session = _make_session(ui=ui)
|
||||
mock_stream = MagicMock()
|
||||
session._cancel_ref.append(mock_stream)
|
||||
assert len(session._cancel_ref) == 1
|
||||
|
||||
@dataclass
|
||||
class FakeChunk:
|
||||
@@ -672,18 +700,87 @@ class TestCancelRef:
|
||||
info_delta: str = ""
|
||||
provider_blocks: list = field(default_factory=list)
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
session,
|
||||
"_create_stream_with_retry",
|
||||
return_value=iter([FakeChunk(content_delta="hi", finish_reason="stop")]),
|
||||
),
|
||||
patch.object(session, "_full_messages", return_value=[]),
|
||||
):
|
||||
session.send("test")
|
||||
_armed_provider(session, iter([FakeChunk(content_delta="hi", finish_reason="stop")]))
|
||||
session.send("test")
|
||||
|
||||
# After send() completes, _cancel_ref is cleared in the finally block
|
||||
assert len(session._cancel_ref) == 0
|
||||
# During the stream the armed handle was registered; after send()
|
||||
# completes the finally block cleared it.
|
||||
assert session._cancel_stream is None
|
||||
|
||||
|
||||
class TestOnFirstAppendHook:
|
||||
"""``_CancelRef.on_first_append`` — the request-accepted observation
|
||||
point (#832): health success, the creation-vs-midstream classifier,
|
||||
and the usage-slot resets all key on it."""
|
||||
|
||||
def test_fires_once_and_arms(self, tmp_db):
|
||||
session = _make_session()
|
||||
fired = []
|
||||
ref = _CancelRef(session, on_first_append=lambda: fired.append(1))
|
||||
assert not ref.armed
|
||||
ref.append(MagicMock())
|
||||
ref.append(MagicMock())
|
||||
assert fired == [1]
|
||||
assert ref.armed
|
||||
|
||||
def test_superseded_append_neither_fires_nor_registers(self, tmp_db):
|
||||
"""A superseded ref's zombie stream is closed on arrival and must
|
||||
neither fire the hook (an orphan must not record health or reset
|
||||
the successor's usage slots) nor hijack ``_cancel_stream`` from
|
||||
the successor generation's live stream."""
|
||||
session = _make_session()
|
||||
session._generation = 5
|
||||
fired = []
|
||||
ref = _CancelRef(session, 4, on_first_append=lambda: fired.append(1))
|
||||
zombie = MagicMock()
|
||||
ref.append(zombie)
|
||||
assert fired == []
|
||||
assert not ref.armed
|
||||
zombie.close.assert_called_once()
|
||||
assert session._cancel_stream is None
|
||||
|
||||
|
||||
class TestForceCancelOrphanNoReissue:
|
||||
"""A force-cancelled generation's mid-stream death is never re-issued
|
||||
on the orphan's behalf (#832 regression pin: the old shared gen-0 ref
|
||||
read ``aborted`` False after a force-cancel — the successor installs a
|
||||
fresh unset event and gen 0 never reads superseded — so the retry
|
||||
machinery would have spent tokens for a generation nobody owns)."""
|
||||
|
||||
def test_orphan_death_not_reissued_no_ui_finalize(self, tmp_db):
|
||||
ui = NullUI()
|
||||
session = _make_session(ui=ui)
|
||||
|
||||
@dataclass
|
||||
class FakeChunk:
|
||||
content_delta: str = ""
|
||||
reasoning_delta: str = ""
|
||||
tool_call_deltas: list = field(default_factory=list)
|
||||
usage: None = None
|
||||
finish_reason: str = ""
|
||||
info_delta: str = ""
|
||||
provider_blocks: list = field(default_factory=list)
|
||||
|
||||
def dying_orphan_stream():
|
||||
yield FakeChunk(content_delta="old ")
|
||||
# Force-cancel: a successor claims the generation (bumped
|
||||
# counter + fresh UNSET event) while this stream is mid-body —
|
||||
# exactly the state where the old shared ref read aborted=False.
|
||||
session._claim_generation()
|
||||
raise ConnectionError("connection reset")
|
||||
|
||||
provider = _armed_provider(session, dying_orphan_stream())
|
||||
my_generation = session._claim_generation()
|
||||
ends_before = ui.stream_ends
|
||||
|
||||
with pytest.raises(GenerationCancelled):
|
||||
session._stream_response(my_generation)
|
||||
|
||||
# ONE dispatch — the death was not re-issued for the orphan.
|
||||
assert provider.create_streaming.call_count == 1
|
||||
# The orphan touched no UI finalize (a stream_end here would
|
||||
# clobber the successor generation's in-flight stream state).
|
||||
assert ui.stream_ends == ends_before
|
||||
|
||||
|
||||
class TestForceCancelGeneration:
|
||||
@@ -733,15 +830,8 @@ class TestForceCancelGeneration:
|
||||
|
||||
original_event = session._cancel_event
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
session,
|
||||
"_create_stream_with_retry",
|
||||
return_value=iter([FakeChunk(content_delta="hi", finish_reason="stop")]),
|
||||
),
|
||||
patch.object(session, "_full_messages", return_value=[]),
|
||||
):
|
||||
session.send("test")
|
||||
_armed_provider(session, iter([FakeChunk(content_delta="hi", finish_reason="stop")]))
|
||||
session.send("test")
|
||||
|
||||
# After send() completes, _cancel_event should be a NEW Event
|
||||
# (not the same object as before the call).
|
||||
@@ -778,10 +868,8 @@ class TestForceCancelThreaded:
|
||||
yield FakeChunk(content_delta=" more", finish_reason="stop")
|
||||
|
||||
# Start generation 1 (will get stuck)
|
||||
with (
|
||||
patch.object(session, "_create_stream_with_retry", return_value=slow_stream()),
|
||||
patch.object(session, "_full_messages", return_value=[]),
|
||||
):
|
||||
_armed_provider(session, slow_stream())
|
||||
if True:
|
||||
|
||||
def run_old():
|
||||
with contextlib.suppress(Exception):
|
||||
@@ -832,24 +920,17 @@ class TestForceCancelThreaded:
|
||||
yield FakeChunk(content_delta=" end", finish_reason="stop")
|
||||
|
||||
# Start stuck generation
|
||||
with (
|
||||
patch.object(session, "_create_stream_with_retry", return_value=stuck_stream()),
|
||||
patch.object(session, "_full_messages", return_value=[]),
|
||||
):
|
||||
t = threading.Thread(target=lambda: session.send("old"), daemon=True)
|
||||
t.start()
|
||||
assert barrier.wait(timeout=5), "stream did not start"
|
||||
_armed_provider(session, stuck_stream())
|
||||
t = threading.Thread(target=lambda: session.send("old"), daemon=True)
|
||||
t.start()
|
||||
assert barrier.wait(timeout=5), "stream did not start"
|
||||
|
||||
# Force cancel
|
||||
session.cancel()
|
||||
|
||||
# New generation should work
|
||||
fresh_stream = iter([FakeChunk(content_delta="Fresh response")])
|
||||
with (
|
||||
patch.object(session, "_create_stream_with_retry", return_value=fresh_stream),
|
||||
patch.object(session, "_full_messages", return_value=[]),
|
||||
):
|
||||
session.send("new message")
|
||||
_armed_provider(session, iter([FakeChunk(content_delta="Fresh response")]))
|
||||
session.send("new message")
|
||||
|
||||
# The new generation should have completed successfully
|
||||
assert "idle" in ui.states
|
||||
|
||||
Reference in New Issue
Block a user