amanpay / api /dependencies.py
MHamdan's picture
CI deploy f7ddbcd
a12d188 verified
Raw
History Blame Contribute Delete
16.1 kB
"""Shared model state and helpers for the API (dependency injection)."""
from __future__ import annotations
import base64
import binascii
import logging
import os
from io import BytesIO
from typing import Optional
import numpy as np
import torch
from fastapi import HTTPException
from PIL import Image
from amanpay.config import AmanPayConfig, load_config
from amanpay.data.preprocessing import FacePreprocessor, FingerprintPreprocessor
from amanpay.models.authenticator import AmanPayAuthenticator
logger = logging.getLogger(__name__)
class ModelState:
"""Holds the loaded authenticator and preprocessors for the app lifetime."""
def __init__(self) -> None:
self.model: Optional[AmanPayAuthenticator] = None
self.unified = None
self.face_pre: Optional[FacePreprocessor] = None
self.fp_pre: Optional[FingerprintPreprocessor] = None
self.voice_pre = None
self.deepfake = None
self.deepfake_threshold: float = 0.5
from amanpay.security.liveness import LivenessRegistry
from amanpay.security.passkey import PasskeyRegistry
from amanpay.security.voice import VoiceRegistry
from amanpay.security.flash_liveness import FlashLivenessRegistry
from amanpay.banking.wallet import WalletRegistry
from amanpay.banking.risk import RiskEngine
from amanpay.banking.geo import GeoRegistry
from api import services
self.passkeys = PasskeyRegistry()
self.voice = VoiceRegistry()
self.liveness = LivenessRegistry()
self.flash = FlashLivenessRegistry()
self.wallet = WalletRegistry()
self.risk = RiskEngine()
self.geo = GeoRegistry()
# Share the torch-free notification singleton (the /notify router uses the same
# instance) and wire durable-persistence + audit hooks into it.
self.notify = services.notifications
services.wire(persist=self.persist, audit=self.audit)
# Datastore: a real SQL DB (SQLite/Postgres) when DATABASE_URL is set —
# per-user rows, audit log, cascade erasure (P1). Otherwise the HF Dataset
# store (keeps the zero-config HF Space demo working).
if os.getenv("DATABASE_URL") or os.getenv("AMANPAY_DB", "").lower() in (
"sql", "sqlite", "postgres", "1", "true"):
from amanpay.storage.repository import SqlRepository
self.store = SqlRepository()
else:
from amanpay.storage import HFStore
self.store = HFStore()
# Provider-independent, non-custodial Payment Core (Saudi-first). Mock provider
# until a licensed real provider is reviewed + enabled (see PROVIDER_EVALUATION.md).
from amanpay.payments import PaymentCore
self.payments = PaymentCore(audit_sink=self.audit)
self.webauthn = None # lazily created (needs rp_id/origin config)
self.report_card: Optional[dict] = None
self.trained: dict = {} # per-modality trained-weights flags (set in load())
self.device: str = "cpu"
def biometric_status(self) -> dict:
"""Per-modality capability for the UI (delegates to a torch-free builder)."""
from amanpay.biometric_status import build_biometric_status
return build_biometric_status(
device=self.device, model_loaded=self.loaded, trained=self.trained or {},
voice=getattr(self, "voice", None) is not None,
unified=getattr(self, "unified", None) is not None,
deepfake=getattr(self, "deepfake", None) is not None,
liveness=getattr(self, "liveness", None) is not None,
flash=getattr(self, "flash", None) is not None,
oob=getattr(self, "notify", None) is not None)
def load(self, config_path: Optional[str] = None,
checkpoint: Optional[str] = None) -> None:
"""Instantiate the model (and load a checkpoint if provided)."""
self.device = "cuda" if torch.cuda.is_available() else "cpu"
config: AmanPayConfig = load_config(config_path)
config.auth.match_mode = os.getenv("AMANPAY_MATCH_MODE", "score")
# Resolve checkpoint paths and, if any are missing, try the layered HF
# fallbacks (tier 2: auto-download our hosted weights; tier 3: third-party
# backbone; tier 4: ImageNet init inside the encoders).
from amanpay.models.hf_backbone import (
_truthy, ensure_local_weights, maybe_init_encoders_from_hf)
face = os.getenv("AMANPAY_FACE_CKPT", "checkpoints/face_encoder_best.pt")
fp = os.getenv("AMANPAY_FP_CKPT", "checkpoints/fp_socofing_on.pt")
voice = os.getenv("AMANPAY_VOICE_CKPT", "checkpoints/voice_encoder_best.pt")
df_path = os.getenv("AMANPAY_DEEPFAKE_CKPT", "checkpoints/deepfake_detector.pkl")
ensure_local_weights(
[face, fp, voice, df_path],
repo=os.getenv("AMANPAY_HF_REPO", "MHamdan/amanpay-encoders"),
token=os.getenv("HF_TOKEN"),
enabled=_truthy(os.getenv("AMANPAY_HF_AUTO_DOWNLOAD", "1")))
trained = {"face": os.path.exists(face), "fingerprint": os.path.exists(fp),
"voice": os.path.exists(voice)}
self.trained = trained # per-modality: real trained weights vs ImageNet-untrained
model = AmanPayAuthenticator(config)
if checkpoint:
try:
model.load(checkpoint, map_location=self.device)
logger.info("Loaded checkpoint %s", checkpoint)
except FileNotFoundError:
logger.warning("Checkpoint %s not found; using initialized weights", checkpoint)
else:
model.load_pretrained(
face_path=face if trained["face"] else None,
fp_path=fp if trained["fingerprint"] else None,
map_location=self.device,
)
# Tier 3: for any encoder still without trained weights, optionally pull a
# pretrained MobileNetV3 backbone from the Hub (AMANPAY_HF_BACKBONE).
maybe_init_encoders_from_hf(
{"face": model.face_encoder, "fingerprint": model.fingerprint_encoder},
trained)
self.model = model.to(self.device).eval()
# MTCNN face detection is the biggest CPU cost; on weak hosts set
# AMANPAY_FACE_DETECT=0 to use the fast center-crop path (inputs from the
# app are already framed faces).
use_mtcnn = _truthy(os.getenv("AMANPAY_FACE_DETECT", "1"))
self.face_pre = FacePreprocessor(device=self.device, use_mtcnn=use_mtcnn)
self.fp_pre = FingerprintPreprocessor()
# Unified tri-modal authenticator (face + fingerprint + voice), trained ckpts.
from amanpay.data.preprocessing import VoicePreprocessor
from amanpay.models.unified import UnifiedAuthenticator
self.voice_pre = VoicePreprocessor()
uni = UnifiedAuthenticator(config).to(self.device).eval()
uni.load_pretrained(
face_path=face if trained["face"] else None,
fp_path=fp if trained["fingerprint"] else None,
voice_path=voice if trained["voice"] else None,
map_location=self.device)
maybe_init_encoders_from_hf(dict(uni.encoders), trained)
self.unified = uni
self.deepfake_threshold = config.auth.deepfake_threshold
# Passive deepfake/injection detector (optional — loaded if fitted).
if os.path.exists(df_path):
from amanpay.models.deepfake_detector import DeepfakeDetector
self.deepfake = DeepfakeDetector.load(df_path)
logger.info("Deepfake detector loaded from %s", df_path)
else:
logger.info("No deepfake detector at %s; deepfake_score disabled", df_path)
# Report card: serve a precomputed snapshot instantly (the live ISO/24745
# build is ~slow on weak CPUs), then refresh it live in the background.
self._prime_report_card(config)
# Restore enrolled identities from the HF datastore (survives Space rebuilds).
from amanpay.models.hf_backbone import _truthy
if _truthy(os.getenv("AMANPAY_PERSIST", "1")):
try:
self.restore(self.store.load())
except Exception as exc:
logger.warning("Enrollment restore failed (%s)", exc)
logger.info("Model ready on %s", self.device)
# -- persistence: enrolled templates + passkeys + wallet + prefs ----- #
def _ensure_webauthn(self):
if self.webauthn is None:
from amanpay.security.webauthn_server import WebAuthnServer
self.webauthn = WebAuthnServer()
return self.webauthn
def snapshot(self) -> dict:
st: dict = {"unified": {}, "webauthn": {}, "wallet": {}, "notify": {},
"passkeys": {}}
if self.unified is not None:
for uid, tmpl in self.unified.enrolled.items():
st["unified"][uid] = {m: t.flatten().tolist() for m, t in tmpl.items()}
if self.webauthn is not None:
st["webauthn"] = self.webauthn.snapshot()
st["wallet"] = self.wallet.snapshot()
st["notify"] = self.notify.snapshot()
return st
def restore(self, st: dict) -> None:
if not st:
return
import torch
if self.unified is not None:
for uid, tmpl in (st.get("unified") or {}).items():
self.unified.enrolled[uid] = {
m: torch.tensor(v, dtype=torch.float32).reshape(1, -1)
for m, v in tmpl.items()}
if st.get("webauthn"):
self._ensure_webauthn().restore(st["webauthn"])
self.wallet.restore(st.get("wallet") or {})
self.notify.restore(st.get("notify") or {})
def ensure_user_loaded(self, user_id: str) -> None:
"""Read-through credential cache: if this replica doesn't have the user's
durable state in memory (e.g. they enrolled on another replica), load just
that user from the datastore. Bounds cross-replica propagation lag to one
cache-miss DB read instead of waiting for a full reload."""
if not user_id:
return
present = ((self.unified is not None and user_id in self.unified.enrolled)
or user_id in self.passkeys._users
or user_id in self.wallet._users
or (self.webauthn is not None and user_id in self.webauthn._users))
if present:
return
fn = getattr(self.store, "load_user", None)
if fn is None:
return
try:
st = fn(user_id)
if st:
self.restore(st)
except Exception as exc:
logger.info("ensure_user_loaded(%s) failed (%s)", user_id, exc)
def audit(self, user_id: str, action: str, detail: Optional[dict] = None) -> None:
"""Append an audit-log entry if the datastore supports it (SQL backend)."""
fn = getattr(self.store, "append_audit", None)
if fn is not None:
try:
fn(user_id, action, detail or {})
except Exception:
pass
def erase(self, user_id: str) -> bool:
"""Cascade-erase a user across memory + datastore (GDPR/BIPA erasure)."""
if self.unified is not None:
self.unified.enrolled.pop(user_id, None)
if self.model is not None:
self.model.enrolled_templates.pop(user_id, None)
self.wallet._users.pop(user_id, None)
if self.webauthn is not None:
self.webauthn._users.pop(user_id, None)
fn = getattr(self.store, "delete_user", None)
ok = bool(fn(user_id)) if fn is not None else False
self.persist()
return ok
def persist(self) -> None:
"""Snapshot to the datastore (best-effort, off the request path)."""
from amanpay.models.hf_backbone import _truthy
if not self.store.enabled or not _truthy(os.getenv("AMANPAY_PERSIST", "1")):
return
import threading
snap = self.snapshot()
threading.Thread(target=lambda: self.store.save(snap), daemon=True).start()
def _prime_report_card(self, config: "AmanPayConfig") -> None:
import json
import threading
from amanpay.models.hf_backbone import _truthy
snap = os.path.join("results", "report_card.json")
if self.report_card is None and os.path.exists(snap):
try:
with open(snap) as fh:
self.report_card = json.load(fh)
logger.info("Report card loaded from snapshot %s", snap)
except Exception as exc:
logger.info("Report-card snapshot unreadable (%s)", exc)
# Live refresh is opt-in — it burns CPU and spikes memory, which can OOM a
# small host (e.g. a free ~512 MB CPU tier). The snapshot is authoritative;
# ?refresh=true still recomputes on demand.
if _truthy(os.getenv("AMANPAY_REPORTCARD_REFRESH", "0")):
def _refresh() -> None:
try:
from amanpay.evaluation.report_card import build_report_card
self.report_card = build_report_card(
fusion_dim=config.fusion.output_dim,
protection_bits=config.auth.protection_bits)
logger.info("Report card refreshed live")
except Exception as exc:
logger.info("Live report-card refresh failed (%s)", exc)
threading.Thread(target=_refresh, daemon=True).start()
def deepfake_score(self, face_rgb) -> Optional[float]:
"""P(attack) for a decoded RGB face image, or None if detector absent."""
if self.deepfake is None:
return None
return self.deepfake.score(face_rgb)
@property
def loaded(self) -> bool:
return self.model is not None
# Singleton shared across requests.
state = ModelState()
def get_model() -> AmanPayAuthenticator:
"""Dependency: return the loaded model or raise 503 if unavailable."""
if not state.loaded or state.model is None:
raise HTTPException(status_code=503, detail="Model not loaded")
return state.model
# Max accepted payload for a single base64 image/audio field (bytes) — bounds
# memory/DoS from unbounded uploads. Override with AMANPAY_MAX_UPLOAD_MB.
MAX_UPLOAD_BYTES = int(float(os.getenv("AMANPAY_MAX_UPLOAD_MB", "8")) * 1024 * 1024)
def decode_image(b64: str) -> np.ndarray:
"""Decode a base64 (optionally data-URI-prefixed) image to an RGB array."""
try:
if "," in b64 and b64.strip().startswith("data:"):
b64 = b64.split(",", 1)[1]
if len(b64) > MAX_UPLOAD_BYTES * 4 // 3 + 4:
raise HTTPException(status_code=413, detail="image payload too large")
raw = base64.b64decode(b64)
img = Image.open(BytesIO(raw)).convert("RGB")
return np.array(img)
except HTTPException:
raise
except (binascii.Error, ValueError, OSError) as exc:
raise HTTPException(status_code=400, detail=f"Invalid image data: {exc}") from exc
def preprocess_face(image: np.ndarray) -> torch.Tensor:
"""Preprocess a face image to a model-ready tensor on the active device."""
assert state.face_pre is not None
tensor = state.face_pre.process(image)
if tensor is None:
raise HTTPException(status_code=422, detail="No face detected in image")
return tensor.to(state.device)
def preprocess_fingerprint(image: np.ndarray) -> torch.Tensor:
"""Preprocess a fingerprint image to a model-ready tensor on the active device."""
assert state.fp_pre is not None
return state.fp_pre.process(image).to(state.device)
def preprocess_voice(wav_bytes: bytes) -> torch.Tensor:
"""Preprocess WAV audio to a mel-spectrogram tensor on the active device."""
assert state.voice_pre is not None
return state.voice_pre.process(wav_bytes).to(state.device)