Check the fallback range, guard PID reuse, and keep writing studio.pid

This commit is contained in:
Nilay Yadav 2026-07-29 02:17:39 +05:30
commit 306915c031
3 changed files with 158 additions and 51 deletions

View file

@ -733,12 +733,21 @@ def _find_free_port(
host: str,
start: int,
max_attempts: int = 20,
avoid_own_studio: bool = False,
) -> int:
"""Find a free port from `start`, trying up to max_attempts ports."""
"""Find a free port from `start`, trying up to max_attempts ports.
``avoid_own_studio`` aborts rather than skipping past one of our own servers
in the fallback range, which would start a duplicate on a later port.
"""
for offset in range(max_attempts):
candidate = start + offset
if _is_port_free(host, candidate):
return candidate
if avoid_own_studio:
blocker = _get_pid_on_port(candidate)
if _blocker_is_own_studio(blocker):
_abort_already_running(blocker[0], candidate)
raise RuntimeError(f"Could not find a free port in range {start}-{start + max_attempts - 1}")
@ -750,16 +759,45 @@ PID_FILE_GLOB = "studio-*.pid"
def _pid_file_for_port(port: int) -> Path:
return _studio_root() / f"studio-{port}.pid"
# PID in the name: 127.0.0.1 and ::1 can share a port, and one file per port
# would let the second bind overwrite the first.
return _studio_root() / f"studio-{port}-{os.getpid()}.pid"
def _pid_alive(pid: int) -> bool:
try:
import psutil
return psutil.pid_exists(pid)
except ImportError:
pass
try:
os.kill(pid, 0)
except ProcessLookupError:
return False
except OSError:
return True
return True
def _pid_is_studio_backend(pid: int) -> bool:
"""Guard against PID reuse: a stale record must not block an unrelated process."""
try:
import psutil
cmdline = " ".join(psutil.Process(pid).cmdline()).lower()
except Exception:
return True
return "run.py" in cmdline or "unsloth" in cmdline
def _blocker_is_own_studio(blocker: "tuple[int, str] | None") -> bool:
"""True when the process holding the port is a server we recorded."""
return bool(blocker) and blocker[0] in _recorded_studio_pids()
if not blocker:
return False
return blocker[0] in _recorded_studio_pids() and _pid_is_studio_backend(blocker[0])
def _recorded_studio_pids() -> "set[int]":
"""PIDs recorded under this STUDIO_HOME."""
"""Live PIDs recorded under this STUDIO_HOME; prunes dead records."""
pids: "set[int]" = set()
try:
paths = list(_studio_root().glob(PID_FILE_GLOB)) + [_PID_FILE]
@ -770,11 +808,26 @@ def _recorded_studio_pids() -> "set[int]":
text = path.read_text(encoding = "utf-8").strip()
except (OSError, UnicodeDecodeError):
continue
if text.isdigit():
pids.add(int(text))
if not text.isdigit():
continue
pid = int(text)
if _pid_alive(pid):
pids.add(pid)
elif path != _PID_FILE:
path.unlink(missing_ok = True)
return pids
def _abort_already_running(pid: int, port: int) -> "NoReturn":
print(
f"Error: Unsloth Studio is already running on port {port} (PID {pid}). Run "
"`unsloth studio stop` first, or start this one on a different --port.",
file = sys.stderr,
flush = True,
)
sys.exit(1)
# Direct backend launches bypass the CLI's env re-export; do it here for
# real custom roots so unsloth-zoo's import-time LLAMA_CPP_DEFAULT_DIR
# picks up the custom build. Skip legacy-default to avoid flipping
@ -809,22 +862,23 @@ def _write_pid_file(port: int):
try:
path.parent.mkdir(parents = True, exist_ok = True)
path.write_text(str(os.getpid()), encoding = "utf-8")
# An older CLI's `stop` only reads this one.
_PID_FILE.write_text(str(os.getpid()), encoding = "utf-8")
except OSError:
return
_OWN_PID_FILE = path
def _remove_pid_file():
"""Remove the PID file if it belongs to this process."""
"""Remove the PID files that belong to this process."""
if _OWN_PID_FILE is None:
return
try:
if _OWN_PID_FILE.is_file():
stored = _OWN_PID_FILE.read_text(encoding = "utf-8").strip()
if stored == str(os.getpid()):
_OWN_PID_FILE.unlink(missing_ok = True)
except (OSError, UnicodeDecodeError):
pass
for path in (_OWN_PID_FILE, _PID_FILE):
try:
if path.is_file() and path.read_text(encoding = "utf-8").strip() == str(os.getpid()):
path.unlink(missing_ok = True)
except (OSError, UnicodeDecodeError):
pass
def _graceful_shutdown(server = None):
@ -1572,15 +1626,8 @@ def run_server(
blocker = _get_pid_on_port(port)
# Falling back past our own server is what creates the orphan.
if _blocker_is_own_studio(blocker):
print(
f"Error: Unsloth Studio is already running on port {port} "
f"(PID {blocker[0]}). Run `unsloth studio stop` first, or start this "
"one on a different --port.",
file = sys.stderr,
flush = True,
)
sys.exit(1)
port = _find_free_port(host, port + 1)
_abort_already_running(blocker[0], port)
port = _find_free_port(host, port + 1, avoid_own_studio = True)
if not silent:
print("")
print("=" * 50)

View file

@ -26,67 +26,127 @@ def isolated_root(tmp_path, monkeypatch):
monkeypatch.setattr(run, "_studio_root", lambda: tmp_path)
monkeypatch.setattr(run, "_PID_FILE", tmp_path / "studio.pid")
monkeypatch.setattr(run, "_OWN_PID_FILE", None)
monkeypatch.setattr(run, "_pid_alive", lambda pid: True)
monkeypatch.setattr(run, "_pid_is_studio_backend", lambda pid: True)
yield
def test_write_pid_file_is_per_port(tmp_path):
def _files(tmp_path):
return sorted(p.name for p in tmp_path.glob("studio-*.pid"))
def test_write_pid_file_records_port_and_pid(tmp_path):
run._write_pid_file(8901)
path = tmp_path / "studio-8901.pid"
assert path.read_text(encoding = "utf-8") == str(os.getpid())
assert _files(tmp_path) == [f"studio-8901-{os.getpid()}.pid"]
assert (tmp_path / f"studio-8901-{os.getpid()}.pid").read_text() == str(os.getpid())
def test_write_pid_file_also_updates_the_legacy_file(tmp_path):
# An older CLI's `stop` only reads studio.pid.
run._write_pid_file(8901)
assert (tmp_path / "studio.pid").read_text(encoding = "utf-8") == str(os.getpid())
def test_second_port_does_not_clobber_the_first(tmp_path):
(tmp_path / "studio-8901.pid").write_text("8550", encoding = "utf-8")
(tmp_path / "studio-8901-8550.pid").write_text("8550", encoding = "utf-8")
run._write_pid_file(8902)
assert (tmp_path / "studio-8901.pid").read_text(encoding = "utf-8") == "8550"
assert (tmp_path / "studio-8902.pid").read_text(encoding = "utf-8") == str(os.getpid())
assert (tmp_path / "studio-8901-8550.pid").read_text(encoding = "utf-8") == "8550"
assert (tmp_path / f"studio-8902-{os.getpid()}.pid").exists()
def test_same_port_on_two_binds_does_not_clobber(tmp_path):
# 127.0.0.1:8888 and ::1:8888 can both listen; one file per port would lose one.
(tmp_path / "studio-8888-8550.pid").write_text("8550", encoding = "utf-8")
run._write_pid_file(8888)
assert len(_files(tmp_path)) == 2
def test_remove_pid_file_only_removes_our_own(tmp_path):
run._write_pid_file(8901)
(tmp_path / "studio-8902.pid").write_text("8600", encoding = "utf-8")
(tmp_path / "studio-8902-8600.pid").write_text("8600", encoding = "utf-8")
run._remove_pid_file()
assert not (tmp_path / "studio-8901.pid").exists()
assert (tmp_path / "studio-8902.pid").exists()
assert _files(tmp_path) == ["studio-8902-8600.pid"]
assert not (tmp_path / "studio.pid").exists()
def test_remove_pid_file_leaves_a_reused_entry_alone(tmp_path):
run._write_pid_file(8901)
(tmp_path / "studio-8901.pid").write_text("999999", encoding = "utf-8")
own = tmp_path / f"studio-8901-{os.getpid()}.pid"
own.write_text("999999", encoding = "utf-8")
run._remove_pid_file()
assert (tmp_path / "studio-8901.pid").read_text(encoding = "utf-8") == "999999"
assert own.read_text(encoding = "utf-8") == "999999"
def test_recorded_studio_pids_reads_per_port_and_legacy_files(tmp_path):
(tmp_path / "studio-8901.pid").write_text("8550", encoding = "utf-8")
(tmp_path / "studio-8902.pid").write_text("8600", encoding = "utf-8")
(tmp_path / "studio-8901-8550.pid").write_text("8550", encoding = "utf-8")
(tmp_path / "studio-8902-8600.pid").write_text("8600", encoding = "utf-8")
(tmp_path / "studio.pid").write_text("4242", encoding = "utf-8")
assert run._recorded_studio_pids() == {8550, 8600, 4242}
def test_recorded_studio_pids_ignores_corrupt_files(tmp_path):
(tmp_path / "studio-8901.pid").write_text("not-a-pid", encoding = "utf-8")
(tmp_path / "studio-8901-x.pid").write_text("not-a-pid", encoding = "utf-8")
assert run._recorded_studio_pids() == set()
def test_recorded_studio_pids_prunes_dead_records(tmp_path, monkeypatch):
monkeypatch.setattr(run, "_pid_alive", lambda pid: False)
(tmp_path / "studio-8901-8550.pid").write_text("8550", encoding = "utf-8")
assert run._recorded_studio_pids() == set()
assert not (tmp_path / "studio-8901-8550.pid").exists()
def test_own_studio_blocking_the_port_is_recognised(tmp_path):
(tmp_path / "studio-8901.pid").write_text("8550", encoding = "utf-8")
(tmp_path / "studio-8901-8550.pid").write_text("8550", encoding = "utf-8")
assert run._blocker_is_own_studio((8550, "python")) is True
def test_a_foreign_blocker_still_falls_back(tmp_path):
# jupyter-lab on 8888 must keep the fallback, not abort the launch.
(tmp_path / "studio-8901.pid").write_text("8550", encoding = "utf-8")
(tmp_path / "studio-8901-8550.pid").write_text("8550", encoding = "utf-8")
assert run._blocker_is_own_studio((117, "jupyter-lab")) is False
assert run._blocker_is_own_studio(None) is False
def test_a_reused_pid_is_not_treated_as_our_studio(tmp_path, monkeypatch):
# Stale record + the OS handing that PID to something else must not abort.
monkeypatch.setattr(run, "_pid_is_studio_backend", lambda pid: False)
(tmp_path / "studio-8901-8550.pid").write_text("8550", encoding = "utf-8")
assert run._blocker_is_own_studio((8550, "postgres")) is False
def test_fallback_aborts_on_our_own_server_further_up_the_range(tmp_path, monkeypatch):
# jupyter holds 8888, our server holds 8889: skipping to 8890 is the duplicate.
(tmp_path / "studio-8889-8550.pid").write_text("8550", encoding = "utf-8")
monkeypatch.setattr(run, "_is_port_free", lambda host, p: p >= 8890)
monkeypatch.setattr(
run, "_get_pid_on_port", lambda p: (8550, "python") if p == 8889 else (117, "jupyter-lab")
)
with pytest.raises(SystemExit) as excinfo:
run._find_free_port("127.0.0.1", 8889, avoid_own_studio = True)
assert excinfo.value.code == 1
def test_fallback_still_skips_foreign_processes(tmp_path, monkeypatch):
monkeypatch.setattr(run, "_is_port_free", lambda host, p: p >= 8890)
monkeypatch.setattr(run, "_get_pid_on_port", lambda p: (117, "jupyter-lab"))
assert run._find_free_port("127.0.0.1", 8889, avoid_own_studio = True) == 8890

View file

@ -67,8 +67,8 @@ def _run_stop(studio_mod):
def test_stop_kills_every_recorded_server(monkeypatch, tmp_path):
studio_mod, _live, killed = _install(monkeypatch, tmp_path, alive = {8550, 8600})
_write_pid(tmp_path, "studio-8901.pid", 8550)
_write_pid(tmp_path, "studio-8902.pid", 8600)
_write_pid(tmp_path, "studio-8901-8550.pid", 8550)
_write_pid(tmp_path, "studio-8902-8600.pid", 8600)
result = _run_stop(studio_mod)
@ -80,8 +80,8 @@ def test_stop_kills_every_recorded_server(monkeypatch, tmp_path):
def test_stop_does_not_leave_the_older_instance_running(monkeypatch, tmp_path):
# The reported symptom: stop claimed success while instance A kept serving.
studio_mod, live, _killed = _install(monkeypatch, tmp_path, alive = {8550, 8600})
_write_pid(tmp_path, "studio-8901.pid", 8550)
_write_pid(tmp_path, "studio-8902.pid", 8600)
_write_pid(tmp_path, "studio-8901-8550.pid", 8550)
_write_pid(tmp_path, "studio-8902-8600.pid", 8600)
result = _run_stop(studio_mod)
@ -111,13 +111,13 @@ def test_stop_reports_nothing_running_without_pid_files(monkeypatch, tmp_path):
def test_stop_cleans_stale_pid_files_without_claiming_a_stop(monkeypatch, tmp_path):
studio_mod, _live, killed = _install(monkeypatch, tmp_path, alive = set())
_write_pid(tmp_path, "studio-8901.pid", 8550)
_write_pid(tmp_path, "studio-8901-8550.pid", 8550)
result = _run_stop(studio_mod)
assert result.exit_code == 0, result.output
assert killed == []
assert not (tmp_path / "studio-8901.pid").exists()
assert not (tmp_path / "studio-8901-8550.pid").exists()
assert "stopped" not in result.output.lower()
@ -130,14 +130,14 @@ def test_stop_does_not_claim_a_stop_while_a_server_is_still_alive(monkeypatch, t
monkeypatch.setattr(studio_mod, "_pid_alive", lambda pid: True)
monkeypatch.setattr(studio_mod.os, "kill", lambda pid, sig: None)
monkeypatch.setattr(sys, "platform", "linux")
_write_pid(tmp_path, "studio-8901.pid", 8550)
_write_pid(tmp_path, "studio-8901-8550.pid", 8550)
result = _run_stop(studio_mod)
assert result.exit_code == 0, result.output
assert "shutting down" in result.output.lower()
assert "stopped" not in result.output.lower()
assert (tmp_path / "studio-8901.pid").exists()
assert (tmp_path / "studio-8901-8550.pid").exists()
def test_stop_continues_after_one_server_fails_to_stop(monkeypatch, tmp_path):
@ -155,8 +155,8 @@ def test_stop_continues_after_one_server_fails_to_stop(monkeypatch, tmp_path):
monkeypatch.setattr(studio_mod.os, "kill", fake_kill)
monkeypatch.setattr(sys, "platform", "linux")
_write_pid(tmp_path, "studio-8901.pid", 8550)
_write_pid(tmp_path, "studio-8902.pid", 8600)
_write_pid(tmp_path, "studio-8901-8550.pid", 8550)
_write_pid(tmp_path, "studio-8902-8600.pid", 8600)
result = _run_stop(studio_mod)
@ -168,9 +168,9 @@ def test_stop_continues_after_one_server_fails_to_stop(monkeypatch, tmp_path):
def test_stop_discards_a_corrupt_pid_file(monkeypatch, tmp_path):
studio_mod, _live, _killed = _install(monkeypatch, tmp_path, alive = set())
(tmp_path / "studio-8901.pid").write_text("not-a-pid", encoding = "utf-8")
(tmp_path / "studio-8901-8550.pid").write_text("not-a-pid", encoding = "utf-8")
result = _run_stop(studio_mod)
assert result.exit_code == 0, result.output
assert not (tmp_path / "studio-8901.pid").exists()
assert not (tmp_path / "studio-8901-8550.pid").exists()