Studio: make Stop and stall deadlines interrupt a wedged stream portably (#7236)

* Studio: make Stop and stall deadlines interrupt a wedged stream portably

The cancel watcher unblocks a stalled read by shutting the socket down from
another thread, which works on POSIX but not reliably on native Windows, where
Winsock does not dependably wake a recv() already in progress on another thread.
Wrap the httpcore network stream so the reader loops each read in short slices
and polls the cancel event itself. Stop and the stall deadlines now interrupt a
wedged mid-stream read without any cross-thread socket teardown, and a slow but
still-alive stream is never torn down. The POSIX shutdown path is preserved.

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

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

* Honor the post-first-token stall timeout in the cancel-aware read

httpcore snapshots request.extensions timeout read once when the body
starts, so lowering it to the stall timeout after the first token never
reached the socket read and a one-token-then-silent server hung for the
full prefill window. Re-read the live extensions timeout per call and
bound each read by it, falling back to the httpcore-passed timeout when
absent so prefill and normal completion are unchanged.

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

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

* studio: tighten comments in the llama.cpp stall timeout path

* Tighten comments in the stream stall cancel path

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: danielhanchen <unslothshared@gmail.com>
This commit is contained in:
Daniel Han 2026-07-20 05:29:18 -07:00 committed by GitHub
commit bdf51525ea
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 199 additions and 0 deletions

View file

@ -9907,6 +9907,75 @@ class LlamaCppBackend:
except Exception:
logger.debug("Could not close httpx client", exc_info = True)
@staticmethod
def _install_cancel_aware_read(
client: "httpx.Client",
cancel_event: threading.Event,
response: Optional["httpx.Response"] = None,
poll_s: float = 0.2,
) -> None:
"""Wrap the httpcore stream so the reader interrupts its own blocked recv() on cancel.
A cross-thread socket shutdown wakes a parked recv() on POSIX but not on
Windows (Winsock), so read in short slices and poll cancel_event between them
(plain or TLS); slice timeouts are swallowed so a slow-but-alive stream survives.
httpcore snapshots request.extensions["timeout"]["read"] once at body start, so
given ``response`` we re-read the live value per call to honor the post-first-token
stall timeout instead of the long prefill timeout."""
import httpcore
def _live_read_timeout() -> Optional[float]:
if response is None:
return None
try:
ext = response.request.extensions.get("timeout")
if isinstance(ext, dict):
value = ext.get("read")
if isinstance(value, (int, float)):
return float(value)
except Exception:
pass
return None
try:
pool = getattr(getattr(client, "_transport", None), "_pool", None)
for connection in list(getattr(pool, "_connections", []) or []):
inner = getattr(connection, "_connection", None)
stream = getattr(inner, "_network_stream", None)
if stream is None or getattr(stream, "_unsloth_cancel_wrapped", False):
continue
orig_read = stream.read
def read(
max_bytes,
timeout = None,
_orig = orig_read,
):
live = _live_read_timeout()
effective = live if live is not None else timeout
deadline = None if effective is None else time.monotonic() + effective
while True:
if cancel_event.is_set():
raise httpcore.ReadError("stream cancelled by user")
if deadline is None:
step = poll_s
else:
remaining = deadline - time.monotonic()
if remaining <= 0:
raise httpcore.ReadTimeout("read operation timed out")
step = min(poll_s, remaining)
try:
return _orig(max_bytes, timeout = step)
except httpcore.ReadTimeout:
if deadline is not None and time.monotonic() >= deadline:
raise
continue # slow but alive: keep reading
stream.read = read
stream._unsloth_cancel_wrapped = True
except Exception:
logger.debug("Could not install cancel-aware read", exc_info = True)
@staticmethod
@contextlib.contextmanager
def _stream_with_retry(
@ -9964,6 +10033,11 @@ class LlamaCppBackend:
headers = headers,
) as response:
_response_ref[0] = response
if cancel_event is not None:
# Portable mid-stream cancel: the reader polls cancel itself, so
# Stop interrupts a stalled read where the watcher's Windows socket
# shutdown does not. Pass response to honor the live stall timeout.
LlamaCppBackend._install_cancel_aware_read(client, cancel_event, response)
if cancel_event is not None and cancel_event.is_set():
raise _LlamaStreamCancelled
yield response

View file

@ -0,0 +1,125 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Regression test for the post-first-token stall timeout in the cancel-aware read.
httpcore snapshots ``request.extensions["timeout"]["read"]`` once at body start, so
when ``_iter_text_cancellable`` lowers it after the first token, a one-token-then-silent
server hangs for the full prefill window. The fix re-reads the live extensions timeout
per call; a fake clock and always-silent stream check the read gives up after the live
stall timeout, not the stale prefill one.
"""
from __future__ import annotations
import inspect
import sys
import threading
import types as _types
from pathlib import Path
import pytest
_BACKEND_DIR = str(Path(__file__).resolve().parent.parent)
if _BACKEND_DIR not in sys.path:
sys.path.insert(0, _BACKEND_DIR)
# Mirror sibling tests' stubbing so the module imports without fastapi.
_loggers_stub = _types.ModuleType("loggers")
_loggers_stub.get_logger = lambda name: __import__("logging").getLogger(name)
sys.modules.setdefault("loggers", _loggers_stub)
sys.modules.setdefault("structlog", _types.ModuleType("structlog"))
import httpcore # noqa: E402
from core.inference import llama_cpp as llama_cpp_mod # noqa: E402
from core.inference.llama_cpp import LlamaCppBackend # noqa: E402
_PREFILL_TIMEOUT = 1200.0 # what httpcore snapshots from the prefill timeout
_STALL_TIMEOUT = 120.0 # the post-first-token stall timeout the wrapper must honor
class _Obj:
pass
def _install(response, clock, silent_stream):
"""Wire fake client/pool so _install_cancel_aware_read finds the stream; return the wrapped stream.read."""
inner = _Obj()
inner._network_stream = silent_stream
connection = _Obj()
connection._connection = inner
pool = _Obj()
pool._connections = [connection]
transport = _Obj()
transport._pool = pool
client = _Obj()
client._transport = transport
cancel_event = threading.Event() # never set: we test the stall path, not cancel
sig = inspect.signature(LlamaCppBackend._install_cancel_aware_read)
if "response" in sig.parameters:
# Fixed signature: wrapper reads the live extensions timeout.
LlamaCppBackend._install_cancel_aware_read(client, cancel_event, response)
else:
# Pre-fix signature: no response, so the stall assertion fails (proves the bug).
LlamaCppBackend._install_cancel_aware_read(client, cancel_event)
return silent_stream.read
def test_stall_timeout_honored_after_first_token(monkeypatch):
clock = {"t": 0.0}
monkeypatch.setattr(llama_cpp_mod.time, "monotonic", lambda: clock["t"])
# One token then silence: every read times out, advancing fake time by its timeout.
def silent_read(max_bytes, timeout = None):
clock["t"] += timeout if timeout is not None else 0.0
raise httpcore.ReadTimeout("slice timed out on silence")
stream = _Obj()
stream.read = silent_read
# First token seen: the live read timeout is lowered to the stall timeout.
request = _Obj()
request.extensions = {"timeout": {"read": _STALL_TIMEOUT}}
response = _Obj()
response.request = request
wrapped_read = _install(response, clock, stream)
# httpcore still passes the stale prefill timeout it snapshotted at body start.
with pytest.raises(httpcore.ReadTimeout):
wrapped_read(65536, timeout = _PREFILL_TIMEOUT)
# Must give up ~stall timeout after the last token, not the prefill window.
assert clock["t"] <= _STALL_TIMEOUT * 1.5, (
f"stall timeout not honored: waited {clock['t']}s "
f"(expected ~{_STALL_TIMEOUT}s, not {_PREFILL_TIMEOUT}s)"
)
assert clock["t"] >= _STALL_TIMEOUT * 0.5
def test_prefill_timeout_used_when_no_live_override(monkeypatch):
"""Without a lowered live timeout, the wrapper honors the passed prefill timeout, so the normal first-token wait is unchanged."""
clock = {"t": 0.0}
monkeypatch.setattr(llama_cpp_mod.time, "monotonic", lambda: clock["t"])
def silent_read(max_bytes, timeout = None):
clock["t"] += timeout if timeout is not None else 0.0
raise httpcore.ReadTimeout("slice timed out on silence")
stream = _Obj()
stream.read = silent_read
# No timeout extension: wrapper falls back to httpcore's passed timeout.
request = _Obj()
request.extensions = {}
response = _Obj()
response.request = request
wrapped_read = _install(response, clock, stream)
with pytest.raises(httpcore.ReadTimeout):
wrapped_read(65536, timeout = _PREFILL_TIMEOUT)
assert clock["t"] >= _PREFILL_TIMEOUT * 0.9