Keep a durable run alive when no model is loaded for PR #7219
A durable run is claimable within the supervisor's poll interval of startup (main.py starts it in the lifespan, and claim_next takes any 'running' run whose lease expired), Studio has no startup model auto-load, and the browser is not connected yet. So restarting Studio mid-run reliably lands the next model call on the local endpoint's HTTP 400 "No model loaded". That 400 is not retryable: _completion retries only >= 500, and _stream_completion, which serves both planning and synthesis, has no retry at all. The run is marked failed, and the only recovery is retry, which sets report_text NULL and deletes every research_plan_step, research_source and research_document_source. Up to an hour of scraping and synthesis is lost on a plain restart, on the feature whose whole point is surviving one. Treat only that refusal as transient: wait up to the run's own modelTimeoutSeconds for a model to come back, then re-send. Any other 400 still fails immediately, so no behaviour changes on the happy path. The wait polls _check_active, so cancellation and lease loss are still honoured, and the model probe fails open, so a probe error can only send a request, never withhold one. Each wait is bounded by the run timeout and the number of waits per call is capped, so a model that keeps disappearing cannot re-send forever. Deliberately not pinning or restoring the model, which the review comment also suggested. Auto-switch is opt-in, default off, and GGUF-only, so restoring would silently evict the model the user just loaded from a background worker, and comparing the configured name to the loaded id is fragile across variant suffixes and advertised aliases, so it would break working runs. Verified: 853 passed across the research/web/sandbox/chat-history/rag/inference backend suites. Eight of the nine new tests fail without the fix.
This commit is contained in:
parent
b78cb3ba0b
commit
689b06535c
2 changed files with 334 additions and 21 deletions
|
|
@ -114,6 +114,13 @@ _AUTO_SCRAPE_TOP_K = 3
|
|||
_AUTO_SCRAPE_TOTAL_CHARS = 6_000
|
||||
_WEB_RAG_TOP_N = 6
|
||||
_WEB_RAG_MIN_SCORE = 0.30
|
||||
# Poll interval while a run waits for a local model to be (re)loaded, and the detail
|
||||
# routes.inference returns when nothing is loaded (its 400 is transient, not a bad request).
|
||||
_MODEL_WAIT_POLL_SECONDS = 2.0
|
||||
# Each wait is bounded by modelTimeoutSeconds, but a model that keeps disappearing would
|
||||
# otherwise re-send forever, so cap how many times one call may wait.
|
||||
_MAX_MODEL_WAITS = 3
|
||||
_NO_MODEL_LOADED_DETAIL = "No model loaded"
|
||||
|
||||
|
||||
def _auto_scrape_default() -> int:
|
||||
|
|
@ -496,6 +503,42 @@ def _loaded_context_length() -> int | None:
|
|||
return None
|
||||
|
||||
|
||||
async def _model_unloaded(response: httpx.Response) -> bool:
|
||||
"""Whether the local endpoint refused because no model is loaded (routes.inference). That is
|
||||
transient for a durable run -- the model can be loaded again -- unlike any other 400."""
|
||||
if response.status_code != 400:
|
||||
return False
|
||||
try:
|
||||
body = await response.aread()
|
||||
except Exception:
|
||||
return False
|
||||
return _NO_MODEL_LOADED_DETAIL in body.decode("utf-8", "replace")
|
||||
|
||||
|
||||
def _local_model_ready() -> bool:
|
||||
"""Whether the local chat-completions path has a model to serve, using the same two checks
|
||||
routes.inference.openai_chat_completions makes before it 400s. Fails open when neither
|
||||
backend can be probed, so a probe failure can only run a request, never withhold one."""
|
||||
probed = False
|
||||
try:
|
||||
from routes.inference import get_llama_cpp_backend
|
||||
|
||||
if getattr(get_llama_cpp_backend(), "is_loaded", False):
|
||||
return True
|
||||
probed = True
|
||||
except Exception:
|
||||
logger.debug("research.model_probe_llama_failed", exc_info = True)
|
||||
try:
|
||||
from core.inference import get_inference_backend
|
||||
|
||||
if getattr(get_inference_backend(), "active_model_name", None):
|
||||
return True
|
||||
probed = True
|
||||
except Exception:
|
||||
logger.debug("research.model_probe_failed", exc_info = True)
|
||||
return not probed
|
||||
|
||||
|
||||
def _synthesis_evidence_budget() -> int:
|
||||
"""Char budget for synthesis evidence, sized to fit the loaded context (falls back to the
|
||||
full cap when the context is unknown)."""
|
||||
|
|
@ -1106,6 +1149,22 @@ class ResearchSupervisor:
|
|||
raise RuntimeError("Research is waiting for the Studio server port")
|
||||
return f"http://127.0.0.1:{port}/v1/chat/completions"
|
||||
|
||||
async def _wait_for_local_model(self, run: dict) -> bool:
|
||||
"""Wait, up to the run's model timeout, for a model to be loaded again; True if one was.
|
||||
|
||||
A durable run resumes after a Studio restart and is approved long after it was created,
|
||||
so the model it was started with can be gone. Waiting keeps the run alive instead of
|
||||
ending it on a non-retryable 400 that discards every step and source it gathered."""
|
||||
loop = asyncio.get_running_loop()
|
||||
deadline = loop.time() + float(run["config"]["budgets"]["modelTimeoutSeconds"])
|
||||
logger.info("research.waiting_for_local_model run_id=%s", run["id"])
|
||||
while loop.time() < deadline:
|
||||
await self._check_active(run["id"])
|
||||
await asyncio.sleep(_MODEL_WAIT_POLL_SECONDS)
|
||||
if _local_model_ready():
|
||||
return True
|
||||
return False
|
||||
|
||||
async def _completion(
|
||||
self,
|
||||
run: dict,
|
||||
|
|
@ -1144,7 +1203,9 @@ class ResearchSupervisor:
|
|||
try:
|
||||
timeout = httpx.Timeout(float(config["budgets"]["modelTimeoutSeconds"]))
|
||||
async with httpx.AsyncClient(timeout = timeout, trust_env = False) as client:
|
||||
for attempt in range(3):
|
||||
attempt = 0
|
||||
model_waits = 0
|
||||
while True:
|
||||
await self._check_active(run["id"])
|
||||
try:
|
||||
post_task = asyncio.create_task(
|
||||
|
|
@ -1169,6 +1230,17 @@ class ResearchSupervisor:
|
|||
body = response.json()
|
||||
break
|
||||
except (httpx.TransportError, httpx.HTTPStatusError) as exc:
|
||||
# Nothing loaded (restart, eject): wait for a model and re-send without
|
||||
# spending an attempt, so the run survives instead of failing here.
|
||||
if isinstance(exc, httpx.HTTPStatusError) and await _model_unloaded(
|
||||
exc.response
|
||||
):
|
||||
model_waits += 1
|
||||
if model_waits <= _MAX_MODEL_WAITS and await self._wait_for_local_model(
|
||||
run
|
||||
):
|
||||
continue
|
||||
raise
|
||||
retryable = (
|
||||
not isinstance(exc, httpx.HTTPStatusError)
|
||||
or exc.response.status_code >= 500
|
||||
|
|
@ -1176,6 +1248,7 @@ class ResearchSupervisor:
|
|||
if not retryable or attempt == 2:
|
||||
raise
|
||||
await asyncio.sleep(2**attempt)
|
||||
attempt += 1
|
||||
message = body["choices"][0]["message"]
|
||||
thought = message.get("reasoning_content")
|
||||
if isinstance(thought, str) and thought.strip():
|
||||
|
|
@ -1342,26 +1415,43 @@ class ResearchSupervisor:
|
|||
try:
|
||||
timeout = httpx.Timeout(float(config["budgets"]["modelTimeoutSeconds"]))
|
||||
async with httpx.AsyncClient(timeout = timeout, trust_env = False) as client:
|
||||
request = client.build_request(
|
||||
"POST",
|
||||
self._endpoint(),
|
||||
json = payload,
|
||||
headers = {"Authorization": f"Bearer {token}"},
|
||||
)
|
||||
response: httpx.Response | None = None
|
||||
send_task = asyncio.create_task(client.send(request, stream = True))
|
||||
send_task: asyncio.Task | None = None
|
||||
model_waits = 0
|
||||
try:
|
||||
while not send_task.done():
|
||||
await asyncio.wait({send_task}, timeout = 0.2)
|
||||
if self._cancel_event(run["id"]).is_set():
|
||||
send_task.cancel()
|
||||
try:
|
||||
await send_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
await self._check_active(run["id"])
|
||||
response = await send_task
|
||||
response.raise_for_status()
|
||||
while True:
|
||||
request = client.build_request(
|
||||
"POST",
|
||||
self._endpoint(),
|
||||
json = payload,
|
||||
headers = {"Authorization": f"Bearer {token}"},
|
||||
)
|
||||
send_task = asyncio.create_task(client.send(request, stream = True))
|
||||
while not send_task.done():
|
||||
await asyncio.wait({send_task}, timeout = 0.2)
|
||||
if self._cancel_event(run["id"]).is_set():
|
||||
send_task.cancel()
|
||||
try:
|
||||
await send_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
await self._check_active(run["id"])
|
||||
response = await send_task
|
||||
try:
|
||||
response.raise_for_status()
|
||||
break
|
||||
except httpx.HTTPStatusError as exc:
|
||||
# Nothing loaded (restart, eject): wait for a model and re-send.
|
||||
# Nothing has streamed yet, so this cannot duplicate report text.
|
||||
if not await _model_unloaded(exc.response):
|
||||
raise
|
||||
model_waits += 1
|
||||
if model_waits > _MAX_MODEL_WAITS:
|
||||
raise
|
||||
if not await self._wait_for_local_model(run):
|
||||
raise
|
||||
await response.aclose()
|
||||
response = None
|
||||
async for line in self._iter_stream_lines(run["id"], response):
|
||||
if self._cancel_event(run["id"]).is_set():
|
||||
await self._check_active(run["id"])
|
||||
|
|
@ -1396,13 +1486,18 @@ class ResearchSupervisor:
|
|||
):
|
||||
await flush_progress()
|
||||
finally:
|
||||
if not send_task.done():
|
||||
if send_task is not None and not send_task.done():
|
||||
send_task.cancel()
|
||||
try:
|
||||
await send_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
if response is None and send_task.done() and not send_task.cancelled():
|
||||
if (
|
||||
response is None
|
||||
and send_task is not None
|
||||
and send_task.done()
|
||||
and not send_task.cancelled()
|
||||
):
|
||||
try:
|
||||
response = send_task.result()
|
||||
except Exception:
|
||||
|
|
|
|||
|
|
@ -3,9 +3,19 @@
|
|||
|
||||
"""Regression tests for Deep Research query/prompt/citation/config hardening."""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import sys
|
||||
import time
|
||||
from types import SimpleNamespace
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from core import research_runs
|
||||
from core.research_runs import (
|
||||
ResearchSupervisor,
|
||||
RunCancelled,
|
||||
_citation_title,
|
||||
_escape_link_destination,
|
||||
_sanitize_public_query,
|
||||
|
|
@ -282,3 +292,211 @@ def test_dropped_raw_url_does_not_unbalance_prose():
|
|||
# An uncataloged URL is still removed, but the paren it swallowed belongs to the prose.
|
||||
out = _validate_report_sources("Claim (https://nope.com/x) here.", [])
|
||||
assert out == "Claim () here."
|
||||
|
||||
|
||||
def _install_probe_backends(monkeypatch, llama, native) -> None:
|
||||
"""Stand in for the two backend modules _local_model_ready probes, so the check can be
|
||||
exercised without importing the ML stack. Pass an exception to make a probe raise."""
|
||||
|
||||
def _getter(value):
|
||||
def _get():
|
||||
if isinstance(value, Exception):
|
||||
raise value
|
||||
return value
|
||||
|
||||
return _get
|
||||
|
||||
monkeypatch.setitem(
|
||||
sys.modules, "routes.inference", SimpleNamespace(get_llama_cpp_backend = _getter(llama))
|
||||
)
|
||||
monkeypatch.setitem(
|
||||
sys.modules, "core.inference", SimpleNamespace(get_inference_backend = _getter(native))
|
||||
)
|
||||
|
||||
|
||||
def test_local_model_ready_mirrors_the_chat_endpoint_checks(monkeypatch):
|
||||
# Same two checks routes.inference.openai_chat_completions makes before it 400s.
|
||||
unloaded = SimpleNamespace(is_loaded = False)
|
||||
idle = SimpleNamespace(active_model_name = None)
|
||||
_install_probe_backends(monkeypatch, SimpleNamespace(is_loaded = True), idle)
|
||||
assert research_runs._local_model_ready() is True
|
||||
_install_probe_backends(monkeypatch, unloaded, SimpleNamespace(active_model_name = "m"))
|
||||
assert research_runs._local_model_ready() is True
|
||||
_install_probe_backends(monkeypatch, unloaded, idle)
|
||||
assert research_runs._local_model_ready() is False
|
||||
|
||||
|
||||
def test_local_model_ready_fails_open_when_neither_backend_can_be_probed(monkeypatch):
|
||||
# A broken probe must not withhold a request; the endpoint stays the decider.
|
||||
_install_probe_backends(monkeypatch, RuntimeError("boom"), RuntimeError("boom"))
|
||||
assert research_runs._local_model_ready() is True
|
||||
|
||||
|
||||
def _response(status: int, *, detail: str = "", body: str = "") -> httpx.Response:
|
||||
request = httpx.Request("POST", "http://127.0.0.1:1/v1/chat/completions")
|
||||
if detail:
|
||||
return httpx.Response(status, json = {"detail": detail}, request = request)
|
||||
return httpx.Response(status, text = body, request = request)
|
||||
|
||||
|
||||
_NO_MODEL = "No model loaded. Call POST /inference/load first."
|
||||
|
||||
|
||||
def test_model_unloaded_only_matches_the_no_model_refusal():
|
||||
assert asyncio.run(research_runs._model_unloaded(_response(400, detail = _NO_MODEL))) is True
|
||||
# Any other 400 is a real bad request and must stay non-retryable.
|
||||
assert (
|
||||
asyncio.run(research_runs._model_unloaded(_response(400, detail = "Invalid 'tools'")))
|
||||
is False
|
||||
)
|
||||
assert asyncio.run(research_runs._model_unloaded(_response(500, body = _NO_MODEL))) is False
|
||||
|
||||
|
||||
def _make_supervisor(check_active = None) -> ResearchSupervisor:
|
||||
supervisor = ResearchSupervisor(
|
||||
SimpleNamespace(state = SimpleNamespace(server_port = 1)),
|
||||
)
|
||||
if check_active is not None:
|
||||
supervisor._check_active = check_active
|
||||
return supervisor
|
||||
|
||||
|
||||
def _waiting_run(timeout_seconds: float) -> dict:
|
||||
return {
|
||||
"id": "run-1",
|
||||
"ownerSubject": "user-1",
|
||||
"config": {"budgets": {"modelTimeoutSeconds": timeout_seconds}},
|
||||
}
|
||||
|
||||
|
||||
def test_wait_for_local_model_polls_until_a_model_is_loaded(monkeypatch):
|
||||
monkeypatch.setattr(research_runs, "_MODEL_WAIT_POLL_SECONDS", 0.01)
|
||||
states = iter([False, True])
|
||||
monkeypatch.setattr(research_runs, "_local_model_ready", lambda: next(states, True))
|
||||
checked: list[str] = []
|
||||
|
||||
async def _check_active(run_id: str) -> None:
|
||||
checked.append(run_id)
|
||||
|
||||
supervisor = _make_supervisor(_check_active)
|
||||
assert asyncio.run(supervisor._wait_for_local_model(_waiting_run(30.0))) is True
|
||||
# Cancellation/lease are re-checked before every poll.
|
||||
assert checked == ["run-1", "run-1"]
|
||||
|
||||
|
||||
def test_wait_for_local_model_gives_up_at_the_run_timeout(monkeypatch):
|
||||
monkeypatch.setattr(research_runs, "_MODEL_WAIT_POLL_SECONDS", 0.01)
|
||||
monkeypatch.setattr(research_runs, "_local_model_ready", lambda: False)
|
||||
|
||||
async def _check_active(run_id: str) -> None:
|
||||
return None
|
||||
|
||||
supervisor = _make_supervisor(_check_active)
|
||||
started = time.monotonic()
|
||||
assert asyncio.run(supervisor._wait_for_local_model(_waiting_run(0.05))) is False
|
||||
assert time.monotonic() - started < 5
|
||||
|
||||
|
||||
def test_wait_for_local_model_still_honors_cancellation(monkeypatch):
|
||||
monkeypatch.setattr(research_runs, "_MODEL_WAIT_POLL_SECONDS", 0.01)
|
||||
monkeypatch.setattr(research_runs, "_local_model_ready", lambda: False)
|
||||
|
||||
async def _check_active(run_id: str) -> None:
|
||||
raise RunCancelled()
|
||||
|
||||
supervisor = _make_supervisor(_check_active)
|
||||
with pytest.raises(RunCancelled):
|
||||
asyncio.run(supervisor._wait_for_local_model(_waiting_run(30.0)))
|
||||
|
||||
|
||||
def _install_fake_client(monkeypatch, responses: list) -> list:
|
||||
"""Serve ``responses`` in order to both completion paths and record the sends."""
|
||||
sent: list = []
|
||||
|
||||
class _FakeClient:
|
||||
def __init__(self, **kwargs):
|
||||
pass
|
||||
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *exc_info):
|
||||
return False
|
||||
|
||||
def build_request(self, method, url, **kwargs):
|
||||
return (method, url)
|
||||
|
||||
async def post(self, url, **kwargs):
|
||||
sent.append(url)
|
||||
return responses.pop(0)
|
||||
|
||||
async def send(self, request, *, stream = False):
|
||||
sent.append(request)
|
||||
return responses.pop(0)
|
||||
|
||||
monkeypatch.setattr(research_runs.httpx, "AsyncClient", _FakeClient)
|
||||
monkeypatch.setattr(
|
||||
research_runs.auth_storage, "create_api_key", lambda **kwargs: ("token", {"id": 1})
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
research_runs.auth_storage, "revoke_internal_api_key", lambda key_id: None
|
||||
)
|
||||
return sent
|
||||
|
||||
|
||||
def _ready_after_first_poll(monkeypatch) -> None:
|
||||
monkeypatch.setattr(research_runs, "_MODEL_WAIT_POLL_SECONDS", 0.01)
|
||||
monkeypatch.setattr(research_runs, "_local_model_ready", lambda: True)
|
||||
|
||||
|
||||
def test_completion_retries_after_the_model_is_loaded_again(monkeypatch):
|
||||
# A durable run resumes after a Studio restart and is approved long after creation, so the
|
||||
# model can be unloaded when it calls. That 400 used to end the run and its gathered work.
|
||||
_ready_after_first_poll(monkeypatch)
|
||||
reply = {"choices": [{"message": {"content": "answer"}}]}
|
||||
sent = _install_fake_client(
|
||||
monkeypatch,
|
||||
[_response(400, detail = _NO_MODEL), _response(200, body = json.dumps(reply))],
|
||||
)
|
||||
|
||||
async def _check_active(run_id: str) -> None:
|
||||
return None
|
||||
|
||||
supervisor = _make_supervisor(_check_active)
|
||||
result = asyncio.run(supervisor._completion(_waiting_run(30.0), [{"role": "user"}]))
|
||||
assert result == "answer"
|
||||
assert len(sent) == 2
|
||||
|
||||
|
||||
def test_completion_still_fails_fast_on_a_real_bad_request(monkeypatch):
|
||||
_ready_after_first_poll(monkeypatch)
|
||||
sent = _install_fake_client(monkeypatch, [_response(400, detail = "Invalid 'tools'")])
|
||||
|
||||
async def _check_active(run_id: str) -> None:
|
||||
return None
|
||||
|
||||
supervisor = _make_supervisor(_check_active)
|
||||
with pytest.raises(httpx.HTTPStatusError):
|
||||
asyncio.run(supervisor._completion(_waiting_run(30.0), [{"role": "user"}]))
|
||||
assert len(sent) == 1
|
||||
|
||||
|
||||
def test_stream_completion_retries_after_the_model_is_loaded_again(monkeypatch):
|
||||
_ready_after_first_poll(monkeypatch)
|
||||
chunk = json.dumps({"choices": [{"delta": {"content": "report"}, "finish_reason": "stop"}]})
|
||||
stream = f"data: {chunk}\n\ndata: [DONE]\n\n"
|
||||
sent = _install_fake_client(
|
||||
monkeypatch, [_response(400, detail = _NO_MODEL), _response(200, body = stream)]
|
||||
)
|
||||
|
||||
async def _check_active(run_id: str) -> None:
|
||||
return None
|
||||
|
||||
supervisor = _make_supervisor(_check_active)
|
||||
report, reasoning, finish_reason = asyncio.run(
|
||||
supervisor._stream_completion(
|
||||
_waiting_run(30.0), [{"role": "user"}], report_progress = False
|
||||
)
|
||||
)
|
||||
assert (report, reasoning, finish_reason) == ("report", "", "stop")
|
||||
assert len(sent) == 2
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue