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:
parent
c759597ad0
commit
df82068aeb
83
api/main.py
83
api/main.py
@ -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
59
tests/test_auth.py
Normal 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("<script>", escaped)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Loading…
x
Reference in New Issue
Block a user