Update CSRF middleware example

Commit 4a5e20a · patx · 2026-07-26T22:21:32-04:00

Changeset
4a5e20a935403edbf6f20244da65474d6b3e7db4
Parents
0493dcd5006d3a612a47d6e3ea01f6472c1e4fde

View source at this commit

Comments

No comments yet.

Log in to comment

Diff

diff --git a/examples/middleware/csrf.py b/examples/middleware/csrf.py
index 698d8a0..ec398ea 100644
--- a/examples/middleware/csrf.py
+++ b/examples/middleware/csrf.py
@@ -1,139 +1,173 @@
-from typing import Optional, Dict, List, Tuple, Any
-from urllib.parse import urlparse
-import uuid
-from itsdangerous import URLSafeTimedSerializer, BadSignature, SignatureExpired
+"""Synchronizer-token CSRF protection for a MicroPie application.
+
+Run this example from its directory with::
+
+    pip install "micropie[all]"
+    uvicorn app:app --reload
+
+Then open http://127.0.0.1:8000. The built-in in-memory session backend is
+convenient for this single-process demo. A real deployment should use a
+shared, persistent ``SessionBackend`` when running more than one process.
+"""
+
+from __future__ import annotations
+
+import html
+import re
+import secrets
+from collections.abc import Iterable
+from typing import Any
+from urllib.parse import urlsplit
+
 from micropie import App, HttpMiddleware, Request
 
 
 class CSRFMiddleware(HttpMiddleware):
-    """
-    MicroPie-ready CSRF middleware using itsdangerous + session binding.
-
-    - Verifies on POST/PUT/PATCH/DELETE
-    - Accepts token from body (form or JSON) or 'X-CSRF-Token' header
-    - For multipart/form-data, strongly prefer header (parser may still be streaming)
-    - Token payload includes the session_id (when present) to bind token to that session
-    - Emits 'X-CSRF-Token' **only** when we create/rotate one during this request
-    - Supports exempt_paths (e.g. webhook endpoints like /sms_order)
+    """Protect unsafe HTTP methods with a session-bound CSRF token.
+
+    A token is created during the first safe request and stored in the
+    server-side session. Clients must echo it in ``X-CSRF-Token`` or in a
+    ``csrf_token`` form/JSON field on every unsafe request.
+
+    Multipart requests must use the header because MicroPie streams their
+    fields while ``before_request`` is running. Only exempt an endpoint when
+    it has separate protection, such as signature verification for a webhook.
     """
 
+    SAFE_METHODS = frozenset({"GET", "HEAD", "OPTIONS", "TRACE"})
+    TOKEN_PATTERN = re.compile(r"[A-Za-z0-9_-]{43}\Z")
+
     def __init__(
         self,
-        app: App,
-        secret_key: str,
         *,
-        max_age: int = 8 * 3600,
-        trusted_origins: Optional[List[str]] = None,
+        trusted_origins: Iterable[str] | None = None,
         body_field: str = "csrf_token",
         header_name: str = "x-csrf-token",
-        require_header_for_multipart: bool = True,
-        exempt_paths: Optional[List[str]] = None,
-    ):
-        self.app = app
-        self.serializer = URLSafeTimedSerializer(secret_key, salt="csrf-token")
-        self.max_age = max_age
-        self.trusted = set(
-            trusted_origins or []
-        )  # e.g. ["https://gardenfresh.vegy.app"]
+        exempt_paths: Iterable[str] = (),
+    ) -> None:
         self.body_field = body_field
         self.header_name = header_name.lower()
-        self.require_header_for_multipart = require_header_for_multipart
-        self.exempt_paths = set(exempt_paths or [])
+        self.exempt_paths = frozenset(exempt_paths)
+        self.trusted_origins = self._normalize_trusted_origins(trusted_origins)
 
-    # ---------- helpers ----------
+    @classmethod
+    def _normalize_trusted_origins(
+        cls, values: Iterable[str] | None
+    ) -> frozenset[str]:
+        if values is None:
+            return frozenset()
 
-    def _is_mutating(self, method: str) -> bool:
-        return method in ("POST", "PUT", "PATCH", "DELETE")
+        origins = set()
+        for value in values:
+            origin = cls._origin(value)
+            if origin is None:
+                raise ValueError(f"Invalid trusted origin: {value!r}")
+            origins.add(origin)
+        return frozenset(origins)
 
-    def _is_multipart(self, ct: str) -> bool:
-        return "multipart/form-data" in (ct or "")
+    @staticmethod
+    def _origin(value: str | None) -> str | None:
+        """Return a normalized HTTP origin from an Origin or Referer value."""
+        if not value:
+            return None
 
-    def _origin_ok(self, headers: Dict[str, str]) -> bool:
-        if not self.trusted:
-            return True
-        origin = headers.get("origin")
-        referer = headers.get("referer")
-        for hdr in (origin, referer):
-            if not hdr:
-                continue
-            try:
-                p = urlparse(hdr)
-                base = f"{p.scheme}://{p.netloc}"
-                if base in self.trusted:
-                    return True
-            except Exception:
-                pass
-        return False
-
-    def _get_session_id(self, request: Request) -> Optional[str]:
-        cookie = request.headers.get("cookie", "")
-        if "session_id=" in cookie:
-            return cookie.split("session_id=", 1)[1].split(";", 1)[0].strip() or None
-        return None
+        try:
+            parsed = urlsplit(value)
+            scheme = parsed.scheme.lower()
+            hostname = parsed.hostname
+            port = parsed.port
+        except (TypeError, ValueError):
+            return None
 
-    def _issue_token(self, session_id: Optional[str]) -> str:
-        payload = {"nonce": str(uuid.uuid4())}
-        if session_id:
-            payload["sid"] = session_id
-        return self.serializer.dumps(payload)
+        if (
+            scheme not in {"http", "https"}
+            or not hostname
+            or parsed.username is not None
+            or parsed.password is not None
+        ):
+            return None
 
-    def _extract_submitted_token(self, request: Request) -> Optional[str]:
-        ct = request.headers.get("content-type", "")
-        if self._is_multipart(ct) and self.require_header_for_multipart:
-            token = request.headers.get(self.header_name)
-            if token:
-                return token
+        host = hostname.lower()
+        if ":" in host:  # Preserve brackets around IPv6 hosts.
+            host = f"[{host}]"
 
-        lst = request.body_params.get(self.body_field)
-        if lst and isinstance(lst, list) and lst:
-            return lst[0]
+        default_port = 80 if scheme == "http" else 443
+        authority = host if port in (None, default_port) else f"{host}:{port}"
+        return f"{scheme}://{authority}"
 
-        j = request.json()
-        if isinstance(j, dict):
-            tok = j.get(self.body_field)
-            if isinstance(tok, str):
-                return tok
+    def _origin_is_trusted(self, request: Request) -> bool:
+        if not self.trusted_origins:
+            return True
 
-        return request.headers.get(self.header_name)
+        # Prefer Origin. Only use Referer when Origin is absent; an invalid
+        # Origin must not be rescued by a valid-looking Referer.
+        origin = request.headers.get("origin")
+        if origin is not None:
+            return self._origin(origin) in self.trusted_origins
 
-    # ---------- middleware hooks ----------
+        referer = request.headers.get("referer")
+        return self._origin(referer) in self.trusted_origins
 
-    async def before_request(self, request: Request) -> Optional[Dict]:
-        path = request.scope.get("path", "")
+    def _submitted_token(self, request: Request) -> str | None:
+        header_token = request.headers.get(self.header_name)
+        if header_token:
+            return header_token
 
-        # Exempt specific paths (e.g. /sms_order webhook)
-        if path in self.exempt_paths:
+        content_type = request.headers.get("content-type", "").lower()
+        if "multipart/form-data" in content_type:
+            # Multipart parsing is concurrent in MicroPie. Requiring the
+            # header avoids waiting for a field that may never arrive.
             return None
 
-        if "csrf_token" not in request.session:
-            sid = self._get_session_id(request)
-            request.session["csrf_token"] = self._issue_token(sid)
-            setattr(request, "_csrf_emit", request.session["csrf_token"])
+        form_token = request.form(self.body_field)
+        if form_token:
+            return form_token
 
-        if not self._is_mutating(request.method):
+        json_token = request.json(self.body_field)
+        return json_token if isinstance(json_token, str) else None
+
+    @classmethod
+    def _valid_session_token(cls, token: Any) -> bool:
+        return (
+            isinstance(token, str)
+            and cls.TOKEN_PATTERN.fullmatch(token) is not None
+        )
+
+    async def before_request(self, request: Request) -> dict[str, Any] | None:
+        if request.scope.get("path", "") in self.exempt_paths:
             return None
 
-        if not self._origin_ok(request.headers):
-            return {"status_code": 403, "body": "Forbidden: invalid origin/referer"}
+        method = (request.method or "").upper()
+        session_token = request.session.get("csrf_token")
 
-        submitted = self._extract_submitted_token(request)
-        if not submitted:
-            return {"status_code": 403, "body": "Missing CSRF token"}
+        if method in self.SAFE_METHODS:
+            if not self._valid_session_token(session_token):
+                session_token = secrets.token_urlsafe(32)
+                request.session["csrf_token"] = session_token
+                request._csrf_token_was_created = session_token
+            return None
 
-        try:
-            data = self.serializer.loads(submitted, max_age=self.max_age)
-        except (BadSignature, SignatureExpired):
-            return {"status_code": 403, "body": "Invalid or expired CSRF token"}
+        if not self._origin_is_trusted(request):
+            return {
+                "status_code": 403,
+                "body": "403 Forbidden: untrusted Origin or Referer",
+            }
 
-        sid = self._get_session_id(request)
-        token_sid = data.get("sid")
-        if token_sid is not None and (sid != token_sid):
-            return {"status_code": 403, "body": "Invalid CSRF token for this session"}
+        submitted_token = self._submitted_token(request)
+        if not submitted_token:
+            return {
+                "status_code": 403,
+                "body": "403 Forbidden: missing CSRF token",
+            }
 
-        sid_now = self._get_session_id(request)
-        new_token = self._issue_token(sid_now)
-        request.session["csrf_token"] = new_token
-        setattr(request, "_csrf_emit", new_token)
+        if (
+            not self._valid_session_token(session_token)
+            or not secrets.compare_digest(submitted_token, session_token)
+        ):
+            return {
+                "status_code": 403,
+                "body": "403 Forbidden: invalid CSRF token",
+            }
 
         return None
 
@@ -142,34 +176,125 @@ class CSRFMiddleware(HttpMiddleware):
         request: Request,
         status_code: int,
         response_body: Any,
-        extra_headers: List[Tuple[str, str]],
-    ) -> Optional[Dict]:
-        to_emit = getattr(request, "_csrf_emit", None)
-        if to_emit:
-            extra_headers.append(("X-CSRF-Token", to_emit))
+        extra_headers: list[tuple[str, str]],
+    ) -> dict[str, Any] | None:
+        # This is useful to clients that first obtain a token with a safe
+        # request. HTML templates can instead read request.session directly.
+        token = getattr(request, "_csrf_token_was_created", None)
+        if token:
+            extra_headers.append(("X-CSRF-Token", token))
         return None
 
 
-class Root(App):
+PAGE = """<!doctype html>
+<html lang="en">
+<head>
+  <meta charset="utf-8">
+  <meta name="viewport" content="width=device-width, initial-scale=1">
+  <title>MicroPie CSRF example</title>
+  <style>
+    body { font: 16px/1.5 system-ui, sans-serif; max-width: 48rem; margin: 3rem auto; padding: 0 1rem; }
+    section { border: 1px solid #ccc; border-radius: .5rem; margin: 1rem 0; padding: 1rem; }
+    input, button { font: inherit; padding: .45rem .65rem; }
+    code, pre { background: #f4f4f4; border-radius: .25rem; padding: .15rem .3rem; }
+    pre { min-height: 1.5rem; overflow-wrap: anywhere; white-space: pre-wrap; }
+  </style>
+</head>
+<body>
+  <h1>MicroPie CSRF protection</h1>
+  <p>
+    This page receives a random token that is stable for the session and never
+    used as the session cookie itself. Unsafe requests must echo it back.
+  </p>
+
+  <section>
+    <h2>HTML form</h2>
+    <p>The synchronizer token travels in a hidden form field.</p>
+    <form action="/submit_form" method="post">
+      <input type="hidden" name="csrf_token" value="__CSRF_TOKEN__">
+      <label>Message <input name="message" value="Sent by a protected form"></label>
+      <button type="submit">Submit form</button>
+    </form>
+  </section>
+
+  <section>
+    <h2>JSON request</h2>
+    <p>JavaScript sends the same token in the <code>X-CSRF-Token</code> header.</p>
+    <button id="protected-request">Send protected request</button>
+    <button id="unprotected-request">Try without a token</button>
+    <pre id="result" aria-live="polite"></pre>
+  </section>
+
+  <script>
+    const csrfToken = "__CSRF_TOKEN__";
+    const result = document.querySelector("#result");
+
+    async function send(includeToken) {
+      const headers = {"Content-Type": "application/json"};
+      if (includeToken) headers["X-CSRF-Token"] = csrfToken;
+
+      const response = await fetch("/submit_json", {
+        method: "POST",
+        headers,
+        body: JSON.stringify({message: "Sent by a protected JSON request"})
+      });
+      result.textContent = `${response.status} ${await response.text()}`;
+    }
+
+    document.querySelector("#protected-request").onclick = () => send(true);
+    document.querySelector("#unprotected-request").onclick = () => send(false);
+  </script>
+</body>
+</html>
+"""
+
+
+class CSRFExample(App):
     async def index(self):
-        csrf_token = self.request.session.get("csrf_token", "")
-        return f"""<form method="POST" action="/submit">
-            <input type="hidden" name="csrf_token" value="{csrf_token}">
-            <input type="text" name="name">
-            <button type="submit">Submit</button>
-            </form>"""
+        if self.request.method not in {"GET", "HEAD"}:
+            return 405, "405 Method Not Allowed"
+
+        token = html.escape(self.request.session["csrf_token"], quote=True)
+        return (
+            200,
+            PAGE.replace("__CSRF_TOKEN__", token),
+            [("Cache-Control", "no-store")],
+        )
+
+    async def submit_form(self):
+        if self.request.method != "POST":
+            return 405, {"error": "Method Not Allowed"}
+
+        return {
+            "ok": True,
+            "message": self.request.form("message", ""),
+            "token_source": "form field",
+        }
+
+    async def submit_json(self):
+        if self.request.method != "POST":
+            return 405, {"error": "Method Not Allowed"}
 
-    async def submit(self):
-        if self.request.method == "POST":
-            name = self.request.form("name", "World")
-            return f"Hello {name}"
+        token_source = (
+            "X-CSRF-Token header"
+            if self.request.headers.get("x-csrf-token")
+            else "JSON field"
+        )
+        return {
+            "ok": True,
+            "message": self.request.json("message", ""),
+            "token_source": token_source,
+        }
 
 
-app = Root()
+# Local HTTP needs Secure=False so the browser will return the session cookie.
+# Keep MicroPie's default Secure=True when deploying this application over HTTPS.
+app = CSRFExample(session_cookie_secure=False)
 app.middlewares.append(
     CSRFMiddleware(
-        app=app,
-        secret_key="my-secret-key",
-        exempt_paths=["/sms_order"],
+        trusted_origins={
+            "http://127.0.0.1:8000",
+            "http://localhost:8000",
+        }
     )
 )