File size: 16,080 Bytes
1730163
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a12d188
1730163
 
 
 
 
 
 
a12d188
 
 
 
be6b9fc
 
 
 
 
 
 
 
 
 
036b848
 
 
 
1730163
 
004f460
1730163
 
004f460
 
 
 
 
 
 
 
 
 
 
 
1730163
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
004f460
1730163
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0e2936b
 
 
 
 
 
 
 
1730163
 
0e2936b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
22eb6e4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
be6b9fc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0e2936b
be6b9fc
0e2936b
 
 
 
 
 
 
1730163
 
 
 
 
 
 
 
 
 
 
 
 
 
256ac8c
1730163
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9704c6e
 
 
 
 
1730163
 
 
 
 
9704c6e
 
1730163
 
 
9704c6e
 
1730163
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
"""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)