"""Sign-in, session cookies, and the per-trace allow list. The OAuth round trip to Hugging Face is stubbed; what's tested here is the part that would actually let the wrong person in — signature forgery, expiry, open redirects, CSRF on the callback, and the allow-list check itself. """ import os import tempfile import time os.environ["TRACE_DIR"] = tempfile.mkdtemp() os.environ.pop("UPLOAD_TOKEN", None) os.environ["OAUTH_CLIENT_ID"] = "cid" os.environ["OAUTH_CLIENT_SECRET"] = "csecret" import pytest # noqa: E402 from fastapi.testclient import TestClient # noqa: E402 import app as app_module # noqa: E402 import auth # noqa: E402 client = TestClient(app_module.app, raise_server_exceptions=False, follow_redirects=False) @pytest.fixture(autouse=True) def fresh(): app_module._hits.clear() app_module.UPLOAD_TOKEN = "" client.cookies.clear() yield app_module._hits.clear() client.cookies.clear() def sign_in_as(username): client.cookies.set(auth.SESSION_COOKIE, auth.make_session(username)) def create(trace="hello", allowed=None, title="t"): body = {"title": title, "trace": trace} if allowed is not None: body["allowed_users"] = allowed response = client.post("/api/traces", json=body) assert response.status_code == 201, response.text return response.json() # ---------------------------------------------------------------- allow list def test_normalize_users(): assert auth.normalize_users("Alice, @Bob carol") == ["alice", "bob", "carol"] assert auth.normalize_users(["A", "a", "b"]) == ["a", "b"] # deduplicated assert auth.normalize_users(None) == [] assert auth.normalize_users(123) == [] assert len(auth.normalize_users([f"u{i}" for i in range(500)])) == 100 # capped def test_normalize_users_rejects_junk(): # Anything with characters an HF username can't contain is dropped entirely. assert "bad!name" not in auth.normalize_users("bad!name good-name") assert "good-name" in auth.normalize_users("bad!name good-name") def test_is_allowed(): assert auth.is_allowed(None, None) is True # unrestricted assert auth.is_allowed(None, []) is True assert auth.is_allowed("alice", ["Alice"]) is True # case-insensitive assert auth.is_allowed("mallory", ["alice"]) is False assert auth.is_allowed(None, ["alice"]) is False # ---------------------------------------------------------------- sessions def test_session_roundtrip(): assert auth.unsign(auth.make_session("alice"))["u"] == "alice" def test_tampered_session_rejected(): good = auth.make_session("alice") body, _, mac = good.partition(".") forged = auth.sign({"u": "mallory", "exp": time.time() + 999}) assert auth.unsign(f"{forged.split('.')[0]}.{mac}") is None # swapped payload assert auth.unsign(body + ".") is None assert auth.unsign("garbage") is None assert auth.unsign(None) is None def test_expired_session_rejected(): assert auth.unsign(auth.sign({"u": "alice", "exp": int(time.time()) - 1})) is None def test_session_secret_actually_gates(monkeypatch): token = auth.make_session("alice") monkeypatch.setattr(auth, "SESSION_SECRET", "a-different-secret") assert auth.unsign(token) is None # ---------------------------------------------------------------- access def test_unrestricted_trace_needs_no_signin(): created = create("public trace") assert client.get(created["url"]).status_code == 200 def test_restricted_trace_redirects_anonymous_to_login(): created = create("secret", allowed=["alice"]) response = client.get(created["url"]) assert response.status_code == 302 assert response.headers["location"].startswith("/login?next=") def test_restricted_trace_visible_to_listed_user(): created = create("for alice only", allowed=["alice"]) sign_in_as("alice") response = client.get(created["url"]) assert response.status_code == 200 assert "for alice only" in response.text def test_restricted_trace_404s_for_unlisted_user(): created = create("for alice only", allowed=["alice"]) sign_in_as("mallory") response = client.get(created["url"]) assert response.status_code == 404 assert "for alice only" not in response.text def test_unlisted_user_gets_same_404_as_missing_trace(): created = create("x", allowed=["alice"]) sign_in_as("mallory") restricted = client.get(created["url"]) missing = client.get("/t/" + "A" * 22) assert restricted.status_code == missing.status_code == 404 assert restricted.text == missing.text def test_raw_endpoint_enforces_allow_list(): created = create({"messages": []}, allowed=["alice"]) assert client.get(created["raw_url"]).status_code == 404 # anonymous, no redirect sign_in_as("mallory") assert client.get(created["raw_url"]).status_code == 404 sign_in_as("alice") assert client.get(created["raw_url"]).status_code == 200 def test_download_sets_attachment_header(): created = create({"messages": []}) response = client.get(created["raw_url"] + "?download=1") assert response.status_code == 200 assert "attachment" in response.headers["content-disposition"] assert created["id"] in response.headers["content-disposition"] def test_allowed_users_stored_and_reported(): created = create("x", allowed=["Alice", "bob"]) assert created["allowed_users"] == ["alice", "bob"] def test_restricted_creation_rejected_without_oauth(monkeypatch): monkeypatch.delenv("OAUTH_CLIENT_ID", raising=False) response = client.post("/api/traces", json={"trace": "x", "allowed_users": ["alice"]}) assert response.status_code == 400 assert "OAuth" in response.json()["error"] # ---------------------------------------------------------------- oauth flow FAKE_PROVIDER = { "authorization_endpoint": "https://huggingface.co/oauth/authorize", "token_endpoint": "https://huggingface.co/oauth/token", "userinfo_endpoint": "https://huggingface.co/oauth/userinfo", } def test_login_redirects_to_provider(monkeypatch): monkeypatch.setattr(auth, "_provider", lambda: FAKE_PROVIDER) response = client.get("/login?next=/t/" + "A" * 22) assert response.status_code == 302 assert response.headers["location"].startswith("https://huggingface.co/oauth/authorize") assert "scope=openid+profile" in response.headers["location"] assert auth.STATE_COOKIE in response.cookies def test_login_next_cannot_be_an_open_redirect(): assert app_module.safe_next("https://evil.test") == "/" assert app_module.safe_next("//evil.test") == "/" assert app_module.safe_next("/t/abc") == "/t/abc" assert app_module.safe_next(None) == "/" def test_callback_rejects_missing_or_forged_state(monkeypatch): monkeypatch.setattr(auth, "_provider", lambda: FAKE_PROVIDER) assert client.get("/auth/callback?code=c&state=s").status_code == 400 client.cookies.set(auth.STATE_COOKIE, auth.sign( {"s": "realstate", "n": "/", "exp": int(time.time()) + 600})) assert client.get("/auth/callback?code=c&state=wrong").status_code == 400 assert client.get("/auth/callback?state=realstate").status_code == 400 def test_callback_sets_session_and_returns_to_next(monkeypatch): monkeypatch.setattr(auth, "exchange_code", lambda code, uri: "access-token") monkeypatch.setattr(auth, "fetch_username", lambda token: "alice") client.cookies.set(auth.STATE_COOKIE, auth.sign( {"s": "st", "n": "/t/" + "A" * 22, "exp": int(time.time()) + 600})) response = client.get("/auth/callback?code=c&state=st") assert response.status_code == 302 assert response.headers["location"] == "/t/" + "A" * 22 assert auth.unsign(response.cookies[auth.SESSION_COOKIE])["u"] == "alice" def test_callback_survives_provider_failure(monkeypatch): def boom(*a, **k): raise RuntimeError("provider down") monkeypatch.setattr(auth, "exchange_code", boom) client.cookies.set(auth.STATE_COOKIE, auth.sign( {"s": "st", "n": "/", "exp": int(time.time()) + 600})) assert client.get("/auth/callback?code=c&state=st").status_code == 502 def test_logout_clears_session(): sign_in_as("alice") response = client.get("/logout") assert response.status_code == 302 assert not response.cookies.get(auth.SESSION_COOKIE) def test_session_cookie_is_httponly_and_lax(monkeypatch): monkeypatch.setattr(auth, "exchange_code", lambda code, uri: "tok") monkeypatch.setattr(auth, "fetch_username", lambda token: "alice") client.cookies.set(auth.STATE_COOKIE, auth.sign( {"s": "st", "n": "/", "exp": int(time.time()) + 600})) response = client.get("/auth/callback?code=c&state=st") header = response.headers["set-cookie"] assert "httponly" in header.lower() assert "samesite=lax" in header.lower()