Improve unsloth start runtime lifecycle

This commit is contained in:
oobabooga 2026-07-21 13:57:36 -03:00
commit 314627cd6a
4 changed files with 620 additions and 32 deletions

View file

@ -14,12 +14,13 @@ import signal
import subprocess
import sys
import tempfile
import threading
import time
import urllib.error
import urllib.request
from pathlib import Path
from typing import NamedTuple, NoReturn, Optional
from urllib.parse import urlparse
from urllib.parse import urlencode, urlparse
import click
import typer
@ -105,8 +106,8 @@ _SERVE_OPTION = typer.Option(
True,
"--serve/--no-serve",
help = (
"If no Unsloth server is running, auto-start one for --model and stop it when the "
"agent exits. --no-serve keeps the old behavior of erroring out."
"If no Unsloth server is running, auto-start one for --model and keep it available "
"after the agent exits. --no-serve keeps the old behavior of erroring out."
),
)
# Model-load knobs mirrored from `unsloth run`; only used when --model triggers a
@ -378,6 +379,265 @@ def _http_json(
_auto_served_server: Optional[subprocess.Popen] = None
# Model download + load can be slow; give the auto-started server room before giving up.
_SERVER_START_TIMEOUT_S = 900
_DOWNLOAD_POLL_INTERVAL_S = 1.0
_START_API_KEY_PREFIX = "UNSLOTH_START_API_KEY: "
def _format_download_bytes(value: int) -> str:
value = max(0, int(value))
for unit in ("B", "KiB", "MiB", "GiB", "TiB"):
if value < 1024 or unit == "TiB":
precision = 0 if unit in ("B", "KiB") else 1
return f"{value:.{precision}f} {unit}"
value /= 1024
return "0 B"
def _format_download_eta(seconds: float) -> str:
seconds = max(0, int(seconds))
if seconds < 60:
return f"{seconds}s"
minutes, seconds = divmod(seconds, 60)
if minutes < 60:
return f"{minutes}m {seconds:02d}s"
hours, minutes = divmod(minutes, 60)
return f"{hours}h {minutes:02d}m"
class _DownloadProgressDisplay:
"""Render download progress without making redirected output noisy."""
def __init__(self) -> None:
self._samples: list[tuple[float, int]] = []
self._shown = False
self._last_bucket = -1
self._last_line_length = 0
self._last_expected = 0
self._interactive = bool(getattr(sys.stdout, "isatty", lambda: False)())
def update(self, progress: dict) -> None:
downloaded = max(0, int(progress.get("downloaded_bytes") or 0))
completed = max(0, int(progress.get("completed_bytes") or 0))
expected = max(0, int(progress.get("expected_bytes") or 0))
self._last_expected = max(self._last_expected, expected)
fraction = float(progress.get("progress") or 0)
if downloaded <= 0:
return
# The hub endpoint can report a fully cached snapshot as 99% when an
# older synchronous load has no download manifest. No incomplete bytes
# means no transfer is occurring, so do not label model startup as a download.
if completed >= downloaded > 0:
return
now = time.monotonic()
if self._samples and downloaded < self._samples[-1][1]:
self._samples.clear()
self._samples.append((now, downloaded))
cutoff = now - 15.0
while len(self._samples) > 2 and self._samples[0][0] < cutoff:
self._samples.pop(0)
rate = 0.0
if len(self._samples) >= 2:
elapsed = self._samples[-1][0] - self._samples[0][0]
delta = self._samples[-1][1] - self._samples[0][1]
if elapsed >= 1.0 and delta > 0:
rate = delta / elapsed
if expected > 0:
# Trust the endpoint's capped value. It deliberately reports at
# most 99% while bytes still live in an incomplete file, even when
# that sparse file's logical size already equals the final blob.
fraction = min(1.0, max(0.0, fraction))
percent = min(100, max(0, int(fraction * 100)))
filled = min(24, int(fraction * 24))
bar = "=" * filled + ">" + "." * max(0, 23 - filled) if filled < 24 else "=" * 24
line = (
f"Downloading model [{bar}] {percent:3d}% "
f"{_format_download_bytes(downloaded)} / {_format_download_bytes(expected)}"
)
bucket = percent // 10
if rate > 0:
line += f" | {_format_download_bytes(rate)}/s"
if downloaded < expected:
line += f" | ETA {_format_download_eta((expected - downloaded) / rate)}"
else:
line = f"Downloading model: {_format_download_bytes(downloaded)}"
bucket = downloaded // (1024**3)
if rate > 0:
line += f" | {_format_download_bytes(rate)}/s"
if self._interactive:
padding = " " * max(0, self._last_line_length - len(line))
typer.echo(f"\r{line}{padding}", nl = False)
sys.stdout.flush()
self._last_line_length = len(line)
elif not self._shown or bucket > self._last_bucket:
typer.echo(line)
self._last_bucket = bucket
self._shown = True
def close(self) -> None:
if self._interactive and self._shown:
typer.echo()
self._last_line_length = 0
def complete(self) -> None:
"""Finish a displayed transfer after the model load confirms success."""
if not self._shown:
return
downloaded = self._samples[-1][1] if self._samples else 0
expected = max(downloaded, getattr(self, "_last_expected", 0))
self.update(
{
"downloaded_bytes": expected,
"expected_bytes": expected,
"progress": 1.0,
}
)
def _normalized_variant(value: object) -> str:
return re.sub(r"[^a-z0-9]", "", str(value or "").lower())
class _ModelDownloadProgress:
"""Best-effort polling of the model download endpoints."""
def __init__(self, base: str, key: str, model: str, variant: Optional[str]) -> None:
self._base = base
self._key = key
self._model = model
self._variant = variant or ""
self._expected_bytes = 0
self._display = _DownloadProgressDisplay()
self._configured = False
self._disabled = not _is_hub_model_id(model)
self._progress_prefix = "/api/hub"
def _configure(self) -> None:
self._configured = True
if self._disabled:
return
# GGUF repos need the selected quant's size. The generic repo endpoint
# totals every quant in the repository and would report a misleading
# percentage, so resolve the variant first and otherwise show bytes only.
if self._variant or "gguf" in self._model.lower():
try:
params = urlencode({"repo_id": self._model})
try:
info = _http_json(
"GET",
f"{self._base}/api/hub/gguf-variants?{params}",
self._key,
timeout = 10,
)
except urllib.error.HTTPError as exc:
if exc.code != 404:
raise
self._progress_prefix = "/api/models"
info = _http_json(
"GET",
f"{self._base}/api/models/gguf-variants?{params}",
self._key,
timeout = 10,
)
self._variant = self._variant or str(info.get("default_variant") or "")
wanted = _normalized_variant(self._variant)
for item in info.get("variants") or []:
quant = _normalized_variant(item.get("quant"))
filename = _normalized_variant(item.get("filename"))
if wanted and (wanted == quant or wanted in filename):
self._expected_bytes = int(
item.get("download_size_bytes") or item.get("size_bytes") or 0
)
break
except Exception:
# Older servers may not expose the variant endpoint. Byte progress
# is still useful, and load errors remain owned by the load request.
pass
def poll(self) -> None:
if not self._configured:
self._configure()
if self._disabled:
return
try:
if self._variant or "gguf" in self._model.lower():
params = urlencode(
{
"repo_id": self._model,
"variant": self._variant,
"expected_bytes": self._expected_bytes,
}
)
url = f"{self._base}{self._progress_prefix}/gguf-download-progress?{params}"
else:
url = (
f"{self._base}{self._progress_prefix}/download-progress?"
f"{urlencode({'repo_id': self._model})}"
)
try:
reading = _http_json("GET", url, self._key, timeout = 10)
except urllib.error.HTTPError as exc:
if exc.code != 404 or self._progress_prefix == "/api/models":
raise
self._progress_prefix = "/api/models"
self.poll()
return
self._display.update(reading)
except Exception:
# Progress is an enhancement. Never turn an unsupported endpoint or
# a transient polling failure into a model-load failure.
self._disabled = True
def close(self) -> None:
self._display.close()
def complete(self) -> None:
self._display.complete()
def _load_model_with_progress(
base: str,
key: str,
model: str,
load: LoadOptions,
payload: dict,
) -> dict:
"""Run the blocking load request while polling its download progress."""
result: list[tuple[bool, object]] = []
done = threading.Event()
def _load() -> None:
try:
value = _http_json(
"POST",
f"{base}/api/inference/load",
key,
payload,
timeout = 3600,
error = "Model load failed",
)
result.append((True, value))
except BaseException as exc:
result.append((False, exc))
finally:
done.set()
threading.Thread(target = _load, name = "unsloth-model-load", daemon = True).start()
progress = _ModelDownloadProgress(base, key, model, load.gguf_variant)
try:
while not done.wait(_DOWNLOAD_POLL_INTERVAL_S):
progress.poll()
ok, value = result[0]
if not ok:
assert isinstance(value, BaseException)
raise value
progress.complete()
return value if isinstance(value, dict) else {}
finally:
progress.close()
def _studio_healthy(base: str, timeout: float = 3.0) -> bool:
@ -438,6 +698,14 @@ def _shutdown_auto_served() -> None:
_shutdown_server(server)
def _keep_auto_served() -> bool:
"""Release ownership so a successfully started server survives this CLI."""
global _auto_served_server
server, _auto_served_server = _auto_served_server, None
atexit.unregister(_shutdown_auto_served)
return server is not None and server.poll() is None
def _start_studio_server(base: str, model: str, load: LoadOptions) -> subprocess.Popen:
"""Spawn `unsloth run` for `model`, wait until it is fully ready, and return it."""
global _auto_served_server
@ -456,6 +724,10 @@ def _start_studio_server(base: str, model: str, load: LoadOptions) -> subprocess
"--no-cloudflare",
"--model",
model,
# The child writes this marker only to our private 0600 log as soon as
# it creates the API key. That lets us authenticate progress polling
# while the child is still blocked loading the model.
"--start-api-key-marker",
]
if load.gguf_variant:
command += ["--gguf-variant", load.gguf_variant]
@ -491,17 +763,44 @@ def _start_studio_server(base: str, model: str, load: LoadOptions) -> subprocess
atexit.register(_shutdown_auto_served)
deadline = time.monotonic() + _SERVER_START_TIMEOUT_S
while time.monotonic() < deadline:
if server.poll() is not None:
tail = _log_tail(log_path)
_shutdown_auto_served()
_fail(f"The Unsloth server stopped before it was ready. Last log lines:\n{tail}")
# `unsloth run` prints the minted key only after the server is up AND the model is
# loaded, so it is the fully-ready signal (same contract serve-unsloth-run.sh uses).
if _studio_healthy(base) and "sk-unsloth-" in _log_tail(log_path, lines = 400):
typer.echo(f"Unsloth server ready at {base}.")
return server
time.sleep(2.0)
progress: Optional[_ModelDownloadProgress] = None
early_key_seen = False
try:
while time.monotonic() < deadline:
if server.poll() is not None:
tail = _log_tail(log_path)
_shutdown_auto_served()
_fail(f"The Unsloth server stopped before it was ready. Last log lines:\n{tail}")
tail = _log_tail(log_path, lines = 400)
if progress is None:
marker = re.search(
rf"^{re.escape(_START_API_KEY_PREFIX)}(sk-unsloth-[^\s]+)$",
tail,
flags = re.MULTILINE,
)
if marker:
early_key_seen = True
progress = _ModelDownloadProgress(
base,
marker.group(1),
model,
load.gguf_variant,
)
if progress is not None:
progress.poll()
# New children emit an early key marker, so wait for the final model
# banner. The fallback preserves compatibility with an older child
# that only prints its API key after loading has completed.
ready_signal = "Model loaded:" in tail if early_key_seen else "sk-unsloth-" in tail
if _studio_healthy(base) and ready_signal:
if progress is not None:
progress.complete()
typer.echo(f"Unsloth server ready at {base}.")
return server
time.sleep(2.0)
finally:
if progress is not None:
progress.close()
_shutdown_auto_served()
_fail(
f"The Unsloth server didn't become ready within {_SERVER_START_TIMEOUT_S}s. See {log_path}."
@ -825,6 +1124,17 @@ def _resolve_model(
)
)
if requested and match is None:
active = next((m for m in models if m.get("loaded") is not False), None)
active_id = active.get("id") if active else None
if active_id and not _model_id_matches(
active_id,
requested,
allow_casefold = allow_casefold,
):
typer.echo(
f"Switching the Unsloth server from {active_id} to {requested}. "
"This unloads the current model for every attached session."
)
typer.echo(
f"Loading {requested} - please wait…"
if load_has_overrides
@ -841,14 +1151,7 @@ def _resolve_model(
payload["load_in_4bit"] = False
if load.tensor_parallel:
payload["tensor_parallel"] = True
loaded = _http_json(
"POST",
f"{base}/api/inference/load",
key,
payload,
timeout = 3600,
error = "Model load failed",
)
loaded = _load_model_with_progress(base, key, requested, load, payload)
# Unsloth registers the model under a canonical id (resolved identifier,
# casing) that /v1/models echoes but which may differ from the path we
# passed; match on the id the load reports so we don't silently fall
@ -881,7 +1184,13 @@ def _resolve_model(
"No model is loaded in Unsloth. Load one from the model dropdown in "
"the UI, or pass --model <hf-id-or-path> to load it from here."
)
return models[0]
resident = next((m for m in models if m.get("loaded") is not False), None)
if resident is None:
_fail(
"No model is currently resident in Unsloth. Pass --model <hf-id-or-path> "
"to reload one, or load it from the model dropdown in the UI."
)
return resident
def _require_gguf_for_codex(base: str, key: str, model_id: str) -> None:
@ -902,6 +1211,11 @@ def _require_gguf_for_codex(base: str, key: str, model_id: str) -> None:
_DYNAMIC_SECTIONS_FLAG = "--exclude-dynamic-system-prompt-sections"
_GEMMA_CLAUDE_NUDGE = (
"When the user asks to create an Unsloth fine-tuning run, begin with the next "
"concrete action. Do not wait silently. If action is blocked, state the exact "
"permission or input needed."
)
def _claude_settings_overlay(model_id: str) -> str:
@ -940,7 +1254,10 @@ def _claude_flags(model_id: str) -> list:
version = _claude_version()
if version is not None and version < (2, 1, 98):
return []
return [_DYNAMIC_SECTIONS_FLAG, "--settings", _claude_settings_overlay(model_id)]
flags = [_DYNAMIC_SECTIONS_FLAG, "--settings", _claude_settings_overlay(model_id)]
if "gemma" in model_id.casefold():
flags += ["--append-system-prompt", _GEMMA_CLAUDE_NUDGE]
return flags
def _merge_codex_config(existing: str, base: str) -> str:
@ -1356,7 +1673,7 @@ def _launch(
env: dict,
install_hint: str,
unset_env: tuple = (),
) -> NoReturn:
) -> int:
# Resolve well-known install dirs (e.g. ~/.local/bin) first, so an already-installed
# agent not yet on PATH is found instead of prompting a needless reinstall.
_augment_path_with_install_dirs()
@ -1382,7 +1699,7 @@ def _launch(
finally:
signal.signal(signal.SIGINT, previous)
# Negative returncode means killed by signal N; shells expect 128+N.
raise typer.Exit(code = code if code >= 0 else 128 - code)
return code if code >= 0 else 128 - code
def _connect(
@ -1438,12 +1755,32 @@ def _run(
if not launch:
env, wsl_env_bridge = _wsl_shim_env(command, env, unset_env)
_print_env(env, command, unset_env = unset_env, wsl_env_bridge = wsl_env_bridge)
if _keep_auto_served():
typer.echo(
f"Unsloth Studio is still running at {base}. "
"Stop it with `unsloth studio stop`."
)
return
try:
_launch(command, env, install_hint = install_hint, unset_env = unset_env)
finally:
# Tear down a server we auto-started once the agent session ends (no-op otherwise).
code = _launch(command, env, install_hint = install_hint, unset_env = unset_env)
except BaseException:
# Startup succeeded but the agent itself could not launch. In that failure
# path, retain the old cleanup behavior instead of orphaning a surprise server.
_shutdown_auto_served()
raise
auto_started = _auto_served_server is not None
kept = _keep_auto_served()
if auto_started and not kept:
typer.echo(f"The auto-started Unsloth server at {base} stopped during the session.")
raise typer.Exit(code = code)
if is_loopback_url(base):
typer.echo(
f"Unsloth Studio is still running at {base}. "
"Stop it with `unsloth studio stop`."
)
else:
typer.echo(f"The remote Unsloth server is still running at {base}.")
raise typer.Exit(code = code)
def _agents_config_root() -> Path:

View file

@ -1760,6 +1760,12 @@ def run(
"decode speed, MoE usually don't."
),
),
start_api_key_marker: bool = typer.Option(
False,
"--start-api-key-marker",
hidden = True,
help = "Emit an early API key marker for the unsloth start parent process.",
),
password: str = typer.Option(
"",
"--password",
@ -1985,6 +1991,8 @@ def run(
args.append("--no-cloudflare")
args.append("--secure" if secure else "--no-secure")
args.append("--tensor-parallel" if tensor_parallel else "--no-tensor-parallel")
if start_api_key_marker:
args.append("--start-api-key-marker")
if verbose:
args.append("--verbose")
# llama-server pass-through extras → child ctx.args → load payload.
@ -2045,6 +2053,11 @@ def run(
# 4. Create API key in-process.
api_key = _create_api_key_inprocess(api_key_name)
if start_api_key_marker:
# `unsloth start` redirects this process to a private 0600 log and
# uses the key to authenticate download-progress polling before the
# blocking load returns. The normal `unsloth run` output is unchanged.
typer.echo(f"UNSLOTH_START_API_KEY: {api_key}")
# 5. Load model via HTTP.
if not silent:

View file

@ -20,6 +20,7 @@ if str(_REPO_ROOT) not in sys.path:
import pytest
import typer
from typer.testing import CliRunner
import unsloth_cli.commands.start as start
@ -92,6 +93,8 @@ def test_claude_flags_passed_to_supported_claude(monkeypatch):
"--exclude-dynamic-system-prompt-sections",
"--settings",
start._claude_settings_overlay(MODEL["id"]),
"--append-system-prompt",
start._GEMMA_CLAUDE_NUDGE,
]
@ -100,6 +103,15 @@ def test_claude_flags_skipped_on_old_claude(monkeypatch):
assert start._claude_flags(MODEL["id"]) == []
def test_claude_nudge_is_gemma_only(monkeypatch):
_fake_claude(monkeypatch, "2.1.215 (Claude Code)\n")
flags = start._claude_flags("unsloth/Qwen3.5-9B-GGUF")
assert "--append-system-prompt" not in flags
assert start._GEMMA_CLAUDE_NUDGE not in flags
def test_claude_flags_skipped_on_unparseable_version(monkeypatch):
_fake_claude(monkeypatch, "weird build string\n")
assert start._claude_flags(MODEL["id"]) == []
@ -113,6 +125,8 @@ def test_claude_flags_detected_when_version_not_first_token(monkeypatch):
"--exclude-dynamic-system-prompt-sections",
"--settings",
start._claude_settings_overlay(MODEL["id"]),
"--append-system-prompt",
start._GEMMA_CLAUDE_NUDGE,
]
@ -426,6 +440,8 @@ def test_claude_flags_detects_supported_agent_only_in_install_dir(monkeypatch, t
"--exclude-dynamic-system-prompt-sections",
"--settings",
start._claude_settings_overlay(MODEL["id"]),
"--append-system-prompt",
start._GEMMA_CLAUDE_NUDGE,
]
@ -931,6 +947,25 @@ def test_resolve_model_attaches_to_loaded_catalog_hit_without_reload(monkeypatch
assert not any(u.endswith("/api/inference/load") for _, u in calls)
def test_resolve_model_without_request_rejects_unloaded_catalog(monkeypatch):
monkeypatch.setattr(
start,
"_http_json",
lambda *a, **k: {
"data": [
{
"id": "unsloth/Gemma-4-GGUF",
"loaded": False,
"context_length": 131072,
}
]
},
)
with pytest.raises(typer.Exit):
start._resolve_model(BASE, "sk-test", None)
def test_resolve_model_remote_studio_does_not_casefold_attach(monkeypatch):
# Against a remote Unsloth the local existence probe cannot see server-side paths,
# so a case-variant loaded id must NOT attach without a load: it could be a distinct
@ -1213,6 +1248,11 @@ def test_connect_model_flag_loads_on_server(fake_studio):
assert loads == [
("POST", f"{BASE}/api/inference/load", {"model_path": "unsloth/Qwen3.5-35B-A3B"})
]
assert (
f"Switching the Unsloth server from {MODEL['id']} to unsloth/Qwen3.5-35B-A3B"
in result.output
)
assert "unloads the current model for every attached session" in result.output
_assert_env_set(result.output, "ANTHROPIC_MODEL", "unsloth/Qwen3.5-35B-A3B")
@ -1761,13 +1801,130 @@ def test_start_studio_server_builds_command_and_waits(monkeypatch):
assert cmd[cmd.index("--gguf-variant") + 1] == "UD-Q4_K_XL"
assert cmd[cmd.index("--context-length") + 1] == "8192"
assert "--tensor-parallel" in cmd
assert "--start-api-key-marker" in cmd
assert cmd[cmd.index("-p") + 1] == "8888"
assert start.LoadOptions().load_in_4bit is True and "--no-load-in-4bit" not in cmd
assert captured["kwargs"].get("start_new_session") is True # own process group
assert server.pid == 4321
def test_auto_serves_when_no_server_then_tears_down(fake_studio, monkeypatch):
def test_start_studio_server_polls_progress_from_early_key(monkeypatch):
class FakePopen:
pid = 4321
def poll(self):
return None
tails = iter(
[
"UNSLOTH_START_API_KEY: sk-unsloth-early\nLoading model...",
"UNSLOTH_START_API_KEY: sk-unsloth-early\nModel loaded: owner/model",
]
)
created = []
class FakeProgress:
def __init__(self, base, key, model, variant):
created.append((base, key, model, variant, "created"))
def poll(self):
created.append("poll")
def close(self):
created.append("close")
def complete(self):
created.append("complete")
monkeypatch.setattr(start.subprocess, "Popen", lambda *a, **k: FakePopen())
monkeypatch.setattr(start, "_studio_healthy", lambda *a, **k: True)
monkeypatch.setattr(start, "_log_tail", lambda *a, **k: next(tails))
monkeypatch.setattr(start, "_ModelDownloadProgress", FakeProgress)
monkeypatch.setattr(start.time, "sleep", lambda _s: None)
server = start._start_studio_server(
BASE,
"owner/model-GGUF",
start.LoadOptions(gguf_variant = "Q4_K_M"),
)
assert server.pid == 4321
assert created[0] == (
BASE,
"sk-unsloth-early",
"owner/model-GGUF",
"Q4_K_M",
"created",
)
assert created.count("poll") == 2
assert created[-2:] == ["complete", "close"]
def test_load_model_with_progress_uses_selected_gguf_size(monkeypatch, capsys):
release = start.threading.Event()
calls = []
def http_json(method, url, token, payload = None, timeout = 30, error = None):
calls.append((method, url, payload))
if url.endswith("/api/inference/load"):
assert release.wait(timeout = 2)
return {"model": "owner/model-GGUF"}
if "/api/hub/gguf-variants?" in url:
return {
"default_variant": "Q8_0",
"variants": [
{
"quant": "UD-Q4_K_XL",
"filename": "model-UD-Q4_K_XL.gguf",
"size_bytes": 4 * 1024**3,
"download_size_bytes": 4 * 1024**3,
}
],
}
if "/api/hub/gguf-download-progress?" in url:
release.set()
return {
"downloaded_bytes": 2 * 1024**3,
"expected_bytes": 4 * 1024**3,
"progress": 0.5,
}
raise AssertionError(f"unexpected request: {method} {url}")
monkeypatch.setattr(start, "_http_json", http_json)
monkeypatch.setattr(start, "_DOWNLOAD_POLL_INTERVAL_S", 0.001)
result = start._load_model_with_progress(
BASE,
"sk-test",
"owner/model-GGUF",
start.LoadOptions(gguf_variant = "UD-Q4_K_XL"),
{"model_path": "owner/model-GGUF", "gguf_variant": "UD-Q4_K_XL"},
)
assert result == {"model": "owner/model-GGUF"}
output = capsys.readouterr().out
assert "Downloading model" in output
assert "100%" in output
progress_url = next(url for method, url, _ in calls if "gguf-download-progress" in url)
assert "variant=UD-Q4_K_XL" in progress_url
assert f"expected_bytes={4 * 1024**3}" in progress_url
def test_download_progress_ignores_fully_cached_bytes(capsys):
display = start._DownloadProgressDisplay()
display.update(
{
"downloaded_bytes": 4 * 1024**3,
"completed_bytes": 4 * 1024**3,
"expected_bytes": 4 * 1024**3,
"progress": 0.99,
}
)
display.close()
assert capsys.readouterr().out == ""
def test_auto_serves_when_no_server_then_keeps_server(fake_studio, monkeypatch):
monkeypatch.setattr(start, "find_studio_server", lambda: None)
started = {}
fake = SimpleNamespace(pid = 999, poll = lambda: None)
@ -1793,8 +1950,82 @@ def test_auto_serves_when_no_server_then_tears_down(fake_studio, monkeypatch):
assert started["model"] == "unsloth/Qwen3-1.7B-GGUF"
assert started["load"].gguf_variant == "UD-Q4_K_XL"
assert started["base"] == BASE
# Torn down after the agent session ended.
assert started.get("down") is fake
# A successful agent exit releases ownership and leaves the server available
# for another terminal. Explicit startup failures still use the cleanup path.
assert "down" not in started
assert start._auto_served_server is None
assert "is still running" in result.output
assert "unsloth studio stop" in result.output
def test_auto_served_agent_launch_failure_stops_server(fake_studio, monkeypatch):
monkeypatch.setattr(start, "find_studio_server", lambda: None)
stopped = []
fake = SimpleNamespace(pid = 999, poll = lambda: None)
def fake_start(*_args):
start._auto_served_server = fake
return fake
monkeypatch.setattr(start, "_start_studio_server", fake_start)
monkeypatch.setattr(start, "_shutdown_server", stopped.append)
monkeypatch.setattr(
start,
"_launch",
lambda *a, **k: (_ for _ in ()).throw(RuntimeError("agent launch failed")),
)
result = CliRunner().invoke(
start.start_app,
["claude", "--model", "unsloth/Qwen3-1.7B-GGUF"],
)
assert result.exit_code == 1
assert stopped == [fake]
assert "is still running" not in result.output
def test_auto_served_server_exit_is_not_reported_as_running(fake_studio, monkeypatch):
monkeypatch.setattr(start, "find_studio_server", lambda: None)
fake = SimpleNamespace(pid = 999, poll = lambda: 1)
def fake_start(*_args):
start._auto_served_server = fake
return fake
monkeypatch.setattr(start, "_start_studio_server", fake_start)
monkeypatch.setattr(start, "_launch", lambda *a, **k: 0)
result = CliRunner().invoke(
start.start_app,
["claude", "--model", "unsloth/Qwen3-1.7B-GGUF"],
)
assert result.exit_code == 0, result.output
assert "stopped during the session" in result.output
assert "is still running" not in result.output
def test_attached_server_prints_stop_hint_after_agent_exits(fake_studio, monkeypatch):
monkeypatch.setattr(start.shutil, "which", lambda _: "/usr/local/bin/claude")
monkeypatch.setattr(start, "_claude_flags", lambda *a, **k: [])
monkeypatch.setattr(
start.subprocess,
"run",
lambda command, env: SimpleNamespace(returncode = 0),
)
result = CliRunner().invoke(start.start_app, ["claude"])
assert result.exit_code == 0, result.output
assert f"Unsloth Studio is still running at {BASE}." in result.output
assert "unsloth studio stop" in result.output
def test_no_launch_recipe_does_not_print_stop_hint(fake_studio):
result = CliRunner().invoke(start.start_app, ["claude", "--no-launch"])
assert result.exit_code == 0, result.output
assert "is still running" not in result.output
def test_codex_preflight_failure_tears_down_auto_served(fake_studio, monkeypatch):

View file

@ -235,6 +235,13 @@ def test_reexec_forwards_parallel_all_aliases(monkeypatch, flag, value):
), f"{flag} {value} was dropped on re-exec; argv = {argv}"
def test_reexec_forwards_start_api_key_marker(monkeypatch):
"""The internal progress marker must survive the studio-venv re-exec."""
result, captured = _invoke_run(monkeypatch, _BASE + ["--start-api-key-marker"])
assert len(captured) == 1, result.output
assert "--start-api-key-marker" in captured[0]["argv"]
@pytest.mark.parametrize("platform", ["linux", "darwin", "win32"])
def test_reexec_argv_is_consistent_across_platforms(monkeypatch, platform):
"""Linux/Darwin (execvp) and Windows (Popen) must build the same argv."""