unsloth/tests/studio/test_stream_cancel_registration_timing.py
Anmol Mishra 9f694ab750
fix(studio): Windows GGUF cancel hang + CPU spinlock overhead (#5692) (#5749)
* fix(studio): Windows GGUF cancel hang + CPU spinlock overhead (#5692)

Two fixes for Windows-native GGUF inference via llama-server:

**Issue 1 — GPU/CUDA Hang on Stream Cancellation:**
- Add `Connection: close` header to all httpx requests proxying to
  llama-server, preventing Keep-Alive from masking downstream socket
  closure.
- Introduce `_await_disconnect_then_close` background watcher that
  polls `request.is_disconnected()` every 100ms and calls
  `resp.aclose()` immediately when the client disconnects. This runs
  alongside the existing cancel-POST watcher and covers client aborts
  that never reach the /cancel endpoint (tab close, proxy aborts,
  Colab, mobile navigation, etc.).
- Change all StreamingResponse `Connection: keep-alive` headers to
  `Connection: close`.

**Issue 2 — High CPU Spinlock & KV Cache Backup Overhead:**
- Set OMP_WAIT_POLICY=PASSIVE and OMP_NUM_THREADS=2 in the
  llama-server subprocess environment on Windows to prevent OpenMP
  from spin-waiting on all logical cores while the GPU decodes.
- Limit `--threads` to 2 on Windows when the model is fully
  GPU-offloaded (`-ngl -1`). Auto-detect otherwise.
- Pass `--cache-ram 0 --ctx-checkpoints 0 --no-cache-prompt
  --checkpoint-every-n-tokens -1` on Windows to disable prompt-cache
  snapshots that copy KV cache to system RAM over the WDDM/PCI-E bus.

Closes #5692.

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

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

* fix: use local import to avoid ruff F823 (sys used before assignment)

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

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

* review: address gemini review feedback

- Simplify _fully_gpu_offloaded init: default to False, only set True
  in the gpu_indices branch, drop redundant else.
- Log exceptions in _await_disconnect_then_close at debug level instead
  of silent pass, per review suggestion.

* Adjust review feedback for PR #5749

- _await_disconnect_then_close: set cancel_event before resp.aclose() so
  the streamer's RemoteProtocolError handler treats the watcher-driven
  close as cancellation, not an upstream error. Both call sites pass
  cancel_event through.
- Windows --cache-ram / --no-cache-prompt / --ctx-checkpoints block: gate
  on _fully_gpu_offloaded so CPU and partial-offload Windows runs keep
  prompt-cache reuse across turns.
- Windows OMP_WAIT_POLICY / OMP_NUM_THREADS env: same gate so CPU and
  partial-offload Windows runs keep default OpenMP parallelism.

* Shorten code comments touched by PR #5749

* Clean up local imports and rename underscore locals in PR #5749

- Drop the function-local `import sys as _sys` introduced as an F823
  workaround; remove the redundant in-function `import os`/`import sys`
  block so module-level imports resolve sys/os instead. F823 no longer
  triggers because no shadowing import remains inside load_model.
- Rename `_fully_gpu_offloaded` and `_t` to `fully_gpu_offloaded` and
  `threads_arg`. Underscore-prefixed names usually mean private/module-
  level; plain locals match Python style for in-function temporaries.

No behavior change. ruff clean, py_compile clean, 35 studio cancel-
infra tests + 13 launch-gating AST locks + 6 disconnect-watcher locks
+ 4 spoof live-import tests all pass.

* Fix Windows GGUF follow-ups for PR #5749

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

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

* Fix cache flag gating for PR #5749

* Fix Python 3.9 annotations for PR #5749

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

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

---------

Co-authored-by: Anmol Mishra <anmolx.work@gmail.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Daniel Han <danielhanchen@gmail.com>
Co-authored-by: wasimysaid <wasimysdev@gmail.com>
Co-authored-by: Lee Jackson <130007945+Imagineer99@users.noreply.github.com>
2026-06-15 10:32:10 +01:00

707 lines
23 KiB
Python

"""Tests that the cancel tracker is registered BEFORE StreamingResponse is
returned and that cleanup runs in a `finally` inside each async generator.
Zombie scenario: Stop during prefill/warmup/proxy buffering (before the first
SSE chunk). If _tracker.__enter__ ran inside the generator body, the registry
would be empty when /api/inference/cancel lands, so cancel returns 0 and decode
runs to completion. The fix registers in the sync body of
openai_chat_completions and cleans up in each generator's `finally` -- a
BackgroundTask would be skipped when stream_response raises.
Structural verifies: no `_tracker.__enter__()` inside the async generators;
each of the four generators has `_tracker.__exit__(...)` in a try/finally; no
StreamingResponse passes `background=`. Behavioral verifies (running the
extracted `_TrackedCancel`): finally cleanup on normal completion, mid-stream
OSError, and aclose() from ClientDisconnect; and a pre-set cancel_event lets
the GGUF loop break cleanly emitting 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:2240 `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():
# Simulates OSError / BrokenPipeError from Starlette send() 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 the async generator when the client
# disconnects mid-stream. The generator's 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 set cancel_event via __enter__. The
# generator must break cleanly and emit final_chunk + [DONE] rather than
# calling next(gen) and propagating GeneratorExit out of 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 cancel_event is 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 out on the next iteration
# rather than draining the stub 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():` statement anywhere inside a
# `while`/`for` loop body is sufficient -- without it, a cancel POST
# cannot interrupt the loop because Colab-style proxies do not
# propagate 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 regression back to `for chunk_text in
# audio_input_generate():` -- which blocks the event loop on each
# whisper chunk and prevents POST /api/inference/cancel from being
# serviced until the chunk yields.
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 Unsloth cancel branch must call backend.reset_generation_state() to
# flush GPU/KV-cache state, else a cancel-via-POST leaves the subprocess
# dirty for the next request.
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():
# Regression guard: replace the post-fix loop with the pre-fix
# `for chunk in audio_input_generate()` pattern and assert it 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 was pre-set before iterating, so the
# top-of-loop check must short-circuit iteration 1 (zero tokens). Catches a
# regression that moves the check below next() (leaks one extra token).
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's top-of-loop
# cancel check must skip the asyncio.to_thread(next, ...) call when
# cancel_event was already 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"
)