Rhinox13/chatapi
0
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 