fix: CSRF protection, cookie security flags, rate limiting, XSS (#23,#24,#25,#28,#33)

Add CSRF tokens to all cookie-based POST endpoints.
Set Secure and SameSite=Strict on auth cookie.
Rate limit login to 5 attempts per 15min per IP.
Escape HTML in signing error messages (XSS fix).
Remove duplicate get_user_from_cookie definition.
This commit is contained in:
Jarian Cottingham 2026-07-04 04:51:56 +00:00
parent c759597ad0
commit df82068aeb
2 changed files with 126 additions and 16 deletions

View File

@ -1,4 +1,4 @@
import os, sqlite3, datetime, secrets, hashlib, subprocess, json
import os, sqlite3, datetime, secrets, hashlib, subprocess, json, time, functools
from fastapi import FastAPI, Request, Depends, HTTPException, Form
from fastapi.responses import HTMLResponse, FileResponse, RedirectResponse, JSONResponse, PlainTextResponse, StreamingResponse
from fastapi.staticfiles import StaticFiles
@ -14,12 +14,42 @@ from cryptography.hazmat.primitives import serialization
app = FastAPI(title="CertAuth Key Vault")
app.mount("/static", StaticFiles(directory="/opt/certauth/api/static"), name="static")
_login_attempts = {}
_LOGIN_MAX_ATTEMPTS = 5
_LOGIN_WINDOW_SECONDS = 900
_csrf_secret = secrets.token_hex(32)
def _check_rate_limit(client_ip: str) -> bool:
now = time.time()
if client_ip not in _login_attempts:
_login_attempts[client_ip] = []
_login_attempts[client_ip] = [
t for t in _login_attempts[client_ip] if now - t < _LOGIN_WINDOW_SECONDS
]
if len(_login_attempts[client_ip]) >= _LOGIN_MAX_ATTEMPTS:
return False
_login_attempts[client_ip].append(now)
return True
def _generate_csrf_token() -> str:
return secrets.token_hex(32)
def _verify_csrf_token(request: Request, token: str) -> bool:
stored = request.cookies.get("csrf_token")
if not stored or not token:
return False
return secrets.compare_digest(stored, token)
def _set_csrf_cookie(resp):
token = _generate_csrf_token()
resp.set_cookie("csrf_token", token, httponly=False, samesite="strict", path="/")
return token
def get_user_from_cookie(request: Request):
token = request.cookies.get("token")
if not token: return None
try: return jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
except: return None
jinja_env = Environment(
loader=FileSystemLoader("/opt/certauth/api/templates"),
@ -43,7 +73,10 @@ class LoginRequest(BaseModel):
password: str
@app.post("/api/token")
async def login(req: LoginRequest):
async def login(req: LoginRequest, request: Request):
client_ip = request.client.host
if not _check_rate_limit(client_ip):
raise HTTPException(429, "Too many login attempts. Try again later.")
conn = get_db()
row = conn.execute("SELECT * FROM users WHERE username = ?", (req.username,)).fetchone()
conn.close()
@ -103,7 +136,9 @@ async def sign_cert(cert_id: int, user: str = Depends(get_current_user)):
raise HTTPException(400, "Not found or already signed")
conn.close()
result, err = build_leaf_cert(row["subject"], row["san"], 365)
if err: raise HTTPException(500, f"Signing failed: {err}")
if err:
import html as h
raise HTTPException(500, f"Signing failed: {h.escape(str(err))}")
cf = f"/etc/ssl/ca/issued/cert-{result['serial']}.crt"
kf = f"/etc/ssl/ca/issued/cert-{result['serial']}.key"
open(cf, "w").write(result["cert_pem"])
@ -203,7 +238,11 @@ async def login_page(request: Request):
return render("login.html", {"request": request, "error": None})
@app.post("/login")
async def login_post(username: str = Form(...), password: str = Form(...)):
async def login_post(username: str = Form(...), password: str = Form(...),
request: Request = None):
client_ip = request.client.host if request else "unknown"
if not _check_rate_limit(client_ip):
return render("login.html", {"request": None, "error": "Too many login attempts. Try again later."})
conn = get_db()
row = conn.execute("SELECT * FROM users WHERE username = ?", (username,)).fetchone()
conn.close()
@ -211,26 +250,31 @@ async def login_post(username: str = Form(...), password: str = Form(...)):
return render("login.html", {"request": None, "error": "Invalid credentials"})
token = create_access_token({"sub": username})
resp = RedirectResponse("/", status_code=302)
resp.set_cookie("token", token, httponly=True, samesite="lax", path="/")
resp.set_cookie("token", token, httponly=True, samesite="strict", secure=True, path="/")
_set_csrf_cookie(resp)
return resp
@app.post("/api/certs/{cert_id}/sign/web")
async def sign_cert_web(cert_id: int, request: Request = None):
async def sign_cert_web(cert_id: int, csrf_token: str = Form(""), request: Request = None):
user = get_user_from_cookie(request)
if not user:
return RedirectResponse("/login", status_code=302)
if not _verify_csrf_token(request, csrf_token):
raise HTTPException(403, "Invalid CSRF token")
conn = get_db()
row = conn.execute("SELECT * FROM certificates WHERE id = ?", (cert_id,)).fetchone()
if not row or row["status"] != "pending":
conn.close()
raise HTTPException(400, "Not found or already issued")
conn.close()
import html as html_lib
try:
result, err = build_leaf_cert(row["subject"], row["san"], 365)
if err:
return HTMLResponse(f'<span class="text-red-400">Issue failed: {err}</span>')
safe_err = html_lib.escape(str(err))
return HTMLResponse(f'<span class="text-red-400">Issue failed: {safe_err}</span>')
cf = f"/etc/ssl/ca/issued/cert-{result['serial']}.crt"
kf = f"/etc/ssl/ca/issued/cert-{result['serial']}.key"
open(cf, "w").write(result["cert_pem"])
@ -244,7 +288,8 @@ async def sign_cert_web(cert_id: int, request: Request = None):
conn2.close()
return HTMLResponse(f'<span class="text-green-400">Issued! <a href="/api/certs/{cert_id}/pem" class="underline">PEM</a> | <a href="/api/certs/{cert_id}/pfx" class="underline">PFX</a> | <a href="/certs" class="underline">Refresh</a></span>')
except Exception as ex:
return HTMLResponse(f'<span class="text-red-400">Issue failed: {str(ex)}</span>')
safe_ex = html_lib.escape(str(ex))
return HTMLResponse(f'<span class="text-red-400">Issue failed: {safe_ex}</span>')
@app.get("/logout")
@ -255,10 +300,13 @@ async def logout():
# --- Web API (cookie auth) ---
@app.post("/api/domains/web")
async def create_domain_web(name: str = Form(...), description: str = Form(""), request: Request = None):
async def create_domain_web(name: str = Form(...), description: str = Form(""),
csrf_token: str = Form(""), request: Request = None):
user = get_user_from_cookie(request)
if not user:
return RedirectResponse("/login", status_code=302)
if not _verify_csrf_token(request, csrf_token):
raise HTTPException(403, "Invalid CSRF token")
conn = get_db()
cur = conn.cursor()
cur.execute("INSERT INTO domains (name, description, created_by) VALUES (?,?,?)", (name, description, 1))
@ -267,13 +315,16 @@ async def create_domain_web(name: str = Form(...), description: str = Form(""),
return HTMLResponse('<span class="text-green-400">Domain registered! <a href="/domains" class="underline">Refresh</a></span>')
@app.post("/api/certs/web/request")
async def request_cert_web(cn: str = Form(...), sans: str = Form(""), days: int = Form(365), domain_id: int = Form(0), request: Request = None):
async def request_cert_web(cn: str = Form(...), sans: str = Form(""), days: int = Form(365),
domain_id: int = Form(0), csrf_token: str = Form(""),
request: Request = None):
user = get_user_from_cookie(request)
if not user:
return RedirectResponse("/login", status_code=302)
if not _verify_csrf_token(request, csrf_token):
raise HTTPException(403, "Invalid CSRF token")
conn = get_db()
cur = conn.cursor()
# Look up domain by CN if domain_id not provided
if domain_id == 0:
cur.execute("SELECT id FROM domains WHERE name=?", (cn,))
row = cur.fetchone()

59
tests/test_auth.py Normal file
View File

@ -0,0 +1,59 @@
import os
import sys
import unittest
from unittest import mock
from httpx import AsyncClient, ASGITransport
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "api"))
class TestAuthSecurity(unittest.TestCase):
def test_rate_limiting_exists(self):
import main
self.assertTrue(hasattr(main, '_check_rate_limit'))
self.assertEqual(main._LOGIN_MAX_ATTEMPTS, 5)
self.assertEqual(main._LOGIN_WINDOW_SECONDS, 900)
def test_csrf_token_generation(self):
import main
t1 = main._generate_csrf_token()
t2 = main._generate_csrf_token()
self.assertNotEqual(t1, t2)
self.assertEqual(len(t1), 64)
def test_csrf_verify_rejects_empty(self):
import main
req = mock.MagicMock()
req.cookies.get.return_value = None
self.assertFalse(main._verify_csrf_token(req, "any-token"))
def test_csrf_verify_rejects_mismatch(self):
import main
req = mock.MagicMock()
req.cookies.get.return_value = "stored-token"
self.assertFalse(main._verify_csrf_token(req, "different-token"))
def test_csrf_verify_accepts_match(self):
import main
req = mock.MagicMock()
req.cookies.get.return_value = "matching-token"
self.assertTrue(main._verify_csrf_token(req, "matching-token"))
def test_no_duplicate_get_user_from_cookie(self):
import main, inspect
sources = inspect.getsourcelines(main)[0]
count = sum(1 for line in sources if 'def get_user_from_cookie' in line)
self.assertEqual(count, 1, "get_user_from_cookie must be defined exactly once")
def test_xss_escaped_in_sign_error(self):
"""Check that HTML in error messages gets escaped."""
import html as h
err = '<script>alert("xss")</script>'
escaped = h.escape(err)
self.assertNotIn("<script>", escaped)
self.assertIn("&lt;script&gt;", escaped)
if __name__ == "__main__":
unittest.main()