unsloth/studio/backend/tests/test_middleware.py
Daniel Han 42ff23de73 Studio: fix HTML/SVG preview sanitizer, sandbox, and streaming gaps
Bundle of follow-ups to the HTML/SVG fence renderer landed earlier in
this PR. Each item came out of either the parallel reviewer pass or a
manual Playwright probe against the live Studio with an Anthropic
provider attached.

Sanitizer:
  - filter, mask, and clip-path are now in FORBID_ATTR. They accept
    url(https://...) values and the CSS engine still fetches that URL
    when the SVG renders, which previously slipped past the FORBID
    list.
  - href and xlink:href are no longer blanket-forbidden; they survive
    only when the value is a same-document fragment (href="#id"),
    which is what textPath, gradient, and use refs need. External
    schemes are dropped via a uponSanitizeAttribute hook so a beacon
    href cannot make it through.
  - The hook approach replaces DOMPurify's ALLOWED_URI_REGEXP, which
    also filtered presentation attrs (cx, cy, r, fill, width, height)
    and rendered circles with r=0.

SVG preview:
  - Inner stylesheet caps both max-width AND max-height so a square
    viewBox (200x200) scaled to the container width no longer
    overflows the fixed-height iframe and clips at the bottom.

HTML preview:
  - srcdoc carries a defense-in-depth meta-CSP (default-src 'none',
    connect-src 'none', frame-src 'none', img-src data: blob:,
    script-src 'self' 'unsafe-inline', style-src 'self' 'unsafe-inline').
    The host CSP already blocks inline scripts; this layer also blocks
    network egress, nested iframes, and form submission so a future
    host-CSP relaxation does not silently turn the preview into an
    exfiltration channel.
  - Sandbox grows allow-modals so alert/confirm/prompt are not
    silently no-oped if the host CSP ever permits inline scripts.
  - Pop-out spacer now uses the live HTML iframe height instead of
    hardcoded DEFAULT_PREVIEW_HEIGHT, so popping out a short preview
    does not leave a 500px hole in the chat bubble.
  - autoHeight resets on source change so a long-running session that
    swaps from a tall demo to a short one no longer keeps the previous
    iframe size during the gap before the new doc posts its height.

Streaming and a11y:
  - parseIncompleteCodeFence parses an in-flight open fence (no closing
    backticks yet). markdown-text falls back to it when streaming is
    incomplete, so the advertised isIncomplete -> Code-tab-lock path
    actually runs.
  - Tab buttons gain aria-controls / aria-labelledby wiring and a
    roving tabindex so the WAI-ARIA tab pattern is complete.
  - Pop-out modal gets role="dialog" and aria-modal.

Tooling:
  - vitest now runs in the Studio Frontend CI workflow so sanitizer or
    renderer regressions block the gate.
  - test-setup shims URL.createObjectURL / revokeObjectURL for jsdom in
    case future iframe work needs it.
  - frame-src in the host CSP is now declared explicitly as 'self' so
    a future change that loosens it leaves a visible diff for review.

Tests added: ARIA wiring, SVG height fit, srcdoc meta-CSP shape,
incomplete-fence helper, filter/mask/clip-path attr stripping, safe
fragment-href survival, external-href rejection. Vitest passes 21/21,
tsc -b and vite build are clean.
2026-05-24 16:05:11 +00:00

328 lines
11 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
"""Tests for MaxBodyMiddleware, SecurityHeadersMiddleware, and the /api/health auth gate."""
import asyncio
import importlib.util
import json
import os
import sys
from pathlib import Path
import pytest
from fastapi import FastAPI, HTTPException, Request
from fastapi.responses import Response
from fastapi.testclient import TestClient
_BACKEND_ROOT = Path(__file__).resolve().parents[1]
if str(_BACKEND_ROOT) not in sys.path:
sys.path.insert(0, str(_BACKEND_ROOT))
@pytest.fixture(scope = "module")
def main_module():
import main as _main # noqa: F401
return _main
# =====================================================================
# MaxBodyMiddleware
# =====================================================================
def _make_protected_app(max_bytes: int, main_module):
app = FastAPI()
app.add_middleware(
main_module.MaxBodyMiddleware,
max_bytes = max_bytes,
protected_prefixes = ("/v1/chat/completions", "/api/train"),
)
@app.post("/v1/chat/completions")
async def chat(payload: dict):
return {"ok": True, "n": len(payload.get("text", ""))}
@app.post("/api/other")
async def other(payload: dict):
return {"ok": True, "unprotected": True}
@app.get("/api/train/status")
async def status_get():
return {"ok": True, "get": True}
return app
class TestMaxBodyMiddleware:
def test_small_protected_body_passes(self, main_module):
app = _make_protected_app(1024, main_module)
c = TestClient(app)
r = c.post("/v1/chat/completions", json = {"text": "x" * 100})
assert r.status_code == 200
assert r.json()["n"] == 100
def test_large_declared_content_length_rejected(self, main_module):
app = _make_protected_app(1024, main_module)
c = TestClient(app)
r = c.post("/v1/chat/completions", json = {"text": "x" * 5000})
assert r.status_code == 413
assert "too large" in r.json()["detail"].lower()
def test_unprotected_prefix_passes_large_body(self, main_module):
app = _make_protected_app(1024, main_module)
c = TestClient(app)
r = c.post("/api/other", json = {"text": "x" * 5000})
assert r.status_code == 200
assert r.json()["unprotected"] is True
def test_chunked_upload_over_cap_rejected(self, main_module):
# Regression: declared-Content-Length-only check could be bypassed
# by chunked transfer-encoding.
app = _make_protected_app(1024, main_module)
c = TestClient(app)
def gen():
yield b'{"text":"'
yield b"x" * 800
yield b'"}'
yield b"\n" + b"y" * 500
r = c.post(
"/v1/chat/completions",
content = gen(),
headers = {"content-type": "application/json"},
)
assert r.status_code == 413
assert "too large" in r.json()["detail"].lower()
def test_chunked_upload_under_cap_passes(self, main_module):
app = _make_protected_app(1024, main_module)
c = TestClient(app)
def gen():
yield b'{"text":"'
yield b"x" * 50
yield b'"}'
r = c.post(
"/v1/chat/completions",
content = gen(),
headers = {"content-type": "application/json"},
)
assert r.status_code == 200
assert r.json()["n"] == 50
def test_get_not_subject_to_cap(self, main_module):
app = _make_protected_app(1024, main_module)
c = TestClient(app)
r = c.get("/api/train/status")
assert r.status_code == 200
# =====================================================================
# SecurityHeadersMiddleware / CSP
# =====================================================================
def _make_csp_app(main_module, attach_nonce: str | None = None):
app = FastAPI()
app.add_middleware(main_module.SecurityHeadersMiddleware)
@app.get("/plain")
async def plain():
return {"ok": True}
@app.get("/with-nonce")
async def with_nonce():
headers = {}
if attach_nonce:
headers[main_module._CSP_SCRIPT_NONCE_HEADER] = attach_nonce
return Response(
content = b"<html></html>",
media_type = "text/html",
headers = headers,
)
return app
class TestSecurityHeadersMiddleware:
def test_csp_has_no_unsafe_inline_for_script_src(self, main_module):
app = _make_csp_app(main_module)
c = TestClient(app)
r = c.get("/plain")
assert r.status_code == 200
csp = r.headers["content-security-policy"]
# Parse per-directive so style-src unsafe-inline does not false-match.
directives = {
chunk.strip().split(" ", 1)[0]: chunk.strip()
for chunk in csp.split(";")
if chunk.strip()
}
assert "script-src" in directives
assert "'unsafe-inline'" not in directives["script-src"]
# style-src keeps unsafe-inline for Vite-injected styles.
assert "'unsafe-inline'" in directives["style-src"]
def test_default_security_headers_present(self, main_module):
app = _make_csp_app(main_module)
c = TestClient(app)
r = c.get("/plain")
assert r.headers["x-frame-options"] == "DENY"
assert r.headers["x-content-type-options"] == "nosniff"
assert r.headers["referrer-policy"] == "no-referrer"
assert "camera=()" in r.headers["permissions-policy"]
assert r.headers["server"] == "unsloth-studio"
def test_internal_nonce_header_is_spliced_into_csp_and_stripped(self, main_module):
nonce = "test-nonce-abc"
app = _make_csp_app(main_module, attach_nonce = nonce)
c = TestClient(app)
r = c.get("/with-nonce")
csp = r.headers["content-security-policy"]
assert f"'nonce-{nonce}'" in csp
# Internal handoff header must not leak to clients.
assert main_module._CSP_SCRIPT_NONCE_HEADER not in {
k.lower() for k in r.headers.keys()
}
def test_build_csp_helper_shape(self, main_module):
plain = main_module._build_csp()
assert "script-src 'self';" in plain
assert "'unsafe-inline'" not in plain.split("script-src", 1)[1].split(";", 1)[0]
nonced = main_module._build_csp("XYZ")
assert "script-src 'self' 'nonce-XYZ';" in nonced
def test_frame_src_is_explicitly_self_only(self, main_module):
# The assistant HTML/SVG preview iframe uses srcdoc (no URL fetch),
# so frame-src does not need to permit data: / blob:. Pinning to
# 'self' explicitly is the strictest setting CSP allows here, and
# leaves a visible directive a reviewer can grep for if a future
# change tries to relax it without an audit.
csp = main_module._build_csp()
frame_src = next(
chunk.strip()
for chunk in csp.split(";")
if chunk.strip().startswith("frame-src ")
)
tokens = frame_src.split()
assert tokens[0] == "frame-src"
assert "'self'" in tokens
assert "data:" not in tokens
assert "blob:" not in tokens
def test_img_src_allows_google_favicons(self, main_module):
# sources.tsx fetches https://www.google.com/s2/favicons?... ; without
# this allowlist entry citation favicons fall back to gray initials.
csp = main_module._build_csp()
img_directive = next(
chunk.strip()
for chunk in csp.split(";")
if chunk.strip().startswith("img-src ")
)
# Tokenise and compare with `==` so CodeQL's URL-substring rule does
# not read directive-string `in` membership as URL sanitisation.
img_sources = img_directive.split()
assert any(src == "https://www.google.com" for src in img_sources)
# Pre-existing favicon CDNs stay allowed.
for host in (
"https://t0.gstatic.com",
"https://t1.gstatic.com",
"https://t2.gstatic.com",
"https://t3.gstatic.com",
):
assert any(src == host for src in img_sources)
# =====================================================================
# /api/health auth gate
# =====================================================================
@pytest.fixture
def health_app(tmp_path, monkeypatch):
"""Mount /api/health on a fresh app against an isolated auth db."""
from auth import storage
monkeypatch.setattr(storage, "DB_PATH", tmp_path / "auth.db")
monkeypatch.setattr(storage, "_BOOTSTRAP_PW_PATH", tmp_path / ".bootstrap_password")
monkeypatch.setattr(storage, "_bootstrap_password", None)
import main as _main
app = FastAPI()
app.add_api_route("/api/health", _main.health_check, methods = ["GET"])
import secrets as _secrets
storage.create_initial_user(
username = storage.DEFAULT_ADMIN_USERNAME,
password = "human-password-123",
jwt_secret = _secrets.token_urlsafe(64),
must_change_password = False,
)
return app
class TestHealthAuthGate:
# Launcher / frontend bootstrap fields are available unauth so the Tauri
# watchdog can re-adopt a sibling backend and the SPA can detect chat-only
# mode before any token exists. Version / device_type still require a bearer.
LAUNCHER_BITS = (
"service",
"studio_root_id",
"chat_only",
"desktop_protocol_version",
"desktop_manageability_version",
"supports_desktop_auth",
"supports_desktop_backend_ownership",
"native_path_leases_supported",
)
FINGERPRINT_FIELDS = ("version", "studio_version", "device_type")
def test_no_auth_exposes_launcher_bits(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 field in self.LAUNCHER_BITS:
assert field in body, f"missing launcher bit: {field}"
assert body["service"] == "Unsloth UI Backend"
for forbidden in self.FINGERPRINT_FIELDS:
assert forbidden not in body
def test_invalid_bearer_returns_launcher_bits_only(self, health_app):
# Regression: calling the async dep without await made any Bearer header pass.
c = TestClient(health_app)
r = c.get(
"/api/health",
headers = {"Authorization": "Bearer not-a-real-token"},
)
assert r.status_code == 200
body = r.json()
assert body["status"] == "healthy"
for field in self.LAUNCHER_BITS:
assert field in body
for forbidden in self.FINGERPRINT_FIELDS:
assert forbidden not in body
def test_valid_bearer_returns_full_payload(self, health_app):
from auth import storage
from auth.authentication import create_access_token
token = create_access_token(storage.DEFAULT_ADMIN_USERNAME)
c = TestClient(health_app)
r = c.get(
"/api/health",
headers = {"Authorization": f"Bearer {token}"},
)
assert r.status_code == 200
body = r.json()
assert body["status"] == "healthy"
for field in self.LAUNCHER_BITS + self.FINGERPRINT_FIELDS:
assert field in body, f"missing: {field}"