534 lines
No EOL
19 KiB
Python
534 lines
No EOL
19 KiB
Python
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("\"", """)}"' 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("\"", """)}"' 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 |