From 9d693b35e2519615f143fceb29d80a7f6308a791 Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Wed, 13 May 2026 14:18:37 +0000 Subject: [PATCH 1/3] studio: hardening follow-up (refresh-token race, sandbox alias bypass, /api/health launcher contract, Windows kill path) Addresses the open review findings on PR #5375 plus the Windows Studio UI CI regression that landed on main. ### Refresh-token rotation (auth/storage.py) * DELETE ... RETURNING is SQLite 3.35+. Older system SQLite (Ubuntu 20.04, RHEL 8, some Windows builds) raised OperationalError and turned /api/auth/refresh into a 500 for every user. consume_refresh_token now feature-detects RETURNING at first use, falls back to a transactional SELECT + DELETE, and uses delete_cursor.rowcount as the canonical "did I win the race" signal so two concurrent refreshes still produce exactly one winner. * Added test_refresh_token_consume.py covering single-use rotation, replay -> None, the desktop flag round-trip, and an 8-thread race that asserts exactly one winner on the fallback path. ### /api/auth/logout (routes/auth.py) * Logout was swallowing all exceptions from revoke_user_refresh_tokens and returning 204 even when refresh tokens were not actually invalidated. The endpoint now surfaces a 500 with a generic detail (and logs the exception class for the operator) so a caller cannot be told "you're logged out" while a stolen refresh token stays live. ### Sandbox AST policy (core/inference/tools.py) The PR-5375 visitor only matched calls on literal "requests." / "urllib.request.urlopen" FQ names. That left three bypasses: 1. Module aliases: `import requests as r; r.get("http://169.254.169.254/")`. 2. From-import + alias: `from requests import get as fetch; fetch(...)`. 3. Session-bound variables: `s = requests.Session(); s.get(...)`. 4. Variable URLs: `u = "http://..."; requests.get(u)`. The visitor now tracks imports (Import, ImportFrom) and assignments (Assign, including JoinedStr f-strings that fold to a constant), synthesises canonical FQ names for aliased calls and session methods, and resolves simple variable URLs through the assignment table before policy eval. Genuinely runtime-computed URLs (env vars, user input) are now flagged as "opaque_url_blocked" rather than allowed through silently. _NETWORK_FQ_PREFIXES gained the session-method synthetic prefixes (requests.Session., httpx.Client., httpx.AsyncClient., aiohttp.ClientSession.); _UPLOAD_HTTP_METHODS gained the matching Session.post/put/patch/delete/request entries. Added 12 tests across TestImportAliasResolution and TestSessionObjectMethods plus updated TestUntrustedHostBlock (test_dynamic_url_not_statically_blocked replaced with three sharper tests: variable URL resolved, f-string folded, and opaque-runtime URL flagged). ### Windows process-group kill (core/inference/tools.py) * _kill_process_tree was unconditionally calling os.getpgid / os.killpg, which raised AttributeError on Windows and skipped the kill entirely. The supervisor then leaked runaway tool processes and returned an execution error instead of a clean timeout. The helper now gates on hasattr(os, "getpgid") and hasattr(os, "killpg"), and on Windows falls back to proc.kill() + a best-effort taskkill /F /T. Added test_kill_process_tree_platform.py with a Linux/macOS pgid path test plus two simulated-Windows tests that monkeypatch the attributes off os. ### /api/health launcher contract (main.py) * Stripping every legacy identity field from the unauthenticated payload broke install.sh::_check_health, studio/src-tauri/src/preflight/backend.rs, and the run_studio_browser_test orchestrator, all of which match on service / studio_root_id / desktop_protocol_version without authenticating. The launcher contract (status, timestamp, service, studio_root_id, the four desktop_* capability bits) is now always exposed; the sensitive diagnostic fields (version, studio_version, device_type, chat_only, native_path_leases_supported, desktop_owner) remain gated on a valid bearer. * Added test_health_unauth_contract.py for the contract on both sides, and updated test_middleware.py::TestHealthAuthGate to match. ### CSP (main.py) * connect-src "self" was blocking the frontend's direct Hugging Face searches (use-hf-model-search, use-hf-dataset-search). connect-src now includes huggingface.co + *.huggingface.co + cdn-lfs.huggingface.co + cdn-lfs.hf.co + hf.co + *.hf.co; img-src adds huggingface.co + cdn-avatars.huggingface.co for the search avatar pickers. script-src stays at 'self' + per-response nonce; no 'unsafe-inline' anywhere. ### tool_call_id correlation (models/inference.py, routes/inference.py) * ChatMessage._validate_role_shape was synthesising a random tool_call_id when role="tool" arrived without one. The random id broke correlation with the preceding assistant tool_calls and OpenAI-compatible backends rejected the tool result. The validator now emits a recognisable TOOL_CALL_ID_SYNTH_PREFIX placeholder; _pair_orphan_tool_ids in the route walks the message list before passthrough and rewrites synth ids to the matching announced tool_call id (FIFO, skipping already-consumed ids). When no preceding tool_call is available the synth id stays so the upstream backend can produce an explicit error. * Added test_tool_id_pairing.py covering single rewrite, idempotency, FIFO pairing, no-announce fallthrough, and not double-consuming an explicit match. ### Training cancel cleanup (core/training/{training,worker}.py) * On cancel-no-save the worker emits "complete" with output_dir=None; force_terminate was snapshotting _output_dir (None at that point) and the new _cleanup_cancelled_checkpoints call was skipped, so periodic checkpoint-* dirs stayed on disk. The worker now emits "run_started" with the resolved output_dir immediately after path resolution; force_terminate prefers that value (_active_run_dir) when cleaning up so the cancel-no-save path actually removes the partial checkpoints. ### Windows Studio UI CI test robustness (tests/studio/playwright_extra_ui.py) * The /studio block was looking for "Configure", "Current run", "History" tabs without waiting for runtime hydration -- under the 1.5s timeout the loading placeholder was still rendered and the assertions failed in CI. The probe now waits up to 30s for either the studio tabs or the chat_only redirect, clicks Configure before checking the data-tour anchors, falls back to text-based selectors if Radix tabs do not yet expose role="tab", and adds a 3s grace for the lazy-mounted ParamsSection. chat_only is now read from /api/health with the bearer token (since the field is gated post-hardening); the test falls back to URL-shape detection if the field is absent. ## Cross-platform / cross-browser simulation Before pushing, the patches were exercised in an isolated `uv venv` under workspace/temp/sim_venv/: * 19 cross-platform sim tests pinning _kill_process_tree (Linux / macOS / Windows simulated by monkeypatching os.getpgid/killpg + sys.platform), the refresh-token RETURNING fallback under simulated old SQLite, and the AST policy across all three simulated platforms. * 24 multi-browser Playwright smokes (Chromium, Firefox, WebKit) against all 8 live Studios (ports 18801-18808), verifying /api/health response shape, CSP, X-Frame-Options, X-Content-Type-Options, Referrer-Policy, Permissions-Policy and the Server header on each engine. WebKit skips automatically when libgtk-4 / libgraphene / libavif are not installed system-wide. * 8 Studios were brought up in parallel (2 per GPU across CUDA_VISIBLE_DEVICES=4,5,6,7) and ran the live security probe; all 8 returned PASS=18 FAIL=0 SKIP=1 (skip is the auth-bearer-flow blocked by the in-test rate-limit hit). ## Test plan * Studio backend unit tests: pytest studio/backend/tests/ ignoring the GPU-dependent and KV-cache networked tests -> 816 passed, 10 skipped. * New tests: 12 sandbox AST cases + 8 refresh-token cases + 4 kill_process_tree cases + 8 tool-id pairing cases + 6 health-contract cases all pass. * Live HTTP probe across 8 Studios: 18/18 PASS on every Studio. * Multi-browser Playwright probe (Chromium + Firefox) across all 8 Studios: 16/16 PASS. --- studio/backend/auth/storage.py | 77 +++- studio/backend/core/inference/tools.py | 351 +++++++++++++++--- studio/backend/core/training/training.py | 36 +- studio/backend/core/training/worker.py | 14 + studio/backend/main.py | 62 +++- studio/backend/models/inference.py | 19 +- studio/backend/routes/auth.py | 27 +- studio/backend/routes/inference.py | 74 ++++ .../tests/test_health_unauth_contract.py | 177 +++++++++ .../tests/test_kill_process_tree_platform.py | 160 ++++++++ studio/backend/tests/test_middleware.py | 50 ++- .../tests/test_refresh_token_consume.py | 139 +++++++ studio/backend/tests/test_sandbox_tools.py | 139 ++++++- studio/backend/tests/test_tool_id_pairing.py | 139 +++++++ tests/studio/playwright_extra_ui.py | 101 ++++- 15 files changed, 1475 insertions(+), 90 deletions(-) create mode 100644 studio/backend/tests/test_health_unauth_contract.py create mode 100644 studio/backend/tests/test_kill_process_tree_platform.py create mode 100644 studio/backend/tests/test_refresh_token_consume.py create mode 100644 studio/backend/tests/test_tool_id_pairing.py diff --git a/studio/backend/auth/storage.py b/studio/backend/auth/storage.py index 3233aa05ef..0d873ff8f9 100644 --- a/studio/backend/auth/storage.py +++ b/studio/backend/auth/storage.py @@ -480,11 +480,44 @@ def save_refresh_token( conn.close() +_RETURNING_SUPPORTED: Optional[bool] = None + + +def _supports_returning() -> bool: + """Feature-detect SQLite ``DELETE ... RETURNING`` support (SQLite 3.35+). + + Cached after first probe. Older system SQLite (e.g. Ubuntu 20.04, RHEL 8, + some Windows builds) ships SQLite < 3.35 and raises ``OperationalError`` + when ``RETURNING`` is parsed. We fall back to a transactional + ``SELECT`` + ``DELETE`` whose atomicity is enforced by checking the + rowcount of the ``DELETE`` (a concurrent winner returns rowcount=0). + """ + global _RETURNING_SUPPORTED + if _RETURNING_SUPPORTED is not None: + return _RETURNING_SUPPORTED + try: + import sqlite3 as _sqlite3 + + ver = getattr(_sqlite3, "sqlite_version_info", (0, 0, 0)) + if ver >= (3, 35, 0): + _RETURNING_SUPPORTED = True + return True + except Exception: + pass + _RETURNING_SUPPORTED = False + return False + + def consume_refresh_token(token: str) -> Optional[Tuple[str, bool]]: """Atomically validate-and-delete a refresh token for single-use rotation. - DELETE RETURNING fuses validate and delete into one statement so two - concurrent refresh requests cannot both consume the same token. + On SQLite 3.35+ we use ``DELETE ... RETURNING`` which fuses validate and + delete into one statement so two concurrent refresh requests cannot both + consume the same token. On older SQLite the helper falls back to + ``SELECT`` + ``DELETE`` inside a single transaction; the ``DELETE``'s + ``rowcount`` is then the source of truth (a concurrent winner returns 0, + so the loser correctly sees ``None``). Either way the contract is + "at most one caller succeeds per token". """ token_hash = _hash_token(token) now = datetime.now(timezone.utc).isoformat() @@ -494,17 +527,51 @@ def consume_refresh_token(token: str) -> Optional[Tuple[str, bool]]: "DELETE FROM refresh_tokens WHERE expires_at < ?", (now,), ) + if _supports_returning(): + try: + cur = conn.execute( + """ + DELETE FROM refresh_tokens + WHERE token_hash = ? AND expires_at >= ? + RETURNING username, is_desktop + """, + (token_hash, now), + ) + row = cur.fetchone() + conn.commit() + if row is None: + return None + return row["username"], bool(row["is_desktop"]) + except Exception: + # Some Python builds advertise sqlite_version_info >= 3.35 + # while bundling an older amalgamation; flip the cache so we + # never retry the RETURNING path on this process and let the + # SELECT+DELETE path handle the rest. + global _RETURNING_SUPPORTED + _RETURNING_SUPPORTED = False + conn.rollback() + # Fallback: SELECT then DELETE with rowcount check. The DELETE's + # rowcount is the canonical "did I win the race" signal -- two + # concurrent callers can both SELECT the row, but only one DELETE + # will return rowcount=1; the other gets 0 and must report failure. cur = conn.execute( """ - DELETE FROM refresh_tokens + SELECT id, username, is_desktop FROM refresh_tokens WHERE token_hash = ? AND expires_at >= ? - RETURNING username, is_desktop """, (token_hash, now), ) row = cur.fetchone() - conn.commit() if row is None: + conn.commit() + return None + del_cur = conn.execute( + "DELETE FROM refresh_tokens WHERE id = ?", + (row["id"],), + ) + conn.commit() + if del_cur.rowcount != 1: + # Another caller consumed it between our SELECT and DELETE. return None return row["username"], bool(row["is_desktop"]) finally: diff --git a/studio/backend/core/inference/tools.py b/studio/backend/core/inference/tools.py index 70db5477d4..48bf0a5e26 100644 --- a/studio/backend/core/inference/tools.py +++ b/studio/backend/core/inference/tools.py @@ -1091,6 +1091,16 @@ def _check_signal_escape_patterns(code: str): "requests.head", "requests.request", "requests.Session", + # Session-bound methods. The visitor's _resolve_fq synthesises + # ``requests.Session.`` / ``requests.sessions.Session.`` + # / ``httpx.Client.`` etc. when the receiver is a variable + # bound to a session constructor; the prefix entries below make + # those synthesised names match. + "requests.Session.", + "requests.sessions.Session.", + "httpx.Client.", + "httpx.AsyncClient.", + "aiohttp.ClientSession.", "http.client.HTTPConnection", "http.client.HTTPSConnection", "httpx.get", @@ -1109,11 +1119,31 @@ def _check_signal_escape_patterns(code: str): "requests.patch", "requests.delete", "requests.request", + "requests.Session.post", + "requests.Session.put", + "requests.Session.patch", + "requests.Session.delete", + "requests.Session.request", + "requests.sessions.Session.post", + "requests.sessions.Session.put", + "requests.sessions.Session.patch", + "requests.sessions.Session.delete", + "requests.sessions.Session.request", "httpx.post", "httpx.put", "httpx.patch", "httpx.delete", "httpx.request", + "httpx.Client.post", + "httpx.Client.put", + "httpx.Client.patch", + "httpx.Client.delete", + "httpx.Client.request", + "httpx.AsyncClient.post", + "httpx.AsyncClient.put", + "httpx.AsyncClient.patch", + "httpx.AsyncClient.delete", + "httpx.AsyncClient.request", "urllib.request.urlopen", "urllib.request.Request", ) @@ -1334,16 +1364,204 @@ def _check_signal_escape_patterns(code: str): and node.func.attr in _UPLOAD_HF_METHODS ) + # Modules whose top-level / class-level functions are HTTP egress. + # Used to resolve aliases (``import requests as r``) and named imports + # (``from requests import get``) back to a canonical FQ name so the + # prefix check below still fires. Listed once so both ``visit_Import`` + # and ``visit_ImportFrom`` agree on what's network-relevant. + _NETWORK_MODULES = frozenset( + { + "socket", + "urllib", + "urllib.request", + "urllib3", + "requests", + "httpx", + "aiohttp", + "http", + "http.client", + } + ) + # Symbols imported from a network module whose call is itself an + # egress (``from requests import get; get("http://...")``). + _NETWORK_FROM_NAMES = frozenset( + { + "get", "post", "put", "delete", "patch", "head", "request", + "Session", "urlopen", "urlretrieve", "create_connection", + "getaddrinfo", "HTTPConnection", "HTTPSConnection", + "Client", "AsyncClient", "ClientSession", + } + ) + # Session/Client method names that perform HTTP egress. A literal + # call on any variable bound to a session-shaped constructor + # (``s = requests.Session(); s.get(...)``) becomes egress-equivalent. + _SESSION_METHOD_NAMES = frozenset( + {"get", "post", "put", "delete", "patch", "head", "request", "send"} + ) + # Constructor calls that produce a session/client object whose + # method calls become egress. + _SESSION_CONSTRUCTOR_FQS = frozenset( + { + "requests.Session", + "requests.sessions.Session", + "httpx.Client", + "httpx.AsyncClient", + "aiohttp.ClientSession", + } + ) + class NetworkAndIoVisitor(ast.NodeVisitor): + def __init__(self): + super().__init__() + # alias -> canonical module name. Populated by visit_Import + # so ``import requests as r`` lets us map ``r`` -> ``requests``. + self._module_aliases: dict[str, str] = {} + # local name -> canonical FQ. Populated by visit_ImportFrom + # so ``from requests import get as fetch`` maps ``fetch`` -> + # ``requests.get``. + self._symbol_aliases: dict[str, str] = {} + # variable -> literal-string value, for simple assignments. + # Lets ``u = "http://169.254.169.254"; requests.get(u)`` resolve. + self._string_vars: dict[str, str] = {} + # variable -> session constructor FQ; lets + # ``s = requests.Session(); s.get(url)`` register as egress. + self._session_vars: dict[str, str] = {} + + # ── Import tracking ───────────────────────────────────── + def visit_Import(self, node): + for alias in node.names: + target = alias.asname or alias.name + # ``import requests`` -> aliases["requests"] = "requests"; + # ``import requests as r`` -> aliases["r"] = "requests". + self._module_aliases[target] = alias.name + self.generic_visit(node) + + def visit_ImportFrom(self, node): + mod = node.module or "" + if mod in _NETWORK_MODULES: + for alias in node.names: + target = alias.asname or alias.name + if alias.name in _NETWORK_FROM_NAMES: + # Build the canonical FQ name; ``from requests + # import get`` -> aliases["get"] = "requests.get". + self._symbol_aliases[target] = f"{mod}.{alias.name}" + self.generic_visit(node) + + # ── Variable tracking (literal string + session constructor) ── + def visit_Assign(self, node): + # Track simple ``name = "literal"`` and + # ``name = requests.Session()``. + if len(node.targets) == 1 and isinstance(node.targets[0], ast.Name): + name = node.targets[0].id + value = node.value + if isinstance(value, ast.Constant) and isinstance(value.value, str): + self._string_vars[name] = value.value + elif isinstance(value, ast.JoinedStr): + # f-string: only resolve when every part is a constant + # or a tracked string variable. Otherwise treat as + # opaque so we don't accept a partly-dynamic URL. + parts: list[str] = [] + ok = True + for piece in value.values: + if isinstance(piece, ast.Constant) and isinstance(piece.value, str): + parts.append(piece.value) + elif ( + isinstance(piece, ast.FormattedValue) + and isinstance(piece.value, ast.Name) + and piece.value.id in self._string_vars + ): + parts.append(self._string_vars[piece.value.id]) + else: + ok = False + break + if ok: + self._string_vars[name] = "".join(parts) + elif isinstance(value, ast.Call): + fq = self._resolve_fq(value) + if fq in _SESSION_CONSTRUCTOR_FQS: + self._session_vars[name] = fq + self.generic_visit(node) + + # ── FQ resolution helpers ──────────────────────────────── + def _resolve_fq(self, call_node: ast.Call) -> str: + """Return the canonical FQ name for a Call's target. + + ``requests.get(...)`` -> ``"requests.get"``; + ``r.get(...)`` where ``r=requests`` -> ``"requests.get"``; + ``fetch(...)`` where ``from requests import get as fetch`` + -> ``"requests.get"``; + ``s.get(...)`` where ``s = requests.Session()`` -> + ``"requests.Session.get"`` (a synthetic prefix that the + prefix-check below treats as egress). + """ + func = call_node.func + # Bare name: ``get(...)`` -- check symbol aliases first + # (``from requests import get`` -> ``"requests.get"``). + if isinstance(func, ast.Name): + return self._symbol_aliases.get(func.id, func.id) + if isinstance(func, ast.Attribute): + parts: list[str] = [func.attr] + cur = func.value + while isinstance(cur, ast.Attribute): + parts.insert(0, cur.attr) + cur = cur.value + if isinstance(cur, ast.Name): + head = cur.id + # If head is a session-bound variable, synthesise a + # prefixable FQ (e.g. ``"requests.Session.get"``). + if head in self._session_vars and parts and parts[-1] in _SESSION_METHOD_NAMES: + return f"{self._session_vars[head]}.{parts[-1]}" + # Map the head through module-alias table. + resolved_head = self._module_aliases.get(head, head) + return ".".join([resolved_head, *parts]) + return "" + + def _resolve_url_arg(self, expr: ast.AST) -> "tuple[str | None, str | None]": + """Return (url_string, host_string) for a call argument. + + Resolves literal strings, simple variable assignments to + literals, and f-strings that fold to a constant. Returns + ``(None, None)`` for opaque / dynamic expressions, which the + caller must treat as "host not statically verifiable" rather + than "allowed". + """ + url: "str | None" = None + host: "str | None" = None + if isinstance(expr, ast.Constant) and isinstance(expr.value, str): + url = expr.value + elif isinstance(expr, ast.Name) and expr.id in self._string_vars: + url = self._string_vars[expr.id] + elif isinstance(expr, ast.JoinedStr): + parts: list[str] = [] + ok = True + for piece in expr.values: + if isinstance(piece, ast.Constant) and isinstance(piece.value, str): + parts.append(piece.value) + elif ( + isinstance(piece, ast.FormattedValue) + and isinstance(piece.value, ast.Name) + and piece.value.id in self._string_vars + ): + parts.append(self._string_vars[piece.value.id]) + else: + ok = False + break + if ok: + url = "".join(parts) + elif isinstance(expr, ast.Tuple) and expr.elts: + e0 = expr.elts[0] + if isinstance(e0, ast.Constant) and isinstance(e0.value, str): + host = e0.value + elif isinstance(e0, ast.Name) and e0.id in self._string_vars: + host = self._string_vars[e0.id] + if url and host is None: + m = re.match(r"^\w+://([^/?#]+)", url) + if m: + host = m.group(1) + return url, host + def visit_Call(self, node): - parts: list[str] = [] - cur = node.func - while isinstance(cur, ast.Attribute): - parts.insert(0, cur.attr) - cur = cur.value - if isinstance(cur, ast.Name): - parts.insert(0, cur.id) - fq = ".".join(parts) if parts else "" + fq = self._resolve_fq(node) if _method_call_is_hf_upload(node): network_calls.append( @@ -1360,14 +1578,7 @@ def _check_signal_escape_patterns(code: str): and node.func.attr == "connect" and node.args ): - a0 = node.args[0] - host_lit = None - if isinstance(a0, ast.Tuple) and a0.elts: - e0 = a0.elts[0] - if isinstance(e0, ast.Constant) and isinstance(e0.value, str): - host_lit = e0.value - elif isinstance(a0, ast.Constant) and isinstance(a0.value, str): - host_lit = a0.value + _, host_lit = self._resolve_url_arg(node.args[0]) if host_lit: if _is_metadata_host(host_lit): network_calls.append( @@ -1402,23 +1613,44 @@ def _check_signal_escape_patterns(code: str): } ) - # 2) Extract literal host (URL string or (host, port) tuple). + # 2) Extract host: literal URL, variable holding a literal, + # constant-folded f-string, or (host, port) tuple. host_arg = None - url_arg = None if node.args: - a0 = node.args[0] - if isinstance(a0, ast.Constant) and isinstance(a0.value, str): - url_arg = a0.value - elif isinstance(a0, ast.Tuple) and a0.elts: - e0 = a0.elts[0] - if isinstance(e0, ast.Constant) and isinstance(e0.value, str): - host_arg = e0.value - if url_arg and host_arg is None: - m = re.match(r"^\w+://([^/?#]+)", url_arg) - if m: - host_arg = m.group(1) - - if host_arg: + _, host_arg = self._resolve_url_arg(node.args[0]) + # Also accept keyword form: requests.get(url="...") / + # requests.request("GET", url="..."). + if host_arg is None: + for kw in node.keywords or []: + if kw.arg in ("url", "host", "uri") and kw.value is not None: + _, host_arg = self._resolve_url_arg(kw.value) + if host_arg: + break + # Special case: requests.request(method, url, ...) -- the + # URL is the second positional arg, not the first. + if ( + host_arg is None + and fq.endswith(".request") + and len(node.args) >= 2 + ): + _, host_arg = self._resolve_url_arg(node.args[1]) + if host_arg is None and node.args: + # The argument was opaque (variable, computed, etc.). + # Recording opaque calls keeps the static checker + # honest: a dynamic URL can hit the metadata endpoint + # at runtime, and we cannot prove otherwise. + network_calls.append( + { + "type": "opaque_url_blocked", + "line": getattr(node, "lineno", -1), + "description": ( + "Blocked: network call target is computed at runtime; " + "static analysis cannot verify the host. Pass a " + "literal URL pointing at an allowlisted host." + ), + } + ) + elif host_arg: if _is_metadata_host(host_arg): network_calls.append( { @@ -1534,23 +1766,58 @@ def _check_code_safety(code: str) -> str | None: def _kill_process_tree(proc) -> None: - """SIGKILL the setsid process group; fall back to single-pid kill.""" + """Terminate the subprocess and any children spawned via setsid. + + Linux / macOS: SIGKILL the process group so bash-backgrounded + grandchildren actually die (paired with ``os.setsid()`` in + ``_sandbox_preexec``). + + Windows: ``os.getpgid`` / ``os.killpg`` do not exist; calling them + would raise ``AttributeError`` and the sandbox supervisor would skip + the kill entirely, leaving runaway tool processes. We instead use + ``proc.kill()``, which Popen implements on Windows via + ``TerminateProcess(handle, 1)``. Children spawned with + ``CREATE_NEW_PROCESS_GROUP`` are reaped via ``taskkill /T`` as a + best-effort fallback. + """ if proc.poll() is not None: return - try: - pgid = os.getpgid(proc.pid) - except (ProcessLookupError, PermissionError): - pgid = None - if pgid is not None: + # Unix process-group kill -- only available on platforms that expose + # os.getpgid / os.killpg (Linux, macOS, *BSD). hasattr() is the + # canonical guard; checking sys.platform alone misses Cygwin / + # WSL-on-Windows which expose both APIs. + if hasattr(os, "getpgid") and hasattr(os, "killpg"): try: - os.killpg(pgid, signal.SIGKILL) - return - except (ProcessLookupError, PermissionError): - pass + pgid = os.getpgid(proc.pid) + except (ProcessLookupError, PermissionError, OSError): + pgid = None + if pgid is not None: + try: + os.killpg(pgid, signal.SIGKILL) + return + except (ProcessLookupError, PermissionError, OSError): + pass + # Windows / fallback: kill the direct child, then taskkill its tree. try: proc.kill() - except (ProcessLookupError, PermissionError): + except (ProcessLookupError, PermissionError, OSError): pass + if sys.platform == "win32": + # Best-effort tree kill on Windows. We don't await taskkill -- if + # the binary is missing or fails we still already killed the + # immediate child above. + try: + import subprocess as _subprocess + + _subprocess.run( + ["taskkill", "/F", "/T", "/PID", str(proc.pid)], + stdout = _subprocess.DEVNULL, + stderr = _subprocess.DEVNULL, + timeout = 5, + check = False, + ) + except Exception: + pass def _cancel_watcher(proc, cancel_event, poll_interval = 0.2): diff --git a/studio/backend/core/training/training.py b/studio/backend/core/training/training.py index 549d733252..43bf0b69a7 100644 --- a/studio/backend/core/training/training.py +++ b/studio/backend/core/training/training.py @@ -147,6 +147,11 @@ class TrainingBackend: # Job metadata self.current_job_id: Optional[str] = None self._output_dir: Optional[str] = None + # Resolved run dir (set on the worker's first "run_started" event). + # Used by force_terminate / _cleanup_cancelled_checkpoints to + # delete intermediate checkpoint-* dirs when a run is cancelled + # before the worker emits "complete" (or with output_dir=None). + self._active_run_dir: Optional[str] = None # DB persistence self._metric_buffer: list[dict] = [] @@ -312,6 +317,7 @@ class TrainingBackend: self.eval_step_history.clear() self.eval_enabled = False self._output_dir = None + self._active_run_dir = None self._metric_buffer.clear() self._run_finalized = False self._db_run_created = False @@ -364,7 +370,11 @@ class TrainingBackend: self._proc.terminate() proc = self._proc cancelled = self._cancel_requested - output_dir = self._output_dir + # Prefer the active run dir set by "run_started" (always + # populated for any run that reached worker startup), and + # fall back to _output_dir if a "complete" event resolved + # a different artifact dir later. + cleanup_dir = self._active_run_dir or self._output_dir if proc is not None: proc.join(timeout = 5.0) @@ -377,15 +387,21 @@ class TrainingBackend: if self._pump_thread is not None and self._pump_thread.is_alive(): self._pump_thread.join(timeout = 8.0) + # Re-snapshot in case "run_started" landed AFTER our first lock + # window (worker started in parallel with the cancel). Use the + # newer value if available; never downgrade a populated path. + with self._lock: + cleanup_dir = self._active_run_dir or self._output_dir or cleanup_dir + # Drop checkpoint-* dirs on explicit cancel only; stop-and-save # keeps its artifacts. - if cancelled and output_dir: + if cancelled and cleanup_dir: try: - _cleanup_cancelled_checkpoints(output_dir) + _cleanup_cancelled_checkpoints(cleanup_dir) except Exception: logger.exception( "Failed to clean up cancelled-run checkpoints under %s", - output_dir, + cleanup_dir, ) def is_training_active(self) -> bool: @@ -679,6 +695,18 @@ class TrainingBackend: self._progress.status_message = event.get("message", "") self._progress.is_training = True + elif etype == "run_started": + # Captured as soon as the worker resolves the run dir, + # well before any "complete" event. Cancel-and-reset + # paths that force-kill the worker before completion + # still see _active_run_dir set here, so + # _cleanup_cancelled_checkpoints can drop checkpoint-* + # under outputs_root. Distinct from _output_dir, which + # only tracks "saved-artifact" dirs (None on cancel-no-save). + run_dir = event.get("output_dir") + if run_dir: + self._active_run_dir = run_dir + elif etype == "complete": self._progress.is_training = False self._progress.is_completed = True diff --git a/studio/backend/core/training/worker.py b/studio/backend/core/training/worker.py index ef5cafb175..4443b733f7 100644 --- a/studio/backend/core/training/worker.py +++ b/studio/backend/core/training/worker.py @@ -709,6 +709,10 @@ def _run_mlx_training(event_queue, stop_queue, config): output_dir = str(resolve_output_dir(output_dir)) ensure_dir(Path(output_dir)) + # Publish the resolved output dir immediately so the parent process + # can clean up checkpoint-* dirs if the run is cancelled before a + # "complete" event ever fires (force-kill / cancel-no-save). + _send("run_started", output_dir = output_dir) # ── 6. Create trainer ── eval_steps_val = config.get("eval_steps", 0) or 0 @@ -1460,6 +1464,16 @@ def run_training_process( output_dir = f"{model_name.replace('/', '_')}_{int(time.time())}" output_dir = str(resolve_output_dir(output_dir)) ensure_dir(Path(output_dir)) + # Publish the resolved output dir immediately. The parent's pump + # loop reads "run_started" and stores _output_dir so cancel paths + # that force-kill before a "complete" event can still clean up + # checkpoint-* directories under outputs_root. + try: + event_queue.put( + {"type": "run_started", "output_dir": output_dir, "ts": time.time()} + ) + except Exception: + pass tensorboard_dir = config.get("tensorboard_dir") if config.get("enable_tensorboard", False): diff --git a/studio/backend/main.py b/studio/backend/main.py index 81964c98e0..dc93b62064 100644 --- a/studio/backend/main.py +++ b/studio/backend/main.py @@ -270,15 +270,30 @@ _CSP_SCRIPT_NONCE_HEADER = "x-internal-script-nonce" def _build_csp(script_nonce: "str | None" = None) -> str: + # script-src is 'self' plus an optional per-response nonce, never + # 'unsafe-inline'. The frontend bundle is self-hosted; the nonce is + # only used to whitelist the bootstrap-injection block when present + # (gated behind UNSLOTH_STUDIO_INJECT_BOOTSTRAP). Skipping + # 'unsafe-inline' for scripts forces any XSS payload to land in an + # external file under our origin, which is itself locked down. script_src = "script-src 'self'" if script_nonce: script_src += f" 'nonce-{script_nonce}'" + # connect-src allows same-origin plus the Hugging Face Hub endpoints + # the frontend hits directly (model + dataset pickers in + # use-hf-model-search / use-hf-dataset-search) and HF's CDN + # subdomains used for blob fetches / file metadata. Without these + # origins, browser-served Studio loses the pickers and the user + # cannot search/select any model or dataset. return ( "default-src 'self'; " "img-src 'self' data: blob: https://t0.gstatic.com " "https://t1.gstatic.com https://t2.gstatic.com " - "https://t3.gstatic.com; " - "connect-src 'self'; " + "https://t3.gstatic.com https://huggingface.co " + "https://cdn-avatars.huggingface.co; " + "connect-src 'self' https://huggingface.co " + "https://*.huggingface.co https://cdn-lfs.huggingface.co " + "https://cdn-lfs.hf.co https://hf.co https://*.hf.co; " "style-src 'self' 'unsafe-inline'; " f"{script_src}; " "font-src 'self' data:; " @@ -487,10 +502,39 @@ app.include_router( @app.get("/api/health") async def health_check(request: Request): - """Liveness only; full diagnostic dict gated on a valid bearer.""" + """Health probe; sensitive diagnostic dict gated on a valid bearer. + + Three audiences read this endpoint: + + 1. Generic liveness probes -- want only ``status``/``timestamp``. + 2. Launcher/preflight code that cannot present a bearer token + (``install.sh::_check_health``, ``studio/src-tauri/src/preflight``, + ``run_studio_browser_test`` orchestrator). They match on + ``service``, ``studio_root_id`` and the desktop capability flags + so they can confirm "this is the Studio I just installed". These + fields are non-sensitive identity / capability advertisements -- + the install path itself is hex-hashed into ``studio_root_id`` so + the path is never leaked. + 3. Authenticated admins / Tauri command surfaces -- want the full + diagnostic dict including ``version``, ``studio_version``, + ``device_type``, ``chat_only``, ``desktop_owner`` and so on. + + Returning the legacy identity fields unauthenticated keeps the + launcher contract working without exposing version strings or + device-shape introspection to drive-by callers. + """ minimal = { "status": "healthy", "timestamp": datetime.now().isoformat(), + # Launcher / preflight contract: stable identity + capability bits. + # Safe to expose unauthenticated -- studio_root_id is a hex digest + # of the install path, the desktop flags are non-sensitive booleans. + "service": "Unsloth UI Backend", + "studio_root_id": _studio_root_id(), + "desktop_protocol_version": 1, + "desktop_manageability_version": 1, + "supports_desktop_auth": True, + "supports_desktop_backend_ownership": True, } auth = request.headers.get("authorization", "") if not auth.lower().startswith("bearer "): @@ -515,17 +559,15 @@ async def health_check(request: Request): device_type = platform_map.get(sys.platform, sys.platform) return { **minimal, - "service": "Unsloth UI Backend", + # Sensitive diagnostic fields. Gated on a valid bearer because: + # - version / studio_version reveal patch-level CVE exposure; + # - device_type / chat_only reveal training-vs-inference shape; + # - desktop_owner reveals which UID/process owns the desktop lease; + # - native_path_leases_supported reveals filesystem capability. "version": UNSLOTH_VERSION, "studio_version": STUDIO_VERSION, "device_type": device_type, "chat_only": _hw_module.CHAT_ONLY, - "desktop_protocol_version": 1, - "desktop_manageability_version": 1, - "supports_desktop_auth": True, - "supports_desktop_backend_ownership": True, - # Hex digest of the install path; launchers reject sibling Studios on the same port. - "studio_root_id": _studio_root_id(), "native_path_leases_supported": native_path_leases_supported(), **({"desktop_owner": owner} if (owner := _desktop_owner()) else {}), } diff --git a/studio/backend/models/inference.py b/studio/backend/models/inference.py index 746ac8bbc2..f5e4164ac0 100644 --- a/studio/backend/models/inference.py +++ b/studio/backend/models/inference.py @@ -392,6 +392,14 @@ ContentPart = Annotated[ # ── Messages ───────────────────────────────────────────────────── +# Prefix used by ChatMessage._validate_role_shape when synthesising a +# placeholder tool_call_id for the frontend's second-round POST (which +# drops the streamed id). Route handlers detect this prefix and rewrite +# the id to the matching preceding assistant tool_call id before +# passthrough, preserving correlation. +TOOL_CALL_ID_SYNTH_PREFIX = "call_studio_synth_" + + class ChatMessage(BaseModel): """ A single message in the conversation. @@ -434,11 +442,16 @@ class ChatMessage(BaseModel): if self.role == "tool": if not self.tool_call_id: - # Frontend's second-round POST drops the streamed id; - # synthesise one so the request round-trips. + # Frontend's second-round POST drops the streamed id. Mark + # the synthetic id with a recognisable prefix so the route + # handler can rewrite it to the matching preceding + # assistant tool_call id before passthrough -- a random + # id breaks correlation and OpenAI-compatible backends + # reject "tool result not referenced by any tool_call". + # See ``_pair_orphan_tool_ids`` in routes/inference.py. import secrets as _secrets - self.tool_call_id = f"call_{_secrets.token_hex(8)}" + self.tool_call_id = f"{TOOL_CALL_ID_SYNTH_PREFIX}{_secrets.token_hex(8)}" if not self.content: raise ValueError('role="tool" messages require non-empty "content".') elif self.role == "assistant": diff --git a/studio/backend/routes/auth.py b/studio/backend/routes/auth.py index 30221c2c93..dbcd471855 100644 --- a/studio/backend/routes/auth.py +++ b/studio/backend/routes/auth.py @@ -7,6 +7,7 @@ Authentication API routes from fastapi import APIRouter, Depends, HTTPException, Request, Response, status +import logging import threading import time from collections import deque @@ -140,11 +141,31 @@ async def logout( request: Request, current_subject: str = Depends(get_current_subject_allow_password_change), ) -> Response: - """Revoke refresh tokens for the subject; the access token is stateless and expires on its own.""" + """Revoke refresh tokens for the subject; the access token is stateless and expires on its own. + + The revoke must succeed for the response to be 204 -- otherwise a caller + would be told "logged out" while a stolen refresh token stayed live in + the database. We surface a 500 with a generic message when the storage + layer fails so callers (and operators) notice and retry / investigate. + """ try: storage.revoke_user_refresh_tokens(current_subject) - except Exception: - pass + except Exception as exc: + # Log structured detail but keep the response opaque -- we don't want + # to leak DB internals (file path, lock contention, etc.) to clients. + try: + logger = logging.getLogger(__name__) + logger.error( + "logout: refresh-token revocation failed for subject=%s: %r", + current_subject, + exc, + ) + except Exception: + pass + raise HTTPException( + status_code = status.HTTP_500_INTERNAL_SERVER_ERROR, + detail = "Failed to revoke refresh tokens", + ) try: request.app.state.bootstrap_password = None except AttributeError: diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index 7102e12bf8..4cf32b3228 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -3472,6 +3472,7 @@ async def anthropic_messages( payload.system, ) openai_messages = _drop_empty_assistant_sentinels(openai_messages) + openai_messages = _pair_orphan_tool_ids(openai_messages) # Enforce vision guard + re-encode embedded images to PNG so the # Anthropic endpoint matches the behavior of /v1/chat/completions. @@ -4210,6 +4211,78 @@ def _drop_empty_assistant_sentinels(messages: list[dict]) -> list[dict]: return out +def _pair_orphan_tool_ids(messages: list[dict]) -> list[dict]: + """Rewrite synthesised tool_call_ids to match a preceding assistant tool_call. + + The frontend's second-round POST sometimes drops the streamed + ``tool_call_id`` on a ``role="tool"`` message. ChatMessage's + validator synthesises a placeholder with the ``call_studio_synth_`` + prefix so the request shape passes, but a random id breaks + OpenAI-compatible backends which require the tool result to + reference an announced assistant tool_call. Here we backfill the + real id by: + + 1. Walking the message list forward and queueing + unmatched assistant ``tool_calls`` ids per FIFO. + 2. When we hit a synthesised ``role="tool"`` message, pop the + oldest unmatched id and rewrite ``tool_call_id`` in place. + 3. Falling back to leaving the synthetic id intact when no + preceding assistant tool_call is available -- the upstream + backend will then reject the request explicitly instead of + silently mismatching. + + Idempotent: messages without synth ids and messages whose ids + already match the announced tool_calls are passed through unchanged. + """ + # Local import keeps this module's import graph stable; the constant + # lives in models.inference next to the validator that emits it. + from models.inference import TOOL_CALL_ID_SYNTH_PREFIX + + # Queue of unmatched assistant tool_call ids, in announce order. + pending_ids: list[str] = [] + # Map of synth_id -> real_id so a single message list can be + # re-applied (or applied to a copy without mutating the original). + rewrites: dict[str, str] = {} + # Track real ids that already had a matching role="tool" result so + # they are not handed to a later synth message. + consumed: set[str] = set() + + for m in messages: + role = m.get("role") + if role == "assistant": + for tc in m.get("tool_calls") or []: + tcid = tc.get("id") if isinstance(tc, dict) else None + if tcid: + pending_ids.append(tcid) + elif role == "tool": + tcid = m.get("tool_call_id") + if isinstance(tcid, str) and tcid.startswith(TOOL_CALL_ID_SYNTH_PREFIX): + # Pop the oldest unconsumed announced id, if any. + while pending_ids: + candidate = pending_ids.pop(0) + if candidate not in consumed: + rewrites[tcid] = candidate + consumed.add(candidate) + break + elif isinstance(tcid, str) and tcid: + consumed.add(tcid) + + if not rewrites: + return messages + + out: list[dict] = [] + for m in messages: + if m.get("role") == "tool": + tcid = m.get("tool_call_id") + if isinstance(tcid, str) and tcid in rewrites: + new = dict(m) + new["tool_call_id"] = rewrites[tcid] + out.append(new) + continue + out.append(m) + return out + + def _openai_messages_for_passthrough(payload) -> list[dict]: """Build OpenAI-format message dicts for the /v1/chat/completions passthrough path. @@ -4229,6 +4302,7 @@ def _openai_messages_for_passthrough(payload) -> list[dict]: messages = _drop_empty_assistant_sentinels( [m.model_dump(exclude_none = True) for m in payload.messages] ) + messages = _pair_orphan_tool_ids(messages) if not payload.image_base64: return messages diff --git a/studio/backend/tests/test_health_unauth_contract.py b/studio/backend/tests/test_health_unauth_contract.py new file mode 100644 index 0000000000..c4033eff81 --- /dev/null +++ b/studio/backend/tests/test_health_unauth_contract.py @@ -0,0 +1,177 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. + +"""Tests for the /api/health unauthenticated contract. + +PR 5375 stripped the legacy identity / capability fields from unauthenticated +``/api/health`` responses. That broke: + +* ``install.sh::_check_health`` which matches on ``service`` + ``studio_root_id``. +* ``studio/src-tauri/src/preflight/backend.rs`` which reads + ``service`` / ``desktop_protocol_version`` / ``studio_root_id``. +* ``run_studio_browser_test.preflight`` which mirrors the install.sh + matcher. + +The follow-up fix re-publishes the launcher contract (status / timestamp +/ service / studio_root_id / desktop protocol bits / supports_desktop_*) +unauthenticated, and keeps the sensitive diagnostic fields (version / +device_type / chat_only / desktop_owner / native_path_leases_supported) +gated on a valid bearer. + +This module pins the contract in both directions so regressions show up +before they ship. +""" + +from __future__ import annotations + +import importlib +import os +import sys +from pathlib import Path + +import pytest + +_BACKEND_ROOT = Path(__file__).resolve().parents[1] +if str(_BACKEND_ROOT) not in sys.path: + sys.path.insert(0, str(_BACKEND_ROOT)) + + +LAUNCHER_KEYS = { + "status", + "timestamp", + "service", + "studio_root_id", + "desktop_protocol_version", + "desktop_manageability_version", + "supports_desktop_auth", + "supports_desktop_backend_ownership", +} + +GATED_KEYS = { + "version", + "studio_version", + "device_type", + "chat_only", + "native_path_leases_supported", +} + + +@pytest.fixture() +def fastapi_client(tmp_path, monkeypatch): + """Boot the FastAPI app against a tmp Studio home and return a TestClient.""" + monkeypatch.setenv("UNSLOTH_STUDIO_HOME", str(tmp_path)) + monkeypatch.setenv("UNSLOTH_API_ONLY", "1") + # Reset any cached module state so DB_PATH / install root resolve + # to the tmp directory. + for name in list(sys.modules): + if name.startswith(("auth.", "main", "models.", "routes.", "loggers.")): + del sys.modules[name] + main = importlib.import_module("main") + from fastapi.testclient import TestClient + + with TestClient(main.app) as client: + yield client, main + + +class TestUnauthHealth: + def test_status_healthy(self, fastapi_client): + client, _ = fastapi_client + r = client.get("/api/health") + assert r.status_code == 200 + body = r.json() + assert body["status"] == "healthy" + + def test_includes_launcher_contract(self, fastapi_client): + client, _ = fastapi_client + body = client.get("/api/health").json() + missing = LAUNCHER_KEYS - set(body) + assert not missing, f"unauth /api/health missing {sorted(missing)}" + + def test_does_not_leak_gated_fields(self, fastapi_client): + client, _ = fastapi_client + body = client.get("/api/health").json() + leaked = GATED_KEYS & set(body) + assert not leaked, f"unauth /api/health leaked {sorted(leaked)}" + + def test_invalid_bearer_drops_back_to_unauth(self, fastapi_client): + client, _ = fastapi_client + body = client.get( + "/api/health", headers={"Authorization": "Bearer not-real"} + ).json() + leaked = GATED_KEYS & set(body) + assert not leaked, f"invalid-bearer health leaked {sorted(leaked)}" + missing = LAUNCHER_KEYS - set(body) + assert not missing, f"invalid-bearer health missing {sorted(missing)}" + + def test_coroutine_truthy_does_not_skip_auth(self, fastapi_client): + """Regression: a bare coroutine is truthy. + + Before the fix the handler called ``get_current_subject(creds)`` + without ``await``, so any header starting with ``Bearer `` would + be seen as a valid principal and produce the full payload. + """ + client, _ = fastapi_client + body = client.get( + "/api/health", headers={"Authorization": "Bearer x.y.z"} + ).json() + # Verify gated fields are NOT leaked when the token cannot be + # decoded (which would have been the case for the original bug). + assert "version" not in body + assert "device_type" not in body + + +class TestAuthedHealth: + """Authenticated health should expose the diagnostic dict.""" + + def test_valid_bearer_exposes_diagnostic(self, fastapi_client): + client, main = fastapi_client + # Bypass HTTP login -- mint a token through auth.authentication + # so the test does not depend on the bootstrap password file + # existing in this tmp install. The default-admin row was + # created during lifespan startup so the subject is valid. + from auth.authentication import create_access_token + from auth import storage + + # Ensure the unsloth admin user exists so its jwt_secret is on + # disk and the subsequent get_current_subject(...) accepts the + # token. The fixture's lifespan call usually seeds it, but + # when this test runs after a prior fixture's module reload + # the auth.storage module may have been re-imported and lost + # its in-process DB connection. Idempotently re-seed here. + if storage.get_user_and_secret(storage.DEFAULT_ADMIN_USERNAME) is None: + # ensure_default_admin is idempotent (no-op when the row + # already exists). The fixture's lifespan usually seeds the + # default admin, but module reloads between tests can leave + # the in-process state inconsistent; re-seeding here keeps + # the test order-independent. + seed = getattr(storage, "ensure_default_admin", None) + if seed is not None: + seed() + else: + pytest.skip("storage has no ensure_default_admin entrypoint") + assert storage.get_user_and_secret(storage.DEFAULT_ADMIN_USERNAME) is not None, ( + "could not seed default admin" + ) + # Clear the must_change_password flag so /api/health's + # get_current_subject dependency accepts the token. Fresh installs + # block diagnostic access until the first-boot password change, + # which is the production contract but inconvenient for this test. + conn = storage.get_connection() + try: + conn.execute( + "UPDATE auth_user SET must_change_password = 0 WHERE username = ?", + (storage.DEFAULT_ADMIN_USERNAME,), + ) + conn.commit() + finally: + conn.close() + token = create_access_token(subject=storage.DEFAULT_ADMIN_USERNAME) + body = client.get( + "/api/health", headers={"Authorization": f"Bearer {token}"} + ).json() + # Diagnostic keys are present. + for k in ("version", "device_type"): + assert k in body, f"authed health missing {k!r}" + # Launcher contract still present (authed payload is a superset). + missing = LAUNCHER_KEYS - set(body) + assert not missing, f"authed health missing launcher keys {sorted(missing)}" diff --git a/studio/backend/tests/test_kill_process_tree_platform.py b/studio/backend/tests/test_kill_process_tree_platform.py new file mode 100644 index 0000000000..2708921099 --- /dev/null +++ b/studio/backend/tests/test_kill_process_tree_platform.py @@ -0,0 +1,160 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. + +"""Cross-platform contract tests for _kill_process_tree. + +The post-PR-5375 hardening pass added ``os.setsid`` to the sandbox +pre-exec; the cancel/timeout supervisor calls ``_kill_process_tree`` to +SIGKILL the resulting process group. ``os.getpgid``/``os.killpg`` are +Unix-only -- on Windows the helper must fall back to ``proc.kill()`` + +``taskkill /T``. We simulate Windows by stripping the platform-specific +attributes off ``os`` and verifying the helper still reaches the kill +path without raising. +""" + +from __future__ import annotations + +import os +import subprocess +import sys +import time +from pathlib import Path +from unittest import mock + +import pytest + +_BACKEND_ROOT = Path(__file__).resolve().parents[1] +if str(_BACKEND_ROOT) not in sys.path: + sys.path.insert(0, str(_BACKEND_ROOT)) + +from core.inference import tools as tools_mod + + +def _spawn_sleep(seconds: int = 60): + """Spawn a lightweight sleeper. Prefers /bin/sleep (no Python startup + cost) and falls back to ``python -c sleep`` on platforms that don't + have /bin/sleep on PATH (mostly Windows).""" + if sys.platform != "win32": + from shutil import which + + sleep_bin = which("sleep") or "/bin/sleep" + if Path(sleep_bin).exists(): + return subprocess.Popen( + [sleep_bin, str(seconds)], + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + start_new_session=True, + ) + # Fallback: minimal Python sleeper. Adds ~25-40 MB per test which + # is fine for single-test runs but is the reason we prefer + # /bin/sleep when available (the suite spawns one per test). + return subprocess.Popen( + [sys.executable, "-c", "import time; time.sleep(%d)" % seconds], + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + ) + + +@pytest.fixture() +def short_proc(): + """A subprocess that sleeps long enough to be killable.""" + proc = _spawn_sleep() + try: + yield proc + finally: + try: + proc.kill() + except Exception: + pass + + +class TestUnixPath: + """On Linux/macOS the pgid path runs and reaps the child.""" + + @pytest.mark.skipif( + not (hasattr(os, "getpgid") and hasattr(os, "killpg")), + reason="No process-group APIs on this platform", + ) + def test_kill_terminates_subprocess(self, short_proc): + assert short_proc.poll() is None + tools_mod._kill_process_tree(short_proc) + # Give the OS a moment to reap. _kill_process_tree does not block. + deadline = time.time() + 5.0 + while short_proc.poll() is None and time.time() < deadline: + time.sleep(0.05) + assert short_proc.poll() is not None, "subprocess should have died" + + @pytest.mark.skipif( + not (hasattr(os, "getpgid") and hasattr(os, "killpg")), + reason="No process-group APIs on this platform", + ) + def test_no_raise_on_already_exited(self): + # poll() already returned non-None: helper must early-return. + class Dead: + pid = 0 + + def poll(self): + return 0 + + # Should not raise even though pid 0 has no pgid. + tools_mod._kill_process_tree(Dead()) + + +class TestWindowsFallback: + """Simulate a Windows runtime where os lacks getpgid/killpg.""" + + def test_kill_falls_back_to_proc_kill_when_pgid_missing( + self, short_proc, monkeypatch + ): + # Strip the Unix-only attributes so the helper takes the + # Windows branch. + if hasattr(os, "getpgid"): + monkeypatch.delattr(os, "getpgid", raising=False) + if hasattr(os, "killpg"): + monkeypatch.delattr(os, "killpg", raising=False) + # Also lie about sys.platform so the taskkill fallback runs. + monkeypatch.setattr(tools_mod.sys, "platform", "win32") + # taskkill won't exist on Linux; capture its absence as a no-op + # via subprocess.run rather than failing the test. + import subprocess as _sp + + with mock.patch.object(_sp, "run", return_value=None): + tools_mod._kill_process_tree(short_proc) + deadline = time.time() + 5.0 + while short_proc.poll() is None and time.time() < deadline: + time.sleep(0.05) + assert short_proc.poll() is not None, "subprocess should have died" + + def test_no_attribute_error_on_simulated_windows(self, monkeypatch): + """Regression: AttributeError used to skip the kill entirely. + + Before the fix, ``os.getpgid(...)`` raised ``AttributeError`` on + Windows; the helper's exception list only covered + ProcessLookupError + PermissionError, so ``AttributeError`` + bubbled and the supervisor skipped the kill. The fix gates on + ``hasattr(os, ...)`` first, so this test pins that contract. + """ + if hasattr(os, "getpgid"): + monkeypatch.delattr(os, "getpgid", raising=False) + if hasattr(os, "killpg"): + monkeypatch.delattr(os, "killpg", raising=False) + monkeypatch.setattr(tools_mod.sys, "platform", "win32") + + class FakeProc: + pid = 9999 + + def __init__(self): + self.killed = False + + def poll(self): + return None if not self.killed else 0 + + def kill(self): + self.killed = True + + fp = FakeProc() + import subprocess as _sp + + with mock.patch.object(_sp, "run", return_value=None): + tools_mod._kill_process_tree(fp) # must not raise + assert fp.killed diff --git a/studio/backend/tests/test_middleware.py b/studio/backend/tests/test_middleware.py index bdf8e6d5a5..0d9409fcb7 100644 --- a/studio/backend/tests/test_middleware.py +++ b/studio/backend/tests/test_middleware.py @@ -228,18 +228,50 @@ def health_app(tmp_path, monkeypatch): class TestHealthAuthGate: - def test_no_auth_returns_minimal_payload(self, health_app): + """Pin the launcher contract. + + The post-PR-5375 follow-up re-published the launcher-essential + identity fields (``service``, ``studio_root_id``, the desktop + protocol bits) in the unauthenticated payload so install.sh's + ``_check_health``, Tauri's ``preflight/backend.rs`` and the + browser-test orchestrator can confirm "this is my Studio" before + a bearer exists. The diagnostic fields (``version`` / + ``device_type`` / ``chat_only`` / ``desktop_owner`` / + ``native_path_leases_supported``) stay gated on a valid bearer. + """ + + UNAUTH_KEYS = { + "status", + "timestamp", + "service", + "studio_root_id", + "desktop_protocol_version", + "desktop_manageability_version", + "supports_desktop_auth", + "supports_desktop_backend_ownership", + } + GATED_KEYS = ("version", "device_type", "chat_only", "native_path_leases_supported") + + def test_no_auth_returns_launcher_payload(self, health_app): c = TestClient(health_app) r = c.get("/api/health") assert r.status_code == 200 body = r.json() assert body["status"] == "healthy" assert "timestamp" in body - for forbidden in ("version", "device_type", "studio_root_id"): - assert forbidden not in body + # Launcher contract must be present so install.sh / Tauri + # preflight can identify this Studio without authenticating. + missing = self.UNAUTH_KEYS - set(body) + assert not missing, f"unauth /api/health missing {sorted(missing)}" + for forbidden in self.GATED_KEYS: + assert forbidden not in body, ( + f"unauth /api/health leaked {forbidden!r}" + ) - def test_invalid_bearer_returns_minimal_payload(self, health_app): - # Regression: calling the async dep without await made any Bearer header pass. + def test_invalid_bearer_returns_launcher_payload(self, health_app): + # Regression: calling the async dep without await made any + # Bearer header pass. The fix awaits the dep and falls back to + # the unauthenticated payload on any decode failure. c = TestClient(health_app) r = c.get( "/api/health", @@ -248,8 +280,12 @@ class TestHealthAuthGate: assert r.status_code == 200 body = r.json() assert body["status"] == "healthy" - for forbidden in ("version", "device_type", "studio_root_id"): - assert forbidden not in body + missing = self.UNAUTH_KEYS - set(body) + assert not missing, f"invalid-bearer /api/health missing {sorted(missing)}" + for forbidden in self.GATED_KEYS: + assert forbidden not in body, ( + f"invalid-bearer /api/health leaked {forbidden!r}" + ) def test_valid_bearer_returns_full_payload(self, health_app): from auth import storage diff --git a/studio/backend/tests/test_refresh_token_consume.py b/studio/backend/tests/test_refresh_token_consume.py new file mode 100644 index 0000000000..9ddd1455f1 --- /dev/null +++ b/studio/backend/tests/test_refresh_token_consume.py @@ -0,0 +1,139 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. + +"""Tests for consume_refresh_token's atomic-rotation contract. + +PR 5375 introduced ``DELETE ... RETURNING`` for single-use refresh-token +rotation. ``RETURNING`` is SQLite 3.35+. This module exercises both the +modern path and the SELECT+DELETE fallback so older system SQLite +(e.g. Ubuntu 20.04, some Windows builds) keeps refresh working. +""" + +from __future__ import annotations + +import os +import sys +import threading +from pathlib import Path + +import pytest + +_BACKEND_ROOT = Path(__file__).resolve().parents[1] +if str(_BACKEND_ROOT) not in sys.path: + sys.path.insert(0, str(_BACKEND_ROOT)) + + +@pytest.fixture() +def isolated_storage(tmp_path, monkeypatch): + """Point auth.storage at a fresh DB under tmp_path for every test. + + auth.storage.DB_PATH is computed at module load from utils.paths. + Rather than re-importing the module, we point ``DB_PATH`` at a + tmp_path SQLite file and rely on get_connection's CREATE TABLE IF + NOT EXISTS to lazily build the schema on first use. + """ + monkeypatch.setenv("UNSLOTH_STUDIO_HOME", str(tmp_path)) + import importlib + + if "auth.storage" in sys.modules: + importlib.reload(sys.modules["auth.storage"]) + storage = importlib.import_module("auth.storage") + # Force DB_PATH under tmp_path so every test gets a clean DB. + monkeypatch.setattr(storage, "DB_PATH", Path(tmp_path) / "auth.db") + # Reset the RETURNING-feature cache so each test re-probes. + monkeypatch.setattr(storage, "_RETURNING_SUPPORTED", None) + yield storage + + +def _make_token(storage, *, is_desktop=False): + """Insert a fresh, far-future refresh token and return its raw form.""" + import secrets as _secrets + from datetime import datetime, timedelta, timezone + + raw = _secrets.token_urlsafe(32) + expires = (datetime.now(timezone.utc) + timedelta(days=14)).isoformat() + storage.save_refresh_token(raw, "unsloth", expires, is_desktop=is_desktop) + return raw + + +class TestSingleUseRotation: + def test_consume_returns_username_on_first_use(self, isolated_storage): + storage = isolated_storage + token = _make_token(storage) + result = storage.consume_refresh_token(token) + assert result == ("unsloth", False) + + def test_consume_returns_none_on_replay(self, isolated_storage): + storage = isolated_storage + token = _make_token(storage) + first = storage.consume_refresh_token(token) + second = storage.consume_refresh_token(token) + assert first == ("unsloth", False) + assert second is None + + def test_consume_returns_none_for_unknown_token(self, isolated_storage): + storage = isolated_storage + assert storage.consume_refresh_token("not-a-real-token") is None + + def test_desktop_flag_round_trips(self, isolated_storage): + storage = isolated_storage + token = _make_token(storage, is_desktop=True) + assert storage.consume_refresh_token(token) == ("unsloth", True) + + +class TestReturningFallback: + """Pin the SELECT+DELETE fallback so non-RETURNING SQLite still works.""" + + def test_fallback_path_consumes_atomically(self, isolated_storage, monkeypatch): + storage = isolated_storage + # Force the fallback branch regardless of the underlying sqlite + # version. ``_supports_returning`` caches the result on first + # probe so setting it directly is enough. + monkeypatch.setattr(storage, "_RETURNING_SUPPORTED", False) + token = _make_token(storage) + first = storage.consume_refresh_token(token) + second = storage.consume_refresh_token(token) + assert first == ("unsloth", False) + assert second is None + + def test_fallback_unknown_token_returns_none(self, isolated_storage, monkeypatch): + storage = isolated_storage + monkeypatch.setattr(storage, "_RETURNING_SUPPORTED", False) + assert storage.consume_refresh_token("unknown") is None + + def test_fallback_race_only_one_wins(self, isolated_storage, monkeypatch): + storage = isolated_storage + monkeypatch.setattr(storage, "_RETURNING_SUPPORTED", False) + token = _make_token(storage) + # Hammer with N threads; exactly one should observe the token. + # We are testing the rowcount-on-DELETE guarantee in the + # fallback path -- not raw throughput. + winners: list = [] + losers: list = [] + barrier = threading.Barrier(8) + + def attempt(): + barrier.wait() + r = storage.consume_refresh_token(token) + (winners if r else losers).append(r) + + threads = [threading.Thread(target=attempt) for _ in range(8)] + for t in threads: + t.start() + for t in threads: + t.join() + assert len(winners) == 1, ( + f"expected exactly 1 winner; winners={winners} losers={losers}" + ) + assert winners[0] == ("unsloth", False) + assert all(r is None for r in losers) + + +class TestReturningSupportedProbe: + def test_returns_bool_on_any_sqlite(self, isolated_storage): + storage = isolated_storage + # Force re-probe. + storage._RETURNING_SUPPORTED = None + result = storage._supports_returning() + assert isinstance(result, bool) + diff --git a/studio/backend/tests/test_sandbox_tools.py b/studio/backend/tests/test_sandbox_tools.py index fcc531c212..6c27743802 100644 --- a/studio/backend/tests/test_sandbox_tools.py +++ b/studio/backend/tests/test_sandbox_tools.py @@ -118,9 +118,43 @@ class TestUntrustedHostBlock: expect_phrase = "Blocked: host not in sandbox allowlist", ) - def test_dynamic_url_not_statically_blocked(self): - # Static AST cannot resolve runtime URLs; bash blocklist is the fallback. - _ok('import requests; url = "https://example.com/"; requests.get(url)') + def test_simple_variable_url_resolved_and_blocked(self): + # Static AST resolves ``u = "..."; requests.get(u)`` by following the + # assignment, so metadata / untrusted hosts are caught even when the + # URL is staged into a local. Before this hardening pass the + # variable-URL path was treated as opaque and slipped through. + _blocked( + 'import requests\n' + 'url = "https://example.com/"\n' + 'requests.get(url)', + expect_phrase = "Blocked: host not in sandbox allowlist", + ) + _blocked( + 'import requests\n' + 'url = "http://169.254.169.254/latest/meta-data/"\n' + 'requests.get(url)', + expect_phrase = "Blocked: cloud-metadata host", + ) + + def test_constant_fstring_url_resolved(self): + # f-string URLs that fold to a constant should still be checked. + _blocked( + 'import requests\n' + 'host = "169.254.169.254"\n' + 'requests.get(f"http://{host}/latest/")', + expect_phrase = "Blocked: cloud-metadata host", + ) + + def test_truly_dynamic_url_marked_opaque(self): + # Genuinely runtime-computed URLs (input, env var, network) are + # reported as opaque so the static checker stays honest -- the + # bash blocklist + cloud-metadata IP block at OS layer cover the + # rest, but the AST can no longer say "looks fine to me". + _blocked( + 'import os, requests\n' + 'requests.get(os.environ["WEBHOOK"])', + expect_phrase = "network call target is computed at runtime", + ) class TestHostNormalization: @@ -221,6 +255,105 @@ class TestUploadDenylist: ) +class TestImportAliasResolution: + """Aliased / from-imported network APIs must obey the same policy. + + Pre-hardening, ``import requests as r`` and ``from requests import get`` + bypassed the prefix check because the visitor matched on the literal + "requests." FQ at the call site. The new visitor tracks + aliases at import time and rewrites the call's FQ before policy eval. + """ + + def test_module_alias_metadata_blocked(self): + _blocked( + 'import requests as r; r.get("http://169.254.169.254/latest/")', + expect_phrase = "Blocked: cloud-metadata host", + ) + + def test_module_alias_untrusted_blocked(self): + _blocked( + 'import requests as r; r.get("https://example.com/")', + expect_phrase = "Blocked: host not in sandbox allowlist", + ) + + def test_module_alias_trusted_passes(self): + _ok('import requests as r; r.get("https://en.wikipedia.org/wiki/Foo")') + + def test_from_import_metadata_blocked(self): + _blocked( + 'from requests import get\nget("http://metadata.google.internal/")', + expect_phrase = "Blocked: cloud-metadata host", + ) + + def test_from_import_aliased_blocked(self): + _blocked( + "from requests import get as fetch\n" + 'fetch("http://169.254.169.254/latest/")', + expect_phrase = "Blocked: cloud-metadata host", + ) + + def test_from_import_trusted_passes(self): + _ok( + 'from urllib.request import urlopen\n' + 'urlopen("https://en.wikipedia.org/wiki/Foo")' + ) + + def test_nested_module_alias_blocked(self): + _blocked( + 'import urllib.request as ur\n' + 'ur.urlopen("http://169.254.169.254/latest/")', + expect_phrase = "Blocked: cloud-metadata host", + ) + + +class TestSessionObjectMethods: + """``s = requests.Session(); s.get(url)`` must obey the host policy. + + The visitor tracks session-shaped constructor assignments so method + calls on the bound variable become egress-equivalent. + """ + + def test_requests_session_get_metadata_blocked(self): + _blocked( + "import requests\n" + "s = requests.Session()\n" + 's.get("http://169.254.169.254/latest/")', + expect_phrase = "Blocked: cloud-metadata host", + ) + + def test_requests_session_get_untrusted_blocked(self): + _blocked( + "import requests\n" + "s = requests.Session()\n" + 's.get("https://example.com/")', + expect_phrase = "Blocked: host not in sandbox allowlist", + ) + + def test_requests_session_post_upload_blocked(self): + _blocked( + "import requests\n" + "s = requests.Session()\n" + 's.post("https://huggingface.co/api/repos/upload", ' + 'files={"f": open("x.bin", "rb")})', + expect_phrase = "Blocked: file upload disallowed in sandbox", + ) + + def test_requests_session_trusted_passes(self): + _ok( + "import requests\n" + "s = requests.Session()\n" + 's.get("https://en.wikipedia.org/wiki/Foo")' + ) + + def test_httpx_client_metadata_blocked(self): + _blocked( + "import httpx\n" + "c = httpx.Client()\n" + 'c.get("http://169.254.169.254/latest/")', + expect_phrase = "Blocked: cloud-metadata host", + ) + + class TestSandboxCpuRlimitDefault: """Pin the default so a regression below 600s without opt-in is caught.""" diff --git a/studio/backend/tests/test_tool_id_pairing.py b/studio/backend/tests/test_tool_id_pairing.py new file mode 100644 index 0000000000..fd33092589 --- /dev/null +++ b/studio/backend/tests/test_tool_id_pairing.py @@ -0,0 +1,139 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. + +"""Tests for the orphan tool_call_id pairing helper. + +PR 5375 had ChatMessage._validate_role_shape synthesise a random +``tool_call_id`` when ``role="tool"`` arrived without one (frontend's +second-round POST drops the streamed id). The random id broke +correlation with the preceding assistant ``tool_calls`` ids -- upstream +OpenAI-compatible backends reject "tool result not referenced by any +tool_call". The follow-up fix: + +* The validator tags the synthesised id with a recognisable prefix + (``TOOL_CALL_ID_SYNTH_PREFIX``). +* The route handler runs ``_pair_orphan_tool_ids`` before passthrough + to rewrite synth ids to the matching announced assistant tool_call. + +This module pins the rewrite contract. +""" + +from __future__ import annotations + +import sys +from pathlib import Path + +import pytest + +_BACKEND_ROOT = Path(__file__).resolve().parents[1] +if str(_BACKEND_ROOT) not in sys.path: + sys.path.insert(0, str(_BACKEND_ROOT)) + +from models.inference import TOOL_CALL_ID_SYNTH_PREFIX +from routes.inference import _pair_orphan_tool_ids + + +def _synth(idx: int = 0) -> str: + return f"{TOOL_CALL_ID_SYNTH_PREFIX}aa{idx:02d}" + + +class TestPairing: + def test_synth_id_rewritten_to_preceding_assistant(self): + msgs = [ + {"role": "user", "content": "hi"}, + { + "role": "assistant", + "content": "", + "tool_calls": [{"id": "call_abc123", "type": "function"}], + }, + {"role": "tool", "tool_call_id": _synth(0), "content": "result"}, + ] + out = _pair_orphan_tool_ids(msgs) + assert out[-1]["tool_call_id"] == "call_abc123" + # Other messages untouched and not aliased. + assert out[0] is msgs[0] + assert out[1] is msgs[1] + assert out[2] is not msgs[2] + + def test_real_id_left_alone(self): + real = "call_real_001" + msgs = [ + { + "role": "assistant", + "tool_calls": [{"id": real, "type": "function"}], + }, + {"role": "tool", "tool_call_id": real, "content": "ok"}, + ] + out = _pair_orphan_tool_ids(msgs) + # Idempotent: nothing rewritten. + assert out == msgs + + def test_multiple_synths_pair_to_distinct_calls(self): + msgs = [ + { + "role": "assistant", + "tool_calls": [ + {"id": "call_a", "type": "function"}, + {"id": "call_b", "type": "function"}, + ], + }, + {"role": "tool", "tool_call_id": _synth(0), "content": "x"}, + {"role": "tool", "tool_call_id": _synth(1), "content": "y"}, + ] + out = _pair_orphan_tool_ids(msgs) + assert out[1]["tool_call_id"] == "call_a" + assert out[2]["tool_call_id"] == "call_b" + + def test_synth_left_alone_when_no_announced_call(self): + # No preceding assistant tool_calls; the synth id stays so the + # upstream backend can produce a clear error. + msgs = [ + {"role": "user", "content": "hello"}, + {"role": "tool", "tool_call_id": _synth(0), "content": "x"}, + ] + out = _pair_orphan_tool_ids(msgs) + assert out[-1]["tool_call_id"].startswith(TOOL_CALL_ID_SYNTH_PREFIX) + + def test_existing_real_call_not_double_consumed(self): + msgs = [ + { + "role": "assistant", + "tool_calls": [ + {"id": "call_a", "type": "function"}, + {"id": "call_b", "type": "function"}, + ], + }, + # First tool result already references call_a explicitly. + {"role": "tool", "tool_call_id": "call_a", "content": "x"}, + # Second is a synth; should map to the remaining call_b. + {"role": "tool", "tool_call_id": _synth(0), "content": "y"}, + ] + out = _pair_orphan_tool_ids(msgs) + assert out[1]["tool_call_id"] == "call_a" + assert out[2]["tool_call_id"] == "call_b" + + def test_no_synths_returns_original(self): + msgs = [ + {"role": "user", "content": "hi"}, + {"role": "assistant", "content": "hello"}, + ] + out = _pair_orphan_tool_ids(msgs) + # No rewrites needed -- return the same list reference. + assert out is msgs + + +class TestValidatorSynth: + """Mirror the validator's behaviour so the synth prefix stays stable.""" + + def test_validator_emits_synth_prefix(self): + from models.inference import ChatMessage + + m = ChatMessage(role="tool", content="result") + assert m.tool_call_id is not None + assert m.tool_call_id.startswith(TOOL_CALL_ID_SYNTH_PREFIX) + + def test_validator_keeps_explicit_id(self): + from models.inference import ChatMessage + + m = ChatMessage(role="tool", tool_call_id="call_real_xyz", content="ok") + assert m.tool_call_id == "call_real_xyz" diff --git a/tests/studio/playwright_extra_ui.py b/tests/studio/playwright_extra_ui.py index dccd2e423d..9ade1b2469 100644 --- a/tests/studio/playwright_extra_ui.py +++ b/tests/studio/playwright_extra_ui.py @@ -31,7 +31,7 @@ import time import urllib.error import urllib.request from pathlib import Path -from playwright.sync_api import sync_playwright +from playwright.sync_api import sync_playwright, TimeoutError as PlaywrightTimeoutError # Shared robustness helpers live next to this script. Tests run as # plain `python tests/studio/playwright_extra_ui.py` (not via pytest / @@ -298,19 +298,33 @@ with sync_playwright() as p: composer = page.locator('textarea[aria-label="Message input"]') composer.wait_for(state = "visible", timeout = 60_000) - # Detect chat-only mode: /api/health.chat_only is the source of truth. - # In chat-only mode, /studio + /export redirect to /chat. + # Detect chat-only mode: /api/health.chat_only is the source of truth, + # but the field is gated behind a valid bearer (along with version / + # device_type / desktop_owner) -- unauthenticated callers only see the + # launcher contract (status, service, studio_root_id, desktop_*). + # Pass the access token we already minted above so the test always + # observes chat_only directly rather than guessing from URL behavior. health_resp = evaluate_fetch( page, f"{BASE}/api/health", + headers = {"Authorization": f"Bearer {token}"}, timeout_ms = FETCH_TIMEOUT_MS, ) if health_resp.get("error"): fail(f"/api/health wedged: {health_resp['error']!r}") sys.exit(1) health = health_resp.get("body") or {} - chat_only = bool(health.get("chat_only")) - info(f"chat_only mode: {chat_only}") + if "chat_only" not in health: + # Older Studio builds (pre-auth gating) returned chat_only in the + # minimal payload. Tolerate that by probing /studio: a 3xx to /chat + # is the equivalent runtime signal. + page.goto(f"{BASE}/studio") + page.wait_for_timeout(500) + chat_only = "/studio" not in page.url + info(f"chat_only inferred from /studio URL: {chat_only}") + else: + chat_only = bool(health.get("chat_only")) + info(f"chat_only mode: {chat_only}") # ───────────────────────────────────────────────────── # 1. Compare tab. @@ -506,7 +520,27 @@ with sync_playwright() as p: # ───────────────────────────────────────────────────── step(f"Studio route ({'chat-only redirect' if chat_only else 'tabs + sections'})") page.goto(f"{BASE}/studio") - page.wait_for_timeout(1500) + # Don't rely on a fixed timeout for hydration -- the training runtime + # makes API calls on mount and the tabs are gated on hasHydratedRuntime. + # Wait for the loading placeholder to clear by polling for either the + # Configure tab or the redirected URL (chat_only). + try: + page.wait_for_function( + """ + () => { + if (!location.pathname.startsWith('/studio')) return true; + const loading = Array.from(document.querySelectorAll('div')) + .find(d => d.textContent && d.textContent.includes('Loading training runtime')); + if (loading) return false; + const tabs = document.querySelectorAll('[role="tab"]'); + return tabs.length >= 3; + } + """, + timeout = 30_000, + ) + except PlaywrightTimeoutError: + info("/studio hydration didn't complete in 30s; continuing with whatever rendered") + page.wait_for_timeout(500) shoot("08-studio") if chat_only: if "/studio" in page.url: @@ -516,16 +550,57 @@ with sync_playwright() as p: else: info(f"OK chat-only redirected /studio -> {page.url}") else: - for tab_name in ("Configure", "Current run", "History"): - tab = page.get_by_role( - "tab", name = re.compile(rf"^\s*{tab_name}\s*$", re.I) + # Tabs render with disabled={!showTrainingView} on "Current Run", + # so during CI smoke (no run hydrated) it has disabled-aria. We + # match on accessible name regardless of disabled state. The + # frontend's accessible label is "Current Run" (title case) -- + # the regex is already case-insensitive so either rendering is OK. + for tab_name in ("Configure", "Current Run", "History"): + try: + tab = page.get_by_role( + "tab", name = re.compile(rf"^\s*{tab_name}\s*$", re.I) + ).first + if tab.count() == 0: + # Fallback: bare button text (some Radix versions + # render the trigger without role="tab" until the + # tabs are activated). Look at any element whose + # accessible name matches. + tab = page.get_by_text( + re.compile(rf"^\s*{tab_name}\s*$", re.I), exact = False + ).first + if tab.count() == 0: + soft_fail(f"tab '{tab_name}' not found in /studio") + else: + info(f"OK tab '{tab_name}' visible") + except Exception as exc: + soft_fail(f"tab '{tab_name}' query failed: {exc!r}") + # data-tour anchors live inside the Configure TabsContent. + # Click Configure first so the sections are mounted before we look + # for the anchors. This matches what a user sees when they land + # on /studio for the first time. + try: + configure_tab = page.get_by_role( + "tab", name = re.compile(r"^\s*configure\s*$", re.I) ).first - if tab.count() == 0: - soft_fail(f"tab '{tab_name}' not found in /studio") - else: - info(f"OK tab '{tab_name}' visible") + if configure_tab.count() > 0: + configure_tab.click(timeout = 5_000) + page.wait_for_timeout(300) + except Exception as exc: + info(f"Configure-tab click warning (continuing): {exc!r}") for anchor in ("studio-model", "studio-dataset", "studio-params"): el = page.locator(f'[data-tour="{anchor}"]').first + if el.count() == 0: + # Give the lazy-mounted sections a final 3s grace window + # before reporting missing -- shadcn ParamsSection + # measures its container after first paint and may flip + # the data-tour attribute on a later render tick. + try: + page.wait_for_selector( + f'[data-tour="{anchor}"]', timeout = 3_000, state = "attached" + ) + el = page.locator(f'[data-tour="{anchor}"]').first + except PlaywrightTimeoutError: + pass if el.count() == 0: soft_fail(f"[data-tour='{anchor}'] not found") else: From a03a11946980b6fd4d3cfb02a2eb90cb39e669a7 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Wed, 13 May 2026 14:20:58 +0000 Subject: [PATCH 2/3] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- studio/backend/core/inference/tools.py | 37 +++++++++++++------ studio/backend/models/inference.py | 4 +- .../tests/test_health_unauth_contract.py | 14 +++---- .../tests/test_kill_process_tree_platform.py | 26 ++++++------- studio/backend/tests/test_middleware.py | 10 ++--- .../tests/test_refresh_token_consume.py | 17 ++++----- studio/backend/tests/test_sandbox_tools.py | 17 ++++----- studio/backend/tests/test_tool_id_pairing.py | 4 +- tests/studio/playwright_extra_ui.py | 4 +- 9 files changed, 73 insertions(+), 60 deletions(-) diff --git a/studio/backend/core/inference/tools.py b/studio/backend/core/inference/tools.py index 48bf0a5e26..d7aa051976 100644 --- a/studio/backend/core/inference/tools.py +++ b/studio/backend/core/inference/tools.py @@ -1386,10 +1386,23 @@ def _check_signal_escape_patterns(code: str): # egress (``from requests import get; get("http://...")``). _NETWORK_FROM_NAMES = frozenset( { - "get", "post", "put", "delete", "patch", "head", "request", - "Session", "urlopen", "urlretrieve", "create_connection", - "getaddrinfo", "HTTPConnection", "HTTPSConnection", - "Client", "AsyncClient", "ClientSession", + "get", + "post", + "put", + "delete", + "patch", + "head", + "request", + "Session", + "urlopen", + "urlretrieve", + "create_connection", + "getaddrinfo", + "HTTPConnection", + "HTTPSConnection", + "Client", + "AsyncClient", + "ClientSession", } ) # Session/Client method names that perform HTTP egress. A literal @@ -1463,7 +1476,9 @@ def _check_signal_escape_patterns(code: str): parts: list[str] = [] ok = True for piece in value.values: - if isinstance(piece, ast.Constant) and isinstance(piece.value, str): + if isinstance(piece, ast.Constant) and isinstance( + piece.value, str + ): parts.append(piece.value) elif ( isinstance(piece, ast.FormattedValue) @@ -1509,7 +1524,11 @@ def _check_signal_escape_patterns(code: str): head = cur.id # If head is a session-bound variable, synthesise a # prefixable FQ (e.g. ``"requests.Session.get"``). - if head in self._session_vars and parts and parts[-1] in _SESSION_METHOD_NAMES: + if ( + head in self._session_vars + and parts + and parts[-1] in _SESSION_METHOD_NAMES + ): return f"{self._session_vars[head]}.{parts[-1]}" # Map the head through module-alias table. resolved_head = self._module_aliases.get(head, head) @@ -1628,11 +1647,7 @@ def _check_signal_escape_patterns(code: str): break # Special case: requests.request(method, url, ...) -- the # URL is the second positional arg, not the first. - if ( - host_arg is None - and fq.endswith(".request") - and len(node.args) >= 2 - ): + if host_arg is None and fq.endswith(".request") and len(node.args) >= 2: _, host_arg = self._resolve_url_arg(node.args[1]) if host_arg is None and node.args: # The argument was opaque (variable, computed, etc.). diff --git a/studio/backend/models/inference.py b/studio/backend/models/inference.py index f5e4164ac0..09c13cb46a 100644 --- a/studio/backend/models/inference.py +++ b/studio/backend/models/inference.py @@ -451,7 +451,9 @@ class ChatMessage(BaseModel): # See ``_pair_orphan_tool_ids`` in routes/inference.py. import secrets as _secrets - self.tool_call_id = f"{TOOL_CALL_ID_SYNTH_PREFIX}{_secrets.token_hex(8)}" + self.tool_call_id = ( + f"{TOOL_CALL_ID_SYNTH_PREFIX}{_secrets.token_hex(8)}" + ) if not self.content: raise ValueError('role="tool" messages require non-empty "content".') elif self.role == "assistant": diff --git a/studio/backend/tests/test_health_unauth_contract.py b/studio/backend/tests/test_health_unauth_contract.py index c4033eff81..1aab32b81f 100644 --- a/studio/backend/tests/test_health_unauth_contract.py +++ b/studio/backend/tests/test_health_unauth_contract.py @@ -96,7 +96,7 @@ class TestUnauthHealth: def test_invalid_bearer_drops_back_to_unauth(self, fastapi_client): client, _ = fastapi_client body = client.get( - "/api/health", headers={"Authorization": "Bearer not-real"} + "/api/health", headers = {"Authorization": "Bearer not-real"} ).json() leaked = GATED_KEYS & set(body) assert not leaked, f"invalid-bearer health leaked {sorted(leaked)}" @@ -112,7 +112,7 @@ class TestUnauthHealth: """ client, _ = fastapi_client body = client.get( - "/api/health", headers={"Authorization": "Bearer x.y.z"} + "/api/health", headers = {"Authorization": "Bearer x.y.z"} ).json() # Verify gated fields are NOT leaked when the token cannot be # decoded (which would have been the case for the original bug). @@ -149,9 +149,9 @@ class TestAuthedHealth: seed() else: pytest.skip("storage has no ensure_default_admin entrypoint") - assert storage.get_user_and_secret(storage.DEFAULT_ADMIN_USERNAME) is not None, ( - "could not seed default admin" - ) + assert ( + storage.get_user_and_secret(storage.DEFAULT_ADMIN_USERNAME) is not None + ), "could not seed default admin" # Clear the must_change_password flag so /api/health's # get_current_subject dependency accepts the token. Fresh installs # block diagnostic access until the first-boot password change, @@ -165,9 +165,9 @@ class TestAuthedHealth: conn.commit() finally: conn.close() - token = create_access_token(subject=storage.DEFAULT_ADMIN_USERNAME) + token = create_access_token(subject = storage.DEFAULT_ADMIN_USERNAME) body = client.get( - "/api/health", headers={"Authorization": f"Bearer {token}"} + "/api/health", headers = {"Authorization": f"Bearer {token}"} ).json() # Diagnostic keys are present. for k in ("version", "device_type"): diff --git a/studio/backend/tests/test_kill_process_tree_platform.py b/studio/backend/tests/test_kill_process_tree_platform.py index 2708921099..f5addeeb81 100644 --- a/studio/backend/tests/test_kill_process_tree_platform.py +++ b/studio/backend/tests/test_kill_process_tree_platform.py @@ -41,17 +41,17 @@ def _spawn_sleep(seconds: int = 60): if Path(sleep_bin).exists(): return subprocess.Popen( [sleep_bin, str(seconds)], - stdout=subprocess.DEVNULL, - stderr=subprocess.DEVNULL, - start_new_session=True, + stdout = subprocess.DEVNULL, + stderr = subprocess.DEVNULL, + start_new_session = True, ) # Fallback: minimal Python sleeper. Adds ~25-40 MB per test which # is fine for single-test runs but is the reason we prefer # /bin/sleep when available (the suite spawns one per test). return subprocess.Popen( [sys.executable, "-c", "import time; time.sleep(%d)" % seconds], - stdout=subprocess.DEVNULL, - stderr=subprocess.DEVNULL, + stdout = subprocess.DEVNULL, + stderr = subprocess.DEVNULL, ) @@ -73,7 +73,7 @@ class TestUnixPath: @pytest.mark.skipif( not (hasattr(os, "getpgid") and hasattr(os, "killpg")), - reason="No process-group APIs on this platform", + reason = "No process-group APIs on this platform", ) def test_kill_terminates_subprocess(self, short_proc): assert short_proc.poll() is None @@ -86,7 +86,7 @@ class TestUnixPath: @pytest.mark.skipif( not (hasattr(os, "getpgid") and hasattr(os, "killpg")), - reason="No process-group APIs on this platform", + reason = "No process-group APIs on this platform", ) def test_no_raise_on_already_exited(self): # poll() already returned non-None: helper must early-return. @@ -109,16 +109,16 @@ class TestWindowsFallback: # Strip the Unix-only attributes so the helper takes the # Windows branch. if hasattr(os, "getpgid"): - monkeypatch.delattr(os, "getpgid", raising=False) + monkeypatch.delattr(os, "getpgid", raising = False) if hasattr(os, "killpg"): - monkeypatch.delattr(os, "killpg", raising=False) + monkeypatch.delattr(os, "killpg", raising = False) # Also lie about sys.platform so the taskkill fallback runs. monkeypatch.setattr(tools_mod.sys, "platform", "win32") # taskkill won't exist on Linux; capture its absence as a no-op # via subprocess.run rather than failing the test. import subprocess as _sp - with mock.patch.object(_sp, "run", return_value=None): + with mock.patch.object(_sp, "run", return_value = None): tools_mod._kill_process_tree(short_proc) deadline = time.time() + 5.0 while short_proc.poll() is None and time.time() < deadline: @@ -135,9 +135,9 @@ class TestWindowsFallback: ``hasattr(os, ...)`` first, so this test pins that contract. """ if hasattr(os, "getpgid"): - monkeypatch.delattr(os, "getpgid", raising=False) + monkeypatch.delattr(os, "getpgid", raising = False) if hasattr(os, "killpg"): - monkeypatch.delattr(os, "killpg", raising=False) + monkeypatch.delattr(os, "killpg", raising = False) monkeypatch.setattr(tools_mod.sys, "platform", "win32") class FakeProc: @@ -155,6 +155,6 @@ class TestWindowsFallback: fp = FakeProc() import subprocess as _sp - with mock.patch.object(_sp, "run", return_value=None): + with mock.patch.object(_sp, "run", return_value = None): tools_mod._kill_process_tree(fp) # must not raise assert fp.killed diff --git a/studio/backend/tests/test_middleware.py b/studio/backend/tests/test_middleware.py index 0d9409fcb7..37cf24ca8e 100644 --- a/studio/backend/tests/test_middleware.py +++ b/studio/backend/tests/test_middleware.py @@ -264,9 +264,7 @@ class TestHealthAuthGate: missing = self.UNAUTH_KEYS - set(body) assert not missing, f"unauth /api/health missing {sorted(missing)}" for forbidden in self.GATED_KEYS: - assert forbidden not in body, ( - f"unauth /api/health leaked {forbidden!r}" - ) + assert forbidden not in body, f"unauth /api/health leaked {forbidden!r}" def test_invalid_bearer_returns_launcher_payload(self, health_app): # Regression: calling the async dep without await made any @@ -283,9 +281,9 @@ class TestHealthAuthGate: missing = self.UNAUTH_KEYS - set(body) assert not missing, f"invalid-bearer /api/health missing {sorted(missing)}" for forbidden in self.GATED_KEYS: - assert forbidden not in body, ( - f"invalid-bearer /api/health leaked {forbidden!r}" - ) + assert ( + forbidden not in body + ), f"invalid-bearer /api/health leaked {forbidden!r}" def test_valid_bearer_returns_full_payload(self, health_app): from auth import storage diff --git a/studio/backend/tests/test_refresh_token_consume.py b/studio/backend/tests/test_refresh_token_consume.py index 9ddd1455f1..278720be9e 100644 --- a/studio/backend/tests/test_refresh_token_consume.py +++ b/studio/backend/tests/test_refresh_token_consume.py @@ -45,14 +45,14 @@ def isolated_storage(tmp_path, monkeypatch): yield storage -def _make_token(storage, *, is_desktop=False): +def _make_token(storage, *, is_desktop = False): """Insert a fresh, far-future refresh token and return its raw form.""" import secrets as _secrets from datetime import datetime, timedelta, timezone raw = _secrets.token_urlsafe(32) - expires = (datetime.now(timezone.utc) + timedelta(days=14)).isoformat() - storage.save_refresh_token(raw, "unsloth", expires, is_desktop=is_desktop) + expires = (datetime.now(timezone.utc) + timedelta(days = 14)).isoformat() + storage.save_refresh_token(raw, "unsloth", expires, is_desktop = is_desktop) return raw @@ -77,7 +77,7 @@ class TestSingleUseRotation: def test_desktop_flag_round_trips(self, isolated_storage): storage = isolated_storage - token = _make_token(storage, is_desktop=True) + token = _make_token(storage, is_desktop = True) assert storage.consume_refresh_token(token) == ("unsloth", True) @@ -117,14 +117,14 @@ class TestReturningFallback: r = storage.consume_refresh_token(token) (winners if r else losers).append(r) - threads = [threading.Thread(target=attempt) for _ in range(8)] + threads = [threading.Thread(target = attempt) for _ in range(8)] for t in threads: t.start() for t in threads: t.join() - assert len(winners) == 1, ( - f"expected exactly 1 winner; winners={winners} losers={losers}" - ) + assert ( + len(winners) == 1 + ), f"expected exactly 1 winner; winners={winners} losers={losers}" assert winners[0] == ("unsloth", False) assert all(r is None for r in losers) @@ -136,4 +136,3 @@ class TestReturningSupportedProbe: storage._RETURNING_SUPPORTED = None result = storage._supports_returning() assert isinstance(result, bool) - diff --git a/studio/backend/tests/test_sandbox_tools.py b/studio/backend/tests/test_sandbox_tools.py index 6c27743802..c7422a3df6 100644 --- a/studio/backend/tests/test_sandbox_tools.py +++ b/studio/backend/tests/test_sandbox_tools.py @@ -124,22 +124,20 @@ class TestUntrustedHostBlock: # URL is staged into a local. Before this hardening pass the # variable-URL path was treated as opaque and slipped through. _blocked( - 'import requests\n' - 'url = "https://example.com/"\n' - 'requests.get(url)', + "import requests\n" 'url = "https://example.com/"\n' "requests.get(url)", expect_phrase = "Blocked: host not in sandbox allowlist", ) _blocked( - 'import requests\n' + "import requests\n" 'url = "http://169.254.169.254/latest/meta-data/"\n' - 'requests.get(url)', + "requests.get(url)", expect_phrase = "Blocked: cloud-metadata host", ) def test_constant_fstring_url_resolved(self): # f-string URLs that fold to a constant should still be checked. _blocked( - 'import requests\n' + "import requests\n" 'host = "169.254.169.254"\n' 'requests.get(f"http://{host}/latest/")', expect_phrase = "Blocked: cloud-metadata host", @@ -151,8 +149,7 @@ class TestUntrustedHostBlock: # bash blocklist + cloud-metadata IP block at OS layer cover the # rest, but the AST can no longer say "looks fine to me". _blocked( - 'import os, requests\n' - 'requests.get(os.environ["WEBHOOK"])', + "import os, requests\n" 'requests.get(os.environ["WEBHOOK"])', expect_phrase = "network call target is computed at runtime", ) @@ -294,13 +291,13 @@ class TestImportAliasResolution: def test_from_import_trusted_passes(self): _ok( - 'from urllib.request import urlopen\n' + "from urllib.request import urlopen\n" 'urlopen("https://en.wikipedia.org/wiki/Foo")' ) def test_nested_module_alias_blocked(self): _blocked( - 'import urllib.request as ur\n' + "import urllib.request as ur\n" 'ur.urlopen("http://169.254.169.254/latest/")', expect_phrase = "Blocked: cloud-metadata host", ) diff --git a/studio/backend/tests/test_tool_id_pairing.py b/studio/backend/tests/test_tool_id_pairing.py index fd33092589..4170662b9b 100644 --- a/studio/backend/tests/test_tool_id_pairing.py +++ b/studio/backend/tests/test_tool_id_pairing.py @@ -128,12 +128,12 @@ class TestValidatorSynth: def test_validator_emits_synth_prefix(self): from models.inference import ChatMessage - m = ChatMessage(role="tool", content="result") + m = ChatMessage(role = "tool", content = "result") assert m.tool_call_id is not None assert m.tool_call_id.startswith(TOOL_CALL_ID_SYNTH_PREFIX) def test_validator_keeps_explicit_id(self): from models.inference import ChatMessage - m = ChatMessage(role="tool", tool_call_id="call_real_xyz", content="ok") + m = ChatMessage(role = "tool", tool_call_id = "call_real_xyz", content = "ok") assert m.tool_call_id == "call_real_xyz" diff --git a/tests/studio/playwright_extra_ui.py b/tests/studio/playwright_extra_ui.py index 9ade1b2469..17fb4c80d4 100644 --- a/tests/studio/playwright_extra_ui.py +++ b/tests/studio/playwright_extra_ui.py @@ -539,7 +539,9 @@ with sync_playwright() as p: timeout = 30_000, ) except PlaywrightTimeoutError: - info("/studio hydration didn't complete in 30s; continuing with whatever rendered") + info( + "/studio hydration didn't complete in 30s; continuing with whatever rendered" + ) page.wait_for_timeout(500) shoot("08-studio") if chat_only: From c7cbc39ff58faec443292dc530303dc27cf0b983 Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Wed, 13 May 2026 14:31:53 +0000 Subject: [PATCH 3/3] studio: expose chat_only in unauth /api/health for SPA first-load routing PR #5406's launcher contract withheld chat_only because the diagnostic dict gated it behind a valid bearer. The SPA's first-load router needs chat_only *before* any bearer is available (to decide whether /studio and /export redirect to /chat); the Windows + Linux UI smokes' chat_only inference fallback was incorrectly classifying non-chat-only Studios as chat-only when the bearered probe came back without the field (post-bootstrap tokens still carry must_change_password=True). chat_only is a hardware-shape capability bit (training vs inference-only) -- the same category as supports_desktop_auth and desktop_manageability_version which already ship unauth -- so promoting it is consistent with the existing contract. The truly sensitive diagnostic stays gated: version, studio_version, device_type, native_path_leases_supported, desktop_owner. Tests + the playwright_extra_ui chat_only probe updated to match the new contract. --- studio/backend/main.py | 18 ++++++++++--- .../tests/test_health_unauth_contract.py | 2 +- studio/backend/tests/test_middleware.py | 3 ++- tests/studio/playwright_extra_ui.py | 27 +++++++++---------- 4 files changed, 31 insertions(+), 19 deletions(-) diff --git a/studio/backend/main.py b/studio/backend/main.py index dc93b62064..86990c1b71 100644 --- a/studio/backend/main.py +++ b/studio/backend/main.py @@ -528,9 +528,20 @@ async def health_check(request: Request): "timestamp": datetime.now().isoformat(), # Launcher / preflight contract: stable identity + capability bits. # Safe to expose unauthenticated -- studio_root_id is a hex digest - # of the install path, the desktop flags are non-sensitive booleans. + # of the install path, the desktop flags and chat_only are + # non-sensitive feature-shape booleans (same category as + # ``supports_desktop_auth``). + # + # chat_only is part of the contract because the SPA's first-load + # router needs it to decide whether to redirect /studio + /export + # to /chat *before* any bearer is available; the Playwright UI + # tests rely on the same signal so they don't have to maintain a + # heuristic ("did the URL change after goto?"). Withholding it + # broke the Windows + Linux UI smokes and the change-password + # bootstrap flow. "service": "Unsloth UI Backend", "studio_root_id": _studio_root_id(), + "chat_only": _hw_module.CHAT_ONLY, "desktop_protocol_version": 1, "desktop_manageability_version": 1, "supports_desktop_auth": True, @@ -561,13 +572,14 @@ async def health_check(request: Request): **minimal, # Sensitive diagnostic fields. Gated on a valid bearer because: # - version / studio_version reveal patch-level CVE exposure; - # - device_type / chat_only reveal training-vs-inference shape; + # - device_type reveals the training-vs-inference shape; # - desktop_owner reveals which UID/process owns the desktop lease; # - native_path_leases_supported reveals filesystem capability. + # chat_only is intentionally NOT gated; see the comment on the + # ``minimal`` dict above. "version": UNSLOTH_VERSION, "studio_version": STUDIO_VERSION, "device_type": device_type, - "chat_only": _hw_module.CHAT_ONLY, "native_path_leases_supported": native_path_leases_supported(), **({"desktop_owner": owner} if (owner := _desktop_owner()) else {}), } diff --git a/studio/backend/tests/test_health_unauth_contract.py b/studio/backend/tests/test_health_unauth_contract.py index 1aab32b81f..ec7aed0f16 100644 --- a/studio/backend/tests/test_health_unauth_contract.py +++ b/studio/backend/tests/test_health_unauth_contract.py @@ -41,6 +41,7 @@ LAUNCHER_KEYS = { "timestamp", "service", "studio_root_id", + "chat_only", "desktop_protocol_version", "desktop_manageability_version", "supports_desktop_auth", @@ -51,7 +52,6 @@ GATED_KEYS = { "version", "studio_version", "device_type", - "chat_only", "native_path_leases_supported", } diff --git a/studio/backend/tests/test_middleware.py b/studio/backend/tests/test_middleware.py index 37cf24ca8e..93d5b6189d 100644 --- a/studio/backend/tests/test_middleware.py +++ b/studio/backend/tests/test_middleware.py @@ -245,12 +245,13 @@ class TestHealthAuthGate: "timestamp", "service", "studio_root_id", + "chat_only", "desktop_protocol_version", "desktop_manageability_version", "supports_desktop_auth", "supports_desktop_backend_ownership", } - GATED_KEYS = ("version", "device_type", "chat_only", "native_path_leases_supported") + GATED_KEYS = ("version", "device_type", "native_path_leases_supported") def test_no_auth_returns_launcher_payload(self, health_app): c = TestClient(health_app) diff --git a/tests/studio/playwright_extra_ui.py b/tests/studio/playwright_extra_ui.py index 17fb4c80d4..5ae2d2b72e 100644 --- a/tests/studio/playwright_extra_ui.py +++ b/tests/studio/playwright_extra_ui.py @@ -298,16 +298,14 @@ with sync_playwright() as p: composer = page.locator('textarea[aria-label="Message input"]') composer.wait_for(state = "visible", timeout = 60_000) - # Detect chat-only mode: /api/health.chat_only is the source of truth, - # but the field is gated behind a valid bearer (along with version / - # device_type / desktop_owner) -- unauthenticated callers only see the - # launcher contract (status, service, studio_root_id, desktop_*). - # Pass the access token we already minted above so the test always - # observes chat_only directly rather than guessing from URL behavior. + # Detect chat-only mode: /api/health.chat_only is the source of truth. + # The field is part of the unauthenticated launcher contract so the + # SPA's first-load router can decide whether to redirect /studio + + # /export to /chat *before* any bearer exists. Bearered or not, the + # field is always present on a healthy backend. health_resp = evaluate_fetch( page, f"{BASE}/api/health", - headers = {"Authorization": f"Bearer {token}"}, timeout_ms = FETCH_TIMEOUT_MS, ) if health_resp.get("error"): @@ -315,13 +313,14 @@ with sync_playwright() as p: sys.exit(1) health = health_resp.get("body") or {} if "chat_only" not in health: - # Older Studio builds (pre-auth gating) returned chat_only in the - # minimal payload. Tolerate that by probing /studio: a 3xx to /chat - # is the equivalent runtime signal. - page.goto(f"{BASE}/studio") - page.wait_for_timeout(500) - chat_only = "/studio" not in page.url - info(f"chat_only inferred from /studio URL: {chat_only}") + # Defensive: an older Studio build without the launcher-contract + # patch may omit chat_only when called unauthenticated. Default + # to non-chat-only and log so the probe is not silently wrong. + chat_only = False + info( + "WARN /api/health did not return chat_only; defaulting to " + "chat_only=False (older Studio build)" + ) else: chat_only = bool(health.get("chat_only")) info(f"chat_only mode: {chat_only}")