unsloth/tests/studio/test_stream_cancel_registration_timing.py
Daniel Han a6dc10dad2
Reduce and tighten comments and docstrings across the test suite (#6429)
* Reduce and tighten comments and docstrings in tests

Shorten verbose comments and docstrings across the test suite without
changing any test logic. Remove narration that restates the next line,
collapse long module and test docstrings to a single line, and drop banner
separators. Keep regression context (issue and PR references, run ids),
skip reasons, mocking and timing rationale, license headers, lint and type
directives, and commented-out code.

Comments and docstrings only: an AST signature check confirms no code,
assertions, or string literals changed, and the suite byte-compiles cleanly.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
2026-06-18 01:07:09 -07:00

684 lines
22 KiB
Python

"""Cancel tracker must register BEFORE StreamingResponse returns and clean up
in each async generator's `finally`, else a Stop before the first SSE chunk
leaves a zombie decode (a BackgroundTask would be skipped when stream_response raises).
Structural verifies registration placement and try/finally cleanup; behavioral
verifies the extracted `_TrackedCancel` cleans up across completion/OSError/aclose
and that a pre-set cancel_event breaks the GGUF loop cleanly with final_chunk + [DONE]."""
from __future__ import annotations
import ast
import asyncio
import threading
import time
from pathlib import Path
SOURCE_PATH = Path(__file__).resolve().parents[2] / "studio" / "backend" / "routes" / "inference.py"
SRC = SOURCE_PATH.read_text()
_TREE = ast.parse(SRC)
# ── Structural (AST) helpers ─────────────────────────────────
def _collect_async_functions(tree: ast.AST):
return [n for n in ast.walk(tree) if isinstance(n, ast.AsyncFunctionDef)]
def _has_tracker_enter_call(node: ast.AST) -> bool:
for sub in ast.walk(node):
if not isinstance(sub, ast.Call):
continue
fn = sub.func
if (
isinstance(fn, ast.Attribute)
and fn.attr == "__enter__"
and isinstance(fn.value, ast.Name)
and fn.value.id.startswith("_tracker")
):
return True
return False
def _finalbody_has_tracker_exit(finalbody) -> bool:
for stmt in finalbody:
if not isinstance(stmt, ast.Expr):
continue
call = stmt.value
if not (isinstance(call, ast.Call) and isinstance(call.func, ast.Attribute)):
continue
fn = call.func
if (
fn.attr == "__exit__"
and isinstance(fn.value, ast.Name)
and fn.value.id.startswith("_tracker")
):
return True
return False
def _async_function(name: str) -> ast.AsyncFunctionDef:
for node in ast.walk(_TREE):
if isinstance(node, ast.AsyncFunctionDef) and node.name == name:
return node
raise AssertionError(f"{name} handler missing")
def _calls_name(node: ast.AST, name: str) -> bool:
for sub in ast.walk(node):
if isinstance(sub, ast.Call) and isinstance(sub.func, ast.Name):
if sub.func.id == name:
return True
return False
# ── Structural tests ─────────────────────────────────────────
def test_no_tracker_enter_inside_async_generators():
offenders = []
for fn in _collect_async_functions(_TREE):
if fn.name in {
"gguf_tool_stream",
"gguf_stream_chunks",
"stream_chunks",
"audio_input_stream",
}:
if _has_tracker_enter_call(fn):
offenders.append(fn.name)
assert not offenders, (
f"Cancel tracker registration must live OUTSIDE the async generator "
f"body so a stop POST can find the registry entry before the first "
f"SSE chunk. Offending generators: {offenders}"
)
def test_tracker_enter_exists_in_sync_body_of_chat_completions():
top = None
for n in ast.walk(_TREE):
if isinstance(n, ast.AsyncFunctionDef) and n.name == "openai_chat_completions":
top = n
break
assert top is not None, "openai_chat_completions handler missing"
count = 0
for sub in ast.walk(top):
if not isinstance(sub, ast.Call):
continue
fn = sub.func
if (
isinstance(fn, ast.Attribute)
and fn.attr == "__enter__"
and isinstance(fn.value, ast.Name)
and fn.value.id.startswith("_tracker")
):
count += 1
assert count >= 3, (
f"expected >=3 _tracker.__enter__() calls in openai_chat_completions "
f"(one per streaming path), got {count}"
)
def test_async_generators_cleanup_tracker_in_finally():
required = {
"gguf_tool_stream",
"gguf_stream_chunks",
"stream_chunks",
"audio_input_stream",
}
found: set[str] = set()
for fn in [n for n in ast.walk(_TREE) if isinstance(n, ast.AsyncFunctionDef)]:
if fn.name not in required:
continue
for sub in ast.walk(fn):
if isinstance(sub, ast.Try) and sub.finalbody:
if _finalbody_has_tracker_exit(sub.finalbody):
found.add(fn.name)
break
missing = required - found
assert not missing, (
f"Cleanup must run via `finally: _tracker.__exit__(None, None, None)` "
f"inside each streaming generator so ClientDisconnect / OSError paths "
f"also release registry entries (Starlette skips `background` callbacks "
f"when stream_response raises). Missing in: {sorted(missing)}"
)
def test_streaming_responses_have_no_background_task():
top = None
for n in ast.walk(_TREE):
if isinstance(n, ast.AsyncFunctionDef) and n.name == "openai_chat_completions":
top = n
break
assert top is not None
for sub in ast.walk(top):
if not (isinstance(sub, ast.Call) and isinstance(sub.func, ast.Name)):
continue
if sub.func.id != "StreamingResponse":
continue
kwargs = {kw.arg for kw in sub.keywords if kw.arg}
assert "background" not in kwargs, (
"StreamingResponse in openai_chat_completions must not pass "
"`background=` -- cleanup now lives in the generator's finally "
"block; a BackgroundTask would be skipped on abrupt disconnect"
)
def test_direct_llama_server_streams_install_disconnect_watcher():
required = {
"openai_completions",
"_responses_stream",
"_anthropic_passthrough_stream",
"_openai_passthrough_stream",
}
missing = [
name
for name in sorted(required)
if not _calls_name(_async_function(name), "_await_disconnect_then_close")
]
assert not missing, (
"Direct httpx streams to llama-server must close the upstream response "
"when the downstream client disconnects during prefill. Missing in: "
f"{missing}"
)
# ── Behavioral helpers ───────────────────────────────────────
_WANTED = {
"_CANCEL_REGISTRY",
"_CANCEL_LOCK",
"_PENDING_CANCELS",
"_PENDING_CANCEL_TTL_S",
"_prune_pending",
"_TrackedCancel",
"_cancel_by_keys",
"_cancel_by_cancel_id_or_stash",
}
def _load_registry_module():
chunks = []
for n in _TREE.body:
seg = ast.get_source_segment(SRC, n)
if seg is None:
continue
if isinstance(n, (ast.FunctionDef, ast.ClassDef)) and n.name in _WANTED:
chunks.append(seg)
elif isinstance(n, ast.Assign):
names = [t.id for t in n.targets if isinstance(t, ast.Name)]
if any(name in _WANTED for name in names):
chunks.append(seg)
elif (
isinstance(n, ast.AnnAssign)
and isinstance(n.target, ast.Name)
and n.target.id in _WANTED
):
chunks.append(seg)
mod = {}
exec("import threading, time\n" + "\n\n".join(chunks), mod)
return mod
def _make_stream(tracker, raise_exc):
async def gen():
try:
try:
yield "data: first\n\n"
if raise_exc is not None:
raise raise_exc
yield "data: [DONE]\n\n"
except asyncio.CancelledError:
raise
except Exception:
yield "data: error\n\n"
finally:
tracker.__exit__(None, None, None)
except BaseException:
raise
return gen()
async def _consume(agen):
out = []
try:
async for ch in agen:
out.append(ch)
except BaseException as e:
out.append(type(e).__name__)
return out
def _llama_stub_raises_on_preset_cancel(cancel_event):
# Reproduces llama_cpp.py _stream_with_retry `raise GeneratorExit` when
# cancel_event is already set at entry.
if cancel_event.is_set():
raise GeneratorExit
yield "cumulative-1"
yield "cumulative-2"
async def _post_fix_gguf_loop(cancel_event):
yield "first_chunk"
gen = _llama_stub_raises_on_preset_cancel(cancel_event)
sentinel = object()
while True:
if cancel_event.is_set():
break
cumulative = await asyncio.to_thread(next, gen, sentinel)
if cumulative is sentinel:
break
yield cumulative
yield "final_chunk"
yield "[DONE]"
# ── Behavioral tests ─────────────────────────────────────────
def test_finally_cleanup_on_normal_completion():
m = _load_registry_module()
m["_CANCEL_REGISTRY"].clear()
ev = threading.Event()
tr = m["_TrackedCancel"](ev, "cid-ok", "sid-ok")
tr.__enter__()
assert "cid-ok" in m["_CANCEL_REGISTRY"]
chunks = asyncio.run(_consume(_make_stream(tr, None)))
assert chunks == ["data: first\n\n", "data: [DONE]\n\n"]
assert "cid-ok" not in m["_CANCEL_REGISTRY"]
assert "sid-ok" not in m["_CANCEL_REGISTRY"]
def test_finally_cleanup_on_mid_stream_exception():
# OSError mid-stream: the exact case where pre-fix `background=BackgroundTask(...)`
# was skipped and leaked the registry entry.
m = _load_registry_module()
m["_CANCEL_REGISTRY"].clear()
ev = threading.Event()
tr = m["_TrackedCancel"](ev, "cid-err", "sid-err")
tr.__enter__()
assert "cid-err" in m["_CANCEL_REGISTRY"]
asyncio.run(_consume(_make_stream(tr, OSError("disconnect"))))
assert "cid-err" not in m["_CANCEL_REGISTRY"]
assert "sid-err" not in m["_CANCEL_REGISTRY"]
def test_finally_cleanup_on_aclose():
# Starlette calls aclose() on client disconnect; the finally block must run.
m = _load_registry_module()
m["_CANCEL_REGISTRY"].clear()
ev = threading.Event()
tr = m["_TrackedCancel"](ev, "cid-abort", "sid-abort")
tr.__enter__()
assert "cid-abort" in m["_CANCEL_REGISTRY"]
async def run():
gen = _make_stream(tr, None)
it = gen.__aiter__()
await it.__anext__()
await gen.aclose()
asyncio.run(run())
assert "cid-abort" not in m["_CANCEL_REGISTRY"]
assert "sid-abort" not in m["_CANCEL_REGISTRY"]
def test_preset_cancel_event_exits_cleanly_with_done():
# Pending-replay: a stashed cancel pre-set cancel_event. The loop must break
# cleanly with final_chunk + [DONE], not propagate GeneratorExit from the GGUF wrapper.
ev = threading.Event()
ev.set()
chunks = asyncio.run(_consume(_post_fix_gguf_loop(ev)))
assert "first_chunk" in chunks
assert "final_chunk" in chunks
assert "[DONE]" in chunks
assert "GeneratorExit" not in chunks
assert "cumulative-1" not in chunks
assert "cumulative-2" not in chunks
def test_normal_path_streams_all_tokens():
# Regression: the top-of-loop cancel_event check must not short-circuit when unset.
ev = threading.Event()
chunks = asyncio.run(_consume(_post_fix_gguf_loop(ev)))
assert chunks == ["first_chunk", "cumulative-1", "cumulative-2", "final_chunk", "[DONE]"]
def test_cancel_during_streaming_stops_iteration_promptly():
# Setting cancel_event between yields breaks on the next iteration, not draining the generator.
ev = threading.Event()
async def _run():
gen = _post_fix_gguf_loop(ev)
seen = []
async for ch in gen:
seen.append(ch)
if ch == "cumulative-1":
ev.set()
return seen
seen = asyncio.run(_run())
assert "first_chunk" in seen
assert "cumulative-1" in seen
assert "cumulative-2" not in seen
assert "final_chunk" in seen
assert "[DONE]" in seen
# ── Cancel-event responsiveness in the streaming loops ───────
def _loop_has_cancel_event_check(fn) -> bool:
# An `if cancel_event.is_set():` inside a loop body is sufficient -- without it
# a cancel POST can't interrupt, since Colab-style proxies drop request.is_disconnected().
for sub in ast.walk(fn):
if not isinstance(sub, (ast.While, ast.For, ast.AsyncFor)):
continue
for stmt in ast.walk(sub):
if not isinstance(stmt, ast.If):
continue
t = stmt.test
if (
isinstance(t, ast.Call)
and isinstance(t.func, ast.Attribute)
and t.func.attr == "is_set"
and isinstance(t.func.value, ast.Name)
and t.func.value.id == "cancel_event"
):
return True
return False
def test_streaming_generators_check_cancel_event_in_loop():
required = {
"gguf_tool_stream",
"gguf_stream_chunks",
"stream_chunks",
"audio_input_stream",
}
missing = []
for fn in [n for n in ast.walk(_TREE) if isinstance(n, ast.AsyncFunctionDef)]:
if fn.name not in required:
continue
if not _loop_has_cancel_event_check(fn):
missing.append(fn.name)
assert not missing, (
f"Each streaming generator must check `cancel_event.is_set()` inside "
f"its main loop so `POST /api/inference/cancel` can interrupt the "
f"stream through proxies that do not forward fetch aborts. "
f"Missing in: {sorted(missing)}"
)
def test_audio_input_stream_offloads_blocking_next_to_thread():
# Guards against regressing to `for chunk_text in audio_input_generate():`, which
# blocks the event loop per whisper chunk and stalls POST /api/inference/cancel.
audio = None
for fn in ast.walk(_TREE):
if isinstance(fn, ast.AsyncFunctionDef) and fn.name == "audio_input_stream":
audio = fn
break
assert audio is not None, "audio_input_stream generator missing"
for sub in ast.walk(audio):
if isinstance(sub, (ast.For, ast.AsyncFor)):
it_src = ast.unparse(sub.iter)
assert "audio_input_generate" not in it_src, (
"audio_input_stream must not iterate audio_input_generate() "
"directly -- that blocks the event loop. Use "
"`await asyncio.to_thread(next, gen, _DONE)` inside a "
"`while True` loop instead"
)
found_to_thread_next = False
for sub in ast.walk(audio):
if not isinstance(sub, ast.Call):
continue
fn_expr = sub.func
if not (
isinstance(fn_expr, ast.Attribute)
and fn_expr.attr == "to_thread"
and isinstance(fn_expr.value, ast.Name)
and fn_expr.value.id == "asyncio"
):
continue
if sub.args and isinstance(sub.args[0], ast.Name) and sub.args[0].id == "next":
found_to_thread_next = True
break
assert found_to_thread_next, (
"audio_input_stream must call `asyncio.to_thread(next, gen, ...)` "
"to keep the event loop free while whisper yields the next chunk"
)
def test_stream_chunks_cancel_branch_resets_backend_state():
# The cancel branch must call backend.reset_generation_state() to flush
# GPU/KV-cache state, else cancel-via-POST leaves the subprocess dirty.
fn = None
top = None
for n in ast.walk(_TREE):
if isinstance(n, ast.AsyncFunctionDef) and n.name == "openai_chat_completions":
top = n
break
assert top is not None
for n in ast.walk(top):
if isinstance(n, ast.AsyncFunctionDef) and n.name == "stream_chunks":
fn = n
break
assert fn is not None, "stream_chunks generator missing"
for sub in ast.walk(fn):
if not isinstance(sub, ast.If):
continue
t = sub.test
if not (
isinstance(t, ast.Call)
and isinstance(t.func, ast.Attribute)
and t.func.attr == "is_set"
and isinstance(t.func.value, ast.Name)
and t.func.value.id == "cancel_event"
):
continue
body_src = "\n".join(ast.unparse(s) for s in sub.body)
if "backend.reset_generation_state()" in body_src:
return
raise AssertionError(
"stream_chunks `if cancel_event.is_set():` branch must call "
"backend.reset_generation_state() -- matches the existing "
"request.is_disconnected() / CancelledError cleanup paths and "
"prevents KV-cache drift after cancel-via-POST"
)
# ── Behavioral simulations for the iter-1 fixes ──────────────
def test_unsloth_stream_loop_breaks_on_external_cancel_event():
cancel_event = threading.Event()
reset_calls = [0]
class _Backend:
def reset_generation_state(self):
reset_calls[0] += 1
backend = _Backend()
def _generate():
for i in range(200):
time.sleep(0.005)
yield f"cum-{i}"
async def _loop():
_DONE = object()
loop = asyncio.get_event_loop()
gen = _generate()
seen = []
while True:
if cancel_event.is_set():
backend.reset_generation_state()
break
cumulative = await loop.run_in_executor(None, next, gen, _DONE)
if cumulative is _DONE:
break
seen.append(cumulative)
return seen
async def _fire():
await asyncio.sleep(0.05)
cancel_event.set()
async def _main():
return await asyncio.gather(_loop(), _fire())
seen, _ = asyncio.run(_main())
assert (
len(seen) < 200
), f"loop must not drain the generator after cancel; got {len(seen)} tokens"
assert reset_calls[0] == 1, (
f"backend.reset_generation_state() must be called exactly once on "
f"cancel-via-POST, got {reset_calls[0]}"
)
def test_audio_stream_stays_responsive_under_blocking_next():
# Assert the pre-fix `for chunk in audio_input_generate()` pattern blocks the
# event loop, then confirm the post-fix pattern exits promptly.
cancel_event = threading.Event()
def _audio_gen():
for i in range(8):
time.sleep(0.15)
yield f"chunk-{i}"
async def _prefix_loop():
seen = []
for chunk_text in _audio_gen():
if cancel_event.is_set():
break
seen.append(chunk_text)
return seen
async def _postfix_loop():
_DONE = object()
gen = _audio_gen()
seen = []
while True:
if cancel_event.is_set():
break
chunk_text = await asyncio.to_thread(next, gen, _DONE)
if chunk_text is _DONE:
break
seen.append(chunk_text)
return seen
async def _fire_early():
await asyncio.sleep(0.05)
cancel_event.set()
async def _run(loop_coro):
return await asyncio.gather(loop_coro, _fire_early())
cancel_event.clear()
t0 = time.monotonic()
prefix_seen, _ = asyncio.run(_run(_prefix_loop()))
prefix_elapsed = time.monotonic() - t0
assert prefix_elapsed >= 0.13, (
f"pre-fix pattern should block event loop for >=1 chunk time "
f"(~150ms); got {prefix_elapsed:.3f}s, {len(prefix_seen)} chunks"
)
cancel_event.clear()
t0 = time.monotonic()
postfix_seen, _ = asyncio.run(_run(_postfix_loop()))
postfix_elapsed = time.monotonic() - t0
assert postfix_elapsed < prefix_elapsed, (
f"post-fix pattern must exit faster than pre-fix (blocking) "
f"pattern; post={postfix_elapsed:.3f}s vs pre={prefix_elapsed:.3f}s"
)
assert (
len(postfix_seen) < 8
), f"post-fix loop must not drain all chunks; got {len(postfix_seen)}"
def test_unsloth_stream_loop_emits_zero_tokens_on_preset_cancel():
# Pending-cancel replay: cancel_event pre-set, so the top-of-loop check must
# short-circuit iteration 1 (zero tokens). Catches moving the check below next().
cancel_event = threading.Event()
cancel_event.set()
reset_calls = [0]
class _Backend:
def reset_generation_state(self):
reset_calls[0] += 1
backend = _Backend()
next_calls = [0]
def _generate():
while True:
next_calls[0] += 1
yield f"cum-{next_calls[0]}"
async def _loop():
_DONE = object()
loop = asyncio.get_event_loop()
gen = _generate()
seen = []
while True:
if cancel_event.is_set():
backend.reset_generation_state()
break
cumulative = await loop.run_in_executor(None, next, gen, _DONE)
if cumulative is _DONE:
break
seen.append(cumulative)
return seen
seen = asyncio.run(_loop())
assert seen == [], (
f"loop must emit zero tokens when cancel_event is pre-set "
f"(pending-replay path); got {seen}"
)
assert next_calls[0] == 0, (
f"loop must not call next() at all on pre-set cancel; got " f"{next_calls[0]} calls"
)
assert reset_calls[0] == 1, (
f"backend.reset_generation_state() must still fire exactly once "
f"on pre-set cancel; got {reset_calls[0]}"
)
def test_audio_stream_emits_zero_chunks_on_preset_cancel():
# Symmetric to the Unsloth pre-set test: the audio loop must skip
# asyncio.to_thread(next, ...) when cancel_event was pre-set via pending-replay.
cancel_event = threading.Event()
cancel_event.set()
next_calls = [0]
def _audio_gen():
while True:
next_calls[0] += 1
yield f"chunk-{next_calls[0]}"
async def _loop():
_DONE = object()
gen = _audio_gen()
seen = []
while True:
if cancel_event.is_set():
break
chunk_text = await asyncio.to_thread(next, gen, _DONE)
if chunk_text is _DONE:
break
seen.append(chunk_text)
return seen
seen = asyncio.run(_loop())
assert seen == [], f"audio loop must emit zero chunks on pre-set cancel; got {seen}"
assert next_calls[0] == 0, (
f"audio loop must not call next() on pre-set cancel; got " f"{next_calls[0]} calls"
)