L'admin (vaults: ["*"]) voyait le home de chaque utilisateur dans sa barre latérale : "*" ouvrait tous les vaults, home-* compris. - backend/auth/middleware.py : check_vault_access exige un octroi explicite pour tout vault home-* (nouveau is_home_vault()). - Filtres « * » en dur remplacés par check_vault_access : dashboard, conflits, liens retour, favoris, abonnements push. - /api/search : search_vaults(is_allowed=…) filtre les bruts avant pagination (total et page restent justes). - backend/user_home.py : _grant n'écarte plus les comptes « * » — l'admin reçoit son propre home-admin (auto-réparé au démarrage). - Tests : test_user_home.py +2, assertion API inversée dans test_auth_api.py (admin ne voit plus home-alice).
285 lines
9.8 KiB
Python
285 lines
9.8 KiB
Python
# backend/push.py
|
|
# Push Notifications — Web Push API with VAPID authentication
|
|
# ROADMAP #67
|
|
|
|
import base64
|
|
import json
|
|
import logging
|
|
import os
|
|
from datetime import datetime, timezone
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
from fastapi import APIRouter, Depends, HTTPException
|
|
from pydantic import BaseModel, Field
|
|
|
|
from backend.auth.middleware import check_vault_access, require_auth
|
|
|
|
logger = logging.getLogger("obsigate.push")
|
|
|
|
router = APIRouter(prefix="/api/push", tags=["push"])
|
|
|
|
# VAPID keys storage
|
|
VAPID_KEYS_FILE = Path(os.environ.get("OBSIGATE_DATA_DIR", "data")) / "vapid_keys.json"
|
|
PUSH_SUBSCRIPTIONS_FILE = Path(os.environ.get("OBSIGATE_DATA_DIR", "data")) / "push_subscriptions.json"
|
|
|
|
# In-memory cache
|
|
_vapid_keys: dict[str, str] = {}
|
|
_push_subscriptions: list[dict[str, Any]] = []
|
|
|
|
|
|
def load_vapid_keys() -> dict[str, str]:
|
|
"""Load VAPID keys from file or generate new ones."""
|
|
global _vapid_keys
|
|
if _vapid_keys:
|
|
return _vapid_keys
|
|
|
|
try:
|
|
if VAPID_KEYS_FILE.exists():
|
|
with open(VAPID_KEYS_FILE, "r") as f:
|
|
_vapid_keys = json.load(f)
|
|
else:
|
|
# Generate new VAPID keys
|
|
from cryptography.hazmat.primitives import serialization
|
|
from cryptography.hazmat.primitives.asymmetric import ec
|
|
|
|
private_key = ec.generate_private_key(ec.SECP256R1())
|
|
public_key = private_key.public_key()
|
|
|
|
private_pem = private_key.private_bytes(
|
|
encoding=serialization.Encoding.PEM,
|
|
format=serialization.PrivateFormat.PKCS8,
|
|
encryption_algorithm=serialization.NoEncryption()
|
|
).decode('utf-8')
|
|
|
|
public_pem = public_key.public_bytes(
|
|
encoding=serialization.Encoding.PEM,
|
|
format=serialization.PublicFormat.SubjectPublicKeyInfo
|
|
).decode('utf-8')
|
|
|
|
# Convert to base64url for Web Push
|
|
public_numbers = public_key.public_numbers()
|
|
|
|
def int_to_base64url(n: int) -> str:
|
|
byte_len = (n.bit_length() + 7) // 8
|
|
return base64.urlsafe_b64encode(n.to_bytes(byte_len, 'big')).decode('utf-8').rstrip('=')
|
|
|
|
_vapid_keys = {
|
|
"private_key": private_pem,
|
|
"public_key": public_pem,
|
|
"public_key_base64url": int_to_base64url(public_numbers.x) + "." + int_to_base64url(public_numbers.y)
|
|
}
|
|
|
|
VAPID_KEYS_FILE.parent.mkdir(parents=True, exist_ok=True)
|
|
with open(VAPID_KEYS_FILE, "w") as f:
|
|
json.dump(_vapid_keys, f)
|
|
except Exception as e:
|
|
logger.error(f"Failed to load/generate VAPID keys: {e}")
|
|
# Fallback: return empty (will use demo keys if needed)
|
|
_vapid_keys = {}
|
|
|
|
return _vapid_keys
|
|
|
|
|
|
def get_vapid_public_key() -> str:
|
|
"""Get VAPID public key in base64url format for client."""
|
|
keys = load_vapid_keys()
|
|
return keys.get("public_key_base64url", "")
|
|
|
|
|
|
def load_push_subscriptions() -> list[dict[str, Any]]:
|
|
"""Load push subscriptions from file."""
|
|
global _push_subscriptions
|
|
if _push_subscriptions:
|
|
return _push_subscriptions
|
|
|
|
try:
|
|
if PUSH_SUBSCRIPTIONS_FILE.exists():
|
|
with open(PUSH_SUBSCRIPTIONS_FILE, "r") as f:
|
|
_push_subscriptions = json.load(f)
|
|
except Exception as e:
|
|
logger.error(f"Failed to load push subscriptions: {e}")
|
|
_push_subscriptions = []
|
|
return _push_subscriptions
|
|
|
|
|
|
def save_push_subscriptions() -> None:
|
|
"""Save push subscriptions to file."""
|
|
try:
|
|
PUSH_SUBSCRIPTIONS_FILE.parent.mkdir(parents=True, exist_ok=True)
|
|
with open(PUSH_SUBSCRIPTIONS_FILE, "w") as f:
|
|
json.dump(_push_subscriptions, f, indent=2)
|
|
except Exception as e:
|
|
logger.error(f"Failed to save push subscriptions: {e}")
|
|
|
|
|
|
# ── Models ──────────────────────────────────────────────────────────────────
|
|
|
|
class PushSubscription(BaseModel):
|
|
"""Push subscription from client (matches Push API)."""
|
|
endpoint: str
|
|
keys: dict[str, str] # { p256dh, auth }
|
|
|
|
|
|
class SubscribeRequest(BaseModel):
|
|
"""Request to subscribe to push notifications."""
|
|
subscription: PushSubscription
|
|
vault: str = Field(description="Vault name this subscription is for")
|
|
|
|
|
|
class SubscribeResponse(BaseModel):
|
|
"""Response to subscription request."""
|
|
success: bool
|
|
subscription_id: str | None = None
|
|
|
|
|
|
class VapidPublicKeyResponse(BaseModel):
|
|
"""VAPID public key for client."""
|
|
public_key: str
|
|
|
|
|
|
class PushPayload(BaseModel):
|
|
"""Payload for sending a push notification."""
|
|
vault: str
|
|
title: str
|
|
body: str
|
|
data: dict[str, Any] = Field(default_factory=dict)
|
|
tag: str = "obsigate-notification"
|
|
|
|
|
|
# ── Endpoints ───────────────────────────────────────────────────────────────
|
|
|
|
@router.get("/vapid-public-key", response_model=VapidPublicKeyResponse)
|
|
async def get_vapid_public_key_endpoint():
|
|
"""Get VAPID public key for client subscription."""
|
|
public_key = get_vapid_public_key()
|
|
if not public_key:
|
|
# Return a demo key if generation failed (for testing)
|
|
return {"public_key": "demo-key-for-testing"}
|
|
return {"public_key": public_key}
|
|
|
|
|
|
@router.post("/subscribe", response_model=SubscribeResponse)
|
|
async def subscribe_push(
|
|
request: SubscribeRequest,
|
|
current_user=Depends(require_auth)
|
|
):
|
|
"""Subscribe to push notifications for a vault."""
|
|
username = current_user.get("username", "unknown")
|
|
|
|
# Check if user has access to this vault (#194 : "*" n'inclut pas les
|
|
# dossiers persos, check_vault_access est la source unique).
|
|
if not check_vault_access(request.vault, current_user):
|
|
raise HTTPException(status_code=403, detail="No access to this vault")
|
|
|
|
# Check if subscription already exists
|
|
for sub in _push_subscriptions:
|
|
if sub["endpoint"] == request.subscription.endpoint and sub["username"] == username:
|
|
return SubscribeResponse(success=True, subscription_id=sub.get("id"))
|
|
|
|
# Add new subscription
|
|
sub_id = base64.urlsafe_b64encode(os.urandom(16)).decode('utf-8').rstrip('=')
|
|
subscription = {
|
|
"id": sub_id,
|
|
"endpoint": request.subscription.endpoint,
|
|
"keys": request.subscription.keys,
|
|
"vault": request.vault,
|
|
"username": username,
|
|
"created_at": datetime.now(timezone.utc).isoformat()
|
|
}
|
|
|
|
_push_subscriptions.append(subscription)
|
|
save_push_subscriptions()
|
|
|
|
logger.info(f"Push subscription added for user={username}, vault={request.vault}")
|
|
return SubscribeResponse(success=True, subscription_id=sub_id)
|
|
|
|
|
|
@router.delete("/subscribe")
|
|
async def unsubscribe_push(
|
|
endpoint: str,
|
|
current_user=Depends(require_auth)
|
|
):
|
|
"""Unsubscribe from push notifications."""
|
|
username = current_user.get("username", "unknown")
|
|
|
|
global _push_subscriptions
|
|
original_len = len(_push_subscriptions)
|
|
_push_subscriptions = [
|
|
s for s in _push_subscriptions
|
|
if not (s["endpoint"] == endpoint and s["username"] == username)
|
|
]
|
|
|
|
if len(_push_subscriptions) < original_len:
|
|
save_push_subscriptions()
|
|
return {"success": True, "message": "Unsubscribed"}
|
|
|
|
return {"success": False, "message": "Subscription not found"}
|
|
|
|
|
|
@router.get("/subscriptions")
|
|
async def list_subscriptions(current_user=Depends(require_auth)):
|
|
"""List current user's push subscriptions."""
|
|
username = current_user.get("username", "unknown")
|
|
user_subs = [
|
|
{
|
|
"id": s["id"],
|
|
"vault": s["vault"],
|
|
"created_at": s["created_at"],
|
|
"endpoint": s["endpoint"][:50] + "..." # Truncate for privacy
|
|
}
|
|
for s in _push_subscriptions if s["username"] == username
|
|
]
|
|
return {"subscriptions": user_subs}
|
|
|
|
|
|
# ── Internal: Send push notification ────────────────────────────────────────
|
|
|
|
async def send_push_notification(vault: str, title: str, body: str, data: dict | None = None, tag: str = "obsigate-notification") -> int:
|
|
"""Send push notification to all subscribers of a vault. Returns count of sent notifications."""
|
|
subscriptions = [s for s in _push_subscriptions if s["vault"] == vault]
|
|
|
|
if not subscriptions:
|
|
return 0
|
|
|
|
payload = {
|
|
"title": title,
|
|
"body": body,
|
|
"data": data or {},
|
|
"tag": tag
|
|
}
|
|
|
|
# Import here to avoid circular dependency
|
|
from pywebpush import WebPushException, webpush
|
|
|
|
keys = load_vapid_keys()
|
|
vapid_private_key = keys.get("private_key")
|
|
vapid_claims = {
|
|
"sub": "mailto:[email protected]"
|
|
}
|
|
|
|
sent = 0
|
|
for sub in subscriptions:
|
|
try:
|
|
webpush(
|
|
subscription_info={
|
|
"endpoint": sub["endpoint"],
|
|
"keys": sub["keys"]
|
|
},
|
|
data=json.dumps(payload),
|
|
vapid_private_key=vapid_private_key,
|
|
vapid_claims=vapid_claims
|
|
)
|
|
sent += 1
|
|
except WebPushException as e:
|
|
logger.warning(f"Push failed for subscription {sub['id']}: {e}")
|
|
# If subscription expired/invalid, remove it
|
|
if e.response and e.response.status_code in (404, 410):
|
|
_push_subscriptions.remove(sub)
|
|
except Exception as e:
|
|
logger.error(f"Push error for {sub['id']}: {e}")
|
|
|
|
if sent < len(subscriptions):
|
|
save_push_subscriptions()
|
|
|
|
return sent |