diff --git a/apps/api/src/auth.rs b/apps/api/src/auth.rs index b63f575..25d232e 100644 --- a/apps/api/src/auth.rs +++ b/apps/api/src/auth.rs @@ -48,8 +48,17 @@ pub struct AuthenticatedUser { pub id: Uuid, } -pub async fn csrf(State(state): State) -> Response { - let token = security::random_token(); +pub async fn csrf(State(state): State, headers: HeaderMap) -> Response { + // Reuse the browser's current token so another tab fetching /auth/csrf + // does not invalidate the token this tab already holds. + let token = security::get_cookie(&headers, state.config.csrf_cookie_name()) + .filter(|value| { + (32..=128).contains(&value.len()) + && value + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || byte == b'-' || byte == b'_') + }) + .unwrap_or_else(security::random_token); let mut response = ( StatusCode::OK, Json(CsrfResponse { @@ -425,4 +434,32 @@ mod tests { let error = validate_identity("User", "user@example.com", "short").unwrap_err(); assert!(matches!(error, AppError::Validation { .. })); } + + #[tokio::test] + async fn csrf_reuses_the_browsers_current_token() { + let Some(state) = crate::test_support::test_app_state().await else { + return; + }; + let token_from = |response: Response| { + let cookie = response.headers()[axum::http::header::SET_COOKIE] + .to_str() + .unwrap() + .to_owned(); + cookie.split(';').next().unwrap().split_once('=').unwrap().1.to_owned() + }; + let first = token_from(csrf(State(state.clone()), HeaderMap::new()).await); + let mut headers = HeaderMap::new(); + headers.insert( + axum::http::header::COOKIE, + format!("{}={first}", state.config.csrf_cookie_name()).parse().unwrap(), + ); + assert_eq!(token_from(csrf(State(state.clone()), headers).await), first); + // A malformed cookie is replaced, not echoed back. + let mut headers = HeaderMap::new(); + headers.insert( + axum::http::header::COOKIE, + format!("{}=bad;token", state.config.csrf_cookie_name()).parse().unwrap(), + ); + assert_ne!(token_from(csrf(State(state), headers).await), "bad"); + } } diff --git a/apps/web/src/lib/api.test.ts b/apps/web/src/lib/api.test.ts new file mode 100644 index 0000000..0371203 --- /dev/null +++ b/apps/web/src/lib/api.test.ts @@ -0,0 +1,42 @@ +import { afterEach, describe, expect, it, vi } from "vitest" + +import { apiRequest, resetCsrfToken } from "@/lib/api" + +function json(status: number, body: unknown) { + return new Response(JSON.stringify(body), { + status, + headers: { "Content-Type": "application/json" }, + }) +} + +describe("apiRequest", () => { + afterEach(() => { + vi.unstubAllGlobals() + resetCsrfToken() + }) + + it("refreshes a stale CSRF token and retries the mutation once", async () => { + const tokens = ["stale", "fresh"] + const sent: (string | null)[] = [] + const fetchMock = vi.fn((url: string, init?: RequestInit) => { + if (url.endsWith("/auth/csrf")) { + return Promise.resolve(json(200, { csrfToken: tokens.shift() })) + } + const token = new Headers(init?.headers).get("X-CSRF-Token") + sent.push(token) + return Promise.resolve( + token === "fresh" + ? json(200, { ok: true }) + : json(403, { + error: { code: "CSRF_INVALID", message: "The CSRF token is invalid." }, + }) + ) + }) + vi.stubGlobal("fetch", fetchMock) + + await expect( + apiRequest("/integrations/knotree-registry/authorize", { method: "POST" }) + ).resolves.toEqual({ ok: true }) + expect(sent).toEqual(["stale", "fresh"]) + }) +}) diff --git a/apps/web/src/lib/api.ts b/apps/web/src/lib/api.ts index 1494b61..076b87f 100644 --- a/apps/web/src/lib/api.ts +++ b/apps/web/src/lib/api.ts @@ -80,6 +80,27 @@ export async function apiBinaryRequest( export async function apiRequest( path: string, options: ApiRequestOptions = {} +): Promise { + try { + return await sendApiRequest(path, options) + } catch (error) { + // The CSRF cookie is per browser; if it changed since this tab cached its + // token (sign-in, logout, another tab), fetch a fresh one and retry once. + if ( + error instanceof ApiError && + error.status === 403 && + (error.code === "CSRF_INVALID" || error.code === "CSRF_REQUIRED") + ) { + resetCsrfToken() + return sendApiRequest(path, options) + } + throw error + } +} + +async function sendApiRequest( + path: string, + options: ApiRequestOptions ): Promise { const method = (options.method ?? "GET").toUpperCase() const headers = new Headers(options.headers)