from __future__ import annotations import hashlib import hmac import json import os import re import secrets import sqlite3 import time from collections import defaultdict, deque from contextlib import contextmanager from datetime import datetime, timedelta, timezone from pathlib import Path from typing import Any, Iterator from uuid import uuid4 from argon2 import PasswordHasher from argon2.exceptions import InvalidHashError, VerificationError, VerifyMismatchError from fastapi import Depends, FastAPI, HTTPException, Request, Response, status from fastapi.middleware.cors import CORSMiddleware from pydantic import BaseModel, Field, field_validator ROOT = Path(__file__).resolve().parents[1] DB_PATH = Path(os.getenv("KEYCHAIN_DB_PATH", str(ROOT / "data" / "keychain.db"))) DB_PATH.parent.mkdir(parents=True, exist_ok=True) SESSION_TTL = int(os.getenv("KEYCHAIN_SESSION_TTL_SECONDS", "1209600")) COOKIE_SECURE = os.getenv("KEYCHAIN_COOKIE_SECURE", "false").lower() == "true" COOKIE_SAMESITE = "lax" SESSION_COOKIE = "kc_session" CSRF_COOKIE = "kc_csrf" configured_origins = [ origin.strip() for origin in os.getenv( "KEYCHAIN_ALLOWED_ORIGINS", "http://localhost:5173,http://127.0.0.1:5173" ).split(",") if origin.strip() ] password_hasher = PasswordHasher() email_pattern = re.compile(r"^[^\s@]+@[^\s@]+\.[^\s@]+$") auth_attempts: dict[str, deque[float]] = defaultdict(deque) def utc_now() -> datetime: return datetime.now(timezone.utc) def iso_now() -> str: return utc_now().isoformat() def hash_token(token: str) -> str: return hashlib.sha256(token.encode("utf-8")).hexdigest() @contextmanager def db() -> Iterator[sqlite3.Connection]: connection = sqlite3.connect(DB_PATH, timeout=10) connection.row_factory = sqlite3.Row connection.execute("PRAGMA foreign_keys = ON") try: yield connection connection.commit() finally: connection.close() def init_db() -> None: with db() as connection: connection.executescript( """ CREATE TABLE IF NOT EXISTS users ( id TEXT PRIMARY KEY, email TEXT NOT NULL UNIQUE COLLATE NOCASE, display_name TEXT NOT NULL, password_hash TEXT NOT NULL, created_at TEXT NOT NULL ); CREATE TABLE IF NOT EXISTS vaults ( user_id TEXT PRIMARY KEY REFERENCES users(id) ON DELETE CASCADE, envelope TEXT, revision INTEGER NOT NULL DEFAULT 0, updated_at TEXT NOT NULL ); CREATE TABLE IF NOT EXISTS sessions ( id TEXT PRIMARY KEY, user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE, token_hash TEXT NOT NULL UNIQUE, csrf_hash TEXT NOT NULL, created_at TEXT NOT NULL, expires_at TEXT NOT NULL, last_seen_at TEXT NOT NULL ); CREATE INDEX IF NOT EXISTS sessions_token_idx ON sessions(token_hash); CREATE INDEX IF NOT EXISTS sessions_user_idx ON sessions(user_id); """ ) init_db() class AuthPayload(BaseModel): email: str = Field(min_length=3, max_length=254) password: str = Field(min_length=8, max_length=128) display_name: str | None = Field(default=None, max_length=80) @field_validator("email") @classmethod def normalize_email(cls, value: str) -> str: normalized = value.strip().lower() if not email_pattern.fullmatch(normalized): raise ValueError("Enter a valid email address") return normalized @field_validator("password") @classmethod def reject_control_chars(cls, value: str) -> str: if any(ord(character) < 32 for character in value): raise ValueError("Password contains unsupported control characters") return value class VaultPayload(BaseModel): envelope: dict[str, Any] = Field(min_length=1) revision: int = Field(ge=0) @field_validator("envelope") @classmethod def validate_size(cls, value: dict[str, Any]) -> dict[str, Any]: serialized = json.dumps(value, separators=(",", ":"), ensure_ascii=False) if len(serialized.encode("utf-8")) > 8 * 1024 * 1024: raise ValueError("Vault envelope is too large") return value def public_user(row: sqlite3.Row) -> dict[str, str]: return { "id": row["id"], "email": row["email"], "display_name": row["display_name"], } def enforce_auth_rate_limit(request: Request) -> None: address = request.client.host if request.client else "unknown" now = time.monotonic() attempts = auth_attempts[address] while attempts and attempts[0] < now - 60: attempts.popleft() if len(attempts) >= 12: raise HTTPException(status_code=429, detail="Too many attempts. Try again in a minute.") attempts.append(now) def issue_session(response: Response, user_id: str) -> str: raw_session = secrets.token_urlsafe(32) raw_csrf = secrets.token_urlsafe(24) now = utc_now() expires = now + timedelta(seconds=SESSION_TTL) with db() as connection: connection.execute( "INSERT INTO sessions (id, user_id, token_hash, csrf_hash, created_at, expires_at, last_seen_at) " "VALUES (?, ?, ?, ?, ?, ?, ?)", ( uuid4().hex, user_id, hash_token(raw_session), hash_token(raw_csrf), now.isoformat(), expires.isoformat(), now.isoformat(), ), ) response.set_cookie( SESSION_COOKIE, raw_session, max_age=SESSION_TTL, httponly=True, secure=COOKIE_SECURE, samesite=COOKIE_SAMESITE, path="/", ) response.set_cookie( CSRF_COOKIE, raw_csrf, max_age=SESSION_TTL, httponly=False, secure=COOKIE_SECURE, samesite=COOKIE_SAMESITE, path="/", ) return raw_csrf def current_session(request: Request) -> sqlite3.Row: raw_token = request.cookies.get(SESSION_COOKIE) if not raw_token: raise HTTPException(status_code=401, detail="Authentication required") now = utc_now() with db() as connection: session = connection.execute( "SELECT * FROM sessions WHERE token_hash = ?", (hash_token(raw_token),) ).fetchone() if session is None: raise HTTPException(status_code=401, detail="Session is invalid") if datetime.fromisoformat(session["expires_at"]) <= now: connection.execute("DELETE FROM sessions WHERE id = ?", (session["id"],)) raise HTTPException(status_code=401, detail="Session has expired") connection.execute( "UPDATE sessions SET last_seen_at = ? WHERE id = ?", (now.isoformat(), session["id"]) ) return session def require_csrf(request: Request, session: sqlite3.Row) -> None: header_token = request.headers.get("x-csrf-token", "") cookie_token = request.cookies.get(CSRF_COOKIE, "") if not header_token or not hmac.compare_digest(header_token, cookie_token): raise HTTPException(status_code=403, detail="CSRF validation failed") if not hmac.compare_digest(hash_token(header_token), session["csrf_hash"]): raise HTTPException(status_code=403, detail="CSRF token is invalid") def user_for_session(session: sqlite3.Row) -> sqlite3.Row: with db() as connection: user = connection.execute("SELECT * FROM users WHERE id = ?", (session["user_id"],)).fetchone() if user is None: raise HTTPException(status_code=401, detail="Account no longer exists") return user app = FastAPI(title="Keychain API", version="0.1.0") app.add_middleware( CORSMiddleware, allow_origins=configured_origins, allow_origin_regex=r"^chrome-extension://[a-zA-Z0-9]+$", allow_credentials=True, allow_methods=["GET", "POST", "PUT", "OPTIONS"], allow_headers=["Content-Type", "X-CSRF-Token"], ) @app.middleware("http") async def security_headers(request: Request, call_next): response = await call_next(request) response.headers["X-Content-Type-Options"] = "nosniff" response.headers["X-Frame-Options"] = "DENY" response.headers["Referrer-Policy"] = "strict-origin-when-cross-origin" response.headers["Permissions-Policy"] = "camera=(), microphone=(), geolocation=()" return response @app.get("/api/health") def health() -> dict[str, str]: return {"status": "ok"} @app.post("/api/auth/register", status_code=status.HTTP_201_CREATED) def register(payload: AuthPayload, request: Request, response: Response) -> dict[str, Any]: enforce_auth_rate_limit(request) user_id = uuid4().hex display_name = (payload.display_name or payload.email.split("@", 1)[0]).strip()[:80] now = iso_now() try: with db() as connection: connection.execute( "INSERT INTO users (id, email, display_name, password_hash, created_at) VALUES (?, ?, ?, ?, ?)", (user_id, payload.email, display_name, password_hasher.hash(payload.password), now), ) connection.execute( "INSERT INTO vaults (user_id, envelope, revision, updated_at) VALUES (?, ?, 0, ?)", (user_id, None, now), ) except sqlite3.IntegrityError as error: if "email" in str(error).lower(): raise HTTPException(status_code=409, detail="An account with this email already exists") from error raise csrf_token = issue_session(response, user_id) with db() as connection: user = connection.execute("SELECT * FROM users WHERE id = ?", (user_id,)).fetchone() return {"user": public_user(user), "csrf_token": csrf_token} @app.post("/api/auth/login") def login(payload: AuthPayload, request: Request, response: Response) -> dict[str, Any]: enforce_auth_rate_limit(request) with db() as connection: user = connection.execute("SELECT * FROM users WHERE email = ?", (payload.email,)).fetchone() if user is None: raise HTTPException(status_code=401, detail="Email or password is incorrect") try: password_hasher.verify(user["password_hash"], payload.password) except (VerifyMismatchError, VerificationError, InvalidHashError) as error: raise HTTPException(status_code=401, detail="Email or password is incorrect") from error if password_hasher.check_needs_rehash(user["password_hash"]): with db() as connection: connection.execute( "UPDATE users SET password_hash = ? WHERE id = ?", (password_hasher.hash(payload.password), user["id"]), ) csrf_token = issue_session(response, user["id"]) return {"user": public_user(user), "csrf_token": csrf_token} @app.get("/api/auth/me") def me(session: sqlite3.Row = Depends(current_session)) -> dict[str, Any]: return {"user": public_user(user_for_session(session))} @app.post("/api/auth/logout") def logout(request: Request, response: Response, session: sqlite3.Row = Depends(current_session)) -> dict[str, str]: require_csrf(request, session) with db() as connection: connection.execute("DELETE FROM sessions WHERE id = ?", (session["id"],)) response.delete_cookie(SESSION_COOKIE, path="/") response.delete_cookie(CSRF_COOKIE, path="/") return {"status": "signed_out"} @app.get("/api/vault") def get_vault(session: sqlite3.Row = Depends(current_session)) -> dict[str, Any]: user = user_for_session(session) with db() as connection: vault = connection.execute("SELECT * FROM vaults WHERE user_id = ?", (user["id"],)).fetchone() if vault is None: raise HTTPException(status_code=404, detail="Vault not found") return { "envelope": json.loads(vault["envelope"]) if vault["envelope"] else None, "revision": vault["revision"], "updated_at": vault["updated_at"], } @app.put("/api/vault") def put_vault( payload: VaultPayload, request: Request, session: sqlite3.Row = Depends(current_session), ) -> dict[str, Any]: require_csrf(request, session) with db() as connection: vault = connection.execute( "SELECT revision FROM vaults WHERE user_id = ?", (session["user_id"],) ).fetchone() if vault is None: raise HTTPException(status_code=404, detail="Vault not found") if vault["revision"] != payload.revision: raise HTTPException(status_code=409, detail="Vault changed on another device; reload before saving") new_revision = payload.revision + 1 updated_at = iso_now() connection.execute( "UPDATE vaults SET envelope = ?, revision = ?, updated_at = ? WHERE user_id = ?", ( json.dumps(payload.envelope, separators=(",", ":"), ensure_ascii=False), new_revision, updated_at, session["user_id"], ), ) return {"revision": new_revision, "updated_at": updated_at}