tracebin / test_auth.py
lvwerra's picture
lvwerra HF Staff
oauth-restricted traces, download + agent-prompt buttons
6b41207 verified
Raw History Blame Contribute Delete
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)
@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()