Download test_auth.py from lvwerra/tracebin: direct link, hf CLI and curl.
- Browser
- Download file 8.87 kB
-
https://huggingface.co/spaces/lvwerra/tracebin/resolve/main/test_auth.py
- Command line
-
hf download hf://spaces/lvwerra/tracebin/test_auth.py
-
curl -L -o test_auth.py https://huggingface.co/spaces/lvwerra/tracebin/resolve/main/test_auth.py
8.87 kB
| """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) | |
| 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() | |