CoolFace
Apppublic

Rhinox13/chatapi

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
csrf.py99 linesDownload Raw Back to services
1from __future__ import annotations
2
3from urllib.parse import urlsplit
4
5from flask import Flask, jsonify, request, session
6
7SAFE_METHODS = {"GET", "HEAD", "OPTIONS", "TRACE"}
8
9
10def _first_header_value(value: str) -> str:
11    return value.split(",", 1)[0].strip()
12
13
14def _normalize_origin(value: str) -> str | None:
15    raw = value.strip()
16    if not raw or raw == "null":
17        return None
18    try:
19        parsed = urlsplit(raw)
20        if not parsed.scheme or not parsed.netloc:
21            return None
22        scheme = parsed.scheme.lower()
23        if scheme not in {"http", "https"}:
24            return None
25        hostname = parsed.hostname
26        if not hostname:
27            return None
28        port = parsed.port
29    except ValueError:
30        return None
31
32    if port is None:
33        port = 443 if scheme == "https" else 80
34    return f"{scheme}://{hostname.lower()}:{port}"
35
36
37def _request_origins() -> set[str]:
38    origins: set[str] = set()
39
40    forwarded_proto = _first_header_value(request.headers.get("X-Forwarded-Proto", ""))
41    scheme = forwarded_proto or request.scheme
42    current_origin = _normalize_origin(f"{scheme}://{request.host}")
43    if current_origin:
44        origins.add(current_origin)
45
46    host_origin = _normalize_origin(request.host_url)
47    if host_origin:
48        origins.add(host_origin)
49
50    return origins
51
52
53def _configured_origins(cors_origins: list[str] | tuple[str, ...]) -> set[str]:
54    origins: set[str] = set()
55    for raw_origin in cors_origins:
56        raw = str(raw_origin).strip()
57        if not raw or raw == "*":
58            continue
59        normalized = _normalize_origin(raw)
60        if normalized:
61            origins.add(normalized)
62    return origins
63
64
65def _source_is_allowed(value: str, allowed_origins: set[str]) -> bool:
66    normalized = _normalize_origin(value)
67    return normalized is not None and normalized in allowed_origins
68
69
70def register_csrf_protection(app: Flask, *, cors_origins: list[str] | tuple[str, ...]) -> None:
71    configured_origins = _configured_origins(cors_origins)
72
73    @app.before_request
74    def reject_cross_site_session_mutations():
75        if request.method in SAFE_METHODS:
76            return None
77        if not request.path.startswith("/api/"):
78            return None
79        allowed_origins = _request_origins() | configured_origins
80        origin = request.headers.get("Origin", "").strip()
81        if origin:
82            if _source_is_allowed(origin, allowed_origins):
83                return None
84            app.logger.warning("Rejected cross-site request from Origin %s", origin)
85            return jsonify({"error": "csrf_origin_mismatch"}), 403
86
87        referer = request.headers.get("Referer", "").strip()
88        if referer:
89            if _source_is_allowed(referer, allowed_origins):
90                return None
91            app.logger.warning("Rejected cross-site request from Referer %s", referer)
92            return jsonify({"error": "csrf_referer_mismatch"}), 403
93
94        if not session.get("user_id"):
95            return None
96
97        app.logger.warning("Rejected session request without trusted Origin or Referer")
98        return jsonify({"error": "csrf_origin_required"}), 403
99