loko/streetup/api/views.py
2026-07-22 14:48:40 +02:00

534 lines
No EOL
19 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

from __future__ import annotations
import os
from typing import Any, Optional
import re
from html import parser as html_parser
import httpx
from urllib.parse import urlsplit, urlunsplit, parse_qsl, urlencode, unquote
from rest_framework import generics, permissions, serializers, status
from rest_framework.permissions import IsAuthenticated, AllowAny
from rest_framework.response import Response
from rest_framework.views import APIView
from django.http import HttpResponse
from .models import PersonalAccessToken
from .auth import HybridTokenAuthentication, collections_for, mint_jwt
# === CONFIG ===
PYGEOAPI_BASE_URL = os.environ.get("PYGEOAPI_BASE_URL")
# === VUES TOKENS ===
class IssueJWTView(APIView):
permission_classes = [IsAuthenticated]
authentication_classes = [HybridTokenAuthentication]
def post(self, request, *args, **kwargs):
token = mint_jwt(request.user, lifetime_minutes=int(request.data.get("lifetime_minutes", 20)))
return Response({"access_token": token, "token_type": "Bearer"})
class SharePATView(APIView):
permission_classes = [IsAuthenticated]
authentication_classes = [HybridTokenAuthentication]
def post(self, request, *args, **kwargs):
req_cols = request.data.get("collections", [])
ttl_h = int(request.data.get("ttl_hours", 720))
allowed = sorted(set(collections_for(request.user)) & set(req_cols)) if req_cols else collections_for(request.user)
raw, obj = PersonalAccessToken.issue(
user=request.user,
collections=allowed,
ttl_hours=ttl_h,
name=request.data.get("name", "share link"),
)
base = request.build_absolute_uri("/api/geo/").rstrip("/")
url_path = f"{base}/{raw}/" # <-- forme /api/geo/pat_xxx/
url_query = request.build_absolute_uri("/api/geo/").rstrip("/") + f"?token={raw}"
return Response({
"token": raw,
"url_path": url_path,
"url_query": url_query,
"expires_at": obj.expires_at,
})
# === PROXY GEO -> PYGEOAPI ===
class GeoProxyView(APIView):
"""
Passerelle /api/geo/... -> pygeoapi.
- Auth: HybridTokenAuthentication (JWT, PAT, ou Session)
- AutZ: filtrage par collection si l'URL cible le détail d'une collection
"""
permission_classes = [IsAuthenticated]
authentication_classes = [HybridTokenAuthentication]
def _strip_prefix(self, full_path: str) -> str:
# /api/geo/... -> /collections/... (chemin attendu par pygeoapi)
path = full_path
if path.startswith("/api"):
path = path[len("/api"):]
if path.startswith("/geo"):
path = path[len("/geo"):]
if not path.startswith("/"):
path = "/" + path
# Si le 1er segment est un PAT (pat_...), on l'enlève du chemin
parts = [p for p in path.split("/") if p]
if parts and parts[0].startswith("pat_"):
parts = parts[1:]
path = "/" + "/".join(parts)
return path
def _check_collection(self, url_or_path: str, allowed: list[str] | None) -> tuple[bool, str | None]:
"""
Retourne:
(True, None) -> accès OK (pas une URL de collection spécifique)
(True, cid) -> accès OK à la collection cid
(False, cid) -> accès refusé à la collection cid
"""
# 1) Extraire le path pur (ignore ?f=json, #hash, etc.)
s = urlsplit(url_or_path)
path = s.path or url_or_path
path = unquote(path)
# 2) Trouver /collections/ n'importe où dans le path
needle = "/collections/"
i = path.lower().find(needle)
if i == -1:
# Pas une URL /collections/{cid} -> OK
return True, None
# 3) Récupérer le segment cid (ce qui suit immédiatement /collections/)
after = path[i + len(needle):]
if not after:
# /collections (listing) -> OK
return True, None
cid = after.split("/", 1)[0]
if not cid:
# /collections/ sans id -> traiter comme listing
return True, None
# 4) Vérifier l'autorisation (si filtrage actif)
if allowed is None:
return True, cid
return (cid in allowed, cid)
def request_upstream(self, request, upstream: str, *, path: str, allowed: list[str] | None = None):
headers = {k: v for k, v in request.headers.items() if k.lower() not in {"host", "content-length"}}
try:
r = httpx.request(
method=request.method,
url=upstream,
headers=headers,
params=request.GET,
content=(getattr(request, "body", None) if request.method in {"POST", "PUT", "PATCH"} else None),
timeout=15.0,
)
except httpx.RequestError as exc:
return Response({"detail": f"Upstream error: {exc}"}, status=status.HTTP_502_BAD_GATEWAY)
# --- NEW: renvoi correct selon le Content-Type ---
content_type = (r.headers.get("content-type") or "").split(";")[0].strip().lower()
pat = _extract_pat_from_request(request)
mode = _choose_pat_mode(request)
allowed_set: set[str] | None
if allowed is None:
allowed_set = None
else:
allowed_set = {c for c in allowed if c}
path_no_query = path.split("?", 1)[0]
path_for_filter = (path_no_query.rstrip("/") or "/") if path_no_query else "/"
# Headers Location (redirections)
def _copy_headers(resp):
# Ne pas copier Content-Encoding pour éviter la double compression (Django GZipMiddleware + pygeoapi)
for h in ("Content-Disposition", "Cache-Control", "ETag", "Last-Modified"):
if h in r.headers:
resp[h] = r.headers[h]
loc = r.headers.get("Location")
if loc and pat:
resp["Location"] = _rewrite_api_geo_everywhere(loc, pat, mode=mode)
return resp
# JSON
try:
data = r.json()
if allowed_set is not None and path_for_filter == "/collections" and isinstance(data, dict):
data = _filter_collections_json(data, allowed_set)
def _walk(o):
if isinstance(o, dict):
return {k: _walk(v) for k, v in o.items()}
if isinstance(o, list):
return [_walk(v) for v in o]
if isinstance(o, str) and "/api/geo" in o:
return _rewrite_api_geo_everywhere(o, pat, mode=mode)
return o
data = _walk(data)
resp = Response(data, status=r.status_code)
return _copy_headers(resp)
except ValueError:
pass
# HTML / texte
if content_type in {"text/html", "text/plain", "application/xml", "text/xml"}:
text = r.text
if allowed_set is not None and path_for_filter == "/collections":
text = _filter_collections_html(text, allowed_set)
if pat and "/api/geo" in text:
text = _rewrite_api_geo_everywhere(text, pat, mode=mode)
resp = HttpResponse(text, status=r.status_code, content_type=r.headers.get("content-type"))
return _copy_headers(resp)
# Binaire inchangé
resp = HttpResponse(r.content, status=r.status_code, content_type=r.headers.get("content-type"))
return _copy_headers(resp)
def get(self, request, *args, **kwargs):
path = self._strip_prefix(request.get_full_path())
ok, cid = self._check_collection(path, getattr(request, "collections", []))
if not ok:
return Response({"detail": f"Forbidden collection: {cid}"}, status=status.HTTP_403_FORBIDDEN)
upstream = f"{PYGEOAPI_BASE_URL}{path}"
return self.request_upstream(request, upstream, path=path, allowed=getattr(request, "collections", None))
# Optionnel: POST/PUT/PATCH/DELETE si tu as des opérations d'écriture pygeoapi
def post(self, request, *args, **kwargs):
path = self._strip_prefix(request.get_full_path())
ok, cid = self._check_collection(path, getattr(request, "collections", []))
if not ok:
return Response({"detail": f"Forbidden collection: {cid}"}, status=status.HTTP_403_FORBIDDEN)
upstream = f"{PYGEOAPI_BASE_URL}{path}"
return self.request_upstream(request, upstream, path=path, allowed=getattr(request, "collections", None))
# === DRF pur ===
from interventions.models import Intervention
class InterventionSerializer(serializers.ModelSerializer):
class Meta:
model = Intervention
fields = ("code","title","status","creation_time")
class InterventionListCreateView(generics.ListCreateAPIView):
"""
GET /api/interventions/?status=to_be_processed&status=in_progress
POST /api/interventions/
"""
permission_classes = [IsAuthenticated]
authentication_classes = [HybridTokenAuthentication]
# TODO: remplacer par le queryset réel:
queryset = []
serializer_class = InterventionSerializer
def get_queryset(self):
qs = super().get_queryset()
statuses = self.request.query_params.getlist("status")
if statuses:
# TODO: adapter au champ status réel
# qs = qs.filter(status__in=statuses)
pass
return qs
class InterventionRetrieveUpdateDestroyView(generics.RetrieveUpdateDestroyAPIView):
permission_classes = [IsAuthenticated]
authentication_classes = [HybridTokenAuthentication]
# TODO: remplacer par le queryset réel:
queryset = []
serializer_class = InterventionSerializer
def _extract_pat_from_request(request):
auth = request.headers.get("Authorization", "")
if auth.lower().startswith("bearer pat_"):
return auth.split(None, 1)[1].strip()
tok = request.GET.get("token")
if tok and tok.startswith("pat_"):
return tok
path = request.path or ""
try:
after_geo = path.split("/geo/", 1)[1]
seg0 = after_geo.split("/", 1)[0]
if seg0.startswith("pat_"):
return seg0
except Exception:
pass
return None
def _choose_pat_mode(request) -> str:
# Forçage manuel
forced = request.GET.get("_pat_mode")
if forced in {"path", "query"}:
return forced
# Si PAT en query -> réécritures en query, sinon path
if request.GET.get("token", "").startswith("pat_"):
return "query"
return "path"
# Capture TOUTE l’URL /api/geo… (absolue ou relative), y compris ?query et #fragment,
# mais en évitant le cas où un /pat_… est déjà juste après /api/geo/
_PATTERN_API_GEO_URL = re.compile(
r'((?:(?:https?://|//)[^/\'"\s>]+)?' # schéma + host éventuel
r'/api/geo' # base
r'(?:/(?!pat_[^/]+)[^\'"\s>]*)?' # suite de chemin (sans /pat_… juste après)
r'(?:\?[^#\'"\s>]*)?' # query facultative
r'(?:#[^\'"\s>]*)?' # fragment facultatif
r')'
)
def _rewrite_api_geo_everywhere(text: str, pat: str, mode: str = "path") -> str:
"""
Réécrit toutes les URLs /api/geo… :
- mode="path" -> injecte /pat_xxx/ après /api/geo
- mode="query" -> ajoute ?token=pat_xxx (ou &token=…)
Ne double pas si déjà présent (/pat_… ou token=…).
"""
if not pat or not text:
return text
def _repl(m):
url = m.group(1)
parts = urlsplit(url)
# parts.path commence par /api/geo…
after = parts.path.split("/api/geo", 1)[1] # "", "/...", etc.
if mode == "path":
# si déjà /pat_… juste après /api/geo, ne rien faire
if after.startswith("/") and after[1:].startswith("pat_"):
return url
new_path = parts.path.replace("/api/geo", f"/api/geo/{pat}", 1)
return urlunsplit((parts.scheme, parts.netloc, new_path, parts.query, parts.fragment))
# mode "query"
# si token déjà présent, ne rien faire
q = parse_qsl(parts.query or "", keep_blank_values=True)
if any(k == "token" for k, _ in q):
return url
q.append(("token", pat))
new_query = urlencode(q, doseq=True)
return urlunsplit((parts.scheme, parts.netloc, parts.path, new_query, parts.fragment))
return _PATTERN_API_GEO_URL.sub(_repl, text)
def _filter_collections_json(data: dict[str, Any], allowed: set[str]) -> dict[str, Any]:
"""Filtre la réponse JSON /collections selon les collections autorisées."""
collections = data.get("collections")
if isinstance(collections, list):
filtered = []
for item in collections:
if not isinstance(item, dict):
continue
cid = item.get("id")
if cid in allowed:
filtered.append(item)
data["collections"] = filtered
links = data.get("links")
if isinstance(links, list):
filtered_links = []
for link in links:
if not isinstance(link, dict):
filtered_links.append(link)
continue
href = link.get("href", "")
cid = _extract_collection_id_from_href(str(href))
if cid is None or cid in allowed:
filtered_links.append(link)
data["links"] = filtered_links
return data
def _extract_collection_id_from_href(href: str) -> str | None:
if not href:
return None
ref = href.strip()
if not ref:
return None
# Normalise la partie chemin
path = ref
if "?" in path:
path = path.split("?", 1)[0]
if "#" in path:
path = path.split("#", 1)[0]
if "collections/" in path:
segment = path.split("collections/", 1)[1]
else:
segment = path
segment = segment.lstrip("./")
while segment.startswith("../"):
segment = segment[3:]
segment = segment.lstrip("/")
if not segment:
return None
cid = segment.split("/", 1)[0]
if not cid:
return None
# on évite de considérer des liens génériques (items, schema, etc.)
if cid in {"items", "query", "schema", "coverage", "tiles", "map"}:
return None
return cid
# -------------------------------------------------------------------
# HTML Filtering
# -------------------------------------------------------------------
_COLLECTION_HREF_RE = re.compile(r"/collections/([^/?#]+)", re.I)
class _CollectionHTMLFilter(html_parser.HTMLParser):
CONTAINER_TAGS = {"tr", "li", "article", "div", "section"}
# HTML void elements (no end tag is emitted by the parser)
VOID_TAGS = {
"area", "base", "br", "col", "embed", "hr", "img", "input",
"link", "meta", "param", "source", "track", "wbr"
}
def __init__(self, allowed: Optional[set[str]]):
super().__init__(convert_charrefs=False)
self.allowed = allowed
self.stack = [{"tag": None, "buffer": [], "ids": set(), "container": False}]
def _push(self, tag: Optional[str], container: bool = False):
self.stack.append({"tag": tag, "buffer": [], "ids": set(), "container": container})
def _pop(self):
return self.stack.pop()
def _append(self, s: str):
self.stack[-1]["buffer"].append(s)
def get_html(self) -> str:
return "".join(self.stack[0]["buffer"])
def handle_decl(self, decl: str):
self._append(f"<!{decl}>")
def handle_starttag(self, tag: str, attrs: list[tuple[str, Optional[str]]]):
t = tag.lower()
# Build attributes
if attrs:
a = " ".join(
f'{k}="{(v or "").replace("\"", "&quot;")}"' if v is not None else f"{k}"
for k, v in attrs
)
open_tag = f"<{tag} {a}>"
else:
open_tag = f"<{tag}>"
# Void elements: append and DO NOT push (no end tag will arrive)
if t in self.VOID_TAGS:
self._append(open_tag)
return
# Non-void: push a new frame and append the start tag inside that frame
container = t in self.CONTAINER_TAGS
self._push(tag, container)
self._append(open_tag)
# Detect collection IDs on <a href=".../collections/{id}">
if t == "a":
href = None
for k, v in attrs:
if k.lower() == "href":
href = v or ""
break
if href:
m = _COLLECTION_HREF_RE.search(href)
if m:
cid = m.group(1)
# mark current and nearest container
self.stack[-1]["ids"].add(cid)
for frame in reversed(self.stack):
if frame["container"]:
frame["ids"].add(cid)
break
def handle_endtag(self, tag: str):
t = tag.lower()
# If the stack top isn't the same tag, avoid popping (robustness)
if not self.stack or self.stack[-1]["tag"] is None or self.stack[-1]["tag"].lower() != t:
# Still emit the closing tag for best-effort rendering
self._append(f"</{tag}>")
return
# Close normally
self._append(f"</{tag}>")
frame = self._pop()
if frame["container"] and self.allowed is not None:
ids = frame["ids"]
keep = (not ids) or any(cid in self.allowed for cid in ids)
if not keep:
# Drop the entire container block
return
# Keep the content (opening+children+closing) by appending to parent
self._append("".join(frame["buffer"]))
def handle_startendtag(self, tag: str, attrs: list[tuple[str, Optional[str]]]):
# Treat as a void-like tag
if attrs:
a = " ".join(
f'{k}="{(v or "").replace("\"", "&quot;")}"' if v is not None else f"{k}"
for k, v in attrs
)
self._append(f"<{tag} {a} />")
else:
self._append(f"<{tag} />")
def handle_data(self, data: str):
self._append(data)
def handle_comment(self, data: str):
self._append(f"<!--{data}-->")
def handle_entityref(self, name: str):
self._append(f"&{name};")
def handle_charref(self, name: str):
self._append(f"&#{name};")
def _filter_collections_html(html: str, allowed: Optional[set[str]]) -> str:
if not html:
return "<!DOCTYPE html>\n"
if allowed is None:
out = html
if "<!DOCTYPE" not in out[:300].upper():
out = "<!DOCTYPE html>\n" + out
return out
parser = _CollectionHTMLFilter(allowed)
parser.feed(html)
parser.close()
print(parser.stack)
out = parser.get_html()
if "<!DOCTYPE" not in out[:300].upper():
out = "<!DOCTYPE html>\n" + out
return out