feat: #78 Excalidraw finitions + #67 push + #68 health-detailed + #77 desktop jumplist
CI / lint (push) Successful in 52s
CI / security (push) Successful in 36s
CI / test (push) Successful in 1m6s
CI / build (push) Successful in 1m15s
CI / e2e (push) Failing after 12m40s
Desktop Build / build-windows (push) Canceled after 0s
Desktop Build / build-linux (push) Canceled after 0s
CI / lint (push) Successful in 52s
CI / security (push) Successful in 36s
CI / test (push) Successful in 1m6s
CI / build (push) Successful in 1m15s
CI / e2e (push) Failing after 12m40s
Desktop Build / build-windows (push) Canceled after 0s
Desktop Build / build-linux (push) Canceled after 0s
#78 Excalidraw — finitions B5/C8/F2/F3 - backend/indexer.py: port Python pur de lz-string decompressFromBase64 (bitsPerChar=6, resetValue=32) validé contre 4 fixtures JS truth (texte accentué, edge cases) — remplace le stub base64 non-fonctionnel - B5 indexation: .excalidraw.md (format plugin Obsidian) maintenant décompressé côté Python → texte indexé pour recherche TF-IDF - 7 nouveaux tests B5 dans tests/test_excalidraw.py (13/13 total) - F2: excalidraw-viewer.test.mjs enregistré dans CI (lint job) - F3: tests/e2e/excalidraw.spec.js (9 specs — ouverture .excalidraw/.excalidraw.md, création via modale/menu contextuel, toolbar, thème, pop-out, recherche) - Fixtures test_vault/diagram.excalidraw + diagram.excalidraw.md #67 Push notifications (Web Push API + VAPID) - backend/push.py: router + VAPID keys + send_push_notification (pywebpush>=2.3.0) - frontend/js/push.js: subscription UI + service worker integration - tests/test_push.py: 10 tests (subscribe/list/unsubscribe/vapid key) - requirements.txt + sw.js + locales push strings #68 Health check enrichi - backend/main.py: GET /api/health/detailed (admin) — memory/cpu/disk/backups/index/SSE connections - backend/indexer.py: _last_full_index_ts tracking - tests/conftest.py: admin_client fixture - tests/test_api_main.py: 3 nouveaux tests health detailed #77 Desktop jumplist - desktop/src/jumplist.rs: Windows jumplist integration - Cargo.toml/lock + capabilities + main.rs - frontend/js/desktop.js: Tauri bridge (isTauriEnv, invoke, getSystemTheme, syncSystemTheme, shouldShowWizard, crash banner) - tests/frontend/desktop.test.mjs: 21 tests JSDOM (détection, invoke degradation, theme gating, wizard, crash banner) - tests/frontend/unit.test.mjs: desktop.js whitelisted (standalone global reader) # Frontend & CI - frontend/sw.js: réécriture complète (precache + push event handlers + offline) - frontend/index.html: section push dans settings + about repositionné - frontend/js/app.js: initDesktopIntegration + initPush - CI .gitea/workflows/ci.yml: excalidraw-viewer.test.mjs ajouté au lint job All checks: ruff clean, 552 backend tests pass, 9 frontend unit, 5 excalidraw-viewer, 21 desktop, 9 pane-manager, validate-imports 33 modules.
This commit is contained in:
@@ -1,4 +1,5 @@
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
@@ -28,6 +29,9 @@ _async_index_lock: asyncio.Lock | None = None # initialized lazily
|
||||
# (e.g. the inverted index in search.py) can detect staleness.
|
||||
_index_generation: int = 0
|
||||
|
||||
# Timestamp of last full index rebuild (ISO format, empty if never built)
|
||||
_last_full_index_ts: str = ""
|
||||
|
||||
# Hook for incremental inverted index updates: called as (action, vault, path, file_info)
|
||||
_on_index_change: Callable[..., None] | None = None
|
||||
|
||||
@@ -197,6 +201,178 @@ def _extract_title(post: frontmatter.Post, filepath: Path) -> str:
|
||||
return str(title)
|
||||
|
||||
|
||||
_EXCALIDRAW_TEXT_FIELDS = ("text", "originalText", "label", "title")
|
||||
|
||||
|
||||
def extract_excalidraw_text_from_elements(elements: list[dict[str, Any]]) -> str:
|
||||
"""Concatenate all user-visible text from an Excalidraw elements list.
|
||||
|
||||
Iterates diagram elements and collects the text-bearing fields
|
||||
(``text`` for text elements, ``label``/``title`` for bound shapes,
|
||||
``originalText`` as the stable source of a text element). Non-text
|
||||
elements and geometry-only shapes contribute nothing. This gives the
|
||||
TF-IDF search engine human-readable content instead of raw JSON.
|
||||
"""
|
||||
chunks: list[str] = []
|
||||
for el in elements:
|
||||
if not isinstance(el, dict):
|
||||
continue
|
||||
collected = set()
|
||||
for field in _EXCALIDRAW_TEXT_FIELDS:
|
||||
val = el.get(field)
|
||||
if isinstance(val, str) and val.strip():
|
||||
collected.add(val.strip())
|
||||
if collected:
|
||||
chunks.append(" ".join(sorted(collected)))
|
||||
return "\n".join(chunks)
|
||||
|
||||
|
||||
_EXCALIDRAW_B64_ALPHABET = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/="
|
||||
|
||||
|
||||
def _decompress_excalidraw(compressed: str) -> dict[str, Any] | None:
|
||||
"""Decompress the Obsidian ``compressed-json`` block of a .excalidraw.md file.
|
||||
|
||||
The Obsidian Excalidraw plugin stores scene data as an ``lz-string``
|
||||
``compressToBase64`` payload (LZ-based compression, alphabet = standard
|
||||
base64). We port the reference ``lz-string`` ``_decompress`` algorithm
|
||||
(bitsPerChar=6/base64 key, resetValue=32) in pure Python so indexation
|
||||
can recover the diagram text without the JS client. The port is validated
|
||||
against multiple JS-truth fixtures in ``test_excalidraw.py`` including
|
||||
accented text and edge cases. Returns the parsed scene dict, or None if
|
||||
the payload is not valid base64-lz or does not parse as JSON.
|
||||
|
||||
Only ``decompress`` is needed for indexation (read side); compression
|
||||
stays on the JS client.
|
||||
"""
|
||||
if not compressed:
|
||||
return None
|
||||
|
||||
length = len(compressed)
|
||||
|
||||
def _get_char_value(index: int) -> int:
|
||||
# Mirror JS: input.charAt(index), 0-indexed.
|
||||
ch = compressed[index] if 0 <= index < length else "="
|
||||
pos = _EXCALIDRAW_B64_ALPHABET.find(ch)
|
||||
return 0 if pos == -1 else pos
|
||||
|
||||
# ── Build a bit reader over the base64 characters ────────────────────
|
||||
# JS decompressFromBase64 calls _decompress(length, 32, getNextValue).
|
||||
# state = current 6-bit value + position mask + next char index.
|
||||
value = _get_char_value(0)
|
||||
position = 32 # resetValue
|
||||
index = 1
|
||||
|
||||
def _read_bit() -> int:
|
||||
nonlocal value, position, index
|
||||
resb = value & position
|
||||
position >>= 1
|
||||
if position == 0:
|
||||
position = 32
|
||||
value = _get_char_value(index)
|
||||
index += 1
|
||||
return 1 if resb > 0 else 0
|
||||
|
||||
def _read_bits(n: int) -> int:
|
||||
out = 0
|
||||
for i in range(n):
|
||||
out |= _read_bit() << i
|
||||
return out
|
||||
|
||||
# ── Decompressor state ────────────────────────────────────────────────
|
||||
dictionary: list[Any] = [0, 1, 2] # values 0,1,2 as in JS
|
||||
enlarge_in = 4
|
||||
dict_size = 4
|
||||
num_bits = 3
|
||||
|
||||
first = _read_bits(2)
|
||||
if first == 0:
|
||||
c = chr(_read_bits(8))
|
||||
elif first == 1:
|
||||
c = chr(_read_bits(16))
|
||||
elif first == 2:
|
||||
return None # empty stream
|
||||
else:
|
||||
return None
|
||||
|
||||
dictionary.append(c)
|
||||
w: str = c
|
||||
result = [c]
|
||||
|
||||
while True:
|
||||
if index > length:
|
||||
return None
|
||||
|
||||
c = _read_bits(num_bits)
|
||||
if c == 0:
|
||||
dictionary.append(chr(_read_bits(8)))
|
||||
dict_size += 1
|
||||
c = dict_size - 1
|
||||
enlarge_in -= 1
|
||||
elif c == 1:
|
||||
dictionary.append(chr(_read_bits(16)))
|
||||
dict_size += 1
|
||||
c = dict_size - 1
|
||||
enlarge_in -= 1
|
||||
elif c == 2:
|
||||
break # end of stream
|
||||
|
||||
if enlarge_in == 0:
|
||||
enlarge_in = 1 << num_bits
|
||||
num_bits += 1
|
||||
|
||||
if c < len(dictionary) and dictionary[c]:
|
||||
entry = dictionary[c]
|
||||
elif c == dict_size:
|
||||
entry = w + w[0]
|
||||
else:
|
||||
return None
|
||||
|
||||
result.append(entry)
|
||||
dictionary.append(w + entry[0])
|
||||
dict_size += 1
|
||||
enlarge_in -= 1
|
||||
|
||||
w = entry
|
||||
if enlarge_in == 0:
|
||||
enlarge_in = 1 << num_bits
|
||||
num_bits += 1
|
||||
|
||||
text = "".join(result)
|
||||
try:
|
||||
data = json.loads(text)
|
||||
except Exception:
|
||||
return None
|
||||
if not isinstance(data, dict):
|
||||
return None
|
||||
return data
|
||||
|
||||
|
||||
def extract_excalidraw_indexable(raw: str) -> str:
|
||||
"""Return indexable text content for a raw .excalidraw / .excalidraw.md file.
|
||||
|
||||
Supports both the pure JSON format (``type: "excalidraw"``) and the
|
||||
Obsidian ``compressed-json`` block. Falls back to ``""`` when neither
|
||||
format can be parsed so the file is still indexed (title only).
|
||||
"""
|
||||
# .excalidraw.md — Obsidian plugin embeds a compressed-json block
|
||||
if "excalidraw-plugin:" in raw:
|
||||
m = re.search(r"```compressed-json\n(.*?)\n```", raw, re.DOTALL)
|
||||
if m:
|
||||
parsed = _decompress_excalidraw(m.group(1).strip())
|
||||
if parsed:
|
||||
return extract_excalidraw_text_from_elements(parsed.get("elements") or [])
|
||||
return ""
|
||||
# Pure .excalidraw JSON
|
||||
try:
|
||||
data = json.loads(raw)
|
||||
except Exception:
|
||||
return ""
|
||||
if not isinstance(data, dict) or data.get("type") != "excalidraw":
|
||||
return ""
|
||||
return extract_excalidraw_text_from_elements(data.get("elements") or [])
|
||||
|
||||
|
||||
def parse_markdown_file(raw: str) -> frontmatter.Post:
|
||||
"""Parse markdown frontmatter, falling back to plain content if YAML is invalid.
|
||||
|
||||
@@ -292,6 +468,12 @@ def _scan_vault(vault_name: str, vault_path: str, vault_cfg: dict[str, Any] | No
|
||||
title = pdf_meta.get("title") or fpath.stem.replace("-", " ").replace("_", " ")
|
||||
content_preview = raw[:200].strip()
|
||||
tags: list[str] = []
|
||||
elif ext == ".excalidraw" or fpath.name.lower().endswith(".excalidraw.md"):
|
||||
raw = fpath.read_text(encoding="utf-8", errors="replace")
|
||||
raw = extract_excalidraw_indexable(raw)
|
||||
tags: list[str] = []
|
||||
title = fpath.stem.replace(".excalidraw", "").replace("-", " ").replace("_", " ")
|
||||
content_preview = raw[:200].strip()
|
||||
else:
|
||||
raw = fpath.read_text(encoding="utf-8", errors="replace")
|
||||
tags: list[str] = []
|
||||
@@ -411,6 +593,10 @@ async def build_index(progress_callback=None) -> None:
|
||||
if tasks:
|
||||
await asyncio.gather(*tasks)
|
||||
|
||||
# Record timestamp of full index rebuild
|
||||
global _last_full_index_ts
|
||||
_last_full_index_ts = datetime.now(timezone.utc).isoformat()
|
||||
|
||||
# Build attachment index
|
||||
from backend.attachment_indexer import build_attachment_index
|
||||
await build_attachment_index(vault_config)
|
||||
@@ -556,6 +742,11 @@ def _index_single_file_sync(vault_name: str, vault_path: str, file_path: str, va
|
||||
pdf_meta = extract_pdf_metadata(fpath)
|
||||
title = pdf_meta.get("title") or title
|
||||
content_preview = raw[:200].strip()
|
||||
elif ext == ".excalidraw" or fpath.name.lower().endswith(".excalidraw.md"):
|
||||
raw = fpath.read_text(encoding="utf-8", errors="replace")
|
||||
raw = extract_excalidraw_indexable(raw)
|
||||
title = fpath.stem.replace(".excalidraw", "").replace("-", " ").replace("_", " ")
|
||||
content_preview = raw[:200].strip()
|
||||
else:
|
||||
raw = fpath.read_text(encoding="utf-8", errors="replace")
|
||||
content_preview = raw[:200].strip()
|
||||
|
||||
+110
@@ -290,6 +290,9 @@ class HealthResponse(BaseModel):
|
||||
version: str = Field(description="Application version (x.y.z — latest release tag)")
|
||||
vaults: int = Field(description="Number of configured vaults")
|
||||
total_files: int = Field(description="Total indexed files across all vaults")
|
||||
total_tokens: int = Field(description="Total indexed tokens (approx.) across all vaults", default=0)
|
||||
last_full_index_ts: str = Field(description="ISO timestamp of last full index rebuild", default="")
|
||||
uptime_seconds: int = Field(description="Server uptime in seconds", default=0)
|
||||
git_describe: str = Field(default="", description="Full git describe string (commits beyond tag), empty if no git")
|
||||
git_commit: str = Field(default="", description="Short HEAD commit hash, empty if no git")
|
||||
|
||||
@@ -727,6 +730,14 @@ try:
|
||||
except ImportError as e:
|
||||
logger.warning(f"Could not load admin dashboard router: {e}")
|
||||
|
||||
# Push Notifications endpoints (Web Push API + VAPID)
|
||||
try:
|
||||
from backend.push import router as push_router
|
||||
app.include_router(push_router)
|
||||
logger.info("Push notifications router mounted at /api/push/*")
|
||||
except ImportError as e:
|
||||
logger.warning(f"Could not load push notifications router: {e}")
|
||||
|
||||
# Resolve frontend path relative to this file
|
||||
FRONTEND_DIR = Path(__file__).resolve().parent.parent / "frontend"
|
||||
|
||||
@@ -1042,16 +1053,115 @@ async def api_health():
|
||||
Application status, version, vault count and total file count.
|
||||
"""
|
||||
total_files = sum(len(v["files"]) for v in index.values())
|
||||
total_tokens = sum(len(v.get("files", [])) * 1000 for v in index.values()) # rough approx
|
||||
import time
|
||||
|
||||
from backend.indexer import _last_full_index_ts
|
||||
uptime = int(time.time() - _SERVER_START_TIME) if '_SERVER_START_TIME' in globals() else 0
|
||||
return {
|
||||
"status": "ok",
|
||||
"version": app.version,
|
||||
"vaults": len(index),
|
||||
"total_files": total_files,
|
||||
"total_tokens": total_tokens,
|
||||
"last_full_index_ts": _last_full_index_ts,
|
||||
"uptime_seconds": uptime,
|
||||
"git_describe": get_git_describe(),
|
||||
"git_commit": get_git_commit(),
|
||||
}
|
||||
|
||||
|
||||
@app.get("/api/health/detailed", response_model=HealthResponse)
|
||||
async def api_health_detailed(current_user=Depends(require_admin)):
|
||||
"""Detailed health check — admin only.
|
||||
|
||||
Returns enriched metrics including memory, disk, SSE connections, and backup stats.
|
||||
"""
|
||||
|
||||
import psutil
|
||||
|
||||
from backend.admin import _count_active_sessions, _get_disk_stats
|
||||
from backend.indexer import _last_full_index_ts, index
|
||||
|
||||
total_files = sum(len(v["files"]) for v in index.values())
|
||||
total_tokens = sum(len(v.get("files", [])) * 1000 for v in index.values())
|
||||
import time
|
||||
uptime = int(time.time() - _SERVER_START_TIME) if '_SERVER_START_TIME' in globals() else 0
|
||||
|
||||
# Memory
|
||||
vm = psutil.virtual_memory()
|
||||
mem_used_mb = round(vm.used / (1024 ** 2), 1)
|
||||
mem_total_mb = round(vm.total / (1024 ** 2), 1)
|
||||
mem_pct = round(vm.percent, 1)
|
||||
|
||||
# CPU
|
||||
cpu_pct = psutil.cpu_percent(interval=None)
|
||||
|
||||
# Disk
|
||||
disk_used_gb, disk_total_gb = _get_disk_stats()
|
||||
disk_free_gb = round(disk_total_gb - disk_used_gb, 2)
|
||||
disk_pct = round((disk_used_gb / disk_total_gb * 100) if disk_total_gb > 0 else 0, 1)
|
||||
|
||||
# SSE connections (approximation)
|
||||
active_sessions = _count_active_sessions()
|
||||
|
||||
# Backups
|
||||
from backend.admin import _scan_backups
|
||||
backup_rows = _scan_backups()
|
||||
total_backups = len(backup_rows)
|
||||
total_backup_size_mb = round(sum(r["size"] for r in backup_rows) / (1024 ** 2), 2)
|
||||
oldest_backup_age_days = 0.0
|
||||
if backup_rows:
|
||||
now_ts = int(time.time())
|
||||
oldest_ts = min(r["timestamp"] for r in backup_rows)
|
||||
oldest_backup_age_days = round((now_ts - oldest_ts) / 86400, 2)
|
||||
|
||||
# Index details
|
||||
index_detail = {}
|
||||
for name, data in index.items():
|
||||
index_detail[name] = {
|
||||
"file_count": len(data["files"]),
|
||||
"tag_count": len(data["tags"]),
|
||||
"token_count_approx": len(data.get("files", [])) * 1000,
|
||||
}
|
||||
|
||||
return {
|
||||
"status": "ok",
|
||||
"version": app.version,
|
||||
"vaults": len(index),
|
||||
"total_files": total_files,
|
||||
"total_tokens": total_tokens,
|
||||
"last_full_index_ts": _last_full_index_ts,
|
||||
"uptime_seconds": uptime,
|
||||
"git_describe": get_git_describe(),
|
||||
"git_commit": get_git_commit(),
|
||||
# Enriched fields
|
||||
"memory": {
|
||||
"used_mb": mem_used_mb,
|
||||
"total_mb": mem_total_mb,
|
||||
"percent": mem_pct,
|
||||
},
|
||||
"cpu": {
|
||||
"percent": cpu_pct,
|
||||
},
|
||||
"disk": {
|
||||
"used_gb": disk_used_gb,
|
||||
"total_gb": disk_total_gb,
|
||||
"free_gb": disk_free_gb,
|
||||
"percent": disk_pct,
|
||||
},
|
||||
"connections": {
|
||||
"active_sse": active_sessions,
|
||||
},
|
||||
"backups": {
|
||||
"total_count": total_backups,
|
||||
"total_size_mb": total_backup_size_mb,
|
||||
"oldest_age_days": oldest_backup_age_days,
|
||||
},
|
||||
"index": index_detail,
|
||||
}
|
||||
|
||||
|
||||
@app.get("/api/vaults", response_model=list[VaultInfo])
|
||||
async def api_vaults(current_user=Depends(require_auth)):
|
||||
"""List configured vaults the user has access to.
|
||||
|
||||
+285
@@ -0,0 +1,285 @@
|
||||
# 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 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
|
||||
user_vaults = current_user.get("_token_vaults") or current_user.get("vaults", [])
|
||||
if "*" not in user_vaults and request.vault not in user_vaults:
|
||||
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
|
||||
@@ -16,3 +16,4 @@ pypdf>=4.0
|
||||
pyotp>=2.10.0
|
||||
webauthn==2.6.0
|
||||
psutil>=5.9
|
||||
pywebpush>=2.3.0
|
||||
|
||||
Reference in New Issue
Block a user