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:
parent
6c78147980
commit
bdf51525ea
2 changed files with 199 additions and 0 deletions
|
|
@ -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
|
||||
|
|
|
|||
125
studio/backend/tests/test_llama_cpp_stall_timeout.py
Normal file
125
studio/backend/tests/test_llama_cpp_stall_timeout.py
Normal 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
|
||||
Loading…
Add table
Add a link
Reference in a new issue