Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
38 changes: 38 additions & 0 deletions backend/tests/test_flow.py
Original file line number Diff line number Diff line change
Expand Up @@ -167,6 +167,44 @@ def predict(self, state, questions, model=None):
assert client.get("/internal/api-keys").status_code == 401


@pytest.mark.parametrize("keep_cookies", [True, False], ids=["stale-cookies", "expired-cookies"])
def test_login_after_session_expiry(tmp_path, keep_cookies):
db_path = tmp_path / "db.sqlite3"
settings = Settings("admin", PasswordHasher().hash("test-password"), db_path,
"http://testserver", False, 1, tmp_path / "models", None, 1, tmp_path / "missing-dist")
client = TestClient(create_app(settings, FakePredictor()))
credentials = {"username": "admin", "password": "test-password"}
origin = {"Origin": "http://testserver"}
assert client.post("/internal/auth/login", json=credentials, headers=origin).status_code == 200
old_session = client.cookies["laya_session"]
old_csrf = client.get("/internal/auth/session").json()["csrf_token"]

with sqlite3.connect(db_path) as connection:
connection.execute("UPDATE sessions SET expires_at='2000-01-01T00:00:00+00:00'")
if not keep_cookies:
client.cookies.clear()

assert client.get("/internal/auth/session").status_code == 401
assert client.get("/internal/api-keys").status_code == 401
old_headers = {**origin, "X-CSRF-Token": old_csrf}
assert client.post("/internal/api-keys", json={"name": "expired"}, headers=old_headers).status_code == 401
assert client.post("/internal/auth/logout", headers=old_headers).status_code == 401

# Signing in does not require a successful logout or a valid old CSRF token.
assert client.post("/internal/auth/login", json=credentials, headers=origin).status_code == 200
assert client.cookies["laya_session"] != old_session
session = client.get("/internal/auth/session")
assert session.status_code == 200
new_csrf = session.json()["csrf_token"]
assert new_csrf and new_csrf != old_csrf
assert client.get("/internal/api-keys").status_code == 200
assert client.post("/internal/api-keys", json={"name": "stale-csrf"}, headers=old_headers).status_code == 403
created = client.post("/internal/api-keys", json={"name": "after-login"},
headers={**origin, "X-CSRF-Token": new_csrf})
assert created.status_code == 201
assert [key["name"] for key in client.get("/internal/api-keys").json()] == ["after-login"]


def test_concurrent_inference_requests_are_not_rejected(tmp_path):
class ConcurrentPredictor(FakePredictor):
def __init__(self):
Expand Down
30 changes: 11 additions & 19 deletions frontend/src/App.tsx
Original file line number Diff line number Diff line change
@@ -1,13 +1,13 @@
import { useEffect, useState, type FormEvent, type ReactNode } from "react"
import { useEffect, useState, useSyncExternalStore, type FormEvent, type ReactNode } from "react"
import { Activity, BookOpen, ChartNoAxesColumn, CircleHelp, Copy, KeyRound, Languages, LogOut, Menu, Plus, Send, X } from "lucide-react"
import { Button } from "@/components/ui/button"
import { Input } from "@/components/ui/input"
import { apiErrorMessage, useI18n, type Locale, type MessageKey, type Translate } from "./i18n"
import { api, ApiError, getSession, setSession, subscribeSession, type Session } from "./api"

type ApiKey = { id: number; name: string; mask: string; created_at: string; revoked_at: string | null; last_used_at: string | null }
type Usage = { totals: { requests: number; input_tokens: number; output_tokens: number }; daily: { day: string; requests: number; input_tokens: number; output_tokens: number }[]; sources: { source: string; key_id: number | null; requests: number; input_tokens: number; output_tokens: number }[] }
type Page = "home" | "keys" | "usage" | "playground" | "docs"
type Session = { username: string; csrf_token: string }

const navigation: { id: Page; label: MessageKey; icon: typeof Activity }[] = [
{ id: "home", label: "home", icon: Activity },
Expand All @@ -17,29 +17,20 @@ const navigation: { id: Page; label: MessageKey; icon: typeof Activity }[] = [
{ id: "docs", label: "documentation", icon: BookOpen },
]

class ApiError extends Error {
constructor(readonly code?: string) { super(code || "UNKNOWN_ERROR") }
}

function displayError(error: unknown, t: Translate) {
return apiErrorMessage(error instanceof ApiError ? error.code : undefined, t)
}

async function api<T>(path: string, options: RequestInit = {}, csrf?: string): Promise<T> {
const headers = new Headers(options.headers)
if (options.body) headers.set("Content-Type", "application/json")
if (options.method && options.method !== "GET") headers.set("X-CSRF-Token", csrf || "")
const response = await fetch(path, { credentials: "same-origin", ...options, headers })
const data = await response.json()
if (!response.ok) throw new ApiError(data?.detail?.code)
return data as T
}

function useSession() {
const [session, setSession] = useState<Session | null>(null)
const session = useSyncExternalStore(subscribeSession, getSession)
const [ready, setReady] = useState(false)
useEffect(() => {
api<Session>("/internal/auth/session").then(setSession).catch(() => setSession(null)).finally(() => setReady(true))
const controller = new AbortController()
api<Session>("/internal/auth/session", { signal: controller.signal })
.then(value => { if (!controller.signal.aborted) setSession(value) })
.catch(() => { if (!controller.signal.aborted) setSession(null) })
.finally(() => { if (!controller.signal.aborted) setReady(true) })
return () => controller.abort()
}, [])
return { session, setSession, ready }
}
Expand Down Expand Up @@ -100,8 +91,9 @@ function App() {
async function logout() {
try {
await api("/internal/auth/logout", { method: "POST" }, session!.csrf_token)
setSession(null)
if (getSession() === session) setSession(null)
} catch (cause) {
if (cause instanceof ApiError && cause.status === 401) return
alert(displayError(cause, t))
}
}
Expand Down
44 changes: 44 additions & 0 deletions frontend/src/api.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,44 @@
export type Session = { username: string; csrf_token: string }

let session: Session | null = null
let sessionVersion = 0
const listeners = new Set<() => void>()

export function getSession() { return session }

export function setSession(value: Session | null) {
session = value
sessionVersion += 1
listeners.forEach(listener => listener())
}

export function subscribeSession(listener: () => void) {
listeners.add(listener)
return () => { listeners.delete(listener) }
}

export class ApiError extends Error {
readonly code?: string
readonly status: number

constructor(status: number, code?: string) {
super(code || "UNKNOWN_ERROR")
this.status = status
this.code = code
}
}

export async function api<T>(path: string, options: RequestInit = {}, csrf?: string): Promise<T> {
const requestSessionVersion = sessionVersion
const headers = new Headers(options.headers)
if (options.body) headers.set("Content-Type", "application/json")
if (options.method && options.method !== "GET") headers.set("X-CSRF-Token", csrf || "")
const response = await fetch(path, { credentials: "same-origin", ...options, headers })
if (!response.ok) {
const unauthenticated = response.status === 401 && path.startsWith("/internal/") && path !== "/internal/auth/login"
if (unauthenticated && requestSessionVersion === sessionVersion) setSession(null)
const data = await response.json().catch(() => null)
throw new ApiError(response.status, data?.detail?.code || (unauthenticated ? "UNAUTHENTICATED" : undefined))
}
return response.json() as Promise<T>
}
Loading