diff --git a/api/main.py b/api/main.py index 4bcca57..1daa0e0 100644 --- a/api/main.py +++ b/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'Issue failed: {err}') + safe_err = html_lib.escape(str(err)) + return HTMLResponse(f'Issue failed: {safe_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"]) @@ -244,7 +288,8 @@ async def sign_cert_web(cert_id: int, request: Request = None): conn2.close() return HTMLResponse(f'Issued! PEM | PFX | Refresh') except Exception as ex: - return HTMLResponse(f'Issue failed: {str(ex)}') + safe_ex = html_lib.escape(str(ex)) + return HTMLResponse(f'Issue failed: {safe_ex}') @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('Domain registered! Refresh') @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() diff --git a/tests/test_auth.py b/tests/test_auth.py new file mode 100644 index 0000000..c2220c1 --- /dev/null +++ b/tests/test_auth.py @@ -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 = '' + escaped = h.escape(err) + self.assertNotIn("