diff --git a/.env.example b/.env.example index 98d9fa4..03eda48 100644 --- a/.env.example +++ b/.env.example @@ -17,6 +17,8 @@ PAYGATE_V4_WEBHOOK_SECRET=replace-with-a-long-random-webhook-secret # Online SQLite backup schedule. PAYGATE_V4_BACKUP_HOUR_UTC=3 PAYGATE_V4_BACKUP_RETENTION=30 +# Raw relay notification body retention; normalized evidence is retained. +PAYGATE_V4_RAW_EVENT_RETENTION=168h # paygate-v4-migrate only; not required by the normal server. PAYGATE_V4_ACTIVE_PROFILE=kotak diff --git a/README.md b/README.md index fcf3a7c..758b5a3 100644 --- a/README.md +++ b/README.md @@ -8,7 +8,7 @@ # PayGate -PayGate creates a payment instruction with a unique exact payable amount, watches supported incoming-payment notifications through one trusted Android phone, and marks the matching payment only when the server can do so unambiguously. +PayGate creates a payment instruction with a unique exact payable amount, watches supported incoming-payment notifications through one or more independently revocable Android relay phones, and changes payment state only when the server has an origin-bound confirmation it can trust. The payer sends money directly to the configured UPI account. PayGate does not custody, route, or settle funds; notification matching is an evidence mechanism rather than a bank/acquirer settlement guarantee. @@ -54,6 +54,7 @@ v4.0 currently supports: - **Paytm for Business** payment notifications; - **Kotak** credit notifications delivered by Google Messages on the PayGate phone. +Paytm package-bound notifications are currently the only automatic confirmation source. Kotak/Google Messages and package-agnostic notifications remain signed evidence for web-admin confirmation until an independent sender/provider proof is available. The merchant does not select Paytm or Kotak. PayGate snapshots the active profile and destination when the payment is created, so later profile changes cannot alter an existing payment. @@ -108,7 +109,7 @@ The web UI is embedded in the v4 server image. The Android app remains a separat The PayGate phone uses `NotificationListenerService`, a durable local queue and a P-256 ECDSA key stored in Android Keystore. -Android performs only cheap source allowlisting and notification capture. The server owns source-specific parsing, incoming-credit semantics, profile inference, matching, deduplication and payment mutation. +Android performs only cheap source-agnostic notification capture. The server owns source-specific parsing, incoming-credit semantics, profile inference, source trust, matching, deduplication and payment mutation. The foreground relay is intentionally independent of operator login and is designed to survive screen lock, Doze, process recreation, temporary network loss and normal Battery Saver when the app is exempt from battery optimization. ## Persistence and recovery @@ -150,7 +151,7 @@ CI validates the v4 frontend, all retained Go packages, static analysis and the - plaintext merchant/webhook secrets are never persisted when a verifier/hash is sufficient; - Android private signing keys remain non-exportable in Android Keystore; - payer identity and raw notification detail stay out of unauthenticated/public payment views; -- a false-positive confirmation is considered worse than a delayed/manual outcome, so matching fails closed; +- a false-positive confirmation is considered worse than a delayed/manual outcome, so only origin-bound evidence may transition payment state; generic notifications and Google Messages/SMS remain evidence for operator confirmation; - never run a second PayGate process against the live SQLite volume. ## Documentation diff --git a/cmd/paygate-v4/main.go b/cmd/paygate-v4/main.go index 6e99455..07931e7 100644 --- a/cmd/paygate-v4/main.go +++ b/cmd/paygate-v4/main.go @@ -3,17 +3,21 @@ package main import ( "context" "errors" + "flag" "fmt" + "io" "log" "net/http" "os" "os/signal" + "path/filepath" "strconv" "strings" "syscall" "time" v4runtime "github.com/Phloraxx/payment-api/internal/v4/runtime" + "github.com/Phloraxx/payment-api/internal/v4/storage" ) func main() { @@ -23,8 +27,13 @@ func main() { } func run() error { - if len(os.Args) > 1 && os.Args[1] == "healthcheck" { - return healthcheck() + if len(os.Args) > 1 { + switch os.Args[1] { + case "healthcheck": + return healthcheck() + case "restore-drill": + return restoreDrill() + } } cfg, err := configFromEnv() if err != nil { @@ -88,6 +97,37 @@ func healthcheck() error { } return nil } +func restoreDrill() error { + flags := flag.NewFlagSet("restore-drill", flag.ContinueOnError) + flags.SetOutput(io.Discard) + backup := flags.String("backup", "", "completed standalone SQLite backup") + liveDB := flags.String("live-db", "", "live database path to protect") + expectedSHA256 := flags.String("sha256", "", "expected backup SHA-256") + if err := flags.Parse(os.Args[2:]); err != nil { + return fmt.Errorf("parse restore-drill arguments: %w", err) + } + if flags.NArg() != 0 { + return errors.New("restore-drill does not accept positional arguments") + } + if strings.TrimSpace(*backup) == "" { + return errors.New("restore-drill requires --backup") + } + protectedPath := strings.TrimSpace(*liveDB) + if protectedPath == "" { + dataDir := strings.TrimSpace(os.Getenv("PAYGATE_V4_DATA_DIR")) + if dataDir == "" { + return errors.New("restore-drill requires --live-db or PAYGATE_V4_DATA_DIR") + } + protectedPath = filepath.Join(dataDir, "paygate.db") + } + report, err := storage.RestoreDrill(context.Background(), *backup, protectedPath, *expectedSHA256) + if err != nil { + return err + } + fmt.Printf("restore drill passed: schema=%d payments=%d relay_events=%d webhook_deliveries=%d relay_devices=%d\n", + report.SchemaVersion, report.Payments, report.RelayEvents, report.WebhookDeliveries, report.RelayDevices) + return nil +} func configFromEnv() (v4runtime.Config, error) { hour := 3 @@ -106,6 +146,14 @@ func configFromEnv() (v4runtime.Config, error) { } retention = parsed } + var rawEventRetention time.Duration + if value := strings.TrimSpace(os.Getenv("PAYGATE_V4_RAW_EVENT_RETENTION")); value != "" { + parsed, err := time.ParseDuration(value) + if err != nil { + return v4runtime.Config{}, fmt.Errorf("invalid PAYGATE_V4_RAW_EVENT_RETENTION: %w", err) + } + rawEventRetention = parsed + } origins := splitCSV(os.Getenv("PAYGATE_V4_ALLOWED_ORIGINS")) return v4runtime.Config{ DataDir: os.Getenv("PAYGATE_V4_DATA_DIR"), @@ -119,6 +167,7 @@ func configFromEnv() (v4runtime.Config, error) { BackupDir: os.Getenv("PAYGATE_V4_BACKUP_DIR"), BackupHourUTC: hour, BackupRetention: retention, + RawEventRetention: rawEventRetention, }, nil } diff --git a/cmd/paygate-v4/main_test.go b/cmd/paygate-v4/main_test.go index 3a3fe75..3001104 100644 --- a/cmd/paygate-v4/main_test.go +++ b/cmd/paygate-v4/main_test.go @@ -1,6 +1,9 @@ package main -import "testing" +import ( + "testing" + "time" +) func TestConfigFromEnvUsesExplicitV4BootstrapValues(t *testing.T) { t.Setenv("PAYGATE_V4_DATA_DIR", t.TempDir()) @@ -21,6 +24,16 @@ func TestConfigFromEnvUsesExplicitV4BootstrapValues(t *testing.T) { t.Fatal("explicit v4 webhook secret missing") } } +func TestConfigFromEnvReadsRawEventRetention(t *testing.T) { + t.Setenv("PAYGATE_V4_RAW_EVENT_RETENTION", "48h") + cfg, err := configFromEnv() + if err != nil { + t.Fatal(err) + } + if cfg.RawEventRetention != 48*time.Hour { + t.Fatalf("raw event retention = %s", cfg.RawEventRetention) + } +} func TestConfigFromEnvIgnoresRemovedLegacyBootstrapAliases(t *testing.T) { t.Setenv("PAYGATE_V4_DATA_DIR", t.TempDir()) diff --git a/docs/v4/00_PRODUCT_VISION.md b/docs/v4/00_PRODUCT_VISION.md index 116ad76..071604f 100644 --- a/docs/v4/00_PRODUCT_VISION.md +++ b/docs/v4/00_PRODUCT_VISION.md @@ -114,11 +114,11 @@ The frontend never chooses the collection profile and never constructs a destina ## Android responsibility -The phone is a trustworthy sensor, not the payment engine. +The phone is a transport sensor, not the payment engine or a payment-proof authority. It knows: -- allowed notification packages; +- package identity as evidence, without a payment-app allowlist; - a cheap generic decimal-money prefilter; - notification package/key/text/post time; - durable local queue/retry; diff --git a/docs/v4/01_TARGET_ARCHITECTURE.md b/docs/v4/01_TARGET_ARCHITECTURE.md index e6cbad0..e80e03b 100644 --- a/docs/v4/01_TARGET_ARCHITECTURE.md +++ b/docs/v4/01_TARGET_ARCHITECTURE.md @@ -92,7 +92,7 @@ Owns: - short-lived QR pairing sessions; - enrolled Android public key; - signed request verification; -- one active relay-device policy for v4.0; +- additive relay-device enrollment and independent revocation; - heartbeat/device health; - event deduplication. @@ -243,7 +243,7 @@ not requested amount. ### Android -- package allowlist; +- capture package identity as evidence; do not maintain an Android payment-package allowlist; - cheap generic decimal-money prefilter; - capture notification package/key/text/post time; - stable local event ID; @@ -290,9 +290,8 @@ Oracle/Docker host local paygate.db + WAL/SHM completed-backup exporter -Android phone - one PayGate APK -``` +Android phones + one or more additive PayGate APK relay clients Production invariant remains one PayGate process owning the live SQLite database. @@ -311,11 +310,11 @@ Production invariant remains one PayGate process owning the live SQLite database ## Future-source extension rule -Adding GPay/Slice later should require only: +Adding GPay/Slice or another source later should require only: -1. allowlist a new Android package if necessary; -2. add server parser + sanitized fixtures; -3. map parser to a collection profile/source; +1. add parser + sanitized fixtures while the source-agnostic transport remains unchanged; +2. map parser output to a collection profile/source; +3. independently authenticate the source before allowing automatic confirmation; otherwise retain evidence for operator confirmation; 4. pass the same observation/matching pipeline. It must not require changing merchant payment creation or Android/server ownership boundaries. \ No newline at end of file diff --git a/docs/v4/02_NOTIFICATION_INGESTION_AND_PAIRING.md b/docs/v4/02_NOTIFICATION_INGESTION_AND_PAIRING.md index f1c7a2c..01cbd63 100644 --- a/docs/v4/02_NOTIFICATION_INGESTION_AND_PAIRING.md +++ b/docs/v4/02_NOTIFICATION_INGESTION_AND_PAIRING.md @@ -16,9 +16,9 @@ The server decides: ## Universal package intake -Android does not maintain a payment-app package allowlist. A package name is retained as evidence, but package identity alone never authorizes a payment match. This allows BHIM, Google Pay, PhonePe, bank apps and future sources to work without an Android release when their visible notification text satisfies the generic server parser. +Android does not maintain a payment-app package allowlist. A package name is retained as evidence, but package identity and visible text alone never authorize a payment match. This keeps BHIM, Google Pay, PhonePe, bank apps and future sources capturable without an Android release; automatic payment confirmation still requires an independently trusted source. -Specialized server parsers remain useful for sources with stronger known semantics, but they are an optimization and confidence boundary rather than an Android transport restriction. +Specialized server parsers remain useful for sources with stronger origin-bound semantics. Generic and Google Messages/SMS evidence is retained for diagnostics and operator confirmation, not automatic payment confirmation. ## Generic on-device prefilter @@ -152,12 +152,13 @@ It extracts best-effort: - transaction time when text contains one. Unknown decimal-money Google Messages notifications are ignored/unmatched; they are not treated as Kotak by default. +Even a positively recognized Kotak credit remains evidence-only: Google Messages exposes SMS text, not an authenticated bank/provider assertion, so it is stored as ambiguous when it has a payment candidate and cannot enqueue `payment.paid`. ### Generic incoming-payment notifications Any other package can produce `android_notification` evidence when the bounded visible text positively expresses an incoming payment and contains a PayGate decimal amount. Google Messages messages that are not recognized as Kotak use the equivalent `android_message` source. Real regression fixtures cover generic wallet/bank wording and BHIM notifications, including dotted UPI IDs. -Generic evidence does **not** inherit whichever collection profile happens to be active when the HTTP request arrives. The server first searches historical amount reservations using the exact payable amount and trusted occurrence time. One qualifying profile is used; multiple qualifying profiles are ambiguous and cannot pay either candidate. With no historical candidate, the current active profile is retained only so unmatched diagnostic evidence can still be recorded. +Generic evidence does **not** inherit whichever collection profile happens to be active when the HTTP request arrives. The server first searches historical amount reservations using the exact payable amount and trusted occurrence time. One qualifying profile is used for evidence attribution; multiple qualifying profiles are ambiguous. Even one qualifying generic candidate is stored as ambiguous/manual evidence and cannot change payment state or enqueue a webhook. With no historical candidate, the current active profile is retained only so unmatched diagnostic evidence can still be recorded. The normalized observation schema is intentionally not tied to a fixed app package enum, so adding another wallet or bank notification usually requires only parser fixtures unless its wording needs a specialized parser. @@ -233,10 +234,10 @@ One real UPI credit may create more than one phone notification. Example: a Kota PayGate must not try to collapse those notifications on the phone. Each source event remains independently signed and stored. The **payment** is the dedupe anchor: -1. first safe observation -> `matched`; if needed, transition payment to `paid` and enqueue exactly one `payment.paid` webhook; -2. later independent observation for the same historical reservation -> `corroborated`, attach to the same payment and optionally enrich missing payer fields; -3. exact retry of the same relay event ID -> replay prior result without a second observation; -4. reused amount + insufficient timestamp confidence -> `ambiguous`, never assume it corroborates the newest payment. +1. a trusted origin-bound observation may match and enqueue exactly one `payment.paid` webhook; +2. generic/GPay and Google Messages/Kotak observations remain ambiguous evidence for operator confirmation, even when their amount/profile/time candidate is unique; +3. later trusted evidence for the same historical reservation may be `corroborated` without a second payment transition/webhook; +4. exact retry of the same relay event ID replays the prior result; reused amount + insufficient timestamp confidence remains ambiguous. This makes duplicate handling provider-agnostic and does not require UTR/RRN or fragile notification-text hashes. @@ -246,7 +247,7 @@ Android should not upload unrelated personal notifications. Rules: -- allowlist only required apps; +- capture unrelated package text only when the cheap decimal-money candidate filter passes; - cheap decimal-money filter before persistence/upload; - bounded title/text/big-text sizes; - short local retention; @@ -326,19 +327,17 @@ https://pay.mulearnscet.in/device/pair/ - consumption and device enrollment happen atomically; - failed/expired/used token cannot be replayed. -## One active relay phone in v4.0 +## Additive relay phones -Because the product currently needs one phone, keep the operational model simple: +Relay phones are additive and may remain enabled concurrently: ```text -one active payment relay device +one or more independently revocable payment relay devices ``` -Pairing a second phone should require an explicit **Replace device** flow that revokes/disables the old relay only after the new enrollment succeeds. +Pairing a new phone adds its device-key row. It does not disable or replace existing phones. Operators revoke a specific device explicitly when it is lost, retired or no longer trusted. -This avoids two phones delivering duplicate notifications for the same bank/payment stream. - -Historical device records may remain for audit. +Historical device records remain for audit, and each device's signed events are deduplicated independently. ## Device signing diff --git a/docs/v4/03_PAYMENT_LIFECYCLE_AND_MATCHING.md b/docs/v4/03_PAYMENT_LIFECYCLE_AND_MATCHING.md index 1b830c0..82b28a9 100644 --- a/docs/v4/03_PAYMENT_LIFECYCLE_AND_MATCHING.md +++ b/docs/v4/03_PAYMENT_LIFECYCLE_AND_MATCHING.md @@ -191,7 +191,9 @@ Preference: ## Reuse/collision safety rule -If a payable value has only one historical reservation compatible with the trusted occurrence time, it may match that reservation. +If the payable value has only one historical reservation compatible with the trusted occurrence time, it is a necessary candidate, not sufficient proof. + +Only an origin-bound source may automatically confirm that candidate. The current v4 Paytm notification path is bound to the Paytm Business package; generic package text and Google Messages/Kotak SMS have no authenticated app, sender or provider assertion. Those observations are stored as ambiguous Activity for explicit web-admin confirmation or future independent provider corroboration. If the same value has been reused and the observation time is missing, implausible or too low-confidence to distinguish the reservations: @@ -201,13 +203,12 @@ DO NOT AUTO-MATCH Store the observation as unmatched/ambiguous Activity. The operator can correct the payment directly if necessary. -This is the final defense against extremely delayed SMS/notification delivery. +This is the final defense against extremely delayed SMS/notification delivery and forged visible text. ## Auto-match algorithm For normalized observation `O`: -```text 1. dedupe signed relay event 2. parse source and validate incoming-credit semantics 3. validate amount > 0 and paise != 00 @@ -215,23 +216,26 @@ For normalized observation `O`: 5. resolve collection profile P from source semantics or, for generic evidence, historical exact-amount reservations at occurred_at 6. find historical payments on P with exact payable amount 7. restrict candidates to payments whose lifecycle can contain O.occurred_at -8. if exactly one candidate is safe and not yet paid: +8. if the source is origin-bound and exactly one candidate is safe and not yet paid: mark paid attach observation as `matched` copy payer enrichment when present append payment history enqueue one payment.paid webhook in same transaction -9. if exactly one safe candidate is already paid and already has a confirming observation: - attach this independent observation as `corroborated` +9. if the source is generic or Google Messages/Kotak, retain the candidate as ambiguous evidence: + do not mutate payment state + do not attach a payment confirmation + do not enqueue a payment.paid webhook +10. if exactly one trusted candidate is already paid and already has a confirming observation: + attach this independent trusted observation as `corroborated` optionally enrich missing payer fields do not create another payment transition/history/webhook -10. if zero candidates: +11. if zero candidates: save unmatched Activity -11. if multiple/uncertain candidates: +12. if multiple/uncertain candidates: fail closed; save ambiguous Activity -``` -The **currently active collection profile is irrelevant to a match when historical evidence identifies the profile**. A Paytm observation can still pay an older Paytm payment after the operator switches new payment creation to Kotak. Generic evidence similarly follows a unique historical reservation profile; if the same amount was simultaneously reserved on multiple profiles, it fails closed as ambiguous rather than guessing. +The **currently active collection profile is irrelevant to a trusted match when historical evidence identifies the profile**. A Paytm observation can still pay an older Paytm payment after the operator switches new payment creation to Kotak. Generic/Kotak evidence can still be attributed to the unique historical reservation for operator review, but never changes payment state automatically. ## Relay amount hint is not trusted diff --git a/docs/v4/05_ANDROID_APP.md b/docs/v4/05_ANDROID_APP.md index ab78e5a..ede5a47 100644 --- a/docs/v4/05_ANDROID_APP.md +++ b/docs/v4/05_ANDROID_APP.md @@ -135,14 +135,9 @@ Operator auth controls management UI only. ## NotificationListenerService -Initial allowlist: +The listener captures notifications from any package; package identity is retained as evidence and is not an Android payment allowlist. Paytm and Google Messages have specialized server parsers, while unknown packages remain eligible for the generic decimal-money prefilter. -```text -com.paytm.business -com.google.android.apps.messaging -``` - -GPay/Slice are deferred. +GPay/Slice are deferred as automatic confirmation sources. Their notifications may still be captured as evidence and remain ambiguous/manual when a payment candidate exists. Phone applies only the generic non-`.00` decimal-money prefilter. It does not decide incoming-credit semantics/profile/payment match. @@ -230,16 +225,17 @@ No permanent aggressive wake lock. No manual server URL/device ID/pairing secret in normal UX. -## One active phone +## Additive relay phones -V4.0 assumes one payment-notification phone. +V4.0 supports one or more payment-notification phones concurrently. Each device has its own signed identity, queue and health state. -Pairing another should be a **Replace phone** action: +Pairing another phone is an **Add phone** action: -- new phone enrolls successfully first; -- server activates new device and revokes/disables old one atomically or in one controlled flow; -- old historical device record remains for audit; -- avoid two active devices delivering the same payment stream. +- the new phone enrolls successfully; +- existing enabled phones remain enabled; +- each device can be revoked independently; +- historical device records remain for audit; +- duplicate source events are deduplicated server-side without disabling a healthy peer. ## Existing production phone @@ -313,7 +309,7 @@ Immutable: ```text Paytm payment detected · ₹100.37 · matched to Sourav P Bijoy -Kotak payment detected · ₹500.42 · unmatched +Kotak payment detected · ₹500.42 · ambiguous; operator confirmation required Payment updated · operator Webhook delivered · 200 Phone replaced @@ -325,7 +321,7 @@ Primary UI should not expose implementation words such as `evidence_reference`, ### Notification/queue -- allowlist Paytm + Google Messages only; +- source-agnostic capture from any package; specialized Paytm/Google Messages parser routing; - decimal money with non-zero paise passes; - `.00` is filtered; - unrelated personal message never queues; @@ -340,7 +336,7 @@ Primary UI should not expose implementation words such as `evidence_reference`, - valid one-time App Link enrollment; - expired/used token rejected; - existing key reused after app upgrade; -- Replace phone leaves one active relay device; +- multiple relay phones may remain enabled concurrently; - revoke blocks signed relay requests; - operator logout does not stop relay. diff --git a/docs/v4/06_ADMIN_UI_AND_DESIGN_SYSTEM.md b/docs/v4/06_ADMIN_UI_AND_DESIGN_SYSTEM.md index 243bf37..fc37525 100644 --- a/docs/v4/06_ADMIN_UI_AND_DESIGN_SYSTEM.md +++ b/docs/v4/06_ADMIN_UI_AND_DESIGN_SYSTEM.md @@ -240,7 +240,7 @@ Examples: ```text Payment detected · Paytm · ₹100.37 · matched to Sourav P Bijoy -Payment detected · Kotak · ₹501.42 · unmatched +Payment detected · Kotak · ₹501.42 · ambiguous; operator confirmation required Payment updated · pay_... · operator Webhook delivered · payment.paid · 200 Webhook failed · payment.expired · 404 diff --git a/docs/v4/07_DATA_STORAGE_SECURITY_OPERATIONS.md b/docs/v4/07_DATA_STORAGE_SECURITY_OPERATIONS.md index 6d12c91..28200b2 100644 --- a/docs/v4/07_DATA_STORAGE_SECURITY_OPERATIONS.md +++ b/docs/v4/07_DATA_STORAGE_SECURITY_OPERATIONS.md @@ -272,7 +272,7 @@ device_model TEXT android_version TEXT ``` -v4.0 should prefer **one active payment-relay device** to avoid duplicate phone streams. Pairing a replacement phone should be an explicit operator action. +Multiple payment-relay devices may remain enabled concurrently. Pairing adds a device row; explicit operator revocation disables only the selected device, while historical device records remain for audit. ### `pairing_sessions` @@ -489,7 +489,7 @@ Plus application checks: - payment count plausible; - latest payments/history readable; - collection profile state valid; -- one active relay device invariant valid; +- every enabled relay device is independently valid; no singleton active-device invariant; - pending webhook counts readable. Do not run a full integrity scan on every request/health check. @@ -511,20 +511,25 @@ Restore acceptance: ## Sensitive-data retention -Keep payment/history records according to business/audit needs, but aggressively bound notification content. +Keep payment/history records according to business/audit needs, but bound relay +notification bodies. The runtime redactor runs once at startup and hourly, +processing at most 500 rows per pass. Its default boundary is seven days, +configurable with `PAYGATE_V4_RAW_EVENT_RETENTION` between `1h` and `30d`. -Recommended policy to validate during implementation: +Redaction clears `relay_events.title`, `text`, and `big_text` only. It retains +the event identity, source package/event ID, timestamps, amount hint, payload +hash, status/error, and any normalized observation/payment/history records. +Rows still marked `received` at expiry become `ignored` with an expiry reason, +so expired bodies cannot later trigger parsing or payment transitions. -- relay raw title/text/bigText: days, not permanent; -- normalized payer fields: retained with payment when operationally useful; -- unmatched raw notification data: short retention; -- delivered local Android queue: short retention; -- failed local rows: longer but bounded diagnostics retention; -- pairing sessions: purge quickly after use/expiry; -- admin sessions: purge after expiry/revocation; -- delivered webhook bodies: bounded retention after operational window. +The Android relay intentionally has no automatic age, row-count, or retry +deletion. Settings provides an explicit confirmation action that clears +delivered/local-only/candidate rows while preserving pending, retry, and failed +evidence. App-data deletion remains the user's separate Android control. -If we need stronger deletion hygiene for notification text, evaluate `PRAGMA secure_delete=FAST` against write cost and WAL/backup behavior rather than assuming row deletion instantly removes every forensic copy. +Backups and WAL files may retain prior SQLite pages. If stronger deletion +hygiene is required, evaluate `PRAGMA secure_delete=FAST` against write cost +and backup behavior rather than assuming row updates erase every forensic copy. Reference: https://www.sqlite.org/pragma.html#pragma_secure_delete @@ -533,6 +538,10 @@ Reference: https://www.sqlite.org/pragma.html#pragma_secure_delete ### Admin One password, hashed with Argon2id. Password change invalidates existing admin sessions. +Password-only login is bounded to two concurrent Argon2 verifications and +throttled per observed remote address after five failures in fifteen minutes. +The temporary block is one minute; `Retry-After` is returned and password +content is never logged. Web session cookie: @@ -547,6 +556,11 @@ Path=/admin Keep ECDSA signing and Android Keystore private key. Pairing token is short-lived and single-use; server stores public key only. +The sole device-authenticated mutation, active collection-destination update, +is checked inside the same immediate transaction against both `enabled=1` and +the exact enrollment epoch authenticated for the request. Revocation or +re-pairing therefore wins before an in-flight phone write can commit. + ### Merchant API Separate merchant API key. Never reuse the admin password or Android device credential. diff --git a/docs/v4/08_MIGRATION_AND_IMPLEMENTATION_PLAN.md b/docs/v4/08_MIGRATION_AND_IMPLEMENTATION_PLAN.md index eab2c80..71cf64d 100644 --- a/docs/v4/08_MIGRATION_AND_IMPLEMENTATION_PLAN.md +++ b/docs/v4/08_MIGRATION_AND_IMPLEMENTATION_PLAN.md @@ -17,8 +17,8 @@ Before implementation, confirm these product rules: - `name` = merchant-supplied person/payee identifier, not event title; - `external_id` = merchant/event ID and is allowed to repeat across many payments; - idempotency uses `Idempotency-Key`, not `external_id`; -- Paytm + Kotak only for v4.0; -- GPay, Amazon Pay and Slice active matching deferred; +- Paytm package-bound automatic confirmation; Kotak/Google Messages and generic package-agnostic evidence remain manual until an independent provider/source proof exists; +- GPay, Amazon Pay and Slice may be captured as evidence, but active automatic confirmation is deferred; - merchant never selects collection profile; - PayGate returns UPI URI string; frontend renders QR; - no hosted checkout requirement; @@ -28,7 +28,7 @@ Before implementation, confirm these product rules: - 5m active + 5m grace + 5m hard quarantine; - soft recent-use avoidance after hard release; - direct SQLite only; no PocketBase runtime; -- one active Android relay device for v4.0; +- one or more additive Android relay devices with independent revocation; - Razorpay-inspired dark navy/blue UI. Capture sanitized real notification fixtures before parser implementation. @@ -182,7 +182,7 @@ Implement/verify: - bounded text payload; - server event dedupe; - timestamp confidence/source; -- one active relay phone policy; +- additive relay phones with independent revocation; - QR/App-Link pairing session. ### Relay edge cases @@ -266,14 +266,13 @@ unique relay event - grace match; - delayed relay during quarantine with original post time; - event delivered after amount release but occurrence belongs to old reservation; -- same amount reused later + high-confidence old timestamp -> old payment only; +- same amount reused later + high-confidence old timestamp -> trusted old payment only; - same amount reused later + ambiguous/low-confidence timestamp -> no auto-match; - unknown decimal payment -> unmatched Activity; - duplicate relay event -> exact replay, no second observation or paid webhook; -- different source event for the same already-paid reservation -> `corroborated`, no second transition/webhook; -- future Kotak SMS + GPay/Amazon Pay observations converge on the payment as dedupe anchor; -- reused amount + weak timing never becomes false corroboration; -- inactive historical profile still matchable; +- generic, Kotak and future untrusted-source observations -> ambiguous evidence, never automatic state transition/webhook; +- multiple independently trusted observations for the same already-paid reservation -> `corroborated`, no second transition/webhook; +- inactive historical profile still matchable for trusted evidence; - active profile switch never affects match decision. ## Phase 9 — Admin auth and web dashboard @@ -314,7 +313,7 @@ Order: 5. Settings; 6. password-only login; 7. QR/App-Link Connect/Replace phone UX; -8. Paytm + Google Messages package allowlist; +8. source-agnostic Android capture; server parser and independent source-trust gate; 9. generic decimal prefilter; 10. queue/heartbeat/Doze regression tests. diff --git a/docs/v4/10_EDGE_CASES_AND_INVARIANTS.md b/docs/v4/10_EDGE_CASES_AND_INVARIANTS.md index 4345b61..09f2ace 100644 --- a/docs/v4/10_EDGE_CASES_AND_INVARIANTS.md +++ b/docs/v4/10_EDGE_CASES_AND_INVARIANTS.md @@ -170,9 +170,9 @@ If trusted occurrence time proves money arrived before/around cancellation accor Exact retry of the same signed source event -> return the prior result idempotently. -A different source event can describe the same underlying credit. Example: Kotak SMS + GPay, or later Kotak SMS + Amazon Pay. Do not dedupe these on Android and do not hash text in an attempt to prove they are identical. +Different source events can describe the same underlying credit. Example: Kotak SMS + GPay, or later Kotak SMS + Amazon Pay. Do not dedupe these on Android and do not hash text in an attempt to prove they are identical. -The first safe observation is `matched`. A later independent observation that resolves to the same historical payment reservation is `corroborated`. It attaches to the same payment and may fill missing payer information, but it creates **no second payment transition and no second merchant webhook**. +Only an independently authenticated/origin-bound source event may become `matched` or `corroborated`. Generic and Google Messages/Kotak text remains signed evidence for operator confirmation, even when it resolves to the same historical reservation. If the payable amount has been reused and the later source has only low-confidence timing, fail closed as `ambiguous`; never call it corroboration merely because the amount is equal. @@ -218,13 +218,13 @@ Server Paytm parser rejects it as non-incoming credit. No payment mutation. ### Google Messages balance message with decimals -Cheap filter may pass. Kotak parser must positively recognize incoming Kotak credit; otherwise ignore/unmatched. +Cheap filter may pass. Kotak parser must positively recognize incoming Kotak credit; otherwise ignore/unmatched. Recognized Google Messages/Kotak evidence remains ambiguous/manual because the SMS sender/provider is not authenticated. ### Multiple amounts in one notification Example may contain transaction amount and balance. -Phone does not choose authority. Server source parser must know which token is the credited amount. If parser cannot determine reliably, no auto-match. +Phone does not choose authority. Server source parser must know which token is the credited amount. If parser cannot determine reliably, or the source has no authenticated origin, no auto-match. ### Currency comma separators @@ -285,7 +285,7 @@ Token consumption and device enrollment are atomic so failure cannot consume tok ### Replace phone -New phone must enroll before old phone is disabled. Final state has exactly one active relay device. +New phone must enroll without disabling old phones. Final state may contain multiple concurrently enabled relay devices, each independently revocable. ### Old phone comes online later @@ -325,7 +325,7 @@ Persist server/domain times as UTC Unix milliseconds. Parse displayed bank times ### Exactly one safe candidate -Auto-match. +Auto-match only when the source is independently origin-bound. Generic package text and Google Messages/Kotak SMS with a unique candidate remain ambiguous/manual evidence because visible notification/SMS content is forgeable. ### Zero candidates @@ -503,9 +503,10 @@ After meaningful v4 payments exist, rollback cannot simply restore old v3 databa ## 14. Security edge cases -### Admin password brute force - -Rate-limit/throttle password-only login and log security events without password content. +Rate-limit/throttle password-only login and log security events without +password content. PayGate enforces two concurrent Argon2 verifications, +five failures per remote address within fifteen minutes, and a one-minute +temporary block. ### Merchant API key leaked @@ -538,8 +539,8 @@ Always render as escaped text. Never inject raw notification content into dashbo 11. One PayGate process owns live SQLite. 12. No second maintenance/backup PayGate process opens live DB. 13. UTR/RRN is not required. -14. GPay/Amazon Pay/Slice do not auto-match in v4.0. -15. Multiple independent notification sources can corroborate one payment but can never create multiple `payment.paid` transitions/webhooks. +14. GPay/Amazon Pay/Slice and other untrusted generic sources are evidence-only in v4.0; they never auto-match without independent origin/provider proof. +15. Multiple independently trusted evidence sources can corroborate one payment but can never create multiple `payment.paid` transitions/webhooks; untrusted notification/SMS evidence remains ambiguous/manual. 16. PocketBase/libgm are absent from final v4 runtime. 17. Every risky operator correction is visible in immutable history. 18. If PayGate cannot prove which payment owns money, it records Activity and does not guess. \ No newline at end of file diff --git a/docs/v4/README.md b/docs/v4/README.md index 16b448e..aabba86 100644 --- a/docs/v4/README.md +++ b/docs/v4/README.md @@ -33,9 +33,10 @@ A create request contains context such as: 2. v4.0 supports Paytm and Kotak. GPay and Slice are deferred. 3. PayGate snapshots profile/destination onto each payment, so switching the active profile never changes existing sessions. 4. PayGate returns the canonical `upi://pay?...` string. The frontend renders the QR; PayGate does not need SVG/PNG/hosted-checkout output. -5. Android does not know the active profile or expected payment. It relays a minimal signed notification snapshot from allowlisted packages when it contains a plausible non-`.00` money value. +5. Android does not know the active profile or expected payment. It relays a minimal signed notification snapshot from any package when it contains a plausible non-`.00` money value; the server applies the source-trust gate. 6. Server parses source-specific wording, infers Paytm/Kotak, validates incoming-credit semantics and performs matching. 7. Target v4 has no server-side Google Messages/libgm connector. Kotak arrives through the phone's Google Messages notification. +- Paytm package-bound notifications are the current automatic confirmation source; generic and Google Messages/Kotak evidence remains manual until independently authenticated. 8. UTR/RRN is not part of v4 matching. 9. Payable amounts use **ordered random buckets**: for a ₹N request, PayGate randomly chooses among free `₹N.01…₹N.99` values first. It only considers `₹(N+1).01…₹(N+1).99` when the entire base-rupee bucket is unavailable. 10. The default v4.0 capacity is therefore two 99-value buckets (maximum adjustment `₹1.99`), always skipping `.00`. Randomness applies **inside the current bucket**, never across both buckets at once. diff --git a/internal/v4/adminpayments/service.go b/internal/v4/adminpayments/service.go index eaa1b55..d6f805e 100644 --- a/internal/v4/adminpayments/service.go +++ b/internal/v4/adminpayments/service.go @@ -50,6 +50,9 @@ type Payment struct { PayerName string PayerUPIID string InternalNote string + EvidencePackage string + EvidenceDeviceName string + EvidenceReceivedAt *time.Time } type ListInput struct { @@ -151,9 +154,62 @@ func (s *Service) List(ctx context.Context, input ListInput) (ListResult, error) if err := rows.Err(); err != nil { return ListResult{}, fmt.Errorf("iterate payments: %w", err) } + if err := s.attachLatestEvidence(ctx, items); err != nil { + return ListResult{}, err + } return ListResult{Items: items, Total: total, Limit: limit, Offset: offset}, nil } +func (s *Service) attachLatestEvidence(ctx context.Context, items []Payment) error { + if len(items) == 0 { + return nil + } + placeholders := make([]string, len(items)) + args := make([]any, len(items)) + byID := make(map[string]int, len(items)) + for i := range items { + placeholders[i] = "?" + args[i] = items[i].ID + byID[items[i].ID] = i + } + query := `SELECT po.matched_payment_id,re.package_name,COALESCE(rd.name,''),po.received_at + FROM payment_observations po + JOIN relay_events re ON re.id=po.relay_event_id + LEFT JOIN relay_devices rd ON rd.id=re.device_id + WHERE po.matched_payment_id IN (` + strings.Join(placeholders, ",") + `) + AND po.match_result IN ('matched','corroborated') + ORDER BY po.received_at DESC,po.rowid DESC` + rows, err := s.DB.SQL.QueryContext(ctx, query, args...) + if err != nil { + return fmt.Errorf("load payment evidence: %w", err) + } + defer rows.Close() + seen := make(map[string]bool, len(items)) + for rows.Next() { + var paymentID, packageName, deviceName string + var receivedAt int64 + if err := rows.Scan(&paymentID, &packageName, &deviceName, &receivedAt); err != nil { + return fmt.Errorf("scan payment evidence: %w", err) + } + if seen[paymentID] { + continue + } + i, ok := byID[paymentID] + if !ok { + continue + } + at := time.UnixMilli(receivedAt).UTC() + items[i].EvidencePackage = packageName + items[i].EvidenceDeviceName = deviceName + items[i].EvidenceReceivedAt = &at + seen[paymentID] = true + } + if err := rows.Err(); err != nil { + return fmt.Errorf("iterate payment evidence: %w", err) + } + return nil +} + func buildWhere(input ListInput) (string, []any, error) { var clauses []string var args []any @@ -251,6 +307,11 @@ func (s *Service) Get(ctx context.Context, id string) (Detail, error) { if err != nil { return Detail{}, fmt.Errorf("get payment: %w", err) } + paymentsWithEvidence := []Payment{payment} + if err := s.attachLatestEvidence(ctx, paymentsWithEvidence); err != nil { + return Detail{}, err + } + payment = paymentsWithEvidence[0] history, err := loadHistory(ctx, s.DB.SQL, id) if err != nil { return Detail{}, err diff --git a/internal/v4/adminpayments/service_test.go b/internal/v4/adminpayments/service_test.go index 426c3bf..7cd97b4 100644 --- a/internal/v4/adminpayments/service_test.go +++ b/internal/v4/adminpayments/service_test.go @@ -106,6 +106,50 @@ func TestDetailIncludesHistoryAndWebhookTimeline(t *testing.T) { } } +func TestListAndDetailIncludeLatestMatchedEvidence(t *testing.T) { + f := newFixture(t) + payment := f.create(t, 100, "Sourav", "evt_1", "evidence") + if _, err := f.db.SQL.Exec(`INSERT INTO relay_devices(id,name,public_key_pem,enabled,enrolled_at) VALUES('device-evidence','Motorola Edge 60 Stylus','pem',1,?)`, f.now.Add(-time.Hour).UnixMilli()); err != nil { + t.Fatal(err) + } + for i, tc := range []struct { + id string + pkg string + status string + offset time.Duration + }{ + {id: "old", pkg: "in.org.npci.upiapp", status: "matched", offset: -time.Minute}, + {id: "new", pkg: "com.google.android.apps.nbu.paisa.user", status: "corroborated", offset: 0}, + } { + at := f.now.Add(tc.offset).UnixMilli() + relayID := "relay-evidence-" + tc.id + if _, err := f.db.SQL.Exec(`INSERT INTO relay_events(id,device_id,source_event_id,package_name,posted_at,received_at,status) VALUES(?,?,?,?,?,?,?)`, + relayID, "device-evidence", "source-"+tc.id, tc.pkg, at, at, "matched"); err != nil { + t.Fatalf("relay %d: %v", i, err) + } + if _, err := f.db.SQL.Exec(`INSERT INTO payment_observations(id,relay_event_id,source,collection_profile_id,amount_paise,occurred_at,occurred_at_source,received_at,matched_payment_id,match_result) VALUES(?,?,?,?,?,?,?,?,?,?)`, + "obs-evidence-"+tc.id, relayID, "android_notification", "paytm", payment.PayableAmountPaise, at, "notification_posted_at", at, payment.ID, tc.status); err != nil { + t.Fatalf("observation %d: %v", i, err) + } + } + + list, err := f.admin.List(context.Background(), ListInput{Query: payment.ID}) + if err != nil || len(list.Items) != 1 { + t.Fatalf("list=%+v err=%v", list, err) + } + got := list.Items[0] + if got.EvidencePackage != "com.google.android.apps.nbu.paisa.user" || got.EvidenceDeviceName != "Motorola Edge 60 Stylus" || got.EvidenceReceivedAt == nil || !got.EvidenceReceivedAt.Equal(*f.now) { + t.Fatalf("list evidence = %+v", got) + } + detail, err := f.admin.Get(context.Background(), payment.ID) + if err != nil { + t.Fatal(err) + } + if detail.Payment.EvidencePackage != got.EvidencePackage || detail.Payment.EvidenceDeviceName != got.EvidenceDeviceName || detail.Payment.EvidenceReceivedAt == nil { + t.Fatalf("detail evidence = %+v", detail.Payment) + } +} + func TestEditMerchantVisibleFieldsCreatesUpdatedWebhook(t *testing.T) { f := newFixture(t) payment := f.create(t, 100, "Sourav", "evt_1", "edit-fields") diff --git a/internal/v4/auth/service.go b/internal/v4/auth/service.go index 59eee79..d54af97 100644 --- a/internal/v4/auth/service.go +++ b/internal/v4/auth/service.go @@ -129,24 +129,32 @@ func (s *Service) ChangePassword(ctx context.Context, currentPassword, newPasswo if err := s.ready(); err != nil { return err } + var encodedCurrent string + if err := s.DB.SQL.QueryRowContext(ctx, `SELECT password_hash FROM admin_credentials WHERE singleton=1`).Scan(&encodedCurrent); errors.Is(err, sql.ErrNoRows) { + return ErrNotInitialized + } else if err != nil { + return fmt.Errorf("read admin password: %w", err) + } + ok, err := verifyPassword(encodedCurrent, currentPassword) + if err != nil { + return fmt.Errorf("verify admin password: %w", err) + } + if !ok { + return ErrInvalidCredentials + } encodedNew, err := s.hashPassword(newPassword) if err != nil { return err } now := s.now().UnixMilli() return s.DB.WithImmediateTx(ctx, func(tx *storage.ImmediateTx) error { - var encodedCurrent string - if err := tx.QueryRowContext(ctx, `SELECT password_hash FROM admin_credentials WHERE singleton=1`).Scan(&encodedCurrent); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return ErrNotInitialized - } - return fmt.Errorf("read admin password: %w", err) - } - ok, err := verifyPassword(encodedCurrent, currentPassword) - if err != nil { - return fmt.Errorf("verify admin password: %w", err) + var current string + if err := tx.QueryRowContext(ctx, `SELECT password_hash FROM admin_credentials WHERE singleton=1`).Scan(¤t); errors.Is(err, sql.ErrNoRows) { + return ErrNotInitialized + } else if err != nil { + return fmt.Errorf("re-read admin password: %w", err) } - if !ok { + if current != encodedCurrent { return ErrInvalidCredentials } if _, err := tx.ExecContext(ctx, `UPDATE admin_credentials SET password_hash=?,updated_at=? WHERE singleton=1`, encodedNew, now); err != nil { @@ -164,11 +172,9 @@ func (s *Service) CreateAdminSession(ctx context.Context, password string) (Admi return AdminSession{}, err } var encoded string - err := s.DB.SQL.QueryRowContext(ctx, `SELECT password_hash FROM admin_credentials WHERE singleton=1`).Scan(&encoded) - if errors.Is(err, sql.ErrNoRows) { + if err := s.DB.SQL.QueryRowContext(ctx, `SELECT password_hash FROM admin_credentials WHERE singleton=1`).Scan(&encoded); errors.Is(err, sql.ErrNoRows) { return AdminSession{}, ErrNotInitialized - } - if err != nil { + } else if err != nil { return AdminSession{}, fmt.Errorf("read admin password: %w", err) } ok, err := verifyPassword(encoded, password) @@ -178,23 +184,40 @@ func (s *Service) CreateAdminSession(ctx context.Context, password string) (Admi if !ok { return AdminSession{}, ErrInvalidCredentials } - token, err := s.randomToken("pg_admin_", 32) - if err != nil { - return AdminSession{}, fmt.Errorf("generate admin session: %w", err) - } - hash := sha256.Sum256([]byte(token)) - now := s.now() - ttl := s.SessionTTL - if ttl <= 0 { - ttl = 24 * time.Hour - } - expires := now.Add(ttl) - _, err = s.DB.SQL.ExecContext(ctx, `INSERT INTO admin_sessions(token_hash,created_at,expires_at,last_seen_at) VALUES(?,?,?,?)`, - hash[:], now.UnixMilli(), expires.UnixMilli(), now.UnixMilli()) + + var session AdminSession + err = s.DB.WithImmediateTx(ctx, func(tx *storage.ImmediateTx) error { + var current string + if err := tx.QueryRowContext(ctx, `SELECT password_hash FROM admin_credentials WHERE singleton=1`).Scan(¤t); errors.Is(err, sql.ErrNoRows) { + return ErrNotInitialized + } else if err != nil { + return fmt.Errorf("re-read admin password: %w", err) + } + if current != encoded { + return ErrInvalidCredentials + } + token, err := s.randomToken("pg_admin_", 32) + if err != nil { + return fmt.Errorf("generate admin session: %w", err) + } + hash := sha256.Sum256([]byte(token)) + now := s.now() + ttl := s.SessionTTL + if ttl <= 0 { + ttl = 24 * time.Hour + } + expires := now.Add(ttl) + if _, err := tx.ExecContext(ctx, `INSERT INTO admin_sessions(token_hash,created_at,expires_at,last_seen_at) VALUES(?,?,?,?)`, + hash[:], now.UnixMilli(), expires.UnixMilli(), now.UnixMilli()); err != nil { + return fmt.Errorf("store admin session: %w", err) + } + session = AdminSession{Token: token, ExpiresAt: expires} + return nil + }) if err != nil { - return AdminSession{}, fmt.Errorf("store admin session: %w", err) + return AdminSession{}, err } - return AdminSession{Token: token, ExpiresAt: expires}, nil + return session, nil } func (s *Service) AuthenticateAdminSession(ctx context.Context, token string) error { @@ -228,11 +251,19 @@ func (s *Service) RevokeAdminSession(ctx context.Context, token string) error { return err } hash := sha256.Sum256([]byte(strings.TrimSpace(token))) - result, err := s.DB.SQL.ExecContext(ctx, `UPDATE admin_sessions SET revoked_at=? WHERE token_hash=? AND revoked_at IS NULL`, s.now().UnixMilli(), hash[:]) + var rowsAffected int64 + err := s.DB.WithImmediateTx(ctx, func(tx *storage.ImmediateTx) error { + result, err := tx.ExecContext(ctx, `UPDATE admin_sessions SET revoked_at=? WHERE token_hash=? AND revoked_at IS NULL`, s.now().UnixMilli(), hash[:]) + if err != nil { + return fmt.Errorf("revoke admin session: %w", err) + } + rowsAffected, _ = result.RowsAffected() + return nil + }) if err != nil { - return fmt.Errorf("revoke admin session: %w", err) + return err } - if rows, _ := result.RowsAffected(); rows != 1 { + if rowsAffected != 1 { return ErrInvalidSession } return nil @@ -247,19 +278,20 @@ func (s *Service) BootstrapAPIKey(ctx context.Context, label, secret string) err if label == "" || len([]rune(label)) > 120 || len(secret) < 32 { return fmt.Errorf("%w: bootstrap API key requires a label and at least 32 characters", ErrInvalidInput) } - var count int - if err := s.DB.SQL.QueryRowContext(ctx, `SELECT COUNT(*) FROM api_keys`).Scan(&count); err != nil { - return fmt.Errorf("read API key bootstrap state: %w", err) - } - if count > 0 { - return nil - } hash := sha256.Sum256([]byte(secret)) - _, err := s.DB.SQL.ExecContext(ctx, `INSERT INTO api_keys(id,label,secret_hash,enabled,created_at) VALUES('key_legacy_v3',?,?,1,?)`, label, hash[:], s.now().UnixMilli()) - if err != nil { - return fmt.Errorf("bootstrap API key: %w", err) - } - return nil + return s.DB.WithImmediateTx(ctx, func(tx *storage.ImmediateTx) error { + var count int + if err := tx.QueryRowContext(ctx, `SELECT COUNT(*) FROM api_keys`).Scan(&count); err != nil { + return fmt.Errorf("read API key bootstrap state: %w", err) + } + if count > 0 { + return nil + } + if _, err := tx.ExecContext(ctx, `INSERT INTO api_keys(id,label,secret_hash,enabled,created_at) VALUES('key_legacy_v3',?,?,1,?)`, label, hash[:], s.now().UnixMilli()); err != nil { + return fmt.Errorf("bootstrap API key: %w", err) + } + return nil + }) } func (s *Service) CreateAPIKey(ctx context.Context, label string) (APIKey, error) { @@ -282,9 +314,15 @@ func (s *Service) CreateAPIKey(ctx context.Context, label string) (APIKey, error secret := "pg_live_" + id + "_" + base64.RawURLEncoding.EncodeToString(secretPart) hash := sha256.Sum256([]byte(secret)) now := s.now() - _, err = s.DB.SQL.ExecContext(ctx, `INSERT INTO api_keys(id,label,secret_hash,enabled,created_at) VALUES(?,?,?,1,?)`, id, label, hash[:], now.UnixMilli()) + err = s.DB.WithImmediateTx(ctx, func(tx *storage.ImmediateTx) error { + _, err := tx.ExecContext(ctx, `INSERT INTO api_keys(id,label,secret_hash,enabled,created_at) VALUES(?,?,?,1,?)`, id, label, hash[:], now.UnixMilli()) + if err != nil { + return fmt.Errorf("store API key: %w", err) + } + return nil + }) if err != nil { - return APIKey{}, fmt.Errorf("store API key: %w", err) + return APIKey{}, err } return APIKey{ID: id, Label: label, Secret: secret, CreatedAt: now}, nil } @@ -321,11 +359,19 @@ func (s *Service) RevokeAPIKey(ctx context.Context, id string) error { return err } id = strings.TrimSpace(id) - result, err := s.DB.SQL.ExecContext(ctx, `UPDATE api_keys SET enabled=0 WHERE id=? AND enabled=1`, id) + var rowsAffected int64 + err := s.DB.WithImmediateTx(ctx, func(tx *storage.ImmediateTx) error { + result, err := tx.ExecContext(ctx, `UPDATE api_keys SET enabled=0 WHERE id=? AND enabled=1`, id) + if err != nil { + return fmt.Errorf("revoke API key: %w", err) + } + rowsAffected, _ = result.RowsAffected() + return nil + }) if err != nil { - return fmt.Errorf("revoke API key: %w", err) + return err } - if rows, _ := result.RowsAffected(); rows != 1 { + if rowsAffected != 1 { return ErrInvalidAPIKey } return nil diff --git a/internal/v4/httpapi/admin.go b/internal/v4/httpapi/admin.go index 64c2f28..a021702 100644 --- a/internal/v4/httpapi/admin.go +++ b/internal/v4/httpapi/admin.go @@ -2,11 +2,14 @@ package httpapi import ( "bytes" + "context" "errors" "io" + "net" "net/http" "strconv" "strings" + "sync" "time" "github.com/Phloraxx/payment-api/internal/v4/adminpayments" @@ -14,14 +17,41 @@ import ( "github.com/Phloraxx/payment-api/internal/v4/operator" "github.com/Phloraxx/payment-api/internal/v4/profiles" "github.com/Phloraxx/payment-api/internal/v4/relay" + "github.com/Phloraxx/payment-api/internal/v4/storage" "github.com/Phloraxx/payment-api/internal/v4/webhooks" ) const ( - adminCookieName = "paygate_admin" - adminLoginConcurrency = 2 + adminCookieName = "paygate_admin" + adminLoginConcurrency = 2 + adminLoginFailureWindow = 15 * time.Minute + adminLoginFailureLimit = 5 + adminLoginBlockDuration = time.Minute + adminLoginThrottleEntries = 1024 ) +type adminLoginFailure struct { + failures int + firstFailure time.Time + blockedUntil time.Time + lastSeen time.Time +} + +type adminDeviceAuthorizationContextKey struct{} + +type adminDeviceAuthorization struct { + ID string + EnrolledAt time.Time +} + +func adminDeviceAuthorizationFromContext(ctx context.Context) (adminDeviceAuthorization, bool) { + if ctx == nil { + return adminDeviceAuthorization{}, false + } + value, ok := ctx.Value(adminDeviceAuthorizationContextKey{}).(adminDeviceAuthorization) + return value, ok && value.ID != "" && !value.EnrolledAt.IsZero() +} + type AdminHandler struct { Auth *auth.Service Payments *adminpayments.Service @@ -33,6 +63,8 @@ type AdminHandler struct { PairingBaseURL string SecureCookies bool loginSlots chan struct{} + loginMu sync.Mutex + loginFailures map[string]adminLoginFailure mux *http.ServeMux } @@ -40,7 +72,7 @@ func NewAdminHandler(authService *auth.Service, paymentService *adminpayments.Se settingsService *operator.SettingsService, profileService *profiles.Service, relayService *relay.Service, webhookService *webhooks.Service) *AdminHandler { h := &AdminHandler{Auth: authService, Payments: paymentService, Operator: operatorService, Settings: settingsService, Profiles: profileService, Relay: relayService, Webhooks: webhookService, SecureCookies: true, - loginSlots: make(chan struct{}, adminLoginConcurrency), mux: http.NewServeMux()} + loginSlots: make(chan struct{}, adminLoginConcurrency), loginFailures: make(map[string]adminLoginFailure), mux: http.NewServeMux()} h.registerRoutes() return h } @@ -84,45 +116,50 @@ func (h *AdminHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { h.mux.ServeHTTP(w, r) return } - deviceID, ok := h.deviceAuthorization(w, r) + requireEpoch := r.Method == http.MethodPatch && r.URL.Path == "/admin/profiles/active/destination" + deviceID, enrolledAt, ok := h.deviceAuthorization(w, r, requireEpoch) if !ok { writeError(w, http.StatusUnauthorized, "unauthorized", "Admin or connected-device authentication is required") return } if !deviceOperationalRoute(r.Method, r.URL.Path) { - _ = deviceID writeError(w, http.StatusForbidden, "admin_required", "This setting requires the web admin session") return } + r = r.WithContext(context.WithValue(r.Context(), adminDeviceAuthorizationContextKey{}, adminDeviceAuthorization{ + ID: deviceID, EnrolledAt: enrolledAt, + })) h.mux.ServeHTTP(w, r) } const adminDeviceBodyLimit = 256 << 10 -func (h *AdminHandler) deviceAuthorization(w http.ResponseWriter, r *http.Request) (string, bool) { +func (h *AdminHandler) deviceAuthorization(w http.ResponseWriter, r *http.Request, requireEpoch bool) (string, time.Time, bool) { if h.Relay == nil { - return "", false + return "", time.Time{}, false } deviceID := strings.TrimSpace(r.Header.Get("X-PayGate-Relay-Device")) timestamp := strings.TrimSpace(r.Header.Get("X-PayGate-Relay-Time")) signature := strings.TrimSpace(r.Header.Get("X-PayGate-Relay-Signature")) - if deviceID == "" || timestamp == "" || signature == "" { - return "", false + enrollmentEpoch := strings.TrimSpace(r.Header.Get("X-PayGate-Relay-Epoch")) + if deviceID == "" || timestamp == "" || signature == "" || (requireEpoch && enrollmentEpoch == "") { + return "", time.Time{}, false } var body []byte if r.Body != nil { raw, err := io.ReadAll(io.LimitReader(r.Body, adminDeviceBodyLimit+1)) if err != nil || len(raw) > adminDeviceBodyLimit { - return "", false + return "", time.Time{}, false } body = raw r.Body = io.NopCloser(bytes.NewReader(raw)) } target := r.URL.RequestURI() - id, err := h.Relay.AuthenticateDevice(r.Context(), relay.RequestAuth{ - DeviceID: deviceID, Timestamp: timestamp, Signature: signature, Method: r.Method, Path: target, + id, enrolledAt, err := h.Relay.AuthenticateDeviceWithEpoch(r.Context(), relay.RequestAuth{ + DeviceID: deviceID, Timestamp: timestamp, Signature: signature, EnrollmentEpoch: enrollmentEpoch, + Method: r.Method, Path: target, }, body) - return id, err == nil + return id, enrolledAt, err == nil } func deviceOperationalRoute(method, path string) bool { @@ -158,6 +195,13 @@ func (h *AdminHandler) login(w http.ResponseWriter, r *http.Request) { writeError(w, http.StatusBadRequest, "invalid_request", err.Error()) return } + key := loginRemoteKey(r) + now := time.Now().UTC() + if allowed, retry := h.loginAttemptAllowed(key, now); !allowed { + w.Header().Set("Retry-After", strconv.Itoa(retryAfterSeconds(retry))) + writeError(w, http.StatusTooManyRequests, "login_rate_limited", "Too many failed login attempts; try again later") + return + } if !h.acquireLoginSlot() { w.Header().Set("Retry-After", "1") writeError(w, http.StatusTooManyRequests, "login_busy", "Too many login attempts are already being verified") @@ -167,12 +211,17 @@ func (h *AdminHandler) login(w http.ResponseWriter, r *http.Request) { session, err := h.Auth.CreateAdminSession(r.Context(), input.Password) if err != nil { if errors.Is(err, auth.ErrInvalidCredentials) || errors.Is(err, auth.ErrNotInitialized) { + h.recordLoginFailure(key, now) writeError(w, http.StatusUnauthorized, "invalid_credentials", "Password is incorrect") return } + if writeAdminLoginServiceError(w, err) { + return + } writeError(w, http.StatusInternalServerError, "internal_error", "PayGate could not create the admin session") return } + h.clearLoginFailures(key) h.setAdminCookie(w, session.Token, session.ExpiresAt) response := adminLoginResponse{ExpiresAt: session.ExpiresAt} if strings.EqualFold(strings.TrimSpace(input.Client), "android") { @@ -181,6 +230,15 @@ func (h *AdminHandler) login(w http.ResponseWriter, r *http.Request) { writeJSON(w, http.StatusOK, response) } +func writeAdminLoginServiceError(w http.ResponseWriter, err error) bool { + if !errors.Is(err, storage.ErrBusy) { + return false + } + w.Header().Set("Retry-After", "1") + writeError(w, http.StatusServiceUnavailable, "login_retryable", "PayGate is temporarily busy; try again shortly") + return true +} + func (h *AdminHandler) acquireLoginSlot() bool { if h == nil || h.loginSlots == nil { return true @@ -199,11 +257,109 @@ func (h *AdminHandler) releaseLoginSlot() { } <-h.loginSlots } +func loginRemoteKey(r *http.Request) string { + if r == nil { + return "unknown" + } + remote := strings.TrimSpace(r.RemoteAddr) + if host, _, err := net.SplitHostPort(remote); err == nil && host != "" { + return host + } + if remote != "" { + return remote + } + return "unknown" +} + +func retryAfterSeconds(duration time.Duration) int { + if duration <= 0 { + return 1 + } + return int((duration + time.Second - 1) / time.Second) +} + +func (h *AdminHandler) loginAttemptAllowed(key string, now time.Time) (bool, time.Duration) { + if h == nil || h.loginFailures == nil { + return true, 0 + } + h.loginMu.Lock() + defer h.loginMu.Unlock() + h.pruneLoginFailuresLocked(now) + state, ok := h.loginFailures[key] + if !ok { + return true, 0 + } + if now.Before(state.blockedUntil) { + state.lastSeen = now + h.loginFailures[key] = state + return false, state.blockedUntil.Sub(now) + } + if state.firstFailure.IsZero() || now.Sub(state.firstFailure) >= adminLoginFailureWindow { + delete(h.loginFailures, key) + } + return true, 0 +} + +func (h *AdminHandler) recordLoginFailure(key string, now time.Time) { + if h == nil || h.loginFailures == nil { + return + } + h.loginMu.Lock() + defer h.loginMu.Unlock() + h.pruneLoginFailuresLocked(now) + state := h.loginFailures[key] + if state.firstFailure.IsZero() || now.Sub(state.firstFailure) >= adminLoginFailureWindow { + state = adminLoginFailure{firstFailure: now} + } + state.failures++ + state.lastSeen = now + if state.failures >= adminLoginFailureLimit { + state.blockedUntil = now.Add(adminLoginBlockDuration) + } + h.loginFailures[key] = state +} + +func (h *AdminHandler) clearLoginFailures(key string) { + if h == nil || h.loginFailures == nil { + return + } + h.loginMu.Lock() + delete(h.loginFailures, key) + h.loginMu.Unlock() +} + +func (h *AdminHandler) pruneLoginFailuresLocked(now time.Time) { + for key, state := range h.loginFailures { + if now.After(state.blockedUntil) && now.Sub(state.lastSeen) >= adminLoginFailureWindow { + delete(h.loginFailures, key) + } + } + for len(h.loginFailures) >= adminLoginThrottleEntries { + var oldestKey string + var oldest time.Time + for key, state := range h.loginFailures { + if oldestKey == "" || state.lastSeen.Before(oldest) { + oldestKey, oldest = key, state.lastSeen + } + } + if oldestKey == "" { + return + } + delete(h.loginFailures, oldestKey) + } +} func (h *AdminHandler) logout(w http.ResponseWriter, r *http.Request) { token, err := h.adminToken(r) if err == nil { - _ = h.Auth.RevokeAdminSession(r.Context(), token) + err = h.Auth.RevokeAdminSession(r.Context(), token) + if err != nil && !errors.Is(err, auth.ErrInvalidSession) { + if writeStorageBusyError(w, err) { + return + } + writeError(w, http.StatusInternalServerError, "internal_error", "Could not revoke admin session") + return + } } h.clearAdminCookie(w) w.WriteHeader(http.StatusNoContent) @@ -259,6 +415,9 @@ func (h *AdminHandler) changePassword(w http.ResponseWriter, r *http.Request) { return } if err := h.Auth.ChangePassword(r.Context(), input.CurrentPassword, input.NewPassword); err != nil { + if writeStorageBusyError(w, err) { + return + } switch { case errors.Is(err, auth.ErrInvalidCredentials): writeError(w, http.StatusUnauthorized, "invalid_credentials", "Current password is incorrect") diff --git a/internal/v4/httpapi/admin_device.go b/internal/v4/httpapi/admin_device.go index 9bed559..48ec8df 100644 --- a/internal/v4/httpapi/admin_device.go +++ b/internal/v4/httpapi/admin_device.go @@ -9,7 +9,6 @@ import ( ) type pairingSessionRequest struct { - ReplaceExisting bool `json:"replace_existing,omitempty"` } func (h *AdminHandler) getDevice(w http.ResponseWriter, r *http.Request) { @@ -43,14 +42,16 @@ func (h *AdminHandler) createPairingSession(w http.ResponseWriter, r *http.Reque writeError(w, http.StatusBadRequest, "invalid_request", err.Error()) return } - session, err := h.Relay.CreatePairing(r.Context(), input.ReplaceExisting) + session, err := h.Relay.CreatePairing(r.Context()) if err != nil { + if writeStorageBusyError(w, err) { + return + } writeError(w, http.StatusInternalServerError, "internal_error", "Could not create pairing session") return } response := map[string]any{ "token": session.Token, "expires_at": session.ExpiresAt, - "replace_existing": session.ReplaceExisting, } if base := strings.TrimRight(strings.TrimSpace(h.PairingBaseURL), "/"); base != "" { response["pairing_url"] = base + "/device/pair/" + session.Token @@ -68,6 +69,9 @@ func (h *AdminHandler) revokeDevice(w http.ResponseWriter, r *http.Request) { writeError(w, http.StatusNotFound, "device_not_found", "PayGate device not found or already revoked") return } + if writeStorageBusyError(w, err) { + return + } writeError(w, http.StatusInternalServerError, "internal_error", "Could not revoke PayGate device") return } @@ -80,6 +84,9 @@ func (h *AdminHandler) retryWebhook(w http.ResponseWriter, r *http.Request) { return } if err := h.Webhooks.RetryOne(r.Context(), r.PathValue("id")); err != nil { + if writeStorageBusyError(w, err) { + return + } writeError(w, http.StatusConflict, "webhook_not_retryable", "Webhook is not retryable") return } diff --git a/internal/v4/httpapi/admin_keys.go b/internal/v4/httpapi/admin_keys.go index 79cb142..743fd4c 100644 --- a/internal/v4/httpapi/admin_keys.go +++ b/internal/v4/httpapi/admin_keys.go @@ -15,6 +15,9 @@ type apiKeyCreateRequest struct { func (h *AdminHandler) listAPIKeys(w http.ResponseWriter, r *http.Request) { items, err := h.Auth.ListAPIKeys(r.Context()) if err != nil { + if writeStorageBusyError(w, err) { + return + } writeError(w, http.StatusInternalServerError, "internal_error", "Could not load API keys") return } @@ -37,6 +40,9 @@ func (h *AdminHandler) createAPIKey(w http.ResponseWriter, r *http.Request) { writeError(w, http.StatusBadRequest, "invalid_api_key", err.Error()) return } + if writeStorageBusyError(w, err) { + return + } writeError(w, http.StatusInternalServerError, "internal_error", "Could not create API key") return } @@ -53,6 +59,9 @@ func (h *AdminHandler) revokeAPIKey(w http.ResponseWriter, r *http.Request) { writeError(w, http.StatusNotFound, "api_key_not_found", "API key not found or already revoked") return } + if writeStorageBusyError(w, err) { + return + } writeError(w, http.StatusInternalServerError, "internal_error", "Could not revoke API key") return } diff --git a/internal/v4/httpapi/admin_payments.go b/internal/v4/httpapi/admin_payments.go index 1a55da3..6e12e75 100644 --- a/internal/v4/httpapi/admin_payments.go +++ b/internal/v4/httpapi/admin_payments.go @@ -33,6 +33,9 @@ type adminPaymentResponse struct { PayerName string `json:"payer_name,omitempty"` PayerUPIID string `json:"payer_upi_id,omitempty"` InternalNote string `json:"internal_note,omitempty"` + EvidencePackage string `json:"evidence_package,omitempty"` + EvidenceDeviceName string `json:"evidence_device_name,omitempty"` + EvidenceReceivedAt *time.Time `json:"evidence_received_at,omitempty"` } func adminPayment(p adminpayments.Payment) adminPaymentResponse { @@ -44,6 +47,7 @@ func adminPayment(p adminpayments.Payment) adminPaymentResponse { TransactionNote: payments.TransactionNote(p.ID), Status: p.Status, CreatedAt: p.CreatedAt, ExpiresAt: p.ExpiresAt, GraceUntil: p.GraceUntil, ReuseAfter: p.ReuseAfter, PaidAt: p.PaidAt, PayerName: p.PayerName, PayerUPIID: p.PayerUPIID, InternalNote: p.InternalNote, + EvidencePackage: p.EvidencePackage, EvidenceDeviceName: p.EvidenceDeviceName, EvidenceReceivedAt: p.EvidenceReceivedAt, } } @@ -142,6 +146,9 @@ func (h *AdminHandler) editAdminPayment(w http.ResponseWriter, r *http.Request) } func writeAdminPaymentError(w http.ResponseWriter, err error) { + if writeStorageBusyError(w, err) { + return + } switch { case errors.Is(err, adminpayments.ErrPaymentNotFound): writeError(w, http.StatusNotFound, "payment_not_found", "Payment not found") diff --git a/internal/v4/httpapi/admin_settings.go b/internal/v4/httpapi/admin_settings.go index e9b53f4..790cbc6 100644 --- a/internal/v4/httpapi/admin_settings.go +++ b/internal/v4/httpapi/admin_settings.go @@ -62,6 +62,9 @@ func (h *AdminHandler) updateWebhookSettings(w http.ResponseWriter, r *http.Requ writeError(w, http.StatusBadRequest, "invalid_webhook", err.Error()) return } + if writeStorageBusyError(w, err) { + return + } writeError(w, http.StatusInternalServerError, "internal_error", "Could not update webhook settings") return } @@ -132,17 +135,28 @@ func (h *AdminHandler) updateProfileDestination(w http.ResponseWriter, r *http.R return } id := strings.TrimSpace(r.PathValue("id")) - if id == "active" { - active, err := h.Profiles.Active(r.Context()) - if err != nil { - writeProfileError(w, err) - return + var profile profiles.Profile + var err error + deviceAuth, hasDevice := adminDeviceAuthorizationFromContext(r.Context()) + if hasDevice { + if id == "active" { + profile, err = h.Profiles.UpdateActiveDestinationForRelay(r.Context(), deviceAuth.ID, deviceAuth.EnrolledAt, profiles.DestinationInput{ + UPIID: input.UPIID, PayeeName: input.PayeeName, + }) + } else { + profile, err = h.Profiles.UpdateDestinationForRelay(r.Context(), deviceAuth.ID, deviceAuth.EnrolledAt, id, profiles.DestinationInput{ + UPIID: input.UPIID, PayeeName: input.PayeeName, + }) } - id = active.ID + } else if id == "active" { + profile, err = h.Profiles.UpdateActiveDestination(r.Context(), profiles.DestinationInput{ + UPIID: input.UPIID, PayeeName: input.PayeeName, + }) + } else { + profile, err = h.Profiles.UpdateDestination(r.Context(), id, profiles.DestinationInput{ + UPIID: input.UPIID, PayeeName: input.PayeeName, + }) } - profile, err := h.Profiles.UpdateDestination(r.Context(), id, profiles.DestinationInput{ - UPIID: input.UPIID, PayeeName: input.PayeeName, - }) if err != nil { writeProfileError(w, err) return @@ -164,6 +178,9 @@ func (h *AdminHandler) activateProfile(w http.ResponseWriter, r *http.Request) { writeJSON(w, http.StatusOK, map[string]any{"profile": profile}) } func writeProfileError(w http.ResponseWriter, err error) { + if writeStorageBusyError(w, err) { + return + } switch { case errors.Is(err, profiles.ErrProfileNotFound): writeError(w, http.StatusNotFound, "profile_not_found", "Collection profile not found") @@ -171,6 +188,8 @@ func writeProfileError(w http.ResponseWriter, err error) { writeError(w, http.StatusConflict, "profile_disabled", "Disabled collection profile cannot be activated") case errors.Is(err, profiles.ErrCannotDisableActiveProfile): writeError(w, http.StatusConflict, "active_profile", "Activate another profile before disabling this one") + case errors.Is(err, profiles.ErrRelayDeviceNotAuthorized): + writeError(w, http.StatusUnauthorized, "relay_device_revoked", "Relay device is no longer authorized") case errors.Is(err, profiles.ErrInvalidProfile): writeError(w, http.StatusBadRequest, "invalid_profile", err.Error()) default: diff --git a/internal/v4/httpapi/admin_test.go b/internal/v4/httpapi/admin_test.go index a09991e..3c94813 100644 --- a/internal/v4/httpapi/admin_test.go +++ b/internal/v4/httpapi/admin_test.go @@ -8,9 +8,11 @@ import ( "crypto/rand" "crypto/sha256" "crypto/x509" + "database/sql" "encoding/base64" "encoding/json" "encoding/pem" + "fmt" "net/http" "net/http/httptest" "path/filepath" @@ -195,6 +197,189 @@ func TestAdminLoginRejectsWhenVerificationCapacityIsSaturated(t *testing.T) { } _ = loginAdmin(t, f.handler, false) } +func TestAdminLoginDatabaseBusyIsRetryable(t *testing.T) { + f := newAdminHTTPFixture(t) + ctx := context.Background() + f.db.SQL.SetMaxOpenConns(2) + first, err := f.db.SQL.Conn(ctx) + if err != nil { + t.Fatal(err) + } + second, err := f.db.SQL.Conn(ctx) + if err != nil { + first.Close() + t.Fatal(err) + } + for _, conn := range []*sql.Conn{first, second} { + if _, err := conn.ExecContext(ctx, `PRAGMA busy_timeout=1`); err != nil { + first.Close() + second.Close() + t.Fatal(err) + } + } + first.Close() + second.Close() + + locked := make(chan struct{}) + release := make(chan struct{}) + txErr := make(chan error, 1) + go func() { + txErr <- f.db.WithImmediateTx(ctx, func(*storage.ImmediateTx) error { + close(locked) + <-release + return nil + }) + }() + select { + case <-locked: + case err := <-txErr: + t.Fatalf("lock transaction failed before starting: %v", err) + case <-time.After(2 * time.Second): + t.Fatal("lock transaction did not start") + } + + req := httptest.NewRequest(http.MethodPost, "/admin/session", strings.NewReader(`{"password":"correct horse battery staple"}`)) + req.Header.Set("Content-Type", "application/json") + rr := httptest.NewRecorder() + f.handler.ServeHTTP(rr, req) + close(release) + if err := <-txErr; err != nil { + t.Fatal(err) + } + if rr.Code != http.StatusServiceUnavailable || rr.Header().Get("Retry-After") != "1" { + t.Fatalf("busy status=%d retry=%q body=%s", rr.Code, rr.Header().Get("Retry-After"), rr.Body.String()) + } + body := rr.Body.String() + if !strings.Contains(body, `"code":"login_retryable"`) || strings.Contains(strings.ToLower(body), "sqlite") { + t.Fatalf("unexpected busy response: %s", body) + } +} +func TestAdminLogoutDatabaseBusyPreservesSession(t *testing.T) { + f := newAdminHTTPFixture(t) + ctx := context.Background() + f.db.SQL.SetMaxOpenConns(2) + first, err := f.db.SQL.Conn(ctx) + if err != nil { + t.Fatal(err) + } + second, err := f.db.SQL.Conn(ctx) + if err != nil { + first.Close() + t.Fatal(err) + } + for _, conn := range []*sql.Conn{first, second} { + if _, err := conn.ExecContext(ctx, `PRAGMA busy_timeout=1`); err != nil { + first.Close() + second.Close() + t.Fatal(err) + } + } + first.Close() + second.Close() + + locked := make(chan struct{}) + release := make(chan struct{}) + txErr := make(chan error, 1) + go func() { + txErr <- f.db.WithImmediateTx(ctx, func(*storage.ImmediateTx) error { + close(locked) + <-release + return nil + }) + }() + select { + case <-locked: + case err := <-txErr: + t.Fatalf("lock transaction failed before starting: %v", err) + case <-time.After(2 * time.Second): + t.Fatal("lock transaction did not start") + } + + failed := adminRequest(t, f, http.MethodDelete, "/admin/session", nil, false) + close(release) + if err := <-txErr; err != nil { + t.Fatal(err) + } + if failed.Code != http.StatusServiceUnavailable || failed.Header().Get("Retry-After") != "1" { + t.Fatalf("busy status=%d retry=%q body=%s", failed.Code, failed.Header().Get("Retry-After"), failed.Body.String()) + } + if strings.Contains(failed.Body.String(), `"code":"retryable_busy"`) == false || + strings.Contains(strings.ToLower(failed.Body.String()), "sqlite") { + t.Fatalf("unexpected busy response: %s", failed.Body.String()) + } + if failed.Header().Get("Set-Cookie") != "" { + t.Fatalf("failed logout cleared cookie: %s", failed.Header().Get("Set-Cookie")) + } + + stillValid := adminRequest(t, f, http.MethodGet, "/admin/overview", nil, false) + if stillValid.Code != http.StatusOK { + t.Fatalf("session after retryable logout status=%d body=%s", stillValid.Code, stillValid.Body.String()) + } + + succeeded := adminRequest(t, f, http.MethodDelete, "/admin/session", nil, false) + if succeeded.Code != http.StatusNoContent { + t.Fatalf("successful logout status=%d body=%s", succeeded.Code, succeeded.Body.String()) + } + cleared := false + for _, cookie := range succeeded.Result().Cookies() { + if cookie.Name == adminCookieName && cookie.MaxAge < 0 { + cleared = true + break + } + } + if !cleared { + t.Fatalf("successful logout did not clear cookie: %s", succeeded.Header().Get("Set-Cookie")) + } + afterLogout := adminRequest(t, f, http.MethodGet, "/admin/overview", nil, false) + if afterLogout.Code != http.StatusUnauthorized { + t.Fatalf("revoked session status=%d body=%s", afterLogout.Code, afterLogout.Body.String()) + } +} +func TestStorageBusyHTTPMappingsAreRetryable(t *testing.T) { + t.Parallel() + err := fmt.Errorf("transaction failed: %w", storage.ErrBusy) + cases := map[string]func(http.ResponseWriter, error){ + "merchant payment": writePaymentError, + "relay pairing": writeRelayPairError, + "relay request": writeRelayError, + "profile": writeProfileError, + "admin payment": writeAdminPaymentError, + } + for name, write := range cases { + t.Run(name, func(t *testing.T) { + recorder := httptest.NewRecorder() + write(recorder, err) + if recorder.Code != http.StatusServiceUnavailable || recorder.Header().Get("Retry-After") != "1" { + t.Fatalf("status=%d retry=%q body=%s", recorder.Code, recorder.Header().Get("Retry-After"), recorder.Body.String()) + } + body := recorder.Body.String() + if !strings.Contains(body, `"code":"retryable_busy"`) || strings.Contains(strings.ToLower(body), "sqlite") { + t.Fatalf("unexpected busy response: %s", body) + } + }) + } +} +func TestAdminLoginThrottlesRepeatedFailuresByRemoteAddress(t *testing.T) { + f := newAdminHTTPFixture(t) + for attempt := 0; attempt < adminLoginFailureLimit; attempt++ { + req := httptest.NewRequest(http.MethodPost, "/admin/session", strings.NewReader(`{"password":"wrong password"}`)) + req.RemoteAddr = "198.51.100.17:4567" + req.Header.Set("Content-Type", "application/json") + rr := httptest.NewRecorder() + f.handler.ServeHTTP(rr, req) + if rr.Code != http.StatusUnauthorized { + t.Fatalf("failed attempt %d status=%d body=%s", attempt+1, rr.Code, rr.Body.String()) + } + } + req := httptest.NewRequest(http.MethodPost, "/admin/session", strings.NewReader(`{"password":"wrong password"}`)) + req.RemoteAddr = "198.51.100.17:4567" + req.Header.Set("Content-Type", "application/json") + rr := httptest.NewRecorder() + f.handler.ServeHTTP(rr, req) + if rr.Code != http.StatusTooManyRequests || rr.Header().Get("Retry-After") == "" || !strings.Contains(rr.Body.String(), `"code":"login_rate_limited"`) { + t.Fatalf("throttled login status=%d retry=%q body=%s", rr.Code, rr.Header().Get("Retry-After"), rr.Body.String()) + } +} func TestAdminPaymentsFilterDetailAndEdit(t *testing.T) { f := newAdminHTTPFixture(t) @@ -313,7 +498,7 @@ func TestDeviceOperationalRouteKeepsPhoneEvidenceOnly(t *testing.T) { } } -func pairAdminTestDevice(t *testing.T, f adminHTTPFixture) (*ecdsa.PrivateKey, string) { +func pairAdminTestDevice(t *testing.T, f adminHTTPFixture) (*ecdsa.PrivateKey, string, int64) { t.Helper() privateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) if err != nil { @@ -324,7 +509,7 @@ func pairAdminTestDevice(t *testing.T, f adminHTTPFixture) (*ecdsa.PrivateKey, s t.Fatal(err) } publicKeyPEM := pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: der}) - session, err := f.handler.Relay.CreatePairing(context.Background(), false) + session, err := f.handler.Relay.CreatePairing(context.Background()) if err != nil { t.Fatal(err) } @@ -334,13 +519,26 @@ func pairAdminTestDevice(t *testing.T, f adminHTTPFixture) (*ecdsa.PrivateKey, s if err != nil { t.Fatal(err) } - return privateKey, paired.DeviceID + return privateKey, paired.DeviceID, paired.EnrolledAtMS } func signedDeviceAdminRequest(t *testing.T, f adminHTTPFixture, privateKey *ecdsa.PrivateKey, deviceID, method, path string, body []byte) *httptest.ResponseRecorder { t.Helper() - timestamp := strconv.FormatInt(time.Now().UTC().UnixMilli(), 10) + return signedDeviceAdminRequestAt(t, f, privateKey, deviceID, method, path, body, 0) +} + +func signedDeviceAdminRequestWithEpoch(t *testing.T, f adminHTTPFixture, privateKey *ecdsa.PrivateKey, deviceID, method, path string, body []byte, enrollmentEpoch int64) *httptest.ResponseRecorder { + t.Helper() + return signedDeviceAdminRequestAt(t, f, privateKey, deviceID, method, path, body, enrollmentEpoch) +} + +func signedDeviceAdminRequestAt(t *testing.T, f adminHTTPFixture, privateKey *ecdsa.PrivateKey, deviceID, method, path string, body []byte, enrollmentEpoch int64) *httptest.ResponseRecorder { + t.Helper() + timestamp := strconv.FormatInt(f.handler.Relay.Now().UTC().UnixMilli(), 10) canonical := relay.CanonicalRequest(method, path, timestamp, body) + if enrollmentEpoch > 0 { + canonical = relay.CanonicalRequestWithEpoch(method, path, timestamp, enrollmentEpoch, body) + } digest := sha256.Sum256([]byte(canonical)) signature, err := ecdsa.SignASN1(rand.Reader, privateKey, digest[:]) if err != nil { @@ -353,17 +551,86 @@ func signedDeviceAdminRequest(t *testing.T, f adminHTTPFixture, privateKey *ecds req.Header.Set("X-PayGate-Relay-Device", deviceID) req.Header.Set("X-PayGate-Relay-Time", timestamp) req.Header.Set("X-PayGate-Relay-Signature", base64.StdEncoding.EncodeToString(signature)) + if enrollmentEpoch > 0 { + req.Header.Set(relayEpochHeader, strconv.FormatInt(enrollmentEpoch, 10)) + } rr := httptest.NewRecorder() f.handler.ServeHTTP(rr, req) return rr } +func TestDeviceDestinationRejectsMissingAndStaleEnrollmentEpoch(t *testing.T) { + f := newAdminHTTPFixture(t) + current := time.Date(2026, 9, 8, 0, 0, 0, 0, time.UTC) + f.handler.Relay.Now = func() time.Time { return current } + privateKey, deviceID, epoch1 := pairAdminTestDevice(t, f) + + legacyBody := []byte(`{"upi_id":"legacy-replay@upi","payee_name":"PayGate"}`) + legacy := signedDeviceAdminRequest(t, f, privateKey, deviceID, http.MethodPatch, "/admin/profiles/active/destination", legacyBody) + if legacy.Code != http.StatusUnauthorized { + t.Fatalf("missing epoch destination status=%d body=%s", legacy.Code, legacy.Body.String()) + } + + oldBody := []byte(`{"upi_id":"stale-replay@upi","payee_name":"PayGate"}`) + oldTimestamp := strconv.FormatInt(current.UnixMilli(), 10) + oldCanonical := relay.CanonicalRequestWithEpoch(http.MethodPatch, "/admin/profiles/active/destination", oldTimestamp, epoch1, oldBody) + oldDigest := sha256.Sum256([]byte(oldCanonical)) + oldSignature, err := ecdsa.SignASN1(rand.Reader, privateKey, oldDigest[:]) + if err != nil { + t.Fatal(err) + } + oldRequest := httptest.NewRequest(http.MethodPatch, "/admin/profiles/active/destination", bytes.NewReader(oldBody)) + oldRequest.Header.Set("Content-Type", "application/json") + oldRequest.Header.Set("X-PayGate-Relay-Device", deviceID) + oldRequest.Header.Set("X-PayGate-Relay-Time", oldTimestamp) + oldRequest.Header.Set("X-PayGate-Relay-Signature", base64.StdEncoding.EncodeToString(oldSignature)) + oldRequest.Header.Set(relayEpochHeader, strconv.FormatInt(epoch1, 10)) + + der, err := x509.MarshalPKIXPublicKey(&privateKey.PublicKey) + if err != nil { + t.Fatal(err) + } + publicKeyPEM := pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: der}) + current = current.Add(time.Second) + session, err := f.handler.Relay.CreatePairing(context.Background()) + if err != nil { + t.Fatal(err) + } + repaired, err := f.handler.Relay.PairDevice(context.Background(), relay.PairDeviceInput{ + Token: session.Token, Name: "evidence-only test phone", PublicKeyPEM: string(publicKeyPEM), AppVersion: "test", + }) + if err != nil { + t.Fatal(err) + } + if repaired.DeviceID != deviceID || repaired.EnrolledAtMS == epoch1 { + t.Fatalf("re-pair identity=%q epoch=%d old_epoch=%d", repaired.DeviceID, repaired.EnrolledAtMS, epoch1) + } + + rr := httptest.NewRecorder() + f.handler.ServeHTTP(rr, oldRequest) + if rr.Code != http.StatusUnauthorized { + t.Fatalf("stale epoch replay status=%d body=%s", rr.Code, rr.Body.String()) + } + var upiID string + if err := f.db.SQL.QueryRow(`SELECT upi_id FROM collection_profiles WHERE active=1`).Scan(&upiID); err != nil { + t.Fatal(err) + } + if upiID != "paygate@paytm" { + t.Fatalf("stale replay changed destination to %q", upiID) + } + + fresh := signedDeviceAdminRequestWithEpoch(t, f, privateKey, deviceID, http.MethodPatch, "/admin/profiles/active/destination", []byte(`{"upi_id":"fresh@upi","payee_name":"PayGate"}`), repaired.EnrolledAtMS) + if fresh.Code != http.StatusOK || !strings.Contains(fresh.Body.String(), `"upi_id":"fresh@upi"`) { + t.Fatalf("fresh epoch destination status=%d body=%s", fresh.Code, fresh.Body.String()) + } +} + func TestPairedDeviceCannotMutatePaymentOrWebhookAuthority(t *testing.T) { f := newAdminHTTPFixture(t) - privateKey, deviceID := pairAdminTestDevice(t, f) + privateKey, deviceID, enrollmentEpoch := pairAdminTestDevice(t, f) payment := createAdminTestPayment(t, f, "Evidence boundary", "evt_device_auth", "device-auth") - rr := signedDeviceAdminRequest(t, f, privateKey, deviceID, http.MethodPatch, "/admin/payments/"+payment.ID, []byte(`{"status":"paid"}`)) + rr := signedDeviceAdminRequestWithEpoch(t, f, privateKey, deviceID, http.MethodPatch, "/admin/payments/"+payment.ID, []byte(`{"status":"paid"}`), enrollmentEpoch) if rr.Code != http.StatusForbidden || !strings.Contains(rr.Body.String(), `"admin_required"`) { t.Fatalf("device payment edit status=%d body=%s", rr.Code, rr.Body.String()) } @@ -375,13 +642,22 @@ func TestPairedDeviceCannotMutatePaymentOrWebhookAuthority(t *testing.T) { t.Fatalf("device changed payment status to %q", status) } - rr = signedDeviceAdminRequest(t, f, privateKey, deviceID, http.MethodPost, "/admin/webhooks/wh_test/retry", nil) + rr = signedDeviceAdminRequestWithEpoch(t, f, privateKey, deviceID, http.MethodPost, "/admin/webhooks/wh_test/retry", nil, enrollmentEpoch) if rr.Code != http.StatusForbidden || !strings.Contains(rr.Body.String(), `"admin_required"`) { t.Fatalf("device webhook retry status=%d body=%s", rr.Code, rr.Body.String()) } - rr = signedDeviceAdminRequest(t, f, privateKey, deviceID, http.MethodPatch, "/admin/profiles/active/destination", []byte(`{"upi_id":"paygate-new@upi","payee_name":"PayGate"}`)) + rr = signedDeviceAdminRequestWithEpoch(t, f, privateKey, deviceID, http.MethodPatch, "/admin/profiles/active/destination", []byte(`{"upi_id":"paygate-new@upi","payee_name":"PayGate"}`), enrollmentEpoch) if rr.Code != http.StatusOK || !strings.Contains(rr.Body.String(), `"upi_id":"paygate-new@upi"`) { t.Fatalf("device destination update status=%d body=%s", rr.Code, rr.Body.String()) } + + rr = adminRequest(t, f, http.MethodDelete, "/admin/device/"+deviceID, nil, false) + if rr.Code != http.StatusNoContent { + t.Fatalf("device revoke status=%d body=%s", rr.Code, rr.Body.String()) + } + rr = signedDeviceAdminRequest(t, f, privateKey, deviceID, http.MethodPatch, "/admin/profiles/active/destination", []byte(`{"upi_id":"revoked@upi"}`)) + if rr.Code != http.StatusUnauthorized || !strings.Contains(rr.Body.String(), `"unauthorized"`) { + t.Fatalf("revoked device destination status=%d body=%s", rr.Code, rr.Body.String()) + } } diff --git a/internal/v4/httpapi/merchant.go b/internal/v4/httpapi/merchant.go index 2f46952..706dd15 100644 --- a/internal/v4/httpapi/merchant.go +++ b/internal/v4/httpapi/merchant.go @@ -12,6 +12,7 @@ import ( "github.com/Phloraxx/payment-api/internal/v4/auth" "github.com/Phloraxx/payment-api/internal/v4/payments" + "github.com/Phloraxx/payment-api/internal/v4/storage" ) const ( @@ -193,7 +194,19 @@ func decodeStrictJSON(w http.ResponseWriter, r *http.Request, dst any) error { } return nil } +func writeStorageBusyError(w http.ResponseWriter, err error) bool { + if !errors.Is(err, storage.ErrBusy) { + return false + } + w.Header().Set("Retry-After", "1") + writeError(w, http.StatusServiceUnavailable, "retryable_busy", "PayGate is temporarily busy; try again shortly") + return true +} + func writePaymentError(w http.ResponseWriter, err error) { + if writeStorageBusyError(w, err) { + return + } switch { case errors.Is(err, payments.ErrInvalidPaymentInput): writeError(w, http.StatusBadRequest, "invalid_payment", "Invalid payment request") diff --git a/internal/v4/httpapi/relay.go b/internal/v4/httpapi/relay.go index 2634a2b..33e8bc0 100644 --- a/internal/v4/httpapi/relay.go +++ b/internal/v4/httpapi/relay.go @@ -15,6 +15,7 @@ const ( relayDeviceHeader = "X-PayGate-Relay-Device" relayTimeHeader = "X-PayGate-Relay-Time" relaySignatureHeader = "X-PayGate-Relay-Signature" + relayEpochHeader = "X-PayGate-Relay-Epoch" ) type RelayHandler struct { @@ -28,7 +29,6 @@ func NewRelayHandler(service *relay.Service) *RelayHandler { h.mux.HandleFunc("POST "+relay.EventPath, h.event) h.mux.HandleFunc("POST "+relay.HeartbeatPath, h.heartbeat) h.mux.HandleFunc("GET "+relay.DevicePath, h.getThisDevice) - h.mux.HandleFunc("DELETE "+relay.DevicePath, h.disconnectDevice) return h } func (h *RelayHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { @@ -73,8 +73,7 @@ func (h *RelayHandler) pair(w http.ResponseWriter, r *http.Request) { return } writeJSON(w, http.StatusOK, map[string]any{ - "device_id": result.DeviceID, "enabled": result.Enabled, - "replaced_device_id": emptyToNil(result.ReplacedDeviceID), + "device_id": result.DeviceID, "enabled": result.Enabled, "enrolled_at_ms": result.EnrolledAtMS, }) } @@ -118,19 +117,6 @@ func (h *RelayHandler) getThisDevice(w http.ResponseWriter, r *http.Request) { writeJSON(w, http.StatusOK, map[string]any{"device": device}) } -func (h *RelayHandler) disconnectDevice(w http.ResponseWriter, r *http.Request) { - deviceID, err := h.Relay.AuthenticateDevice(r.Context(), relayAuth(r, relay.DevicePath), nil) - if err != nil { - writeRelayError(w, err) - return - } - if err := h.Relay.RevokeDevice(r.Context(), deviceID); err != nil { - writeRelayError(w, relayErrorForDeviceRevoke(err)) - return - } - w.WriteHeader(http.StatusNoContent) -} - func relayErrorForDeviceRevoke(err error) error { if errors.Is(err, relay.ErrInvalidDevice) { return &relay.Error{Code: "UNKNOWN_RELAY_DEVICE", Message: "relay device is not enrolled or is disabled", HTTPStatus: http.StatusUnauthorized} @@ -141,7 +127,8 @@ func relayErrorForDeviceRevoke(err error) error { func relayAuth(r *http.Request, path string) relay.RequestAuth { return relay.RequestAuth{ DeviceID: r.Header.Get(relayDeviceHeader), Timestamp: r.Header.Get(relayTimeHeader), - Signature: r.Header.Get(relaySignatureHeader), Method: r.Method, Path: path, + Signature: r.Header.Get(relaySignatureHeader), EnrollmentEpoch: r.Header.Get(relayEpochHeader), + Method: r.Method, Path: path, } } @@ -159,11 +146,12 @@ func readRelayBody(w http.ResponseWriter, r *http.Request) ([]byte, bool) { return raw, true } func writeRelayPairError(w http.ResponseWriter, err error) { + if writeStorageBusyError(w, err) { + return + } switch { case errors.Is(err, relay.ErrPairingTokenInvalid), errors.Is(err, relay.ErrPairingTokenExpired), errors.Is(err, relay.ErrPairingTokenUsed): writeError(w, http.StatusUnauthorized, "invalid_pairing", "Pairing link is invalid or expired") - case errors.Is(err, relay.ErrRelayAlreadyActive): - writeError(w, http.StatusConflict, "device_already_connected", "A PayGate phone is already connected") case errors.Is(err, relay.ErrInvalidDevice): writeError(w, http.StatusBadRequest, "invalid_device", err.Error()) default: @@ -172,6 +160,9 @@ func writeRelayPairError(w http.ResponseWriter, err error) { } func writeRelayError(w http.ResponseWriter, err error) { + if writeStorageBusyError(w, err) { + return + } var relayErr *relay.Error if errors.As(err, &relayErr) { writeError(w, relayErr.HTTPStatus, strings.ToLower(relayErr.Code), relayErr.Message) @@ -179,10 +170,3 @@ func writeRelayError(w http.ResponseWriter, err error) { } writeError(w, http.StatusInternalServerError, "internal_error", "PayGate could not process the relay request") } - -func emptyToNil(value string) any { - if strings.TrimSpace(value) == "" { - return nil - } - return value -} diff --git a/internal/v4/httpapi/relay_test.go b/internal/v4/httpapi/relay_test.go index 93fa1e4..915903e 100644 --- a/internal/v4/httpapi/relay_test.go +++ b/internal/v4/httpapi/relay_test.go @@ -66,9 +66,9 @@ func (f relayHTTPFixture) publicKeyPEM(t *testing.T) string { return string(pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: der})) } -func pairRelayHTTP(t *testing.T, f relayHTTPFixture) { +func pairRelayHTTP(t *testing.T, f relayHTTPFixture) int64 { t.Helper() - session, err := f.service.CreatePairing(context.Background(), false) + session, err := f.service.CreatePairing(context.Background()) if err != nil { t.Fatal(err) } @@ -83,11 +83,25 @@ func pairRelayHTTP(t *testing.T, f relayHTTPFixture) { if rr.Code != http.StatusOK || !strings.Contains(rr.Body.String(), f.device) { t.Fatalf("pair status=%d body=%s", rr.Code, rr.Body.String()) } + var response struct { + EnrolledAtMS int64 `json:"enrolled_at_ms"` + } + if err := json.Unmarshal(rr.Body.Bytes(), &response); err != nil { + t.Fatal(err) + } + var dbEpoch int64 + if err := f.db.SQL.QueryRow(`SELECT enrolled_at FROM relay_devices WHERE id=?`, f.device).Scan(&dbEpoch); err != nil { + t.Fatal(err) + } + if response.EnrolledAtMS != dbEpoch { + t.Fatalf("pair response epoch=%d db epoch=%d", response.EnrolledAtMS, dbEpoch) + } + return response.EnrolledAtMS } -func signedRelayRequest(t *testing.T, f relayHTTPFixture, path string, body []byte) *http.Request { +func signedRelayRequest(t *testing.T, f relayHTTPFixture, path string, body []byte, enrollmentEpoch int64) *http.Request { t.Helper() timestamp := strconv.FormatInt(f.now.UnixMilli(), 10) - canonical := relay.CanonicalRequest(http.MethodPost, path, timestamp, body) + canonical := relay.CanonicalRequestWithEpoch(http.MethodPost, path, timestamp, enrollmentEpoch, body) digest := sha256.Sum256([]byte(canonical)) signature, err := ecdsa.SignASN1(rand.Reader, f.private, digest[:]) if err != nil { @@ -98,34 +112,79 @@ func signedRelayRequest(t *testing.T, f relayHTTPFixture, path string, body []by req.Header.Set(relayDeviceHeader, f.device) req.Header.Set(relayTimeHeader, timestamp) req.Header.Set(relaySignatureHeader, base64.StdEncoding.EncodeToString(signature)) + req.Header.Set(relayEpochHeader, strconv.FormatInt(enrollmentEpoch, 10)) return req } func TestRelayPairHeartbeatAndHealthPersistence(t *testing.T) { f := newRelayHTTPFixture(t) - pairRelayHTTP(t, f) + epoch := pairRelayHTTP(t, f) body := []byte(`{"schema_version":1,"app_version":"0.5.0","android_version":"16","device_model":"motorola edge 60 stylus","notification_access":true,"listener_connected":true,"battery_optimization_exempt":true,"power_save_mode":false,"background_restricted":false,"foreground_service":true,"pending_count":0,"failed_count":2}`) rr := httptest.NewRecorder() - f.handler.ServeHTTP(rr, signedRelayRequest(t, f, relay.HeartbeatPath, body)) + f.handler.ServeHTTP(rr, signedRelayRequest(t, f, relay.HeartbeatPath, body, epoch)) if rr.Code != http.StatusOK { t.Fatalf("heartbeat status=%d body=%s", rr.Code, rr.Body.String()) } - device, err := f.service.ActiveDevice(context.Background()) - if err != nil || device == nil || device.LastHeartbeatAt == nil || device.NotificationAccess == nil || !*device.NotificationAccess || device.FailedCount == nil || *device.FailedCount != 2 { - t.Fatalf("device=%+v err=%v", device, err) + devices, err := f.service.Devices(context.Background()) + if err != nil || len(devices) != 1 || devices[0].LastHeartbeatAt == nil || devices[0].NotificationAccess == nil || !*devices[0].NotificationAccess || devices[0].FailedCount == nil || *devices[0].FailedCount != 2 { + t.Fatalf("devices=%+v err=%v", devices, err) + } +} +func TestRelayHeartbeatDatabaseBusyIsRetryable(t *testing.T) { + f := newRelayHTTPFixture(t) + epoch := pairRelayHTTP(t, f) + ctx := context.Background() + f.db.SQL.SetMaxOpenConns(2) + for i := 0; i < 2; i++ { + conn, err := f.db.SQL.Conn(ctx) + if err != nil { + t.Fatal(err) + } + if _, err := conn.ExecContext(ctx, `PRAGMA busy_timeout=1`); err != nil { + conn.Close() + t.Fatal(err) + } + conn.Close() + } + + locked := make(chan struct{}) + release := make(chan struct{}) + txErr := make(chan error, 1) + go func() { + txErr <- f.db.WithImmediateTx(ctx, func(*storage.ImmediateTx) error { + close(locked) + <-release + return nil + }) + }() + <-locked + + body := []byte(`{"schema_version":1,"app_version":"0.5.0","android_version":"16","device_model":"motorola edge 60 stylus","notification_access":true,"listener_connected":true,"battery_optimization_exempt":true,"power_save_mode":false,"background_restricted":false,"foreground_service":true,"pending_count":0,"failed_count":2}`) + rr := httptest.NewRecorder() + f.handler.ServeHTTP(rr, signedRelayRequest(t, f, relay.HeartbeatPath, body, epoch)) + close(release) + if err := <-txErr; err != nil { + t.Fatal(err) + } + if rr.Code != http.StatusServiceUnavailable || rr.Header().Get("Retry-After") != "1" { + t.Fatalf("busy status=%d retry=%q body=%s", rr.Code, rr.Header().Get("Retry-After"), rr.Body.String()) + } + response := rr.Body.String() + if !strings.Contains(response, `"code":"retryable_busy"`) || strings.Contains(strings.ToLower(response), "sqlite") { + t.Fatalf("unexpected busy response: %s", response) } } func TestRelaySignedEventAndSignatureFailure(t *testing.T) { f := newRelayHTTPFixture(t) - pairRelayHTTP(t, f) + epoch := pairRelayHTTP(t, f) body := []byte(`{"schema_version":1,"event_id":"aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa","package_name":"com.paytm.business","posted_at_ms":1788244200000,"title":"Payment Received on Paytm for Business","text":"₹100.00 Received from Test"}`) rr := httptest.NewRecorder() - f.handler.ServeHTTP(rr, signedRelayRequest(t, f, relay.EventPath, body)) + f.handler.ServeHTTP(rr, signedRelayRequest(t, f, relay.EventPath, body, epoch)) if rr.Code != http.StatusOK || !strings.Contains(rr.Body.String(), `"status":"ignored"`) { t.Fatalf("event status=%d body=%s", rr.Code, rr.Body.String()) } - bad := signedRelayRequest(t, f, relay.EventPath, body) + bad := signedRelayRequest(t, f, relay.EventPath, body, epoch) bad.Header.Set(relaySignatureHeader, base64.StdEncoding.EncodeToString([]byte("bad"))) rr = httptest.NewRecorder() f.handler.ServeHTTP(rr, bad) @@ -144,3 +203,13 @@ func TestRelayRejectsQueryParameters(t *testing.T) { t.Fatalf("query status=%d body=%s", rr.Code, rr.Body.String()) } } + +func TestRelayDeviceCannotSelfRevoke(t *testing.T) { + f := newRelayHTTPFixture(t) + req := httptest.NewRequest(http.MethodDelete, relay.DevicePath, nil) + rr := httptest.NewRecorder() + f.handler.ServeHTTP(rr, req) + if rr.Code != http.StatusMethodNotAllowed { + t.Fatalf("self-revoke status=%d body=%s", rr.Code, rr.Body.String()) + } +} diff --git a/internal/v4/migratev3/migrate_test.go b/internal/v4/migratev3/migrate_test.go index 625abde..f5c4379 100644 --- a/internal/v4/migratev3/migrate_test.go +++ b/internal/v4/migratev3/migrate_test.go @@ -88,6 +88,7 @@ func TestMigratedLegacyIdempotencyKeyFailsClosedOnChangedRequest(t *testing.T) { } defer db.Close() svc := payments.NewService(db) + svc.Now = func() time.Time { return now.Add(time.Hour) } _, err = svc.Create(context.Background(), payments.CreateInput{RequestedAmountPaise: 10000, Name: "Different person", ExternalID: "evt_1", Metadata: []byte(`{"eventId":"evt_1"}`), IdempotencyScope: merchantIdempotencyScope, IdempotencyKey: "idem-late"}) if !errors.Is(err, payments.ErrIdempotencyConflict) { t.Fatalf("changed retry err=%v", err) diff --git a/internal/v4/observations/parser.go b/internal/v4/observations/parser.go index 6eb1339..f9d1975 100644 --- a/internal/v4/observations/parser.go +++ b/internal/v4/observations/parser.go @@ -12,14 +12,33 @@ import ( const ( PaytmBusinessPackage = "com.paytm.business" GoogleMessagesPackage = "com.google.android.apps.messaging" + GmailPackage = "com.google.android.gm" + BHIMPackage = "in.org.npci.upiapp" + GooglePayBusinessPackage = "com.google.android.apps.nbu.paisa.merchant" + GooglePayPackage = "com.google.android.apps.nbu.paisa.user" + PaytmPackage = "net.one97.paytm" + PhonePePackage = "com.phonepe.app" + AmazonPackage = "in.amazon.mShop.android.shopping" + SuperMoneyPackage = "money.super.payments" + Kotak811Package = "com.kotak811mobilebankingapp.instantsavingsupiscanandpayrecharge" GenericNotificationSource = "android_notification" - GenericMessageSource = "android_message" paytmPostTimeRefinementWindow = time.Minute ) +func isTrustedGenericPackage(packageName string) bool { + switch packageName { + case BHIMPackage, GooglePayBusinessPackage, GooglePayPackage, PaytmPackage, + PhonePePackage, AmazonPackage, SuperMoneyPackage, Kotak811Package: + return true + default: + return false + } +} + var ( ErrUnrecognized = errors.New("notification is not a recognized incoming PayGate payment") ErrNonPayGateAmount = errors.New("incoming amount is not a PayGate decimal amount") + ErrAmbiguousAmount = errors.New("notification contains multiple monetary amounts") ) type Snapshot struct { @@ -41,8 +60,9 @@ type Observation struct { } var ( - currencyAmount = `(?:rs\.?|inr|₹)\s*([0-9][0-9,]*(?:\.[0-9]{1,2})?)` - incomingPatterns = []*regexp.Regexp{ + currencyAmount = `(?:rs\.?|inr|₹)\s*([0-9][0-9,]*(?:\.[0-9]+)?)(?:$|[^0-9A-Za-z.]|\.(?:\s|$))` + currencyAmountPattern = regexp.MustCompile(`(?i)` + currencyAmount) + incomingPatterns = []*regexp.Regexp{ regexp.MustCompile(`(?i)\b(?:payment\s+)?received\b.{0,120}?` + currencyAmount), regexp.MustCompile(`(?i)` + currencyAmount + `.{0,80}?\b(?:received|credited|deposited)\b`), regexp.MustCompile(`(?i)\b(?:received|credited|deposited)\b.{0,120}?` + currencyAmount), @@ -57,8 +77,9 @@ var ( } nonPaymentPattern = regexp.MustCompile(`(?i)\b(?:reversal|reversed|refund(?:ed)?|cashback|reward|interest|salary|chargeback|settlement|settled|loan|emi|bill|due|reminder)\b`) debitPattern = regexp.MustCompile(`(?i)\b(?:debited|sent|you\s+paid|paid\s+to|paid\s+for|withdrawn|purchase|spent|transferred\s+to)\b`) + failedPattern = regexp.MustCompile(`(?i)\b(?:failed|failure|declined|decline|unsuccessful|rejected|pending|processing)\b`) upiPattern = regexp.MustCompile(`(?i)[a-z0-9][a-z0-9._-]{0,127}@[a-z0-9][a-z0-9._-]{0,127}`) - fromPattern = regexp.MustCompile(`(?i)\b(?:from|by)\s+(.+?)(?:\s+at\s+\d{1,2}:\d{2}(?:\s*[ap]m)?\b|\s+on\s+|\s+(?:upi\s+)?(?:ref|rrn|utr)|[!|\n]|\.(?:\s|$)|$)`) + fromPattern = regexp.MustCompile(`(?i)\b(?:from|by)\s+(.+?)(?:\s+(?:to|via)\b|\s+at\s+\d{1,2}:\d{2}(?:\s*[ap]m)?\b|\s+on\s+|\s+(?:upi\s+)?(?:ref|rrn|utr)|[!|\n]|\.(?:\s|$)|$)`) paidYouPayerPattern = regexp.MustCompile(`(?i)^(.{1,120}?)\s+paid\s+you\b`) paytmOccurredPattern = regexp.MustCompile(`(?i)\breceived\s+on\s+(\d{1,2}\s+[A-Za-z]{3}\s+\d{4}\s+\d{1,2}:\d{2}\s+(?:AM|PM))\b`) ) @@ -72,21 +93,23 @@ func Parse(snapshot Snapshot) (Observation, error) { if pkg == PaytmBusinessPackage { return parsePaytm(text, snapshot.PostedAt) } - if pkg == GoogleMessagesPackage && strings.Contains(strings.ToLower(text), "kotak") { - return parseKotak(text, snapshot.PostedAt) + if pkg == GoogleMessagesPackage || pkg == GmailPackage { + return Observation{}, ErrUnrecognized } - source := GenericNotificationSource - if pkg == GoogleMessagesPackage { - source = GenericMessageSource + if !isTrustedGenericPackage(pkg) { + return Observation{}, ErrUnrecognized } - return parseGeneric(text, snapshot.PostedAt, source) + return parseGeneric(text, snapshot.PostedAt) } -func parseGeneric(text string, postedAt time.Time, source string) (Observation, error) { +func parseGeneric(text string, postedAt time.Time) (Observation, error) { if rejectedTransactionText(text) { return Observation{}, ErrUnrecognized } - amountText := firstAmount(text, incomingPatterns) + amountText, err := incomingAmount(text, incomingPatterns) + if err != nil { + return Observation{}, err + } if amountText == "" { return Observation{}, ErrUnrecognized } @@ -96,7 +119,7 @@ func parseGeneric(text string, postedAt time.Time, source string) (Observation, } payerName, payerUPI := extractPayer(text) return Observation{ - Source: source, AmountPaise: amount, PayerName: payerName, PayerUPIID: payerUPI, + Source: GenericNotificationSource, AmountPaise: amount, PayerName: payerName, PayerUPIID: payerUPI, OccurredAt: postedAt.UTC(), OccurredAtSource: "notification_posted_at", }, nil } @@ -105,7 +128,10 @@ func parsePaytm(text string, postedAt time.Time) (Observation, error) { if rejectedTransactionText(text) { return Observation{}, ErrUnrecognized } - amountText := firstAmount(text, paytmAmountPatterns) + amountText, err := incomingAmount(text, paytmAmountPatterns) + if err != nil { + return Observation{}, err + } if amountText == "" { return Observation{}, ErrUnrecognized } @@ -127,25 +153,15 @@ func parsePaytm(text string, postedAt time.Time) (Observation, error) { PayerName: payerName, PayerUPIID: payerUPI, OccurredAt: occurredAt, OccurredAtSource: source}, nil } -func parseKotak(text string, postedAt time.Time) (Observation, error) { - if rejectedTransactionText(text) { - return Observation{}, ErrUnrecognized - } - amountText := firstAmount(text, incomingPatterns) - if amountText == "" { - return Observation{}, ErrUnrecognized - } - amount, err := parsePayGateAmount(amountText) - if err != nil { - return Observation{}, err - } - payerName, payerUPI := extractPayer(text) - return Observation{Source: "kotak_sms", CollectionProfileID: "kotak", AmountPaise: amount, - PayerName: payerName, PayerUPIID: payerUPI, OccurredAt: postedAt.UTC(), OccurredAtSource: "notification_posted_at"}, nil +func rejectedTransactionText(text string) bool { + return strings.TrimSpace(text) == "" || nonPaymentPattern.MatchString(text) || debitPattern.MatchString(text) || failedPattern.MatchString(text) } -func rejectedTransactionText(text string) bool { - return strings.TrimSpace(text) == "" || nonPaymentPattern.MatchString(text) || debitPattern.MatchString(text) +func incomingAmount(text string, patterns []*regexp.Regexp) (string, error) { + if len(currencyAmountPattern.FindAllStringSubmatch(text, -1)) > 1 { + return "", ErrAmbiguousAmount + } + return firstAmount(text, patterns), nil } func parsePayGateAmount(value string) (int64, error) { @@ -172,15 +188,27 @@ func firstAmount(text string, patterns []*regexp.Regexp) string { } func extractPayer(text string) (string, string) { - upiID := strings.TrimSpace(upiPattern.FindString(text)) - payerName := "" + payerText := "" if match := paidYouPayerPattern.FindStringSubmatch(text); len(match) > 1 { - payerName = cleanPayer(match[1], upiID) + payerText = match[1] } else if match := fromPattern.FindStringSubmatch(text); len(match) > 1 { - payerName = cleanPayer(match[1], upiID) + payerText = match[1] + } + upiID := firstUPI(payerText) + if upiID == "" { + upiID = upiPattern.FindString(text) } + payerName := cleanPayer(payerText, upiID) return truncateRunes(payerName, 255), truncateRunes(upiID, 255) } + +func firstUPI(value string) string { + matches := upiPattern.FindAllString(value, -1) + if len(matches) == 0 { + return "" + } + return matches[0] +} func cleanPayer(value, upiID string) string { value = strings.Trim(strings.TrimSpace(value), " ,;:-") if upiID != "" { diff --git a/internal/v4/observations/parser_test.go b/internal/v4/observations/parser_test.go index d49e277..7c62791 100644 --- a/internal/v4/observations/parser_test.go +++ b/internal/v4/observations/parser_test.go @@ -92,82 +92,44 @@ func TestPaytmFallsBackToNotificationPostTime(t *testing.T) { t.Fatalf("occurred = %s source=%s", got.OccurredAt, got.OccurredAtSource) } } -func TestParseKotakGoogleMessagesNotification(t *testing.T) { +func TestParseBlocksRetiredMessageAndEmailPackages(t *testing.T) { posted := time.UnixMilli(1_788_200_000_000).UTC() - got, err := Parse(Snapshot{ - PackageName: GoogleMessagesPackage, - PostedAt: posted, - Title: "VM-KOTAKB", - Text: "Kotak: Received Rs. 1,250.50 from Maya UPI Ref No. 123456789012", - }) - if err != nil { - t.Fatal(err) - } - if got.Source != "kotak_sms" || got.CollectionProfileID != "kotak" || got.AmountPaise != 125050 { - t.Fatalf("observation = %+v", got) - } - if got.PayerName != "Maya" || !got.OccurredAt.Equal(posted) || got.OccurredAtSource != "notification_posted_at" { - t.Fatalf("parsed Kotak details = %+v", got) + for _, packageName := range []string{GoogleMessagesPackage, GmailPackage} { + if _, err := Parse(Snapshot{ + PackageName: packageName, + PostedAt: posted, + Title: "Payment received", + Text: "Received Rs. 1,250.50 from Maya", + }); !errors.Is(err, ErrUnrecognized) { + t.Errorf("Parse(%q) error=%v, want %v", packageName, err, ErrUnrecognized) + } } } -func TestKotakExtractsUPIWithoutRequiringReference(t *testing.T) { - posted := time.UnixMilli(1_788_200_000_000).UTC() - got, err := Parse(Snapshot{ - PackageName: GoogleMessagesPackage, - PostedAt: posted, - Title: "KOTAK", - Text: "Payment for Received INR 75.37 from a@upi", - }) - if err != nil { - t.Fatal(err) - } - if got.AmountPaise != 7537 || got.PayerUPIID != "a@upi" { - t.Fatalf("observation = %+v", got) - } -} -func TestMessagesAcceptAnyIncomingBankCreditButStillRejectUnsafeMoneyText(t *testing.T) { +func TestParseRejectsUntrustedGenericPackages(t *testing.T) { posted := time.UnixMilli(1_788_200_000_000).UTC() - accepted := []struct{ title, text string }{ - {"VM-HDFCBK", "Received Rs.100.37 from Rahul"}, - {"JD-SBIUPI-S", "A/c credited Rs.100.37 through UPI Ref 123456789012"}, - {"JK-SBIUPI-S", "A/c credited Rs.100.37 through UPI Ref 123456789012"}, - } - for _, tc := range accepted { - got, err := Parse(Snapshot{PackageName: GoogleMessagesPackage, PostedAt: posted, Title: tc.title, Text: tc.text}) - if err != nil || got.Source != GenericMessageSource { - t.Errorf("Parse(%q,%q)=%+v err=%v", tc.title, tc.text, got, err) - } - } - rejected := []struct { - title, text string - want error - }{ - {"VM-KOTAKB", "Your OTP is 123456 for Rs.100.37", ErrUnrecognized}, - {"VM-KOTAKB", "Rs.100.37 debited from your account", ErrUnrecognized}, - {"VM-KOTAKB", "Cashback of INR 100.37 credited to your account", ErrUnrecognized}, - {"VM-KOTAKB", "Received Rs.100.00 from Rahul", ErrNonPayGateAmount}, - } - for _, tc := range rejected { - _, err := Parse(Snapshot{PackageName: GoogleMessagesPackage, PostedAt: posted, Title: tc.title, Text: tc.text}) - if !errors.Is(err, tc.want) { - t.Errorf("Parse(%q,%q) error=%v want=%v", tc.title, tc.text, err, tc.want) + for _, packageName := range []string{"com.example.wallet", "com.android.shell"} { + if _, err := Parse(Snapshot{ + PackageName: packageName, + PostedAt: posted, + Title: "received", + Text: "₹100.37 received from Rahul", + }); !errors.Is(err, ErrUnrecognized) { + t.Errorf("Parse(%q) error=%v, want %v", packageName, err, ErrUnrecognized) } } } -func TestUnknownPackageSplitTitleAndAmountCanProvideIncomingPaymentEvidence(t *testing.T) { - posted := time.Now().UTC() - got, err := Parse(Snapshot{PackageName: "com.example.wallet", PostedAt: posted, Title: "received", Text: "₹98765.43"}) - if err != nil || got.Source != GenericNotificationSource || got.AmountPaise != 9876543 { - t.Fatalf("split-field generic observation=%+v err=%v", got, err) - } -} - -func TestUnknownPackageCanProvideIncomingPaymentEvidence(t *testing.T) { - got, err := Parse(Snapshot{PackageName: "com.example.wallet", PostedAt: time.Now().UTC(), Text: "₹100.37 received from Rahul"}) - if err != nil || got.Source != GenericNotificationSource || got.AmountPaise != 10037 { - t.Fatalf("generic package observation=%+v err=%v", got, err) +func TestParseRejectsMalformedAmountTokens(t *testing.T) { + posted := time.UnixMilli(1_788_200_000_000).UTC() + for _, text := range []string{ + "₹12.345 received from Rahul", + "₹12.34.56 received from Rahul", + "₹12.34foo received from Rahul", + } { + if _, err := Parse(Snapshot{PackageName: BHIMPackage, PostedAt: posted, Text: text}); err == nil { + t.Errorf("malformed amount %q was accepted", text) + } } } @@ -176,8 +138,8 @@ func TestParseGenericIncomingPaymentApplications(t *testing.T) { cases := []struct{ pkg, text, source string }{ {"in.amazon.mShop.android.shopping", "Payment received: ₹499.37 from Rahul", GenericNotificationSource}, {"com.phonepe.app", "You received INR 250.41 from Maya via UPI", GenericNotificationSource}, + {"com.phonepe.app", "You received INR 12.3 from Maya via UPI", GenericNotificationSource}, {"com.google.android.apps.nbu.paisa.user", "₹99.23 received from user@okaxis", GenericNotificationSource}, - {GoogleMessagesPackage, "HDFC: A/c credited with Rs. 701.19 from Arun UPI Ref 123456789012", GenericMessageSource}, } for _, tc := range cases { got, err := Parse(Snapshot{PackageName: tc.pkg, PostedAt: posted, Text: tc.text}) @@ -266,7 +228,7 @@ func TestGenericPayerCleanupCompatibility(t *testing.T) { } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { - got, err := Parse(Snapshot{PackageName: "example.wallet", PostedAt: posted, Text: tc.text}) + got, err := Parse(Snapshot{PackageName: BHIMPackage, PostedAt: posted, Text: tc.text}) if err != nil { t.Fatalf("Parse() error = %v", err) } @@ -334,3 +296,56 @@ func TestParseGooglePayPaidYouNotification(t *testing.T) { t.Fatalf("outgoing GPay notification error = %v, want %v", err, ErrUnrecognized) } } + +func TestParserRejectsAmbiguousMonetaryAmounts(t *testing.T) { + _, err := Parse(Snapshot{ + PackageName: BHIMPackage, + PostedAt: time.UnixMilli(1_788_200_000_000).UTC(), + Text: "Payment received ₹1.25. Available balance ₹100.37", + }) + if !errors.Is(err, ErrAmbiguousAmount) { + t.Fatalf("ambiguous notification error = %v, want %v", err, ErrAmbiguousAmount) + } +} + +func TestParserRejectsFailedIncomingLanguage(t *testing.T) { + posted := time.UnixMilli(1_788_200_000_000).UTC() + for _, text := range []string{ + "Payment failed but ₹100.37 received", + "UPI payment declined: received ₹100.37", + "Payment pending, amount ₹100.37 received", + } { + if _, err := Parse(Snapshot{PackageName: BHIMPackage, PostedAt: posted, Text: text}); !errors.Is(err, ErrUnrecognized) { + t.Errorf("Parse(%q) error = %v, want %v", text, err, ErrUnrecognized) + } + } +} + +func TestPayerUPIUsesIncomingPayerClause(t *testing.T) { + got, err := Parse(Snapshot{ + PackageName: BHIMPackage, + PostedAt: time.UnixMilli(1_788_200_000_000).UTC(), + Title: "Merchant merchant@upi", + Text: "Received ₹1.25 from Alice (alice@upi)", + }) + if err != nil { + t.Fatalf("Parse() error = %v", err) + } + if got.PayerName != "Alice" || got.PayerUPIID != "alice@upi" { + t.Fatalf("payer = %q / %q", got.PayerName, got.PayerUPIID) + } +} + +func TestPayerUPIPrefersFirstVPAInIncomingClause(t *testing.T) { + got, err := Parse(Snapshot{ + PackageName: BHIMPackage, + PostedAt: time.UnixMilli(1_788_200_000_000).UTC(), + Text: "Received ₹1.25 from Alice (alice@upi) to merchant@upi", + }) + if err != nil { + t.Fatal(err) + } + if got.PayerName != "Alice" || got.PayerUPIID != "alice@upi" { + t.Fatalf("payer = %q / %q", got.PayerName, got.PayerUPIID) + } +} diff --git a/internal/v4/operator/service.go b/internal/v4/operator/service.go index 0e999c8..ed1b848 100644 --- a/internal/v4/operator/service.go +++ b/internal/v4/operator/service.go @@ -24,16 +24,25 @@ type DailyVolume struct { } type Overview struct { - CollectedTodayPaise int64 `json:"collected_today_paise"` - PaymentsToday int `json:"payments_today"` - PaidToday int `json:"paid_today"` - Pending int `json:"pending"` - ExpiredToday int `json:"expired_today"` - StatusCounts map[string]int `json:"status_counts"` - Volume []DailyVolume `json:"volume"` - ActiveProfile *ProfileSummary `json:"active_profile"` - Relay RelaySummary `json:"relay"` - Webhooks WebhookSummary `json:"webhooks"` + CollectedTodayPaise int64 `json:"collected_today_paise"` + PaymentsToday int `json:"payments_today"` + PaidToday int `json:"paid_today"` + Pending int `json:"pending"` + ExpiredToday int `json:"expired_today"` + UnmatchedToday int `json:"unmatched_today"` + ExpiringSoon int `json:"expiring_soon"` + StatusCounts map[string]int `json:"status_counts"` + Volume []DailyVolume `json:"volume"` + ActiveProfile *ProfileSummary `json:"active_profile"` + Relay RelaySummary `json:"relay"` + Webhooks WebhookSummary `json:"webhooks"` + LastObservation *ObservationSummary `json:"last_observation,omitempty"` +} +type ObservationSummary struct { + ReceivedAt time.Time `json:"received_at"` + PackageName string `json:"package_name"` + DeviceName string `json:"device_name,omitempty"` + MatchResult string `json:"match_result"` } type ProfileSummary struct { ID string `json:"id"` @@ -97,6 +106,12 @@ func (s *Service) Overview(ctx context.Context) (Overview, error) { if err := s.DB.SQL.QueryRowContext(ctx, `SELECT COUNT(*) FROM payments WHERE status='expired' AND grace_until>=? AND grace_until=? AND received_at? AND expires_at<=?`, now.UnixMilli(), now.Add(10*time.Minute).UnixMilli()).Scan(&out.ExpiringSoon); err != nil { + return Overview{}, fmt.Errorf("read expiring-soon overview: %w", err) + } rows, err := s.DB.SQL.QueryContext(ctx, `SELECT status,COUNT(*) FROM payments GROUP BY status`) if err != nil { return Overview{}, fmt.Errorf("read status counts: %w", err) @@ -133,6 +148,11 @@ func (s *Service) Overview(ctx context.Context) (Overview, error) { if err := s.loadWebhookSummary(ctx, &out.Webhooks); err != nil { return Overview{}, err } + lastObservation, err := s.lastObservation(ctx) + if err != nil { + return Overview{}, err + } + out.LastObservation = lastObservation return out, nil } func (s *Service) volumeTrend(ctx context.Context, now time.Time, days int) ([]DailyVolume, error) { @@ -186,7 +206,8 @@ func (s *Service) activeProfile(ctx context.Context) (*ProfileSummary, error) { } func (s *Service) loadRelay(ctx context.Context, out *RelaySummary) error { - rows, err := s.DB.SQL.QueryContext(ctx, `SELECT COALESCE(name,''),COALESCE(last_heartbeat_at,last_seen_at),app_version + rows, err := s.DB.SQL.QueryContext(ctx, `SELECT COALESCE(name,''),last_seen_at,last_heartbeat_at,app_version, + notification_access,listener_connected,battery_optimization_exempt,background_restricted,foreground_service FROM relay_devices WHERE enabled=1`) if err != nil { return fmt.Errorf("read relay summary: %w", err) @@ -194,35 +215,48 @@ func (s *Service) loadRelay(ctx context.Context, out *RelaySummary) error { defer rows.Close() var fallbackName, fallbackVersion string - var bestSeen, bestConnected *time.Time + var bestSeenAt, bestConnectedAt, bestConnectedLastSeen *time.Time var bestSeenName, bestSeenVersion, bestConnectedName, bestConnectedVersion string for rows.Next() { var name string - var lastSeen sql.NullInt64 + var lastSeen, lastHeartbeat sql.NullInt64 var appVersion sql.NullString - if err := rows.Scan(&name, &lastSeen, &appVersion); err != nil { + var notificationAccess, listenerConnected, batteryExempt, backgroundRestricted, foregroundService sql.NullInt64 + if err := rows.Scan(&name, &lastSeen, &lastHeartbeat, &appVersion, + ¬ificationAccess, &listenerConnected, &batteryExempt, &backgroundRestricted, &foregroundService); err != nil { return fmt.Errorf("scan relay summary: %w", err) } out.EnabledDevices++ if fallbackName == "" { fallbackName, fallbackVersion = name, appVersion.String } - if !lastSeen.Valid { - continue + diagnosticAt := lastSeen + if !diagnosticAt.Valid { + diagnosticAt = lastHeartbeat } - seen := time.UnixMilli(lastSeen.Int64).UTC() - if bestSeen == nil || seen.After(*bestSeen) { - value := seen - bestSeen = &value - bestSeenName, bestSeenVersion = name, appVersion.String + if diagnosticAt.Valid { + seen := time.UnixMilli(diagnosticAt.Int64).UTC() + if bestSeenAt == nil || seen.After(*bestSeenAt) { + bestSeenAt = &seen + bestSeenName, bestSeenVersion = name, appVersion.String + } } - age := s.now().Sub(seen) - if age >= -5*time.Minute && age <= time.Hour { - out.ConnectedDevices++ - if bestConnected == nil || seen.After(*bestConnected) { - value := seen - bestConnected = &value - bestConnectedName, bestConnectedVersion = name, appVersion.String + ready := lastHeartbeat.Valid && relayHeartbeatReady(notificationAccess, listenerConnected, batteryExempt, backgroundRestricted, foregroundService) + if ready { + heartbeatAt := time.UnixMilli(lastHeartbeat.Int64).UTC() + age := s.now().Sub(heartbeatAt) + if age >= -5*time.Minute && age <= time.Hour { + out.ConnectedDevices++ + if bestConnectedAt == nil || heartbeatAt.After(*bestConnectedAt) { + bestConnectedAt = &heartbeatAt + bestConnectedName, bestConnectedVersion = name, appVersion.String + if lastSeen.Valid { + lastSeenAt := time.UnixMilli(lastSeen.Int64).UTC() + bestConnectedLastSeen = &lastSeenAt + } else { + bestConnectedLastSeen = &heartbeatAt + } + } } } } @@ -231,15 +265,43 @@ func (s *Service) loadRelay(ctx context.Context, out *RelaySummary) error { } out.Connected = out.ConnectedDevices > 0 switch { - case bestConnected != nil: - out.Name, out.AppVersion, out.LastSeenAt = bestConnectedName, bestConnectedVersion, bestConnected - case bestSeen != nil: - out.Name, out.AppVersion, out.LastSeenAt = bestSeenName, bestSeenVersion, bestSeen + case bestConnectedAt != nil: + out.Name, out.AppVersion, out.LastSeenAt = bestConnectedName, bestConnectedVersion, bestConnectedLastSeen + case bestSeenAt != nil: + out.Name, out.AppVersion = bestSeenName, bestSeenVersion + value := *bestSeenAt + out.LastSeenAt = &value default: out.Name, out.AppVersion = fallbackName, fallbackVersion } return nil } +func relayHeartbeatReady(notificationAccess, listenerConnected, batteryExempt, backgroundRestricted, foregroundService sql.NullInt64) bool { + return notificationAccess.Valid && notificationAccess.Int64 == 1 && + listenerConnected.Valid && listenerConnected.Int64 == 1 && + batteryExempt.Valid && batteryExempt.Int64 == 1 && + backgroundRestricted.Valid && backgroundRestricted.Int64 == 0 && + foregroundService.Valid && foregroundService.Int64 == 1 +} + +func (s *Service) lastObservation(ctx context.Context) (*ObservationSummary, error) { + var receivedAt int64 + var out ObservationSummary + err := s.DB.SQL.QueryRowContext(ctx, `SELECT re.received_at,re.package_name,COALESCE(rd.name,''),po.match_result + FROM relay_events re + JOIN payment_observations po ON po.relay_event_id=re.id + LEFT JOIN relay_devices rd ON rd.id=re.device_id + ORDER BY re.received_at DESC LIMIT 1`). + Scan(&receivedAt, &out.PackageName, &out.DeviceName, &out.MatchResult) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + if err != nil { + return nil, fmt.Errorf("read last payment observation: %w", err) + } + out.ReceivedAt = time.UnixMilli(receivedAt).UTC() + return &out, nil +} func (s *Service) loadWebhookSummary(ctx context.Context, out *WebhookSummary) error { if err := s.DB.SQL.QueryRowContext(ctx, `SELECT COUNT(*) FROM webhook_deliveries WHERE status IN ('pending','retry')`).Scan(&out.Pending); err != nil { diff --git a/internal/v4/operator/service_test.go b/internal/v4/operator/service_test.go index bccdc1f..1776b9e 100644 --- a/internal/v4/operator/service_test.go +++ b/internal/v4/operator/service_test.go @@ -68,8 +68,11 @@ func TestOverviewUsesIndiaLocalDayAndShowsOperationalSummary(t *testing.T) { t.Fatal(err) } _ = f.create(t, 200, "pending") - if _, err := f.db.SQL.Exec(`INSERT INTO relay_devices(id,name,public_key_pem,enabled,enrolled_at,last_seen_at,app_version) - VALUES('device-1','Edge 60 Stylus','pem',1,?,?,?)`, f.now.Add(-time.Hour).UnixMilli(), f.now.Add(-time.Minute).UnixMilli(), "0.5.0"); err != nil { + if _, err := f.db.SQL.Exec(`INSERT INTO relay_devices(id,name,public_key_pem,enabled,enrolled_at,last_seen_at,last_heartbeat_at, + notification_access,listener_connected,battery_optimization_exempt,background_restricted,foreground_service,app_version) + VALUES('device-1','Edge 60 Stylus','pem',1,?,?,?,?,?,?,?,?,?)`, + f.now.Add(-time.Hour).UnixMilli(), f.now.Add(-time.Minute).UnixMilli(), f.now.Add(-time.Minute).UnixMilli(), + 1, 1, 1, 0, 1, "0.5.0"); err != nil { t.Fatal(err) } @@ -93,13 +96,62 @@ func TestOverviewUsesIndiaLocalDayAndShowsOperationalSummary(t *testing.T) { t.Fatalf("volume=%+v", overview.Volume) } } +func TestOverviewIncludesAttentionAndLastObservation(t *testing.T) { + f := newOperatorFixture(t) + payment := f.create(t, 100, "attention") + if _, err := f.db.SQL.Exec(`INSERT INTO relay_devices(id,name,public_key_pem,enabled,enrolled_at) VALUES('device-attention','Motorola Edge 60 Stylus','pem',1,?)`, f.now.Add(-time.Hour).UnixMilli()); err != nil { + t.Fatal(err) + } + if _, err := f.db.SQL.Exec(`INSERT INTO relay_events(id,device_id,source_event_id,package_name,posted_at,received_at,status) + VALUES('relay-attention','device-attention','source-attention','com.google.android.apps.nbu.paisa.user',?,?, 'unmatched')`, f.now.Add(-time.Second).UnixMilli(), f.now.Add(-time.Second).UnixMilli()); err != nil { + t.Fatal(err) + } + if _, err := f.db.SQL.Exec(`INSERT INTO payment_observations(id,relay_event_id,source,collection_profile_id,amount_paise,occurred_at,occurred_at_source,received_at,match_result) + VALUES('obs-attention','relay-attention','android_notification','paytm',?,?, 'notification_posted_at',?,'unmatched')`, payment.PayableAmountPaise, f.now.Add(-time.Second).UnixMilli(), f.now.Add(-time.Second).UnixMilli()); err != nil { + t.Fatal(err) + } + + overview, err := f.operator.Overview(context.Background()) + if err != nil { + t.Fatal(err) + } + if overview.UnmatchedToday != 1 || overview.ExpiringSoon != 1 { + t.Fatalf("attention counts = unmatched %d expiring %d", overview.UnmatchedToday, overview.ExpiringSoon) + } + if overview.LastObservation == nil || overview.LastObservation.PackageName != "com.google.android.apps.nbu.paisa.user" || + overview.LastObservation.DeviceName != "Motorola Edge 60 Stylus" || overview.LastObservation.MatchResult != "unmatched" { + t.Fatalf("last observation = %+v", overview.LastObservation) + } +} + +func TestOverviewCountsParserLevelAmbiguousRelayEvidence(t *testing.T) { + f := newOperatorFixture(t) + if _, err := f.db.SQL.Exec(`INSERT INTO relay_devices(id,name,public_key_pem,enabled,enrolled_at) VALUES('device-parser-ambiguous','Relay Phone','pem',1,?)`, f.now.Add(-time.Hour).UnixMilli()); err != nil { + t.Fatal(err) + } + if _, err := f.db.SQL.Exec(`INSERT INTO relay_events(id,device_id,source_event_id,package_name,posted_at,received_at,status,error) + VALUES('relay-parser-ambiguous','device-parser-ambiguous','source-parser-ambiguous','com.google.android.apps.nbu.paisa.user',?,?, 'ambiguous','notification contains multiple monetary amounts')`, f.now.Add(-time.Second).UnixMilli(), f.now.Add(-time.Second).UnixMilli()); err != nil { + t.Fatal(err) + } + + overview, err := f.operator.Overview(context.Background()) + if err != nil { + t.Fatal(err) + } + if overview.UnmatchedToday != 1 { + t.Fatalf("evidence to review = %d, want 1", overview.UnmatchedToday) + } +} + func TestOverviewRelaySummaryUsesAnyHealthyEnabledDevice(t *testing.T) { f := newOperatorFixture(t) - if _, err := f.db.SQL.Exec(`INSERT INTO relay_devices(id,name,public_key_pem,enabled,enrolled_at,last_seen_at,app_version) VALUES - ('stale','Stale phone','pem',1,?,?,?), - ('healthy','Healthy phone','pem',1,?,?,?)`, - f.now.Add(-3*time.Hour).UnixMilli(), f.now.Add(-2*time.Hour).UnixMilli(), "0.6.1", - f.now.Add(-time.Hour).UnixMilli(), f.now.Add(-time.Minute).UnixMilli(), "0.7.0"); err != nil { + if _, err := f.db.SQL.Exec(`INSERT INTO relay_devices( + id,name,public_key_pem,enabled,enrolled_at,last_seen_at,last_heartbeat_at,app_version, + notification_access,listener_connected,battery_optimization_exempt,background_restricted,foreground_service) VALUES + ('stale','Stale phone','pem',1,?,?,?,?,1,1,0,0,1), + ('healthy','Healthy phone','pem',1,?,?,?,?,1,1,1,0,1)`, + f.now.Add(-3*time.Hour).UnixMilli(), f.now.Add(-3*time.Hour).UnixMilli(), f.now.Add(-3*time.Hour).UnixMilli(), "0.6.1", + f.now.Add(-3*time.Hour).UnixMilli(), f.now.Add(-time.Hour).UnixMilli(), f.now.Add(-time.Minute).UnixMilli(), "0.7.0"); err != nil { t.Fatal(err) } overview, err := f.operator.Overview(context.Background()) @@ -109,11 +161,124 @@ func TestOverviewRelaySummaryUsesAnyHealthyEnabledDevice(t *testing.T) { if !overview.Relay.Connected || overview.Relay.EnabledDevices != 2 || overview.Relay.ConnectedDevices != 1 { t.Fatalf("relay summary = %+v", overview.Relay) } - if overview.Relay.Name != "Healthy phone" || overview.Relay.AppVersion != "0.7.0" || overview.Relay.LastSeenAt == nil { + if overview.Relay.Name != "Healthy phone" || overview.Relay.AppVersion != "0.7.0" || overview.Relay.LastSeenAt == nil || + !overview.Relay.LastSeenAt.Equal(f.now.Add(-time.Hour)) { t.Fatalf("representative relay = %+v", overview.Relay) } } +func TestOverviewRelaySummaryRequiresReadyCurrentHeartbeat(t *testing.T) { + f := newOperatorFixture(t) + if _, err := f.db.SQL.Exec(`INSERT INTO relay_devices( + id,name,public_key_pem,enabled,enrolled_at,last_seen_at,last_heartbeat_at,app_version, + notification_access,listener_connected,battery_optimization_exempt,background_restricted,foreground_service) VALUES + ('restricted','Restricted phone','pem',1,?,?,?,?,1,1,0,1,1), + ('healthy','Healthy phone','pem',1,?,?,?,?,1,1,1,0,1)`, + f.now.Add(-time.Hour).UnixMilli(), f.now.Add(-time.Minute).UnixMilli(), f.now.Add(-time.Minute).UnixMilli(), "0.7.2", + f.now.Add(-time.Hour).UnixMilli(), f.now.Add(-2*time.Minute).UnixMilli(), f.now.Add(-2*time.Minute).UnixMilli(), "0.7.2"); err != nil { + t.Fatal(err) + } + overview, err := f.operator.Overview(context.Background()) + if err != nil { + t.Fatal(err) + } + if !overview.Relay.Connected || overview.Relay.ConnectedDevices != 1 || overview.Relay.Name != "Healthy phone" { + t.Fatalf("relay summary = %+v", overview.Relay) + } + + if _, err := f.db.SQL.Exec(`UPDATE relay_devices SET battery_optimization_exempt=0 WHERE id='healthy'`); err != nil { + t.Fatal(err) + } + overview, err = f.operator.Overview(context.Background()) + if err != nil { + t.Fatal(err) + } + if overview.Relay.Connected || overview.Relay.ConnectedDevices != 0 { + t.Fatalf("unready current devices must not be reported online: %+v", overview.Relay) + } +} +func TestOverviewRelaySummaryKeepsLegacyLastSeenAsDiagnosticOnly(t *testing.T) { + f := newOperatorFixture(t) + if _, err := f.db.SQL.Exec(`INSERT INTO relay_devices(id,name,public_key_pem,enabled,enrolled_at,last_seen_at,app_version) + VALUES('legacy','Migrated phone','pem',1,?,?,?)`, f.now.Add(-time.Hour).UnixMilli(), f.now.Add(-time.Minute).UnixMilli(), "0.4.0"); err != nil { + t.Fatal(err) + } + overview, err := f.operator.Overview(context.Background()) + if err != nil { + t.Fatal(err) + } + if overview.Relay.Connected || overview.Relay.ConnectedDevices != 0 || overview.Relay.Name != "Migrated phone" { + t.Fatalf("legacy last-seen diagnostic = %+v", overview.Relay) + } +} + +func TestOverviewRelaySummaryReadinessMatrix(t *testing.T) { + boolInt := func(value bool) int { + if value { + return 1 + } + return 0 + } + tests := []struct { + name string + heartbeat bool + offset time.Duration + notificationAccess bool + listenerConnected bool + batteryExempt bool + powerSaveMode bool + backgroundRestricted bool + foregroundService bool + wantConnected bool + }{ + {name: "current healthy", heartbeat: true, offset: -2 * time.Minute, notificationAccess: true, listenerConnected: true, batteryExempt: true, foregroundService: true, wantConnected: true}, + {name: "battery restricted", heartbeat: true, offset: -2 * time.Minute, notificationAccess: true, listenerConnected: true, batteryExempt: false, foregroundService: true}, + {name: "listener disconnected", heartbeat: true, offset: -2 * time.Minute, notificationAccess: true, listenerConnected: false, batteryExempt: true, foregroundService: true}, + {name: "notification access missing", heartbeat: true, offset: -2 * time.Minute, notificationAccess: false, listenerConnected: true, batteryExempt: true, foregroundService: true}, + {name: "foreground service absent", heartbeat: true, offset: -2 * time.Minute, notificationAccess: true, listenerConnected: true, batteryExempt: true, foregroundService: false}, + {name: "background restricted", heartbeat: true, offset: -2 * time.Minute, notificationAccess: true, listenerConnected: true, batteryExempt: true, backgroundRestricted: true, foregroundService: true}, + {name: "stale heartbeat", heartbeat: true, offset: -61 * time.Minute, notificationAccess: true, listenerConnected: true, batteryExempt: true, foregroundService: true}, + {name: "future within clock tolerance", heartbeat: true, offset: 4 * time.Minute, notificationAccess: true, listenerConnected: true, batteryExempt: true, foregroundService: true, wantConnected: true}, + {name: "future beyond clock tolerance", heartbeat: true, offset: 6 * time.Minute, notificationAccess: true, listenerConnected: true, batteryExempt: true, foregroundService: true}, + {name: "power saver allowed when otherwise ready", heartbeat: true, offset: -2 * time.Minute, notificationAccess: true, listenerConnected: true, batteryExempt: true, powerSaveMode: true, foregroundService: true, wantConnected: true}, + {name: "migrated without heartbeat telemetry", offset: -2 * time.Minute}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + f := newOperatorFixture(t) + seenAt := f.now.Add(test.offset).UnixMilli() + var heartbeat any + if test.heartbeat { + heartbeat = seenAt + } + var notificationAccess, listenerConnected, batteryExempt, powerSaveMode, backgroundRestricted, foregroundService any + if test.heartbeat { + notificationAccess = boolInt(test.notificationAccess) + listenerConnected = boolInt(test.listenerConnected) + batteryExempt = boolInt(test.batteryExempt) + powerSaveMode = boolInt(test.powerSaveMode) + backgroundRestricted = boolInt(test.backgroundRestricted) + foregroundService = boolInt(test.foregroundService) + } + if _, err := f.db.SQL.Exec(`INSERT INTO relay_devices( + id,name,public_key_pem,enabled,enrolled_at,last_seen_at,last_heartbeat_at,app_version, + notification_access,listener_connected,battery_optimization_exempt,power_save_mode,background_restricted,foreground_service) + VALUES('matrix','Matrix phone','pem',1,?,?,?,?,?,?,?,?,?,?)`, + f.now.Add(-time.Hour).UnixMilli(), seenAt, heartbeat, "0.7.2", + notificationAccess, listenerConnected, batteryExempt, powerSaveMode, backgroundRestricted, foregroundService); err != nil { + t.Fatal(err) + } + overview, err := f.operator.Overview(context.Background()) + if err != nil { + t.Fatal(err) + } + if overview.Relay.Connected != test.wantConnected || overview.Relay.ConnectedDevices != boolInt(test.wantConnected) { + t.Fatalf("relay summary = %+v, want connected=%v", overview.Relay, test.wantConnected) + } + }) + } +} + func TestActivityCombinesPaymentObservationAndWebhookEvents(t *testing.T) { f := newOperatorFixture(t) payment := f.create(t, 100, "activity") @@ -182,29 +347,4 @@ func TestWebhookSettingsGenerateHideRotateAndApplyLive(t *testing.T) { if _, _, err := settings.ConfigureWebhook(context.Background(), "", false); err != nil { t.Fatal(err) } - if worker.Enabled() { - t.Fatal("worker remained enabled after webhook was disabled") - } -} - -func TestBootstrapWebhookPreservesLegacySecretAndDoesNotOverwrite(t *testing.T) { - f := newOperatorFixture(t) - worker := webhooks.NewService(f.db, webhooks.Config{}) - settings := NewSettingsService(f.db, worker) - ctx := context.Background() - secret := "legacy-webhook-secret-0123456789abcdef" - if err := settings.BootstrapWebhook(ctx, "https://example.com/legacy-hook", secret); err != nil { - t.Fatal(err) - } - got := worker.ConfigSnapshot() - if got.Endpoint != "https://example.com/legacy-hook" || got.Secret != secret { - t.Fatalf("worker=%+v", got) - } - if err := settings.BootstrapWebhook(ctx, "https://example.com/replacement", "replacement-secret-0123456789abcdef"); err != nil { - t.Fatal(err) - } - got = worker.ConfigSnapshot() - if got.Endpoint != "https://example.com/legacy-hook" || got.Secret != secret { - t.Fatalf("bootstrap overwrote persisted config: %+v", got) - } } diff --git a/internal/v4/payments/allocator.go b/internal/v4/payments/allocator.go index 35ed8c7..a4c3388 100644 --- a/internal/v4/payments/allocator.go +++ b/internal/v4/payments/allocator.go @@ -5,12 +5,15 @@ import ( "crypto/rand" "errors" "fmt" + "math" "math/big" "time" "github.com/Phloraxx/payment-api/internal/v4/storage" ) +const defaultSoftHorizon = 4 * time.Hour + var ErrPaymentCapacity = errors.New("payment capacity temporarily unavailable") type RandomIndex func(max int) (int, error) @@ -24,7 +27,7 @@ type Allocator struct { func NewAllocator() Allocator { return Allocator{ Random: cryptoRandomIndex, - SoftHorizon: 4 * time.Hour, + SoftHorizon: defaultSoftHorizon, Buckets: 2, } } @@ -36,8 +39,8 @@ func (a Allocator) Select(ctx context.Context, tx *storage.ImmediateTx, profileI if profileID == "" { return 0, errors.New("collection profile is required") } - if requestedAmountPaise <= 0 || requestedAmountPaise%100 != 0 { - return 0, errors.New("requested amount must be positive whole INR") + if requestedAmountPaise <= 0 || requestedAmountPaise%100 != 0 || requestedAmountPaise > math.MaxInt64-199 { + return 0, errors.New("requested amount must be positive whole INR within the payable range") } buckets := a.Buckets if buckets <= 0 { @@ -55,11 +58,19 @@ func (a Allocator) Select(ctx context.Context, tx *storage.ImmediateTx, profileI return 0, fmt.Errorf("release due amount reservations: %w", err) } - cutoffMS := now.Add(-a.SoftHorizon).UTC().UnixMilli() + softHorizon := a.SoftHorizon + if softHorizon <= 0 { + softHorizon = defaultSoftHorizon + } + cutoffMS := now.Add(-softHorizon).UTC().UnixMilli() for bucket := 0; bucket < buckets; bucket++ { - start := requestedAmountPaise + int64(bucket*100) + 1 - end := requestedAmountPaise + int64(bucket*100) + 99 - candidates, err := loadBucketCandidates(ctx, tx, profileID, start, end, cutoffMS) + if int64(bucket) > (math.MaxInt64-requestedAmountPaise)/100 { + break + } + offset := int64(bucket) * 100 + start := requestedAmountPaise + offset + 1 + end := start + 98 + candidates, err := loadBucketCandidates(ctx, tx, start, end, cutoffMS) if err != nil { return 0, err } @@ -78,15 +89,14 @@ func (a Allocator) Select(ctx context.Context, tx *storage.ImmediateTx, profileI return 0, ErrPaymentCapacity } -func loadBucketCandidates(ctx context.Context, tx *storage.ImmediateTx, profileID string, start, end, softCutoffMS int64) ([]int64, error) { +func loadBucketCandidates(ctx context.Context, tx *storage.ImmediateTx, start, end, softCutoffMS int64) ([]int64, error) { rows, err := tx.QueryContext(ctx, ` SELECT payable_amount_paise, MAX(CASE WHEN released_at IS NULL THEN 1 ELSE 0 END) AS active, MAX(last_used_at) AS last_used_at FROM amount_reservations -WHERE collection_profile_id = ? - AND payable_amount_paise BETWEEN ? AND ? -GROUP BY payable_amount_paise`, profileID, start, end) +WHERE payable_amount_paise BETWEEN ? AND ? +GROUP BY payable_amount_paise`, start, end) if err != nil { return nil, fmt.Errorf("query amount bucket: %w", err) } @@ -111,9 +121,12 @@ GROUP BY payable_amount_paise`, profileID, start, end) preferred := make([]int64, 0, 99) recent := make([]int64, 0, 99) - for amount := start; amount <= end; amount++ { + for amount := start; ; amount++ { state, seen := used[amount] if seen && state.active { + if amount == end { + break + } continue } if !seen || state.lastUsed < softCutoffMS { @@ -121,6 +134,9 @@ GROUP BY payable_amount_paise`, profileID, start, end) } else { recent = append(recent, amount) } + if amount == end { + break + } } if len(preferred) > 0 { return preferred, nil diff --git a/internal/v4/payments/allocator_test.go b/internal/v4/payments/allocator_test.go index ecf952e..60d8855 100644 --- a/internal/v4/payments/allocator_test.go +++ b/internal/v4/payments/allocator_test.go @@ -5,6 +5,7 @@ import ( "database/sql" "errors" "fmt" + "math" "path/filepath" "testing" "time" @@ -57,6 +58,19 @@ func TestAllocatorRandomizesInsideBaseBucket(t *testing.T) { } } +func TestAllocatorRejectsAmountWhosePayableBucketWouldOverflow(t *testing.T) { + db := openAllocatorDB(t) + now := time.UnixMilli(1_788_200_000_000) + requested := int64(math.MaxInt64) - int64(math.MaxInt64)%100 + err := db.WithImmediateTx(context.Background(), func(tx *storage.ImmediateTx) error { + _, err := NewAllocator().Select(context.Background(), tx, "paytm", requested, now) + return err + }) + if err == nil { + t.Fatal("overflowing payable bucket was accepted") + } +} + func TestAllocatorNeverOverflowsWhileBaseBucketHasOneFreeValue(t *testing.T) { db := openAllocatorDB(t) ctx := context.Background() diff --git a/internal/v4/payments/matching.go b/internal/v4/payments/matching.go index 08ba2c5..cb6d46d 100644 --- a/internal/v4/payments/matching.go +++ b/internal/v4/payments/matching.go @@ -14,8 +14,9 @@ import ( ) var ( - ErrRelayEventNotFound = errors.New("relay event not found") - ErrInvalidObservation = errors.New("invalid payment observation") + ErrRelayEventNotFound = errors.New("relay event not found") + ErrInvalidObservation = errors.New("invalid payment observation") + ErrObservationAmbiguous = errors.New("observation profile changed during matching") ) type MatchResult struct { @@ -26,13 +27,25 @@ type MatchResult struct { } type matchCandidate struct { - PaymentID string - Status string - ReservedAt int64 - ReservedUntil int64 + PaymentID string + Status string + CollectionProfileID string + ReservedAt int64 + ReservedUntil int64 } func (s *Service) ApplyObservation(ctx context.Context, relayEventID string, obs observations.Observation, receivedAt time.Time) (MatchResult, error) { + return s.applyObservation(ctx, relayEventID, obs, receivedAt, "", time.Time{}) +} + +func (s *Service) ApplyObservationForRelay(ctx context.Context, relayEventID string, obs observations.Observation, receivedAt time.Time, deviceID string, enrolledAt time.Time) (MatchResult, error) { + if strings.TrimSpace(deviceID) == "" || enrolledAt.IsZero() { + return MatchResult{}, ErrRelayEventNotFound + } + return s.applyObservation(ctx, relayEventID, obs, receivedAt, strings.TrimSpace(deviceID), enrolledAt.UTC()) +} + +func (s *Service) applyObservation(ctx context.Context, relayEventID string, obs observations.Observation, receivedAt time.Time, deviceID string, enrolledAt time.Time) (MatchResult, error) { if s == nil || s.DB == nil || s.DB.SQL == nil { return MatchResult{}, errors.New("payment storage is required") } @@ -60,6 +73,34 @@ func (s *Service) ApplyObservation(ctx context.Context, relayEventID string, obs var result MatchResult err := s.DB.WithImmediateTx(ctx, func(tx *storage.ImmediateTx) error { + var relayStatus, relayDeviceID string + if deviceID != "" { + var queryErr error + queryErr = tx.QueryRowContext(ctx, `SELECT status,device_id FROM relay_events WHERE id=?`, relayEventID). + Scan(&relayStatus, &relayDeviceID) + if errors.Is(queryErr, sql.ErrNoRows) { + return ErrRelayEventNotFound + } + if queryErr != nil { + return fmt.Errorf("read relay event status: %w", queryErr) + } + if relayStatus != "received" || relayDeviceID != deviceID { + return ErrRelayEventNotFound + } + var enabled int + var currentEnrolledAt int64 + queryErr = tx.QueryRowContext(ctx, `SELECT enabled,enrolled_at FROM relay_devices WHERE id=?`, deviceID). + Scan(&enabled, ¤tEnrolledAt) + if errors.Is(queryErr, sql.ErrNoRows) { + return ErrRelayEventNotFound + } + if queryErr != nil { + return fmt.Errorf("read relay device authorization: %w", queryErr) + } + if enabled != 1 || currentEnrolledAt != enrolledAt.UnixMilli() { + return ErrRelayEventNotFound + } + } replayed, found, err := existingObservationResult(ctx, tx, relayEventID) if err != nil { return err @@ -69,6 +110,19 @@ func (s *Service) ApplyObservation(ctx context.Context, relayEventID string, obs result.Replayed = true return nil } + if deviceID == "" { + err = tx.QueryRowContext(ctx, `SELECT status,device_id FROM relay_events WHERE id=?`, relayEventID). + Scan(&relayStatus, &relayDeviceID) + if errors.Is(err, sql.ErrNoRows) { + return ErrRelayEventNotFound + } + if err != nil { + return fmt.Errorf("read relay event status: %w", err) + } + } + if relayStatus != "received" { + return ErrRelayEventNotFound + } packageName, err := relayPackage(ctx, tx, relayEventID) if err != nil { return err @@ -87,13 +141,8 @@ func (s *Service) ApplyObservation(ctx context.Context, relayEventID string, obs matchResult = "ambiguous" } else if len(candidates) == 1 { candidate := candidates[0] - unsafe, err := reusedLowConfidenceLatest(ctx, tx, obs, candidate) - if err != nil { - return err - } - if unsafe { - matchResult = "ambiguous" - } else { + obs.CollectionProfileID = candidate.CollectionProfileID + if sourceCanAutoConfirm(obs.Source) { matchResult = "matched" if candidate.Status == "paid" { prior, err := hasConfirmedObservation(ctx, tx, candidate.PaymentID) @@ -110,6 +159,17 @@ func (s *Service) ApplyObservation(ctx context.Context, relayEventID string, obs return err } result.Transitioned = transitioned + } else { + matchResult = "ambiguous" + } + } + if strings.TrimSpace(obs.CollectionProfileID) == "" { + if err := tx.QueryRowContext(ctx, `SELECT id FROM collection_profiles WHERE active=1 AND enabled=1 LIMIT 1`). + Scan(&obs.CollectionProfileID); err != nil { + if errors.Is(err, sql.ErrNoRows) { + return fmt.Errorf("%w: no active collection profile is available for observation attribution", ErrInvalidObservation) + } + return fmt.Errorf("read observation collection profile: %w", err) } } observationID, err := idFn("obs") @@ -125,11 +185,11 @@ func (s *Service) ApplyObservation(ctx context.Context, relayEventID string, obs obs.OccurredAtSource, receivedAt.UnixMilli(), nullableString(matchedID), matchResult); err != nil { return fmt.Errorf("insert payment observation: %w", err) } - relayStatus := matchResult - if relayStatus == "corroborated" { - relayStatus = "matched" + relayEventStatus := matchResult + if relayEventStatus == "corroborated" { + relayEventStatus = "matched" } - if _, err := tx.ExecContext(ctx, `UPDATE relay_events SET status=?,error=NULL WHERE id=?`, relayStatus, relayEventID); err != nil { + if _, err := tx.ExecContext(ctx, `UPDATE relay_events SET status=?,error=NULL WHERE id=?`, relayEventStatus, relayEventID); err != nil { return fmt.Errorf("update relay event status: %w", err) } result.Result = matchResult @@ -148,17 +208,10 @@ func validateObservation(obs observations.Observation) error { } switch obs.Source { case "paytm_notification": - if obs.CollectionProfileID != "paytm" { + if strings.TrimSpace(obs.CollectionProfileID) != "paytm" { return fmt.Errorf("%w: Paytm source/profile mismatch", ErrInvalidObservation) } - case "kotak_sms": - if obs.CollectionProfileID != "kotak" { - return fmt.Errorf("%w: Kotak source/profile mismatch", ErrInvalidObservation) - } - case observations.GenericNotificationSource, observations.GenericMessageSource: - if strings.TrimSpace(obs.CollectionProfileID) == "" { - return fmt.Errorf("%w: generic source requires resolved collection profile", ErrInvalidObservation) - } + case observations.GenericNotificationSource: default: return fmt.Errorf("%w: unsupported source %q", ErrInvalidObservation, obs.Source) } @@ -174,11 +227,14 @@ func expectedPackage(source string) string { if source == "paytm_notification" { return observations.PaytmBusinessPackage } - if source == "kotak_sms" { - return observations.GoogleMessagesPackage - } return "" } + +// sourceCanAutoConfirm permits only parsed Paytm or generic app evidence. +// Candidate identity is the exact amount and valid reservation window. +func sourceCanAutoConfirm(source string) bool { + return source == "paytm_notification" || source == observations.GenericNotificationSource +} func existingObservationResult(ctx context.Context, tx *storage.ImmediateTx, relayEventID string) (MatchResult, bool, error) { var matchResult string var paymentID sql.NullString @@ -206,17 +262,24 @@ func relayPackage(ctx context.Context, tx *storage.ImmediateTx, relayEventID str func matchingCandidates(ctx context.Context, tx *storage.ImmediateTx, obs observations.Observation) ([]matchCandidate, error) { occurred := obs.OccurredAt.UnixMilli() - rows, err := tx.QueryContext(ctx, `SELECT p.id,p.status,r.reserved_at,r.reserved_until + query := `SELECT p.id,p.status,r.collection_profile_id,r.reserved_at,r.reserved_until FROM amount_reservations r JOIN payments p ON p.id=r.payment_id - WHERE r.collection_profile_id=? AND r.payable_amount_paise=? AND p.created_at<=? AND r.reserved_until>=? - ORDER BY r.reserved_at`, obs.CollectionProfileID, obs.AmountPaise, occurred, occurred) + WHERE r.payable_amount_paise=? AND p.created_at<=? AND r.reserved_until>=? + AND (r.released_at IS NULL OR r.released_at>=?)` + args := []any{obs.AmountPaise, occurred, occurred, occurred} + if obs.Source == "paytm_notification" { + query += ` AND r.collection_profile_id=?` + args = append(args, obs.CollectionProfileID) + } + query += ` ORDER BY r.reserved_at` + rows, err := tx.QueryContext(ctx, query, args...) if err != nil { return nil, fmt.Errorf("find matching reservations: %w", err) } var raw []matchCandidate for rows.Next() { var c matchCandidate - if err := rows.Scan(&c.PaymentID, &c.Status, &c.ReservedAt, &c.ReservedUntil); err != nil { + if err := rows.Scan(&c.PaymentID, &c.Status, &c.CollectionProfileID, &c.ReservedAt, &c.ReservedUntil); err != nil { rows.Close() return nil, err } @@ -231,32 +294,66 @@ func matchingCandidates(ctx context.Context, tx *storage.ImmediateTx, obs observ candidates := make([]matchCandidate, 0, len(raw)) for _, c := range raw { - if c.Status == "cancelled" { - allowed, err := occurredBeforeCancellation(ctx, tx, c.PaymentID, occurred) - if err != nil { - return nil, err - } - if !allowed { - continue - } + cancelled, err := occurredDuringCancellation(ctx, tx, c.PaymentID, occurred) + if err != nil { + return nil, err + } + if cancelled { + continue } candidates = append(candidates, c) } return candidates, nil } -func occurredBeforeCancellation(ctx context.Context, tx *storage.ImmediateTx, paymentID string, occurredAt int64) (bool, error) { - var cancelledAt int64 - err := tx.QueryRowContext(ctx, `SELECT created_at FROM payment_history WHERE payment_id=? AND type='payment.cancelled' ORDER BY created_at DESC LIMIT 1`, paymentID).Scan(&cancelledAt) - if errors.Is(err, sql.ErrNoRows) { - return false, nil - } +func occurredDuringCancellation(ctx context.Context, tx *storage.ImmediateTx, paymentID string, occurredAt int64) (bool, error) { + rows, err := tx.QueryContext(ctx, `SELECT type,created_at,changes_json FROM payment_history + WHERE payment_id=? ORDER BY created_at,id`, paymentID) if err != nil { - return false, fmt.Errorf("read cancellation time: %w", err) + return false, fmt.Errorf("read payment status history: %w", err) } - return occurredAt <= cancelledAt, nil + defer rows.Close() + var cancelledAt *int64 + for rows.Next() { + var historyType, changes string + var createdAt int64 + if err := rows.Scan(&historyType, &createdAt, &changes); err != nil { + return false, fmt.Errorf("scan payment status history: %w", err) + } + if historyType == "payment.cancelled" { + value := createdAt + cancelledAt = &value + continue + } + var transition struct { + Status struct { + From string `json:"from"` + To string `json:"to"` + } `json:"status"` + } + if err := json.Unmarshal([]byte(changes), &transition); err != nil || transition.Status.To == "" { + continue + } + if transition.Status.To == "cancelled" { + value := createdAt + cancelledAt = &value + continue + } + if cancelledAt != nil { + if occurredAt > *cancelledAt && occurredAt < createdAt { + _ = rows.Close() + return true, nil + } + if transition.Status.From == "cancelled" && transition.Status.To == "pending" { + cancelledAt = nil + } + } + } + if err := rows.Err(); err != nil { + return false, fmt.Errorf("iterate payment status history: %w", err) + } + return cancelledAt != nil && occurredAt > *cancelledAt, nil } - func hasConfirmedObservation(ctx context.Context, tx *storage.ImmediateTx, paymentID string) (bool, error) { var found int err := tx.QueryRowContext(ctx, `SELECT EXISTS(SELECT 1 FROM payment_observations WHERE matched_payment_id=? AND match_result IN ('matched','corroborated'))`, paymentID).Scan(&found) @@ -266,19 +363,6 @@ func hasConfirmedObservation(ctx context.Context, tx *storage.ImmediateTx, payme return found == 1, nil } -func reusedLowConfidenceLatest(ctx context.Context, tx *storage.ImmediateTx, obs observations.Observation, candidate matchCandidate) (bool, error) { - lowConfidence := obs.OccurredAtSource == "server_received_at" || obs.OccurredAtSource == "notification_posted_at" - if !lowConfidence { - return false, nil - } - var count int - var latest sql.NullInt64 - if err := tx.QueryRowContext(ctx, `SELECT COUNT(*),MAX(reserved_at) FROM amount_reservations WHERE collection_profile_id=? AND payable_amount_paise=?`, obs.CollectionProfileID, obs.AmountPaise).Scan(&count, &latest); err != nil { - return false, fmt.Errorf("read reservation reuse history: %w", err) - } - return count > 1 && latest.Valid && candidate.ReservedAt == latest.Int64, nil -} - func applyMatchedPayment(ctx context.Context, tx *storage.ImmediateTx, idFn func(string) (string, error), candidate matchCandidate, obs observations.Observation, transitionAt time.Time) (bool, error) { if candidate.Status == "paid" { if _, err := tx.ExecContext(ctx, `UPDATE payments SET diff --git a/internal/v4/payments/matching_test.go b/internal/v4/payments/matching_test.go index ad3b10d..ac616d8 100644 --- a/internal/v4/payments/matching_test.go +++ b/internal/v4/payments/matching_test.go @@ -2,6 +2,7 @@ package payments import ( "context" + "errors" "testing" "time" @@ -64,6 +65,131 @@ func TestApplyObservationMarksPendingPaidAtomically(t *testing.T) { assertCount(t, db.SQL, "payment_history", 2) assertCount(t, db.SQL, "webhook_deliveries", 2) } + +func TestApplyObservationRejectsStaleRelayEnrollment(t *testing.T) { + ctx := context.Background() + db := openAllocatorDB(t) + createdAt := time.UnixMilli(1_788_200_000_000).UTC() + s := newTestService(t, db, createdAt) + created, err := s.Create(ctx, validCreateInput("match-stale-epoch")) + if err != nil { + t.Fatal(err) + } + occurred := createdAt.Add(2 * time.Minute) + received := occurred.Add(time.Second) + insertRelayEvent(t, db, "relay_stale_epoch", "source_stale_epoch", observations.PaytmBusinessPackage, occurred, received) + var enrolledAt int64 + if err := db.SQL.QueryRow(`SELECT enrolled_at FROM relay_devices WHERE id='device_1'`).Scan(&enrolledAt); err != nil { + t.Fatal(err) + } + if _, err := db.SQL.Exec(`UPDATE relay_devices SET enrolled_at=? WHERE id='device_1'`, enrolledAt+1); err != nil { + t.Fatal(err) + } + _, err = s.ApplyObservationForRelay(ctx, "relay_stale_epoch", paytmObservation(created.Payment.PayableAmountPaise, occurred, "notification_posted_at"), received, "device_1", time.UnixMilli(enrolledAt).UTC()) + if !errors.Is(err, ErrRelayEventNotFound) { + t.Fatalf("stale relay enrollment error = %v", err) + } + var status string + if err := db.SQL.QueryRow(`SELECT status FROM relay_events WHERE id='relay_stale_epoch'`).Scan(&status); err != nil { + t.Fatal(err) + } + if status != "received" { + t.Fatalf("stale relay event status = %q", status) + } + got, err := s.Get(ctx, created.Payment.ID) + if err != nil { + t.Fatal(err) + } + if got.Payment.Status != "pending" { + t.Fatalf("stale relay event changed payment to %q", got.Payment.Status) + } +} +func TestApplyObservationRejectsStaleRelayEnrollmentOnReplay(t *testing.T) { + ctx := context.Background() + db := openAllocatorDB(t) + createdAt := time.UnixMilli(1_788_200_000_000).UTC() + s := newTestService(t, db, createdAt) + created, err := s.Create(ctx, validCreateInput("match-stale-replay")) + if err != nil { + t.Fatal(err) + } + occurred := createdAt.Add(2 * time.Minute) + received := occurred.Add(time.Second) + insertRelayEvent(t, db, "relay_stale_replay", "source_stale_replay", observations.PaytmBusinessPackage, occurred, received) + var enrolledAt int64 + if err := db.SQL.QueryRow(`SELECT enrolled_at FROM relay_devices WHERE id='device_1'`).Scan(&enrolledAt); err != nil { + t.Fatal(err) + } + obs := paytmObservation(created.Payment.PayableAmountPaise, occurred, "notification_posted_at") + if _, err := s.ApplyObservationForRelay(ctx, "relay_stale_replay", obs, received, "device_1", time.UnixMilli(enrolledAt).UTC()); err != nil { + t.Fatal(err) + } + if _, err := db.SQL.Exec(`UPDATE relay_devices SET enrolled_at=? WHERE id='device_1'`, enrolledAt+1); err != nil { + t.Fatal(err) + } + if _, err := s.ApplyObservationForRelay(ctx, "relay_stale_replay", obs, received, "device_1", time.UnixMilli(enrolledAt).UTC()); !errors.Is(err, ErrRelayEventNotFound) { + t.Fatalf("stale replay error = %v", err) + } + assertCount(t, db.SQL, "payment_observations", 1) + got, err := s.Get(ctx, created.Payment.ID) + if err != nil { + t.Fatal(err) + } + if got.Payment.Status != "paid" { + t.Fatalf("stale replay changed payment to %q", got.Payment.Status) + } +} + +func TestGenericAmountMatchesReservationAcrossProfileSwitch(t *testing.T) { + ctx := context.Background() + db := openAllocatorDB(t) + base := time.UnixMilli(1_788_200_000_000).UTC() + s := newTestService(t, db, base) + first, err := s.Create(ctx, validCreateInput("profile-race-first")) + if err != nil { + t.Fatal(err) + } + if _, err := db.SQL.Exec(`INSERT INTO collection_profiles(id,label,upi_id,parser,enabled,active,created_at,updated_at) + VALUES('kotak','Kotak','merchant@kotak','kotak_sms',1,0,?,?)`, base.UnixMilli(), base.UnixMilli()); err != nil { + t.Fatal(err) + } + if _, err := db.SQL.Exec(`UPDATE collection_profiles SET active=0 WHERE id='paytm'`); err != nil { + t.Fatal(err) + } + if _, err := db.SQL.Exec(`UPDATE collection_profiles SET active=1 WHERE id='kotak'`); err != nil { + t.Fatal(err) + } + second, err := s.Create(ctx, validCreateInput("profile-race-second")) + if err != nil { + t.Fatal(err) + } + if first.Payment.PayableAmountPaise == second.Payment.PayableAmountPaise { + t.Fatalf("live payable amount reused across profiles: %d", first.Payment.PayableAmountPaise) + } + occurred := base.Add(time.Minute) + received := occurred.Add(time.Second) + insertRelayEvent(t, db, "relay_profile_race", "source_profile_race", observations.GooglePayPackage, occurred, received) + obs := observations.Observation{Source: observations.GenericNotificationSource, AmountPaise: first.Payment.PayableAmountPaise, + PayerName: "Rahul", OccurredAt: occurred, OccurredAtSource: "notification_posted_at"} + result, err := s.ApplyObservation(ctx, "relay_profile_race", obs, received) + if err != nil { + t.Fatal(err) + } + if result.Result != "matched" || result.PaymentID != first.Payment.ID || !result.Transitioned { + t.Fatalf("amount match result = %+v", result) + } + var firstStatus, secondStatus string + if err := db.SQL.QueryRow(`SELECT status FROM payments WHERE id=?`, first.Payment.ID).Scan(&firstStatus); err != nil { + t.Fatal(err) + } + if err := db.SQL.QueryRow(`SELECT status FROM payments WHERE id=?`, second.Payment.ID).Scan(&secondStatus); err != nil { + t.Fatal(err) + } + if firstStatus != "paid" || secondStatus != "pending" { + t.Fatalf("payment statuses after amount match: %s / %s", firstStatus, secondStatus) + } +} + func TestDifferentRelayEventForAlreadyPaidPaymentDoesNotTransitionTwice(t *testing.T) { ctx := context.Background() db := openAllocatorDB(t) @@ -132,6 +258,7 @@ func TestCancelledPaymentOnlyMatchesMoneyThatOccurredBeforeCancellation(t *testi t.Fatal(err) } cancelAt := createdAt.Add(2 * time.Minute) + s.Now = func() time.Time { return cancelAt } if _, err := s.Cancel(ctx, created.Payment.ID); err != nil { t.Fatal(err) @@ -152,6 +279,46 @@ func TestCancelledPaymentOnlyMatchesMoneyThatOccurredBeforeCancellation(t *testi t.Fatalf("payment status = %s", got.Payment.Status) } } +func TestCancelledReopenedPaymentRejectsEvidenceFromCancellationGap(t *testing.T) { + ctx := context.Background() + db := openAllocatorDB(t) + base := time.UnixMilli(1_788_200_000_000).UTC() + s := newTestService(t, db, base) + created, err := s.Create(ctx, validCreateInput("cancel-reopen-gap")) + if err != nil { + t.Fatal(err) + } + firstCancel := base.Add(2 * time.Minute) + s.Now = func() time.Time { return firstCancel } + if _, err := s.Cancel(ctx, created.Payment.ID); err != nil { + t.Fatal(err) + } + reopenedAt := base.Add(3 * time.Minute) + if _, err := db.SQL.Exec(`UPDATE payments SET status='pending' WHERE id=?`, created.Payment.ID); err != nil { + t.Fatal(err) + } + if _, err := db.SQL.Exec(`INSERT INTO payment_history(id,payment_id,type,actor,summary,changes_json,created_at) + VALUES('hist_reopen_gap',?,'payment.updated','admin','Payment reopened','{"status":{"from":"cancelled","to":"pending"}}',?)`, + created.Payment.ID, reopenedAt.UnixMilli()); err != nil { + t.Fatal(err) + } + secondCancel := base.Add(4 * time.Minute) + s.Now = func() time.Time { return secondCancel } + if _, err := s.Cancel(ctx, created.Payment.ID); err != nil { + t.Fatal(err) + } + occurred := base.Add(2*time.Minute + 30*time.Second) + received := base.Add(5 * time.Minute) + insertRelayEvent(t, db, "relay_cancel_gap", "source_cancel_gap", observations.PaytmBusinessPackage, occurred, received) + s.Now = func() time.Time { return received } + result, err := s.ApplyObservation(ctx, "relay_cancel_gap", paytmObservation(created.Payment.PayableAmountPaise, occurred, "notification_text"), received) + if err != nil { + t.Fatal(err) + } + if result.Result != "unmatched" || result.PaymentID != "" || result.Transitioned { + t.Fatalf("cancellation-gap evidence = %+v", result) + } +} func TestCancelledPaymentRejectsMoneyAfterCancellation(t *testing.T) { ctx := context.Background() @@ -183,6 +350,46 @@ func TestCancelledPaymentRejectsMoneyAfterCancellation(t *testing.T) { t.Fatalf("payment status = %s", got.Payment.Status) } } +func TestLatePreCancellationMatchDoesNotOpenPostCancellationWindow(t *testing.T) { + ctx := context.Background() + db := openAllocatorDB(t) + base := time.UnixMilli(1_788_200_000_000).UTC() + s := newTestService(t, db, base) + created, err := s.Create(ctx, validCreateInput("cancel-late-match")) + if err != nil { + t.Fatal(err) + } + cancelAt := base.Add(2 * time.Minute) + s.Now = func() time.Time { return cancelAt } + if _, err := s.Cancel(ctx, created.Payment.ID); err != nil { + t.Fatal(err) + } + + preOccurred := base.Add(time.Minute) + preReceived := base.Add(5 * time.Minute) + insertRelayEvent(t, db, "relay_cancel_late_pre", "source_cancel_late_pre", observations.PaytmBusinessPackage, preOccurred, preReceived) + s.Now = func() time.Time { return preReceived } + pre, err := s.ApplyObservation(ctx, "relay_cancel_late_pre", paytmObservation(created.Payment.PayableAmountPaise, preOccurred, "notification_posted_at"), preReceived) + if err != nil { + t.Fatal(err) + } + if pre.Result != "matched" || !pre.Transitioned { + t.Fatalf("late pre-cancel result = %+v", pre) + } + + postOccurred := base.Add(4 * time.Minute) + postReceived := base.Add(6 * time.Minute) + insertRelayEvent(t, db, "relay_cancel_late_post", "source_cancel_late_post", observations.PaytmBusinessPackage, postOccurred, postReceived) + s.Now = func() time.Time { return postReceived } + post, err := s.ApplyObservation(ctx, "relay_cancel_late_post", paytmObservation(created.Payment.PayableAmountPaise, postOccurred, "notification_posted_at"), postReceived) + if err != nil { + t.Fatal(err) + } + if post.Result != "unmatched" || post.PaymentID != "" || post.Transitioned { + t.Fatalf("post-cancel result = %+v", post) + } +} + func insertHistoricalReservation(t *testing.T, db *storage.DB, paymentID, profileID string, created time.Time, releasedAt *time.Time, status string) { t.Helper() if profileID == "kotak" { @@ -204,7 +411,7 @@ func insertHistoricalReservation(t *testing.T, db *storage.DB, paymentID, profil } } -func TestKotakLowConfidenceLatestReuseFailsAmbiguous(t *testing.T) { +func TestKotakSMSObservationIsUnsupported(t *testing.T) { ctx := context.Background() db := openAllocatorDB(t) base := time.UnixMilli(1_788_200_000_000).UTC() @@ -216,19 +423,16 @@ func TestKotakLowConfidenceLatestReuseFailsAmbiguous(t *testing.T) { insertRelayEvent(t, db, "relay_kotak_reuse", "source_kotak_reuse", observations.GoogleMessagesPackage, occurred, received) s := newTestService(t, db, received) obs := observations.Observation{Source: "kotak_sms", CollectionProfileID: "kotak", AmountPaise: 10037, OccurredAt: occurred, OccurredAtSource: "notification_posted_at"} - result, err := s.ApplyObservation(ctx, "relay_kotak_reuse", obs, received) - if err != nil { - t.Fatal(err) - } - if result.Result != "ambiguous" || result.PaymentID != "" || result.Transitioned { - t.Fatalf("reuse result = %+v", result) + if _, err := s.ApplyObservation(ctx, "relay_kotak_reuse", obs, received); !errors.Is(err, ErrInvalidObservation) { + t.Fatalf("Kotak SMS observation error = %v, want ErrInvalidObservation", err) } var status string if err := db.SQL.QueryRow(`SELECT status FROM payments WHERE id='new'`).Scan(&status); err != nil || status != "pending" { t.Fatalf("new payment status=%q err=%v", status, err) } } -func TestPaytmPostedTimeLatestReuseFailsAmbiguous(t *testing.T) { + +func TestPaytmPostedTimeLatestReuseMatchesUniqueLiveAmount(t *testing.T) { ctx := context.Background() db := openAllocatorDB(t) base := time.UnixMilli(1_788_200_000_000).UTC() @@ -243,15 +447,39 @@ func TestPaytmPostedTimeLatestReuseFailsAmbiguous(t *testing.T) { if err != nil { t.Fatal(err) } - if result.Result != "ambiguous" || result.PaymentID != "" || result.Transitioned { + if result.Result != "matched" || result.PaymentID != "new_paytm" || !result.Transitioned { t.Fatalf("reuse result = %+v", result) } var status string - if err := db.SQL.QueryRow(`SELECT status FROM payments WHERE id='new_paytm'`).Scan(&status); err != nil || status != "pending" { + if err := db.SQL.QueryRow(`SELECT status FROM payments WHERE id='new_paytm'`).Scan(&status); err != nil || status != "paid" { t.Fatalf("new payment status=%q err=%v", status, err) } } +func TestReleasedReservationDoesNotMatchObservationAfterRelease(t *testing.T) { + ctx := context.Background() + db := openAllocatorDB(t) + base := time.UnixMilli(1_788_200_000_000).UTC() + released := base.Add(5 * time.Minute) + insertHistoricalReservation(t, db, "released_old", "paytm", base, &released, "expired") + occurred := released.Add(time.Second) + received := occurred.Add(time.Second) + insertRelayEvent(t, db, "relay_after_release", "source_after_release", observations.PaytmBusinessPackage, occurred, received) + s := newTestService(t, db, received) + + result, err := s.ApplyObservation(ctx, "relay_after_release", paytmObservation(10037, occurred, "notification_text"), received) + if err != nil { + t.Fatal(err) + } + if result.Result != "unmatched" || result.PaymentID != "" || result.Transitioned { + t.Fatalf("post-release result = %+v", result) + } + var status string + if err := db.SQL.QueryRow(`SELECT status FROM payments WHERE id='released_old'`).Scan(&status); err != nil || status != "expired" { + t.Fatalf("released payment status=%q err=%v", status, err) + } +} + func TestTrustedHistoricalTimeCanMatchOldReservationAfterReuse(t *testing.T) { ctx := context.Background() db := openAllocatorDB(t) @@ -283,6 +511,42 @@ func TestTrustedHistoricalTimeCanMatchOldReservationAfterReuse(t *testing.T) { } } +func TestPaytmObservationCannotConfirmDifferentCollectionProfile(t *testing.T) { + ctx := context.Background() + db := openAllocatorDB(t) + base := time.UnixMilli(1_788_200_000_000).UTC() + insertHistoricalReservation(t, db, "kotak_pending", "kotak", base, nil, "pending") + occurred := base.Add(time.Minute) + received := occurred.Add(time.Second) + insertRelayEvent(t, db, "relay_paytm_wrong_profile", "source_paytm_wrong_profile", observations.PaytmBusinessPackage, occurred, received) + s := newTestService(t, db, received) + + result, err := s.ApplyObservation(ctx, "relay_paytm_wrong_profile", paytmObservation(10037, occurred, "notification_text"), received) + if err != nil { + t.Fatal(err) + } + if result.Result != "unmatched" || result.PaymentID != "" || result.Transitioned { + t.Fatalf("cross-profile Paytm result = %+v", result) + } + var status string + if err := db.SQL.QueryRow(`SELECT status FROM payments WHERE id='kotak_pending'`).Scan(&status); err != nil || status != "pending" { + t.Fatalf("non-Paytm payment status=%q err=%v", status, err) + } + + genericOccurred := occurred.Add(time.Second) + genericReceived := genericOccurred.Add(time.Second) + insertRelayEvent(t, db, "relay_generic_profile", "source_generic_profile", "com.google.android.apps.nbu.paisa.user", genericOccurred, genericReceived) + generic := observations.Observation{Source: observations.GenericNotificationSource, AmountPaise: 10037, OccurredAt: genericOccurred, OccurredAtSource: "notification_posted_at"} + s.Now = func() time.Time { return genericReceived } + result, err = s.ApplyObservation(ctx, "relay_generic_profile", generic, genericReceived) + if err != nil { + t.Fatal(err) + } + if result.Result != "matched" || result.PaymentID != "kotak_pending" || !result.Transitioned { + t.Fatalf("generic cross-profile result = %+v", result) + } +} + func TestPaytmMatchIgnoresCurrentActiveProfile(t *testing.T) { ctx := context.Background() db := openAllocatorDB(t) diff --git a/internal/v4/payments/service.go b/internal/v4/payments/service.go index 9db4ccb..482e020 100644 --- a/internal/v4/payments/service.go +++ b/internal/v4/payments/service.go @@ -11,6 +11,7 @@ import ( "errors" "fmt" "io" + "math" "net/url" "strconv" "strings" @@ -219,9 +220,8 @@ func normalizeCreateInput(input CreateInput) (CreateInput, [32]byte, [32]byte, e input.ExternalID = strings.TrimSpace(input.ExternalID) input.IdempotencyScope = strings.TrimSpace(input.IdempotencyScope) input.IdempotencyKey = strings.TrimSpace(input.IdempotencyKey) - - if input.RequestedAmountPaise <= 0 || input.RequestedAmountPaise%100 != 0 { - return CreateInput{}, [32]byte{}, [32]byte{}, fmt.Errorf("%w: amount must be positive whole INR", ErrInvalidPaymentInput) + if input.RequestedAmountPaise <= 0 || input.RequestedAmountPaise%100 != 0 || input.RequestedAmountPaise > math.MaxInt64-199 { + return CreateInput{}, [32]byte{}, [32]byte{}, fmt.Errorf("%w: amount must be positive whole INR within the payable range", ErrInvalidPaymentInput) } if input.Name == "" || utf8.RuneCountInString(input.Name) > 120 { return CreateInput{}, [32]byte{}, [32]byte{}, fmt.Errorf("%w: name must contain 1-120 characters", ErrInvalidPaymentInput) diff --git a/internal/v4/payments/service_test.go b/internal/v4/payments/service_test.go index bea683a..c0d7e53 100644 --- a/internal/v4/payments/service_test.go +++ b/internal/v4/payments/service_test.go @@ -6,6 +6,7 @@ import ( "encoding/json" "errors" "fmt" + "math" "net/url" "strings" "testing" @@ -289,8 +290,8 @@ func TestProfileSwitchAffectsOnlyNewPayments(t *testing.T) { if second.Payment.CollectionProfileID != "kotak" || second.Payment.UPIIDSnapshot != "merchant@kotak" { t.Fatalf("second payment did not use Kotak: %+v", second.Payment) } - if second.Payment.PayableAmountPaise != 10037 { - t.Fatalf("profile-scoped amount should allow same exact amount, got %d", second.Payment.PayableAmountPaise) + if second.Payment.PayableAmountPaise == first.Payment.PayableAmountPaise { + t.Fatalf("live payable amount reused across profiles: %d", second.Payment.PayableAmountPaise) } if !strings.Contains(second.UPIURI, "pa=merchant%40kotak") || !strings.Contains(second.UPIURI, "pn=PayGate%20Kotak") { t.Fatalf("Kotak UPI URI = %q", second.UPIURI) @@ -312,6 +313,18 @@ func TestCreateUsesOverflowOnlyAfterBaseBucketExhausted(t *testing.T) { } } +func TestCreateRejectsRequestedAmountThatOverflowsPayableBucket(t *testing.T) { + db := openAllocatorDB(t) + now := time.UnixMilli(1_788_200_000_000).UTC() + s := newTestService(t, db, now) + input := validCreateInput("overflow-boundary") + input.RequestedAmountPaise = math.MaxInt64 - math.MaxInt64%100 + if _, err := s.Create(context.Background(), input); !errors.Is(err, ErrInvalidPaymentInput) { + t.Fatalf("boundary amount error = %v, want ErrInvalidPaymentInput", err) + } + assertCount(t, db.SQL, "payments", 0) +} + func TestCreateFailsClosedWithoutActiveProfile(t *testing.T) { db := openAllocatorDB(t) if _, err := db.SQL.Exec(`UPDATE collection_profiles SET active=0 WHERE id='paytm'`); err != nil { diff --git a/internal/v4/profiles/service.go b/internal/v4/profiles/service.go index ac9baf1..b0204e5 100644 --- a/internal/v4/profiles/service.go +++ b/internal/v4/profiles/service.go @@ -17,6 +17,7 @@ var ( ErrProfileDisabled = errors.New("collection profile is disabled") ErrCannotDisableActiveProfile = errors.New("cannot disable active collection profile") ErrInvalidProfile = errors.New("invalid collection profile") + ErrRelayDeviceNotAuthorized = errors.New("relay device is not authorized") ) type Service struct { @@ -101,6 +102,31 @@ func (s *Service) Upsert(ctx context.Context, in UpsertInput) (Profile, error) { } func (s *Service) UpdateDestination(ctx context.Context, id string, in DestinationInput) (Profile, error) { + return s.updateDestination(ctx, "", time.Time{}, id, in) +} + +func (s *Service) UpdateActiveDestination(ctx context.Context, in DestinationInput) (Profile, error) { + return s.updateDestination(ctx, "", time.Time{}, "active", in) +} + +func (s *Service) UpdateActiveDestinationForRelay(ctx context.Context, deviceID string, enrolledAt time.Time, in DestinationInput) (Profile, error) { + deviceID = strings.TrimSpace(deviceID) + if deviceID == "" || enrolledAt.IsZero() { + return Profile{}, ErrRelayDeviceNotAuthorized + } + return s.updateDestination(ctx, deviceID, enrolledAt, "active", in) +} + +func (s *Service) UpdateDestinationForRelay(ctx context.Context, deviceID string, enrolledAt time.Time, id string, in DestinationInput) (Profile, error) { + deviceID = strings.TrimSpace(deviceID) + if deviceID == "" || enrolledAt.IsZero() { + return Profile{}, ErrRelayDeviceNotAuthorized + } + return s.updateDestination(ctx, deviceID, enrolledAt, id, in) +} + +func (s *Service) updateDestination(ctx context.Context, deviceID string, enrolledAt time.Time, id string, in DestinationInput) (Profile, error) { + if s == nil || s.DB == nil || s.DB.SQL == nil { return Profile{}, errors.New("profile storage is required") } @@ -120,6 +146,31 @@ func (s *Service) UpdateDestination(ctx context.Context, id string, in Destinati now := nowFn().UTC() var out Profile err := s.DB.WithImmediateTx(ctx, func(tx *storage.ImmediateTx) error { + if deviceID != "" { + var enabled int + err := tx.QueryRowContext(ctx, `SELECT enabled FROM relay_devices WHERE id=? AND enabled=1 AND enrolled_at=?`, + deviceID, enrolledAt.UTC().UnixMilli()).Scan(&enabled) + if errors.Is(err, sql.ErrNoRows) { + return ErrRelayDeviceNotAuthorized + } + if err != nil { + return fmt.Errorf("read relay device authorization: %w", err) + } + if enabled != 1 { + return ErrRelayDeviceNotAuthorized + } + } + if id == "active" { + var activeID string + err := tx.QueryRowContext(ctx, `SELECT id FROM collection_profiles WHERE active=1 LIMIT 1`).Scan(&activeID) + if errors.Is(err, sql.ErrNoRows) { + return ErrProfileNotFound + } + if err != nil { + return fmt.Errorf("read active collection profile: %w", err) + } + id = activeID + } if _, err := getWith(ctx, tx, id); err != nil { return err } diff --git a/internal/v4/profiles/service_test.go b/internal/v4/profiles/service_test.go index 2f6422b..f5fc355 100644 --- a/internal/v4/profiles/service_test.go +++ b/internal/v4/profiles/service_test.go @@ -151,6 +151,49 @@ func TestUpdateDestinationPreservesExistingPaymentSnapshot(t *testing.T) { } } +func TestRelayDestinationUpdateRejectsRevokedDeviceAtomically(t *testing.T) { + ctx := context.Background() + db := openTestDB(t) + svc := NewService(db) + if _, err := svc.Upsert(ctx, testProfile("paytm", true)); err != nil { + t.Fatal(err) + } + if _, err := svc.Activate(ctx, "paytm"); err != nil { + t.Fatal(err) + } + enrolledAt := time.UnixMilli(1_788_200_000_000).UTC() + if _, err := db.SQL.Exec(`INSERT INTO relay_devices(id,public_key_pem,enabled,enrolled_at) VALUES('relay-1','pem',1,?)`, enrolledAt.UnixMilli()); err != nil { + t.Fatal(err) + } + if _, err := svc.UpdateDestinationForRelay(ctx, "relay-1", enrolledAt, "paytm", DestinationInput{UPIID: "newdestination@upi"}); err != nil { + t.Fatal(err) + } + if _, err := db.SQL.Exec(`UPDATE relay_devices SET enabled=0 WHERE id='relay-1'`); err != nil { + t.Fatal(err) + } + if _, err := svc.UpdateDestinationForRelay(ctx, "relay-1", enrolledAt, "paytm", DestinationInput{UPIID: "revoked@upi"}); !errors.Is(err, ErrRelayDeviceNotAuthorized) { + t.Fatalf("revoked device error = %v", err) + } + profile, err := svc.Get(ctx, "paytm") + if err != nil { + t.Fatal(err) + } + if profile.UPIID != "newdestination@upi" { + t.Fatalf("revoked device changed destination to %q", profile.UPIID) + } + + newEnrolledAt := enrolledAt.Add(time.Minute) + if _, err := db.SQL.Exec(`UPDATE relay_devices SET enabled=1,enrolled_at=? WHERE id='relay-1'`, newEnrolledAt.UnixMilli()); err != nil { + t.Fatal(err) + } + if _, err := svc.UpdateDestinationForRelay(ctx, "relay-1", enrolledAt, "paytm", DestinationInput{UPIID: "stale-epoch@upi"}); !errors.Is(err, ErrRelayDeviceNotAuthorized) { + t.Fatalf("stale enrollment error = %v", err) + } + if _, err := svc.UpdateDestinationForRelay(ctx, "relay-1", newEnrolledAt, "paytm", DestinationInput{UPIID: "repaired@upi"}); err != nil { + t.Fatal(err) + } +} + func TestProfileValidationAndNotFound(t *testing.T) { ctx := context.Background() db := openTestDB(t) diff --git a/internal/v4/relay/heartbeat.go b/internal/v4/relay/heartbeat.go index 4a62940..6ec657c 100644 --- a/internal/v4/relay/heartbeat.go +++ b/internal/v4/relay/heartbeat.go @@ -5,6 +5,7 @@ import ( "encoding/json" "errors" "fmt" + "github.com/Phloraxx/payment-api/internal/v4/storage" "strings" "time" ) @@ -72,19 +73,27 @@ func (s *Service) HeartbeatSigned(ctx context.Context, auth RequestAuth, rawBody } delivered = t.UnixMilli() } - result, err := s.DB.SQL.ExecContext(ctx, `UPDATE relay_devices SET - last_seen_at=?,last_heartbeat_at=?,app_version=?,device_model=?,android_version=?, - notification_access=?,listener_connected=?,battery_optimization_exempt=?,power_save_mode=?, - background_restricted=?,foreground_service=?,pending_count=?,failed_count=?,last_successful_delivery_at=?,last_client_error=? - WHERE id=? AND enabled=1`, - now.UnixMilli(), now.UnixMilli(), nullableText(input.AppVersion), nullableText(input.DeviceModel), nullableText(input.AndroidVersion), - boolInt(input.NotificationAccess), boolInt(input.ListenerConnected), boolInt(input.BatteryOptimizationExempt), boolInt(input.PowerSaveMode), - boolInt(input.BackgroundRestricted), boolInt(input.ForegroundService), input.PendingCount, input.FailedCount, - delivered, nullableText(input.LastClientError), device.ID) + var rowsAffected int64 + err = s.DB.WithImmediateTx(ctx, func(tx *storage.ImmediateTx) error { + result, err := tx.ExecContext(ctx, `UPDATE relay_devices SET + last_seen_at=?,last_heartbeat_at=?,app_version=?,device_model=?,android_version=?, + notification_access=?,listener_connected=?,battery_optimization_exempt=?,power_save_mode=?, + background_restricted=?,foreground_service=?,pending_count=?,failed_count=?,last_successful_delivery_at=?,last_client_error=? + WHERE id=? AND enabled=1 AND enrolled_at=?`, + now.UnixMilli(), now.UnixMilli(), nullableText(input.AppVersion), nullableText(input.DeviceModel), nullableText(input.AndroidVersion), + boolInt(input.NotificationAccess), boolInt(input.ListenerConnected), boolInt(input.BatteryOptimizationExempt), boolInt(input.PowerSaveMode), + boolInt(input.BackgroundRestricted), boolInt(input.ForegroundService), input.PendingCount, input.FailedCount, + delivered, nullableText(input.LastClientError), device.ID, device.EnrolledAt.UnixMilli()) + if err != nil { + return fmt.Errorf("persist relay heartbeat: %w", err) + } + rowsAffected, _ = result.RowsAffected() + return nil + }) if err != nil { - return HeartbeatResult{}, fmt.Errorf("persist relay heartbeat: %w", err) + return HeartbeatResult{}, err } - if rows, _ := result.RowsAffected(); rows != 1 { + if rowsAffected != 1 { return HeartbeatResult{}, relayError("UNKNOWN_RELAY_DEVICE", "relay device is not enrolled or is disabled", 401) } return HeartbeatResult{ReceivedAt: now}, nil diff --git a/internal/v4/relay/pairing.go b/internal/v4/relay/pairing.go index f6da24a..0c7e486 100644 --- a/internal/v4/relay/pairing.go +++ b/internal/v4/relay/pairing.go @@ -9,6 +9,7 @@ import ( "encoding/hex" "errors" "fmt" + "math" "strings" "time" "unicode/utf8" @@ -20,16 +21,14 @@ var ( ErrPairingTokenInvalid = errors.New("pairing token is invalid") ErrPairingTokenExpired = errors.New("pairing token has expired") ErrPairingTokenUsed = errors.New("pairing token has already been used") - ErrRelayAlreadyActive = errors.New("an active relay device already exists") ErrInvalidDevice = errors.New("invalid relay device") ) const defaultPairingTTL = 2 * time.Minute type PairingSession struct { - Token string - ExpiresAt time.Time - ReplaceExisting bool + Token string + ExpiresAt time.Time } type PairDeviceInput struct { @@ -42,15 +41,16 @@ type PairDeviceInput struct { } type PairDeviceResult struct { - DeviceID string - ReplacedDeviceID string - Enabled bool + DeviceID string + Enabled bool + EnrolledAtMS int64 } type DeviceInfo struct { ID string `json:"id"` Name string `json:"name"` Enabled bool `json:"enabled"` + Operational bool `json:"operational"` EnrolledAt time.Time `json:"enrolled_at"` LastSeenAt *time.Time `json:"last_seen_at,omitempty"` LastHeartbeatAt *time.Time `json:"last_heartbeat_at,omitempty"` @@ -69,7 +69,7 @@ type DeviceInfo struct { LastClientError string `json:"last_client_error,omitempty"` } -func (s *Service) CreatePairing(ctx context.Context, replaceExisting bool) (PairingSession, error) { +func (s *Service) CreatePairing(ctx context.Context) (PairingSession, error) { if s == nil || s.DB == nil || s.DB.SQL == nil { return PairingSession{}, errors.New("relay storage is required") } @@ -103,11 +103,9 @@ func (s *Service) CreatePairing(ctx context.Context, replaceExisting bool) (Pair return PairingSession{}, fmt.Errorf("generate pairing session id: %w", err) } tokenHash := sha256.Sum256([]byte(token)) - // replaceExisting is retained only for wire compatibility with older clients. - // Pairing is additive in v5: every valid QR enrolls one independently revocable device. err = s.DB.WithImmediateTx(ctx, func(tx *storage.ImmediateTx) error { - _, err := tx.ExecContext(ctx, `INSERT INTO pairing_sessions(id,token_hash,replace_existing,created_at,expires_at) - VALUES(?,?,?,?,?)`, sessionID, tokenHash[:], 0, now.UnixMilli(), expiresAt.UnixMilli()) + _, err := tx.ExecContext(ctx, `INSERT INTO pairing_sessions(id,token_hash,created_at,expires_at) + VALUES(?,?,?,?)`, sessionID, tokenHash[:], now.UnixMilli(), expiresAt.UnixMilli()) if err != nil { return fmt.Errorf("create pairing session: %w", err) } @@ -116,7 +114,7 @@ func (s *Service) CreatePairing(ctx context.Context, replaceExisting bool) (Pair if err != nil { return PairingSession{}, err } - return PairingSession{Token: token, ExpiresAt: expiresAt, ReplaceExisting: false}, nil + return PairingSession{Token: token, ExpiresAt: expiresAt}, nil } func (s *Service) PairDevice(ctx context.Context, input PairDeviceInput) (PairDeviceResult, error) { @@ -134,6 +132,7 @@ func (s *Service) PairDevice(ctx context.Context, input PairDeviceInput) (PairDe now := nowFn().UTC() tokenHash := sha256.Sum256([]byte(normalized.Token)) result := PairDeviceResult{DeviceID: deviceID, Enabled: true} + enrollmentEpoch := now.UnixMilli() err = s.DB.WithImmediateTx(ctx, func(tx *storage.ImmediateTx) error { var expiresAt int64 var consumedAt sql.NullInt64 @@ -150,12 +149,23 @@ func (s *Service) PairDevice(ctx context.Context, input PairDeviceInput) (PairDe if now.UnixMilli() >= expiresAt { return ErrPairingTokenExpired } + var previousEpoch int64 + err = tx.QueryRowContext(ctx, `SELECT enrolled_at FROM relay_devices WHERE id=?`, deviceID).Scan(&previousEpoch) + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return fmt.Errorf("read existing relay enrollment epoch: %w", err) + } + if err == nil && previousEpoch >= enrollmentEpoch { + if previousEpoch == math.MaxInt64 { + return errors.New("relay enrollment epoch exhausted") + } + enrollmentEpoch = previousEpoch + 1 + } _, err = tx.ExecContext(ctx, `INSERT INTO relay_devices( - id,name,public_key_pem,enabled,enrolled_at,app_version,device_model,android_version) - VALUES(?,?,?,1,?,?,?,?) + id,name,public_key_pem,enabled,enrolled_at,epoch_required,app_version,device_model,android_version) + VALUES(?,?,?,1,?,1,?,?,?) ON CONFLICT(id) DO UPDATE SET name=excluded.name,public_key_pem=excluded.public_key_pem, - enabled=1,app_version=excluded.app_version,device_model=excluded.device_model,android_version=excluded.android_version`, - deviceID, normalized.Name, normalized.PublicKeyPEM, now.UnixMilli(), nullableText(normalized.AppVersion), + enabled=1,enrolled_at=excluded.enrolled_at,epoch_required=1,app_version=excluded.app_version,device_model=excluded.device_model,android_version=excluded.android_version`, + deviceID, normalized.Name, normalized.PublicKeyPEM, enrollmentEpoch, nullableText(normalized.AppVersion), nullableText(normalized.DeviceModel), nullableText(normalized.AndroidVersion)) if err != nil { return fmt.Errorf("enroll relay device: %w", err) @@ -172,6 +182,7 @@ func (s *Service) PairDevice(ctx context.Context, input PairDeviceInput) (PairDe if err != nil { return PairDeviceResult{}, err } + result.EnrolledAtMS = enrollmentEpoch return result, nil } @@ -180,11 +191,19 @@ func (s *Service) RevokeDevice(ctx context.Context, deviceID string) error { if s == nil || s.DB == nil || s.DB.SQL == nil || deviceID == "" { return ErrInvalidDevice } - result, err := s.DB.SQL.ExecContext(ctx, `UPDATE relay_devices SET enabled=0 WHERE id=? AND enabled=1`, deviceID) + var rowsAffected int64 + err := s.DB.WithImmediateTx(ctx, func(tx *storage.ImmediateTx) error { + result, err := tx.ExecContext(ctx, `UPDATE relay_devices SET enabled=0 WHERE id=? AND enabled=1`, deviceID) + if err != nil { + return fmt.Errorf("revoke relay device: %w", err) + } + rowsAffected, _ = result.RowsAffected() + return nil + }) if err != nil { - return fmt.Errorf("revoke relay device: %w", err) + return err } - if rows, _ := result.RowsAffected(); rows != 1 { + if rowsAffected != 1 { return ErrInvalidDevice } return nil @@ -237,6 +256,11 @@ func (s *Service) Devices(ctx context.Context) ([]DeviceInfo, error) { if s == nil || s.DB == nil || s.DB.SQL == nil { return nil, errors.New("relay storage is required") } + nowFn := s.Now + if nowFn == nil { + nowFn = time.Now + } + now := nowFn().UTC() rows, err := s.DB.SQL.QueryContext(ctx, `SELECT id,COALESCE(name,''),enabled,enrolled_at,last_seen_at,last_heartbeat_at, app_version,device_model,android_version,notification_access,listener_connected,battery_optimization_exempt, power_save_mode,background_restricted,foreground_service,pending_count,failed_count,last_successful_delivery_at,last_client_error @@ -270,6 +294,14 @@ func (s *Service) Devices(ctx context.Context) ([]DeviceInfo, error) { info.BatteryOptimizationExempt = nullableBoolPointer(batteryExempt) info.PowerSaveMode = nullableBoolPointer(powerSave) info.BackgroundRestricted = nullableBoolPointer(backgroundRestricted) + age := now.Sub(time.UnixMilli(lastHeartbeat.Int64).UTC()) + info.Operational = lastHeartbeat.Valid && + age >= -5*time.Minute && age <= time.Hour && + notificationAccess.Valid && notificationAccess.Int64 == 1 && + listenerConnected.Valid && listenerConnected.Int64 == 1 && + batteryExempt.Valid && batteryExempt.Int64 == 1 && + backgroundRestricted.Valid && backgroundRestricted.Int64 == 0 && + foregroundService.Valid && foregroundService.Int64 == 1 info.ForegroundService = nullableBoolPointer(foregroundService) info.PendingCount = nullableIntPointer(pendingCount) info.FailedCount = nullableIntPointer(failedCount) @@ -281,17 +313,6 @@ func (s *Service) Devices(ctx context.Context) ([]DeviceInfo, error) { return items, nil } -func (s *Service) ActiveDevice(ctx context.Context) (*DeviceInfo, error) { - items, err := s.Devices(ctx) - if err != nil { - return nil, err - } - if len(items) == 0 { - return nil, nil - } - return &items[0], nil -} - func (s *Service) Device(ctx context.Context, id string) (*DeviceInfo, error) { id = strings.ToLower(strings.TrimSpace(id)) if id == "" { diff --git a/internal/v4/relay/pairing_test.go b/internal/v4/relay/pairing_test.go index 13451ca..47f584b 100644 --- a/internal/v4/relay/pairing_test.go +++ b/internal/v4/relay/pairing_test.go @@ -40,21 +40,20 @@ func TestCreatePairingStoresOnlyTokenHash(t *testing.T) { token := "abcdefghijklmnopqrstuvwxyz0123456789ABCDEFG" service.NewPairingToken = func() (string, error) { return token, nil } - session, err := service.CreatePairing(context.Background(), false) + session, err := service.CreatePairing(context.Background()) if err != nil { t.Fatal(err) } - if session.Token != token || !session.ExpiresAt.Equal(now.Add(2*time.Minute)) || session.ReplaceExisting { + if session.Token != token || !session.ExpiresAt.Equal(now.Add(2*time.Minute)) { t.Fatalf("pairing session = %+v", session) } var stored []byte - var replace int var expires int64 - if err := db.SQL.QueryRow(`SELECT token_hash,replace_existing,expires_at FROM pairing_sessions`).Scan(&stored, &replace, &expires); err != nil { + if err := db.SQL.QueryRow(`SELECT token_hash,expires_at FROM pairing_sessions`).Scan(&stored, &expires); err != nil { t.Fatal(err) } wantHash := sha256.Sum256([]byte(token)) - if hex.EncodeToString(stored) != hex.EncodeToString(wantHash[:]) || replace != 0 || expires != session.ExpiresAt.UnixMilli() { + if hex.EncodeToString(stored) != hex.EncodeToString(wantHash[:]) || expires != session.ExpiresAt.UnixMilli() { t.Fatal("stored pairing session does not match hashed-token contract") } var rawCount int @@ -70,7 +69,7 @@ func TestPairDeviceConsumesTokenAndEnablesFingerprintDevice(t *testing.T) { now := time.Date(2026, 9, 1, 6, 15, 0, 0, time.UTC) service := NewService(db, payments.NewService(db)) service.Now = func() time.Time { return now } - session, err := service.CreatePairing(context.Background(), false) + session, err := service.CreatePairing(context.Background()) if err != nil { t.Fatal(err) } @@ -82,15 +81,15 @@ func TestPairDeviceConsumesTokenAndEnablesFingerprintDevice(t *testing.T) { if err != nil { t.Fatal(err) } - if result.DeviceID != deviceID || !result.Enabled || result.ReplacedDeviceID != "" { + if result.DeviceID != deviceID || !result.Enabled || result.EnrolledAtMS != now.UnixMilli() { t.Fatalf("pair result = %+v", result) } - active, err := service.ActiveDevice(context.Background()) + devices, err := service.Devices(context.Background()) if err != nil { t.Fatal(err) } - if active == nil || active.ID != deviceID || active.Name != "Motorola Edge 60 Stylus" || !active.Enabled { - t.Fatalf("active device = %+v", active) + if len(devices) != 1 || devices[0].ID != deviceID || devices[0].Name != "Motorola Edge 60 Stylus" || !devices[0].Enabled { + t.Fatalf("devices = %+v", devices) } var consumed sql.NullInt64 if err := db.SQL.QueryRow(`SELECT consumed_at FROM pairing_sessions`).Scan(&consumed); err != nil || !consumed.Valid { @@ -102,13 +101,53 @@ func TestPairDeviceConsumesTokenAndEnablesFingerprintDevice(t *testing.T) { t.Fatalf("replay error = %v", err) } } +func TestRePairRefreshesEnrollmentEpoch(t *testing.T) { + db := openRelayDB(t) + firstAt := time.Date(2026, 9, 1, 6, 20, 0, 0, time.UTC) + service := NewService(db, payments.NewService(db)) + service.Now = func() time.Time { return firstAt } + session, err := service.CreatePairing(context.Background()) + if err != nil { + t.Fatal(err) + } + publicKey, deviceID := newPairingPublicKey(t) + input := PairDeviceInput{Token: session.Token, Name: "Phone", PublicKeyPEM: publicKey} + if _, err := service.PairDevice(context.Background(), input); err != nil { + t.Fatal(err) + } + var firstEpoch int64 + if err := db.SQL.QueryRow(`SELECT enrolled_at FROM relay_devices WHERE id=?`, deviceID).Scan(&firstEpoch); err != nil { + t.Fatal(err) + } + + // Re-pair at the exact same wall-clock millisecond. The enrollment epoch + // must still advance so old epoch-bound signatures become invalid. + secondAt := firstAt + service.Now = func() time.Time { return secondAt } + second, err := service.CreatePairing(context.Background()) + if err != nil { + t.Fatal(err) + } + input.Token = second.Token + repaired, err := service.PairDevice(context.Background(), input) + if err != nil { + t.Fatal(err) + } + var secondEpoch int64 + if err := db.SQL.QueryRow(`SELECT enrolled_at FROM relay_devices WHERE id=?`, deviceID).Scan(&secondEpoch); err != nil { + t.Fatal(err) + } + if secondEpoch != firstEpoch+1 || repaired.EnrolledAtMS != secondEpoch { + t.Fatalf("enrollment epoch first=%d second=%d returned=%d", firstEpoch, secondEpoch, repaired.EnrolledAtMS) + } +} func TestAdditionalDevicePairingKeepsExistingDeviceEnabled(t *testing.T) { db := openRelayDB(t) now := time.Date(2026, 9, 1, 6, 30, 0, 0, time.UTC) _, oldID := enrollTestDevice(t, db, now.Add(-time.Hour)) service := NewService(db, payments.NewService(db)) service.Now = func() time.Time { return now } - session, err := service.CreatePairing(context.Background(), false) + session, err := service.CreatePairing(context.Background()) if err != nil { t.Fatal(err) } @@ -117,7 +156,7 @@ func TestAdditionalDevicePairingKeepsExistingDeviceEnabled(t *testing.T) { if err != nil { t.Fatal(err) } - if result.DeviceID != newID || result.ReplacedDeviceID != "" || !result.Enabled { + if result.DeviceID != newID || !result.Enabled { t.Fatalf("pair result=%+v", result) } var oldEnabled, newEnabled int @@ -136,25 +175,22 @@ func TestAdditionalDevicePairingKeepsExistingDeviceEnabled(t *testing.T) { } } -func TestReplaceFlagIsBackwardCompatibleButAdditive(t *testing.T) { +func TestAdditionalPairingKeepsLegacySchemaCompatible(t *testing.T) { db := openRelayDB(t) now := time.Date(2026, 9, 1, 6, 45, 0, 0, time.UTC) _, oldID := enrollTestDevice(t, db, now.Add(-time.Hour)) service := NewService(db, payments.NewService(db)) service.Now = func() time.Time { return now } - session, err := service.CreatePairing(context.Background(), true) + session, err := service.CreatePairing(context.Background()) if err != nil { t.Fatal(err) } - if session.ReplaceExisting { - t.Fatalf("replace flag should be deprecated in v5: %+v", session) - } publicKey, newID := newPairingPublicKey(t) result, err := service.PairDevice(context.Background(), PairDeviceInput{Token: session.Token, Name: "Another Phone", PublicKeyPEM: publicKey}) if err != nil { t.Fatal(err) } - if result.DeviceID != newID || result.ReplacedDeviceID != "" { + if result.DeviceID != newID { t.Fatalf("pair result=%+v", result) } var oldEnabled, newEnabled int @@ -164,13 +200,13 @@ func TestReplaceFlagIsBackwardCompatibleButAdditive(t *testing.T) { t.Fatalf("enabled states old=%d new=%d", oldEnabled, newEnabled) } } -func TestExpiredReplacementLeavesOldDeviceEnabled(t *testing.T) { +func TestExpiredPairingLeavesOldDeviceEnabled(t *testing.T) { db := openRelayDB(t) now := time.Date(2026, 9, 1, 7, 0, 0, 0, time.UTC) _, oldID := enrollTestDevice(t, db, now.Add(-time.Hour)) service := NewService(db, payments.NewService(db)) service.Now = func() time.Time { return now } - session, err := service.CreatePairing(context.Background(), true) + session, err := service.CreatePairing(context.Background()) if err != nil { t.Fatal(err) } @@ -195,11 +231,11 @@ func TestMultiplePairingSessionsCanCoexistAndAreIndividuallyOneUse(t *testing.T) tokens := []string{"first-pairing-token-abcdefghijklmnopqrstuvwxyz", "second-pairing-token-abcdefghijklmnopqrstuvwxyz"} index := 0 service.NewPairingToken = func() (string, error) { value := tokens[index]; index++; return value, nil } - first, err := service.CreatePairing(context.Background(), false) + first, err := service.CreatePairing(context.Background()) if err != nil { t.Fatal(err) } - second, err := service.CreatePairing(context.Background(), false) + second, err := service.CreatePairing(context.Background()) if err != nil { t.Fatal(err) } @@ -231,21 +267,21 @@ func TestRevokeDeviceDisablesWithoutDeletingHistory(t *testing.T) { if err := service.RevokeDevice(context.Background(), deviceID); err != nil { t.Fatal(err) } - active, err := service.ActiveDevice(context.Background()) - if err != nil || active != nil { - t.Fatalf("active after revoke = %+v err=%v", active, err) + devices, err := service.Devices(context.Background()) + if err != nil || len(devices) != 0 { + t.Fatalf("devices after revoke = %+v err=%v", devices, err) } if countRows(t, db, "relay_devices") != 1 { t.Fatal("revocation deleted relay history") } } -func TestReplacementRollbackKeepsOldDeviceAndTokenUnused(t *testing.T) { +func TestPairingRollbackKeepsOldDeviceAndTokenUnused(t *testing.T) { db := openRelayDB(t) now := time.Date(2026, 9, 1, 7, 45, 0, 0, time.UTC) _, oldID := enrollTestDevice(t, db, now.Add(-time.Hour)) service := NewService(db, payments.NewService(db)) service.Now = func() time.Time { return now } - session, err := service.CreatePairing(context.Background(), true) + session, err := service.CreatePairing(context.Background()) if err != nil { t.Fatal(err) } @@ -254,7 +290,7 @@ func TestReplacementRollbackKeepsOldDeviceAndTokenUnused(t *testing.T) { t.Fatal(err) } if _, err := service.PairDevice(context.Background(), PairDeviceInput{ - Token: session.Token, Name: "Replacement", PublicKeyPEM: publicKey, + Token: session.Token, Name: "New Phone", PublicKeyPEM: publicKey, }); err == nil { t.Fatal("expected forced enrollment failure") } @@ -267,9 +303,9 @@ func TestReplacementRollbackKeepsOldDeviceAndTokenUnused(t *testing.T) { t.Fatal(err) } if consumed.Valid { - t.Fatal("failed replacement consumed the pairing token") + t.Fatal("failed pairing consumed the pairing token") } if countRows(t, db, "relay_devices") != 1 { - t.Fatal("failed replacement persisted a partial new device") + t.Fatal("failed pairing persisted a partial new device") } } diff --git a/internal/v4/relay/security.go b/internal/v4/relay/security.go index c1208b4..141a9c3 100644 --- a/internal/v4/relay/security.go +++ b/internal/v4/relay/security.go @@ -25,11 +25,12 @@ const ( ) type RequestAuth struct { - DeviceID string - Timestamp string - Signature string - Method string - Path string + DeviceID string + Timestamp string + Signature string + EnrollmentEpoch string + Method string + Path string } type Error struct { Code string @@ -48,6 +49,10 @@ func relayError(code, message string, status int) error { return &Error{Code: code, Message: message, HTTPStatus: status} } +func invalidRelaySignature() error { + return relayError("INVALID_RELAY_SIGNATURE", "invalid relay signature", 401) +} + type verifiedDevice struct { ID string EnrolledAt time.Time @@ -58,6 +63,12 @@ func CanonicalRequest(method, path, timestamp string, body []byte) string { return fmt.Sprintf("%s\n%s\n%s\n%s", strings.ToUpper(method), path, strings.TrimSpace(timestamp), hex.EncodeToString(sum[:])) } + +func CanonicalRequestWithEpoch(method, path, timestamp string, enrollmentEpoch int64, body []byte) string { + sum := sha256.Sum256(body) + return fmt.Sprintf("%s\n%s\n%s\n%d\n%s", strings.ToUpper(method), path, + strings.TrimSpace(timestamp), enrollmentEpoch, hex.EncodeToString(sum[:])) +} func verifyRequest(ctx context.Context, db *storage.DB, auth RequestAuth, body []byte, now time.Time) (verifiedDevice, error) { deviceID := strings.ToLower(strings.TrimSpace(auth.DeviceID)) if len(deviceID) != 64 { @@ -81,8 +92,10 @@ func verifyRequest(ctx context.Context, db *storage.DB, auth RequestAuth, body [ var publicKeyPEM string var enabled int var enrolledAt int64 - err = db.SQL.QueryRowContext(ctx, `SELECT public_key_pem,enabled,enrolled_at FROM relay_devices WHERE id=?`, deviceID). - Scan(&publicKeyPEM, &enabled, &enrolledAt) + var epochRequired int + err = db.SQL.QueryRowContext(ctx, `SELECT public_key_pem,enabled,enrolled_at,epoch_required + FROM relay_devices WHERE id=?`, deviceID). + Scan(&publicKeyPEM, &enabled, &enrolledAt, &epochRequired) if errors.Is(err, sql.ErrNoRows) { return verifiedDevice{}, relayError("UNKNOWN_RELAY_DEVICE", "relay device is not enrolled or is disabled", 401) } @@ -92,6 +105,22 @@ func verifyRequest(ctx context.Context, db *storage.DB, auth RequestAuth, body [ if enabled != 1 { return verifiedDevice{}, relayError("UNKNOWN_RELAY_DEVICE", "relay device is not enrolled or is disabled", 401) } + epochHeader := strings.TrimSpace(auth.EnrollmentEpoch) + epochBound := epochHeader != "" + // Devices enrolled before epoch enforcement may continue using the legacy canonical + // form so already-installed relays survive the server rollout. Pairing or + // re-pairing marks the device epoch-required, making that enrollment an explicit + // security boundary for every subsequent signed request. + if !epochBound && epochRequired == 1 { + return verifiedDevice{}, invalidRelaySignature() + } + enrollmentEpoch := int64(0) + if epochBound { + enrollmentEpoch, err = strconv.ParseInt(epochHeader, 10, 64) + if err != nil || enrollmentEpoch <= 0 || enrollmentEpoch != enrolledAt { + return verifiedDevice{}, invalidRelaySignature() + } + } pub, der, err := parsePublicKey(publicKeyPEM) if err != nil { return verifiedDevice{}, relayError("INVALID_RELAY_DEVICE_KEY", "stored relay device key is invalid", 500) @@ -102,12 +131,15 @@ func verifyRequest(ctx context.Context, db *storage.DB, auth RequestAuth, body [ } signature, err := base64.StdEncoding.DecodeString(strings.TrimSpace(auth.Signature)) if err != nil { - return verifiedDevice{}, relayError("INVALID_RELAY_SIGNATURE", "invalid relay signature", 401) + return verifiedDevice{}, invalidRelaySignature() } canonical := CanonicalRequest(auth.Method, auth.Path, auth.Timestamp, body) + if epochBound { + canonical = CanonicalRequestWithEpoch(auth.Method, auth.Path, auth.Timestamp, enrollmentEpoch, body) + } digest := sha256.Sum256([]byte(canonical)) if !ecdsa.VerifyASN1(pub, digest[:], signature) { - return verifiedDevice{}, relayError("INVALID_RELAY_SIGNATURE", "invalid relay signature", 401) + return verifiedDevice{}, invalidRelaySignature() } return verifiedDevice{ID: deviceID, EnrolledAt: time.UnixMilli(enrolledAt).UTC()}, nil } @@ -116,18 +148,33 @@ func verifyRequest(ctx context.Context, db *storage.DB, auth RequestAuth, body [ // It is shared by relay ingestion and the Android operational dashboard; callers // decide which application routes a device identity is authorized to use. func (s *Service) AuthenticateDevice(ctx context.Context, auth RequestAuth, body []byte) (string, error) { + device, err := s.authenticateDevice(ctx, auth, body) + if err != nil { + return "", err + } + return device.ID, nil +} + +// AuthenticateDeviceWithEpoch returns the enrollment epoch used to verify the request. +// Mutation handlers bind their write transaction to this epoch so re-pairing invalidates +// an in-flight request even when the device key itself is reused. +func (s *Service) AuthenticateDeviceWithEpoch(ctx context.Context, auth RequestAuth, body []byte) (string, time.Time, error) { + device, err := s.authenticateDevice(ctx, auth, body) + if err != nil { + return "", time.Time{}, err + } + return device.ID, device.EnrolledAt, nil +} + +func (s *Service) authenticateDevice(ctx context.Context, auth RequestAuth, body []byte) (verifiedDevice, error) { if s == nil || s.DB == nil || s.DB.SQL == nil { - return "", errors.New("relay storage is required") + return verifiedDevice{}, errors.New("relay storage is required") } nowFn := s.Now if nowFn == nil { nowFn = time.Now } - device, err := verifyRequest(ctx, s.DB, auth, body, nowFn().UTC()) - if err != nil { - return "", err - } - return device.ID, nil + return verifyRequest(ctx, s.DB, auth, body, nowFn().UTC()) } func parsePublicKey(value string) (*ecdsa.PublicKey, []byte, error) { diff --git a/internal/v4/relay/service.go b/internal/v4/relay/service.go index 77f17d1..982b191 100644 --- a/internal/v4/relay/service.go +++ b/internal/v4/relay/service.go @@ -1,10 +1,13 @@ package relay import ( + "bytes" "context" "crypto/rand" + "crypto/sha256" "database/sql" "encoding/base32" + "encoding/binary" "encoding/hex" "encoding/json" "errors" @@ -57,6 +60,88 @@ func NewService(db *storage.DB, paymentService *payments.Service) *Service { NewPairingToken: randomPairingToken, PairingTTL: 2 * time.Minute, } } + +// RedactRawEvents removes retained notification bodies while keeping the +// immutable event identity and normalized observation available for audit. +func (s *Service) RedactRawEvents(ctx context.Context, before time.Time, limit int) (int64, error) { + if s == nil || s.DB == nil || s.DB.SQL == nil { + return 0, errors.New("relay storage is required") + } + if before.IsZero() { + return 0, errors.New("raw event retention boundary is required") + } + if limit <= 0 || limit > 5000 { + limit = 500 + } + + rows, err := s.DB.SQL.QueryContext(ctx, `SELECT id,package_name,posted_at,amount_hint_paise, + title,text,big_text,payload_hash + FROM relay_events + WHERE received_at < ? + AND (title IS NOT NULL OR text IS NOT NULL OR big_text IS NOT NULL) + ORDER BY received_at,id LIMIT ?`, before.UTC().UnixMilli(), limit) + if err != nil { + return 0, fmt.Errorf("select relay event bodies for redaction: %w", err) + } + type retainedEvent struct { + id, packageName string + postedAt int64 + amountHint sql.NullInt64 + title, text, bigText sql.NullString + payloadHash []byte + } + events := make([]retainedEvent, 0, limit) + for rows.Next() { + var event retainedEvent + if err := rows.Scan(&event.id, &event.packageName, &event.postedAt, &event.amountHint, + &event.title, &event.text, &event.bigText, &event.payloadHash); err != nil { + rows.Close() + return 0, fmt.Errorf("scan relay event body for redaction: %w", err) + } + events = append(events, event) + } + if err := rows.Err(); err != nil { + rows.Close() + return 0, fmt.Errorf("iterate relay event bodies for redaction: %w", err) + } + if err := rows.Close(); err != nil { + return 0, fmt.Errorf("close relay event bodies for redaction: %w", err) + } + + var redacted int64 + err = s.DB.WithImmediateTx(ctx, func(tx *storage.ImmediateTx) error { + for _, event := range events { + payloadHash := event.payloadHash + if len(payloadHash) == 0 { + // Legacy rows predate payload hashing. Preserve a deterministic + // retained-field fingerprint before clearing their raw bodies. + payloadHash = legacyPayloadFingerprint(event.packageName, event.postedAt, + event.amountHint, event.title, event.text, event.bigText) + } + result, err := tx.ExecContext(ctx, `UPDATE relay_events + SET title=NULL,text=NULL,big_text=NULL,payload_hash=?, + status=CASE WHEN status='received' THEN 'ignored' ELSE status END, + error=CASE WHEN status='received' THEN 'raw notification expired before processing' ELSE error END + WHERE id=? AND received_at < ? + AND (title IS NOT NULL OR text IS NOT NULL OR big_text IS NOT NULL)`, + payloadHash, event.id, before.UTC().UnixMilli()) + if err != nil { + return fmt.Errorf("redact relay event %s: %w", event.id, err) + } + count, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("count redacted relay event %s: %w", event.id, err) + } + redacted += count + } + return nil + }) + if err != nil { + return 0, err + } + return redacted, nil +} + func (s *Service) IngestSigned(ctx context.Context, auth RequestAuth, rawBody []byte) (IngestResult, error) { if s == nil || s.DB == nil || s.DB.SQL == nil || s.Payments == nil { return IngestResult{}, errors.New("relay storage and payment service are required") @@ -83,13 +168,18 @@ func (s *Service) IngestSigned(ctx context.Context, auth RequestAuth, rawBody [] if err := validateEventInput(&input); err != nil { return IngestResult{}, err } + payloadHash := sha256.Sum256(rawBody) postedAt, postedReliable := sanitizePostedAt(input.PostedAtMS, now) - result, inserted, err := s.acceptEvent(ctx, device, input, postedAt, postedReliable, now) + result, inserted, err := s.acceptEvent(ctx, device, input, postedAt, postedReliable, now, payloadHash[:]) if err != nil || !inserted { return result, err } if !postedReliable { - postedAt = now + if err := s.finishIgnored(ctx, result.RelayEventID, errors.New("notification posting time is missing or invalid")); err != nil { + return IngestResult{}, err + } + result.Status = "ignored" + return result, nil } obs, parseErr := observations.Parse(observations.Snapshot{ PackageName: input.PackageName, @@ -99,16 +189,20 @@ func (s *Service) IngestSigned(ctx context.Context, auth RequestAuth, rawBody [] BigText: input.BigText, }) if parseErr != nil { + if errors.Is(parseErr, observations.ErrAmbiguousAmount) { + status, err := s.finishAmbiguous(ctx, result.RelayEventID, parseErr) + if err != nil { + return IngestResult{}, err + } + result.Status = status + return result, nil + } if err := s.finishIgnored(ctx, result.RelayEventID, parseErr); err != nil { return IngestResult{}, err } result.Status = "ignored" return result, nil } - if !postedReliable && obs.OccurredAtSource == "notification_posted_at" { - obs.OccurredAt = now - obs.OccurredAtSource = "server_received_at" - } if obs.OccurredAt.After(now.Add(2 * time.Minute)) { if err := s.finishIgnored(ctx, result.RelayEventID, errors.New("payment occurrence time is implausibly in the future")); err != nil { return IngestResult{}, err @@ -116,32 +210,31 @@ func (s *Service) IngestSigned(ctx context.Context, auth RequestAuth, rawBody [] result.Status = "ignored" return result, nil } - if !device.EnrolledAt.IsZero() && obs.OccurredAt.Before(device.EnrolledAt.Add(-2*time.Minute)) { + if !device.EnrolledAt.IsZero() && obs.OccurredAt.Before(device.EnrolledAt) { if err := s.finishIgnored(ctx, result.RelayEventID, errors.New("notification predates relay enrollment")); err != nil { return IngestResult{}, err } result.Status = "ignored" return result, nil } - if obs.CollectionProfileID == "" { - profileID, ambiguous, err := s.resolveGenericCollectionProfileID(ctx, obs) - if err != nil { - if err := s.finishIgnored(ctx, result.RelayEventID, err); err != nil { - return IngestResult{}, err - } - result.Status = "ignored" - return result, nil + matched, err := s.Payments.ApplyObservationForRelay(ctx, result.RelayEventID, obs, now, device.ID, device.EnrolledAt) + if errors.Is(err, payments.ErrRelayEventNotFound) { + // A retention worker or device revocation may have finalized this + // stale event while parsing was in flight. Do not apply its payload. + if finishErr := s.finishIgnored(ctx, result.RelayEventID, errors.New("relay event was no longer authorized for processing")); finishErr != nil { + return IngestResult{}, finishErr } - if ambiguous { - if err := s.finishAmbiguous(ctx, result.RelayEventID, errors.New("generic notification matches reservations in multiple collection profiles")); err != nil { - return IngestResult{}, err - } - result.Status = "ambiguous" - return result, nil + result.Status = "ignored" + return result, nil + } + if errors.Is(err, payments.ErrObservationAmbiguous) { + status, finishErr := s.finishAmbiguous(ctx, result.RelayEventID, err) + if finishErr != nil { + return IngestResult{}, finishErr } - obs.CollectionProfileID = profileID + result.Status = status + return result, nil } - matched, err := s.Payments.ApplyObservation(ctx, result.RelayEventID, obs, now) if err != nil { return IngestResult{}, err } @@ -151,53 +244,6 @@ func (s *Service) IngestSigned(ctx context.Context, auth RequestAuth, rawBody [] return result, nil } -func (s *Service) resolveGenericCollectionProfileID(ctx context.Context, obs observations.Observation) (string, bool, error) { - occurred := obs.OccurredAt.UnixMilli() - rows, err := s.DB.SQL.QueryContext(ctx, `SELECT DISTINCT r.collection_profile_id - FROM amount_reservations r JOIN payments p ON p.id=r.payment_id - WHERE r.payable_amount_paise=? AND p.created_at<=? AND r.reserved_until>=? - AND (p.status<>'cancelled' OR EXISTS( - SELECT 1 FROM payment_history h - WHERE h.payment_id=p.id AND h.type='payment.cancelled' AND h.created_at>=? - )) - ORDER BY r.collection_profile_id`, obs.AmountPaise, occurred, occurred, occurred) - if err != nil { - return "", false, fmt.Errorf("resolve generic notification profile: %w", err) - } - defer rows.Close() - profiles := make([]string, 0, 2) - for rows.Next() { - var profileID string - if err := rows.Scan(&profileID); err != nil { - return "", false, fmt.Errorf("scan generic notification profile: %w", err) - } - profiles = append(profiles, profileID) - } - if err := rows.Err(); err != nil { - return "", false, fmt.Errorf("iterate generic notification profiles: %w", err) - } - if len(profiles) == 1 { - return profiles[0], false, nil - } - if len(profiles) > 1 { - return "", true, nil - } - profileID, err := s.activeCollectionProfileID(ctx) - return profileID, false, err -} - -func (s *Service) activeCollectionProfileID(ctx context.Context) (string, error) { - var profileID string - err := s.DB.SQL.QueryRowContext(ctx, `SELECT id FROM collection_profiles WHERE active=1 AND enabled=1 LIMIT 1`).Scan(&profileID) - if errors.Is(err, sql.ErrNoRows) { - return "", errors.New("no active collection profile is available for generic notification evidence") - } - if err != nil { - return "", fmt.Errorf("read active collection profile: %w", err) - } - return profileID, nil -} - func validateEventInput(in *EventInput) error { if in.SchemaVersion != SchemaVersion { return relayError("UNSUPPORTED_RELAY_SCHEMA", "schema_version must be 1", 400) @@ -213,6 +259,9 @@ func validateEventInput(in *EventInput) error { if in.PackageName == "" || len(in.PackageName) > 255 { return relayError("INVALID_RELAY_APP", "notification package name is required and must be at most 255 characters", 400) } + if blockedRelayPackage(in.PackageName) { + return relayError("UNSUPPORTED_RELAY_APP", "this notification source is no longer accepted", 400) + } if len(in.Title)+len(in.Text)+len(in.BigText) == 0 || len(in.Title)+len(in.Text)+len(in.BigText) > maxNotificationTextBytes { return relayError("RELAY_EVENT_TOO_LARGE", "notification text is empty or too large", 400) } @@ -222,6 +271,15 @@ func validateEventInput(in *EventInput) error { return nil } +func blockedRelayPackage(packageName string) bool { + switch strings.ToLower(strings.TrimSpace(packageName)) { + case observations.GoogleMessagesPackage, observations.GmailPackage: + return true + default: + return false + } +} + func sanitizePostedAt(ms int64, now time.Time) (time.Time, bool) { if ms <= 0 { return now.UTC(), false @@ -232,7 +290,7 @@ func sanitizePostedAt(ms int64, now time.Time) (time.Time, bool) { } return posted, true } -func (s *Service) acceptEvent(ctx context.Context, device verifiedDevice, in EventInput, postedAt time.Time, postedReliable bool, now time.Time) (IngestResult, bool, error) { +func (s *Service) acceptEvent(ctx context.Context, device verifiedDevice, in EventInput, postedAt time.Time, postedReliable bool, now time.Time, payloadHash []byte) (IngestResult, bool, error) { idFn := s.NewID if idFn == nil { idFn = randomID @@ -240,18 +298,21 @@ func (s *Service) acceptEvent(ctx context.Context, device verifiedDevice, in Eve var result IngestResult needsProcessing := false err := s.DB.WithImmediateTx(ctx, func(tx *storage.ImmediateTx) error { - updated, err := tx.ExecContext(ctx, `UPDATE relay_devices SET last_seen_at=? WHERE id=? AND enabled=1`, now.UnixMilli(), device.ID) + updated, err := tx.ExecContext(ctx, `UPDATE relay_devices SET last_seen_at=? WHERE id=? AND enabled=1 AND enrolled_at=?`, now.UnixMilli(), device.ID, device.EnrolledAt.UnixMilli()) if err != nil { return fmt.Errorf("refresh relay device: %w", err) } if rows, _ := updated.RowsAffected(); rows != 1 { return relayError("UNKNOWN_RELAY_DEVICE", "relay device is not enrolled or is disabled", 401) } - existing, found, err := existingRelayEvent(ctx, tx, device.ID, in.EventID) + existing, found, same, err := existingRelayEvent(ctx, tx, device.ID, in.EventID, in, postedAt, postedReliable, payloadHash) if err != nil { return err } if found { + if !same { + return relayError("RELAY_EVENT_ID_CONFLICT", "source event ID was already used with different notification data", 409) + } result = existing result.Duplicate = true needsProcessing = existing.Status == "received" @@ -263,15 +324,15 @@ func (s *Service) acceptEvent(ctx context.Context, device verifiedDevice, in Eve } status := "received" errorText := any(nil) - if postedReliable && !device.EnrolledAt.IsZero() && postedAt.Before(device.EnrolledAt.Add(-2*time.Minute)) { + if postedReliable && !device.EnrolledAt.IsZero() && postedAt.Before(device.EnrolledAt) { status = "ignored" errorText = "notification predates relay enrollment" } _, err = tx.ExecContext(ctx, `INSERT INTO relay_events( - id,device_id,source_event_id,package_name,posted_at,received_at,amount_hint_paise,title,text,big_text,status,error) - VALUES(?,?,?,?,?,?,?,?,?,?,?,?)`, relayID, device.ID, in.EventID, in.PackageName, + id,device_id,source_event_id,package_name,posted_at,received_at,amount_hint_paise,title,text,big_text,payload_hash,status,error) + VALUES(?,?,?,?,?,?,?,?,?,?,?,?,?)`, relayID, device.ID, in.EventID, in.PackageName, postedAt.UnixMilli(), now.UnixMilli(), nullableAmount(in.AmountHintPaise), nullableText(in.Title), - nullableText(in.Text), nullableText(in.BigText), status, errorText) + nullableText(in.Text), nullableText(in.BigText), payloadHash, status, errorText) if err != nil { return fmt.Errorf("insert relay event: %w", err) } @@ -281,32 +342,120 @@ func (s *Service) acceptEvent(ctx context.Context, device verifiedDevice, in Eve }) return result, needsProcessing, err } -func existingRelayEvent(ctx context.Context, tx *storage.ImmediateTx, deviceID, sourceEventID string) (IngestResult, bool, error) { +func existingRelayEvent(ctx context.Context, tx *storage.ImmediateTx, deviceID, sourceEventID string, in EventInput, postedAt time.Time, postedReliable bool, payloadHash []byte) (IngestResult, bool, bool, error) { var result IngestResult var paymentID sql.NullString - err := tx.QueryRowContext(ctx, `SELECT r.id,COALESCE(o.match_result,r.status),o.matched_payment_id + var packageName string + var storedPostedAt int64 + var amountHint sql.NullInt64 + var title, text, bigText sql.NullString + var storedHash []byte + err := tx.QueryRowContext(ctx, `SELECT r.id,COALESCE(o.match_result,r.status),o.matched_payment_id, + r.package_name,r.posted_at,r.amount_hint_paise,r.title,r.text,r.big_text,r.payload_hash FROM relay_events r LEFT JOIN payment_observations o ON o.relay_event_id=r.id WHERE r.device_id=? AND r.source_event_id=?`, deviceID, sourceEventID). - Scan(&result.RelayEventID, &result.Status, &paymentID) + Scan(&result.RelayEventID, &result.Status, &paymentID, &packageName, &storedPostedAt, &amountHint, &title, &text, &bigText, &storedHash) if errors.Is(err, sql.ErrNoRows) { - return IngestResult{}, false, nil + return IngestResult{}, false, false, nil } if err != nil { - return IngestResult{}, false, fmt.Errorf("read relay event: %w", err) + return IngestResult{}, false, false, fmt.Errorf("read relay event: %w", err) } result.PaymentID = paymentID.String - return result, true, nil + same := false + if len(storedHash) == sha256.Size { + same = bytes.Equal(storedHash, payloadHash) + } else if bytes.HasPrefix(storedHash, []byte("legacy:")) { + comparisonPostedAt := postedAt + if !postedReliable { + comparisonPostedAt = time.UnixMilli(storedPostedAt).UTC() + } + same = bytes.Equal(storedHash, inputPayloadFingerprint(in, comparisonPostedAt)) + } else if len(storedHash) == 0 { + // Legacy rows without a hash retain their fields until this service + // redacts them, so retries remain idempotent during migration. + same = packageName == in.PackageName && + (!postedReliable || storedPostedAt == postedAt.UnixMilli()) && + nullableAmountEqual(amountHint, in.AmountHintPaise) && + nullableTextEqual(title, in.Title) && + nullableTextEqual(text, in.Text) && + nullableTextEqual(bigText, in.BigText) + } + return result, true, same, nil +} +func legacyPayloadFingerprint(packageName string, postedAt int64, amountHint sql.NullInt64, + title, text, bigText sql.NullString) []byte { + data := make([]byte, 0, 128) + data = appendFingerprintField(data, packageName, true) + data = appendFingerprintNumber(data, postedAt, true) + data = appendFingerprintNumber(data, amountHint.Int64, amountHint.Valid) + data = appendFingerprintField(data, title.String, title.Valid) + data = appendFingerprintField(data, text.String, text.Valid) + data = appendFingerprintField(data, bigText.String, bigText.Valid) + digest := sha256.Sum256(data) + fingerprint := make([]byte, len("legacy:")+len(digest)) + copy(fingerprint, "legacy:") + copy(fingerprint[len("legacy:"):], digest[:]) + return fingerprint +} + +func inputPayloadFingerprint(in EventInput, postedAt time.Time) []byte { + return legacyPayloadFingerprint(in.PackageName, postedAt.UnixMilli(), + sql.NullInt64{Int64: in.AmountHintPaise, Valid: in.AmountHintPaise != 0}, + fingerprintText(in.Title), fingerprintText(in.Text), fingerprintText(in.BigText)) +} + +func fingerprintText(value string) sql.NullString { + value = strings.TrimSpace(value) + return sql.NullString{String: value, Valid: value != ""} +} + +func appendFingerprintField(dst []byte, value string, valid bool) []byte { + if !valid { + return append(dst, 0) + } + dst = append(dst, 1) + var length [8]byte + binary.BigEndian.PutUint64(length[:], uint64(len(value))) + dst = append(dst, length[:]...) + return append(dst, value...) +} + +func appendFingerprintNumber(dst []byte, value int64, valid bool) []byte { + if !valid { + return append(dst, 0) + } + dst = append(dst, 1) + var number [8]byte + binary.BigEndian.PutUint64(number[:], uint64(value)) + return append(dst, number[:]...) +} + +func nullableAmountEqual(stored sql.NullInt64, wanted int64) bool { + return stored.Valid == (wanted != 0) && (!stored.Valid || stored.Int64 == wanted) +} + +func nullableTextEqual(stored sql.NullString, wanted string) bool { + wanted = strings.TrimSpace(wanted) + return stored.Valid == (wanted != "") && (!stored.Valid || stored.String == wanted) } -func (s *Service) finishAmbiguous(ctx context.Context, relayEventID string, reason error) error { +func (s *Service) finishAmbiguous(ctx context.Context, relayEventID string, reason error) (string, error) { message := "ambiguous generic notification" if reason != nil { message = trimError(reason.Error(), 512) } - return s.DB.WithImmediateTx(ctx, func(tx *storage.ImmediateTx) error { - _, err := tx.ExecContext(ctx, `UPDATE relay_events SET status='ambiguous',error=? WHERE id=? AND status='received'`, message, relayEventID) - return err + var status string + err := s.DB.WithImmediateTx(ctx, func(tx *storage.ImmediateTx) error { + if _, err := tx.ExecContext(ctx, `UPDATE relay_events SET status='ambiguous',error=? WHERE id=? AND status='received'`, message, relayEventID); err != nil { + return err + } + return tx.QueryRowContext(ctx, `SELECT status FROM relay_events WHERE id=?`, relayEventID).Scan(&status) }) + if err != nil { + return "", err + } + return status, nil } func (s *Service) finishIgnored(ctx context.Context, relayEventID string, reason error) error { diff --git a/internal/v4/relay/service_test.go b/internal/v4/relay/service_test.go index 0442515..e7ff3d2 100644 --- a/internal/v4/relay/service_test.go +++ b/internal/v4/relay/service_test.go @@ -1,16 +1,19 @@ package relay import ( + "bytes" "context" "crypto/ecdsa" "crypto/elliptic" "crypto/rand" "crypto/sha256" "crypto/x509" + "database/sql" "encoding/base64" "encoding/hex" "encoding/json" "encoding/pem" + "errors" "net/http" "path/filepath" "strconv" @@ -81,6 +84,23 @@ func signedAuth(t *testing.T, priv *ecdsa.PrivateKey, deviceID string, now time. } } +func signedAuthWithEpoch(t *testing.T, priv *ecdsa.PrivateKey, deviceID string, now time.Time, body []byte, enrollmentEpoch int64) RequestAuth { + t.Helper() + timestamp := strconv.FormatInt(now.UnixMilli(), 10) + canonical := CanonicalRequestWithEpoch(http.MethodPost, EventPath, timestamp, enrollmentEpoch, body) + digest := sha256.Sum256([]byte(canonical)) + sig, err := ecdsa.SignASN1(rand.Reader, priv, digest[:]) + if err != nil { + t.Fatal(err) + } + return RequestAuth{ + DeviceID: deviceID, Timestamp: timestamp, + Signature: base64.StdEncoding.EncodeToString(sig), + EnrollmentEpoch: strconv.FormatInt(enrollmentEpoch, 10), + Method: http.MethodPost, Path: EventPath, + } +} + func marshalEvent(t *testing.T, input EventInput) []byte { t.Helper() body, err := json.Marshal(input) @@ -116,6 +136,121 @@ func countRows(t *testing.T, db *storage.DB, table string) int { } return count } +func TestEpochBoundAuthenticationRejectsMalformedAndWrongEpoch(t *testing.T) { + ctx := context.Background() + db := openRelayDB(t) + now := time.Date(2026, 9, 1, 3, 30, 0, 0, time.UTC) + priv, deviceID := enrollTestDevice(t, db, now.Add(-time.Minute)) + service := NewService(db, payments.NewService(db)) + service.Now = func() time.Time { return now } + body := []byte(`{"schema_version":1}`) + epoch := now.Add(-time.Minute).UnixMilli() + + if got, _, err := service.AuthenticateDeviceWithEpoch(ctx, signedAuthWithEpoch(t, priv, deviceID, now, body, epoch), body); err != nil || got != deviceID { + t.Fatalf("current epoch authentication device=%q err=%v", got, err) + } + + malformed := signedAuthWithEpoch(t, priv, deviceID, now, body, epoch) + malformed.EnrollmentEpoch = "not-an-integer" + if _, err := service.AuthenticateDevice(ctx, malformed, body); err == nil { + t.Fatal("malformed epoch was accepted") + } + + wrong := signedAuthWithEpoch(t, priv, deviceID, now, body, epoch+1) + if _, err := service.AuthenticateDevice(ctx, wrong, body); err == nil { + t.Fatal("wrong epoch was accepted") + } + + if _, err := service.AuthenticateDevice(ctx, signedAuth(t, priv, deviceID, now, body), body); err != nil { + t.Fatalf("legacy authentication rejected: %v", err) + } +} + +func TestLegacySignatureStopsWorkingAfterSameKeyRepair(t *testing.T) { + ctx := context.Background() + db := openRelayDB(t) + requestTime := time.Date(2026, 9, 8, 6, 0, 0, 0, time.UTC) + priv, deviceID := enrollTestDevice(t, db, requestTime.Add(-time.Hour)) + service := NewService(db, payments.NewService(db)) + service.Now = func() time.Time { return requestTime } + body := []byte(`{"schema_version":1}`) + legacyAuth := signedAuth(t, priv, deviceID, requestTime, body) + if got, err := service.AuthenticateDevice(ctx, legacyAuth, body); err != nil || got != deviceID { + t.Fatalf("legacy device authentication device=%q err=%v", got, err) + } + + der, err := x509.MarshalPKIXPublicKey(&priv.PublicKey) + if err != nil { + t.Fatal(err) + } + publicKeyPEM := string(pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: der})) + pairing, err := service.CreatePairing(ctx) + if err != nil { + t.Fatal(err) + } + repaired, err := service.PairDevice(ctx, PairDeviceInput{ + Token: pairing.Token, Name: "Test Phone", PublicKeyPEM: publicKeyPEM, AppVersion: "test-v2", + }) + if err != nil { + t.Fatal(err) + } + var epochRequired int + if err := db.SQL.QueryRow(`SELECT epoch_required FROM relay_devices WHERE id=?`, deviceID).Scan(&epochRequired); err != nil { + t.Fatal(err) + } + if repaired.DeviceID != deviceID || repaired.EnrolledAtMS <= requestTime.Add(-time.Hour).UnixMilli() || epochRequired != 1 { + t.Fatalf("re-pair device=%q epoch=%d epoch_required=%d", repaired.DeviceID, repaired.EnrolledAtMS, epochRequired) + } + + if _, err := service.AuthenticateDevice(ctx, legacyAuth, body); err == nil { + t.Fatal("legacy signature captured before re-pair was accepted after re-pair") + } + fresh := signedAuthWithEpoch(t, priv, deviceID, requestTime, body, repaired.EnrolledAtMS) + if got, err := service.AuthenticateDevice(ctx, fresh, body); err != nil || got != deviceID { + t.Fatalf("fresh epoch authentication device=%q err=%v", got, err) + } +} + +func TestEpochSignatureIsRejectedAfterSameMillisecondRepair(t *testing.T) { + ctx := context.Background() + db := openRelayDB(t) + now := time.Date(2026, 9, 8, 7, 0, 0, 0, time.UTC) + priv, deviceID := enrollTestDevice(t, db, now.Add(-time.Hour)) + service := NewService(db, payments.NewService(db)) + service.Now = func() time.Time { return now } + der, err := x509.MarshalPKIXPublicKey(&priv.PublicKey) + if err != nil { + t.Fatal(err) + } + publicKeyPEM := string(pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: der})) + pair := func() PairDeviceResult { + session, err := service.CreatePairing(ctx) + if err != nil { + t.Fatal(err) + } + result, err := service.PairDevice(ctx, PairDeviceInput{Token: session.Token, Name: "Test Phone", PublicKeyPEM: publicKeyPEM}) + if err != nil { + t.Fatal(err) + } + return result + } + + first := pair() + body := []byte(`{"schema_version":1}`) + captured := signedAuthWithEpoch(t, priv, deviceID, now, body, first.EnrolledAtMS) + second := pair() // service.Now is unchanged: same wall-clock millisecond. + if second.EnrolledAtMS != first.EnrolledAtMS+1 { + t.Fatalf("same-millisecond re-pair epoch first=%d second=%d", first.EnrolledAtMS, second.EnrolledAtMS) + } + if _, err := service.AuthenticateDevice(ctx, captured, body); err == nil { + t.Fatal("captured epoch-bound signature survived same-millisecond re-pair") + } + fresh := signedAuthWithEpoch(t, priv, deviceID, now, body, second.EnrolledAtMS) + if got, err := service.AuthenticateDevice(ctx, fresh, body); err != nil || got != deviceID { + t.Fatalf("fresh epoch authentication device=%q err=%v", got, err) + } +} + func TestSignedPaytmEventMatchesPaymentAndIsIdempotent(t *testing.T) { ctx := context.Background() db := openRelayDB(t) @@ -170,6 +305,206 @@ func TestSignedPaytmEventMatchesPaymentAndIsIdempotent(t *testing.T) { t.Fatal("duplicate relay event created duplicate state") } } +func TestDuplicateRelayEventIDRejectsChangedPayload(t *testing.T) { + ctx := context.Background() + db := openRelayDB(t) + createdAt := time.Date(2026, 9, 1, 3, 30, 0, 0, time.UTC) + insertProfile(t, db, "paytm", "paytm_notification", "merchant@paytm", true, createdAt) + paymentService, created := createPayment(t, db, createdAt, "duplicate-payload") + priv, deviceID := enrollTestDevice(t, db, createdAt.Add(-time.Minute)) + relayService := NewService(db, paymentService) + receivedAt := createdAt.Add(2 * time.Minute) + relayService.Now = func() time.Time { return receivedAt } + eventID := strings.Repeat("d", 64) + first := marshalEvent(t, EventInput{ + SchemaVersion: 1, EventID: eventID, PackageName: observations.PaytmBusinessPackage, + PostedAtMS: receivedAt.UnixMilli(), Title: "Payment Received on Paytm", + Text: "₹100.37 Received from Rahul", AmountHintPaise: 10037, + }) + if _, err := relayService.IngestSigned(ctx, signedAuth(t, priv, deviceID, receivedAt, first), first); err != nil { + t.Fatal(err) + } + changed := marshalEvent(t, EventInput{ + SchemaVersion: 1, EventID: eventID, PackageName: observations.PaytmBusinessPackage, + PostedAtMS: receivedAt.UnixMilli(), Title: "Payment Received on Paytm", + Text: "₹100.37 Received from Mallory", AmountHintPaise: 10037, + }) + _, err := relayService.IngestSigned(ctx, signedAuth(t, priv, deviceID, receivedAt, changed), changed) + var relayErr *Error + if !errors.As(err, &relayErr) || relayErr.Code != "RELAY_EVENT_ID_CONFLICT" || relayErr.HTTPStatus != http.StatusConflict { + t.Fatalf("changed duplicate error=%v, want relay event conflict", err) + } + if countRows(t, db, "relay_events") != 1 || countRows(t, db, "payment_observations") != 1 { + t.Fatal("changed duplicate mutated stored event state") + } + if created.Payment.ID == "" { + t.Fatal("test payment was not created") + } +} + +func TestAmbiguousRelayNotificationIsVisibleInActivity(t *testing.T) { + ctx := context.Background() + db := openRelayDB(t) + now := time.Date(2026, 9, 1, 4, 0, 0, 0, time.UTC) + priv, deviceID := enrollTestDevice(t, db, now.Add(-time.Minute)) + service := NewService(db, payments.NewService(db)) + service.Now = func() time.Time { return now } + body := marshalEvent(t, EventInput{ + SchemaVersion: 1, EventID: strings.Repeat("b", 64), + PackageName: observations.GooglePayPackage, PostedAtMS: now.UnixMilli(), + Title: "Payment received", Text: "₹100.37 received; balance ₹200.00", + AmountHintPaise: 10037, + }) + + result, err := service.IngestSigned(ctx, signedAuth(t, priv, deviceID, now, body), body) + if err != nil { + t.Fatal(err) + } + if result.Status != "ambiguous" || result.Transitioned || result.PaymentID != "" { + t.Fatalf("ambiguous relay result=%+v", result) + } + var status, reason string + if err := db.SQL.QueryRow(`SELECT status,error FROM relay_events WHERE id=?`, result.RelayEventID).Scan(&status, &reason); err != nil { + t.Fatal(err) + } + if status != "ambiguous" || !strings.Contains(reason, "multiple monetary amounts") { + t.Fatalf("relay event status=%s reason=%q", status, reason) + } + if countRows(t, db, "payment_observations") != 0 { + t.Fatal("ambiguous notification must not create a payment observation") + } +} + +func TestAmbiguousFinalizerReportsRetentionWinner(t *testing.T) { + ctx := context.Background() + db := openRelayDB(t) + now := time.Date(2026, 9, 1, 4, 30, 0, 0, time.UTC) + if _, err := db.SQL.Exec(`INSERT INTO relay_devices(id,public_key_pem,enabled,enrolled_at) VALUES('race-device','pem',1,?)`, now.UnixMilli()); err != nil { + t.Fatal(err) + } + if _, err := db.SQL.Exec(`INSERT INTO relay_events(id,device_id,source_event_id,package_name,posted_at,received_at,status) + VALUES('race-event','race-device','race-source','com.example.wallet',?,?,?)`, + now.UnixMilli(), now.UnixMilli(), "ignored"); err != nil { + t.Fatal(err) + } + service := NewService(db, payments.NewService(db)) + status, err := service.finishAmbiguous(ctx, "race-event", errors.New("ambiguous parser result")) + if err != nil || status != "ignored" { + t.Fatalf("finalizer status=%q err=%v", status, err) + } + var stored string + if err := db.SQL.QueryRow(`SELECT status FROM relay_events WHERE id='race-event'`).Scan(&stored); err != nil { + t.Fatal(err) + } + if stored != "ignored" { + t.Fatalf("retention winner was overwritten: %q", stored) + } +} + +func TestRedactRawEventsPreservesIdentityAndMarksStuckReceived(t *testing.T) { + ctx := context.Background() + db := openRelayDB(t) + now := time.Date(2026, 9, 10, 3, 0, 0, 0, time.UTC) + _, err := db.SQL.Exec(`INSERT INTO relay_devices(id,public_key_pem,enabled,enrolled_at) VALUES('retention-device','pem',1,?)`, now.UnixMilli()) + if err != nil { + t.Fatal(err) + } + old := now.Add(-15 * 24 * time.Hour).UnixMilli() + recent := now.Add(-24 * time.Hour).UnixMilli() + insert := `INSERT INTO relay_events(id,device_id,source_event_id,package_name,posted_at,received_at,title,text,big_text,payload_hash,status) + VALUES(?,?,?,?,?,?,?,?,?,?,?)` + if _, err := db.SQL.Exec(insert, "old-event", "retention-device", "old-source", "com.example.paytm", old, old, + "old title", "old body", "old details", bytes.Repeat([]byte{1}, sha256.Size), "matched"); err != nil { + t.Fatal(err) + } + if _, err := db.SQL.Exec(insert, "pending-event", "retention-device", "pending-source", "com.example.paytm", old, old, + "pending title", "pending body", "pending details", bytes.Repeat([]byte{2}, sha256.Size), "received"); err != nil { + t.Fatal(err) + } + if _, err := db.SQL.Exec(insert, "legacy-event", "retention-device", "legacy-source", "com.example.paytm", old, old, + "legacy title", "legacy body", "legacy details", nil, "matched"); err != nil { + t.Fatal(err) + } + if _, err := db.SQL.Exec(insert, "recent-event", "retention-device", "recent-source", "com.example.paytm", recent, recent, + "recent title", "recent body", "recent details", bytes.Repeat([]byte{3}, sha256.Size), "matched"); err != nil { + t.Fatal(err) + } + service := NewService(db, payments.NewService(db)) + redacted, err := service.RedactRawEvents(ctx, now.Add(-14*24*time.Hour), 10) + if err != nil || redacted != 3 { + t.Fatalf("redacted=%d err=%v", redacted, err) + } + var oldTitle, oldText, pendingTitle, recentTitle sql.NullString + var oldStatus, pendingStatus, pendingError string + var payloadHash []byte + if err := db.SQL.QueryRow(`SELECT title,text,payload_hash,status FROM relay_events WHERE id='old-event'`).Scan(&oldTitle, &oldText, &payloadHash, &oldStatus); err != nil { + t.Fatal(err) + } + if oldTitle.Valid || oldText.Valid || len(payloadHash) != sha256.Size || oldStatus != "matched" { + t.Fatalf("redacted event title=%v text=%v payload_hash=%d status=%s", oldTitle, oldText, len(payloadHash), oldStatus) + } + var legacyHash []byte + if err := db.SQL.QueryRow(`SELECT title,text,payload_hash FROM relay_events WHERE id='legacy-event'`).Scan(&oldTitle, &oldText, &legacyHash); err != nil { + t.Fatal(err) + } + if oldTitle.Valid || oldText.Valid || !bytes.HasPrefix(legacyHash, []byte("legacy:")) { + t.Fatalf("legacy redaction title=%v text=%v payload_hash=%q", oldTitle, oldText, legacyHash) + } + legacyInput := EventInput{ + PackageName: "com.example.paytm", PostedAtMS: old, + Title: "legacy title", Text: "legacy body", BigText: "legacy details", + } + var found, same bool + if err := db.WithImmediateTx(ctx, func(tx *storage.ImmediateTx) error { + var err error + _, found, same, err = existingRelayEvent(ctx, tx, "retention-device", "legacy-source", + legacyInput, time.UnixMilli(old), true, bytes.Repeat([]byte{9}, sha256.Size)) + return err + }); err != nil { + t.Fatal(err) + } + if !found || !same { + t.Fatalf("redacted legacy event did not remain idempotent: found=%v same=%v", found, same) + } + legacyInput.PostedAtMS = old + 1000 + if err := db.WithImmediateTx(ctx, func(tx *storage.ImmediateTx) error { + var err error + _, found, same, err = existingRelayEvent(ctx, tx, "retention-device", "legacy-source", + legacyInput, time.UnixMilli(old+1000), true, bytes.Repeat([]byte{9}, sha256.Size)) + return err + }); err != nil { + t.Fatal(err) + } + if !found || same { + t.Fatalf("redacted legacy event accepted changed posted_at: found=%v same=%v", found, same) + } + + legacyInput.PostedAtMS = 0 + if err := db.WithImmediateTx(ctx, func(tx *storage.ImmediateTx) error { + var err error + _, found, same, err = existingRelayEvent(ctx, tx, "retention-device", "legacy-source", + legacyInput, now, false, bytes.Repeat([]byte{9}, sha256.Size)) + return err + }); err != nil { + t.Fatal(err) + } + if !found || !same { + t.Fatalf("redacted legacy event with unreliable posted_at lost idempotency: found=%v same=%v", found, same) + } + + if err := db.SQL.QueryRow(`SELECT title,status,error FROM relay_events WHERE id='pending-event'`).Scan(&pendingTitle, &pendingStatus, &pendingError); err != nil { + t.Fatal(err) + } + if pendingTitle.Valid || pendingStatus != "ignored" || pendingError == "" { + t.Fatalf("stuck received event title=%v status=%s error=%q", pendingTitle, pendingStatus, pendingError) + } + if err := db.SQL.QueryRow(`SELECT title FROM relay_events WHERE id='recent-event'`).Scan(&recentTitle); err != nil { + t.Fatal(err) + } + if !recentTitle.Valid { + t.Fatal("recent event body was redacted") + } +} func TestPaytmMinuteTimestampDoesNotPreDateMidMinutePayment(t *testing.T) { ctx := context.Background() @@ -335,7 +670,7 @@ func TestPreEnrollmentNotificationIsStoredButIgnored(t *testing.T) { body := marshalEvent(t, EventInput{ SchemaVersion: 1, EventID: strings.Repeat("e", 64), PackageName: observations.PaytmBusinessPackage, - PostedAtMS: now.Add(-10 * time.Minute).UnixMilli(), + PostedAtMS: now.Add(-90 * time.Second).UnixMilli(), Title: "Payment Received on Paytm", Text: "₹100.37 Received from Rahul", }) result, err := service.IngestSigned(context.Background(), signedAuth(t, priv, deviceID, now, body), body) @@ -346,6 +681,32 @@ func TestPreEnrollmentNotificationIsStoredButIgnored(t *testing.T) { t.Fatalf("pre-enrollment result = %+v", result) } } +func TestMissingNotificationTimestampIsStoredButIgnored(t *testing.T) { + db := openRelayDB(t) + now := time.Date(2026, 9, 1, 4, 45, 0, 0, time.UTC) + priv, deviceID := enrollTestDevice(t, db, now.Add(-time.Minute)) + service := NewService(db, payments.NewService(db)) + service.Now = func() time.Time { return now } + body := marshalEvent(t, EventInput{ + SchemaVersion: 1, EventID: strings.Repeat("d", 64), + PackageName: observations.PaytmBusinessPackage, + Title: "Payment Received on Paytm", Text: "₹100.37 Received from Rahul", + }) + result, err := service.IngestSigned(context.Background(), signedAuth(t, priv, deviceID, now, body), body) + if err != nil { + t.Fatal(err) + } + if result.Status != "ignored" || countRows(t, db, "payment_observations") != 0 { + t.Fatalf("missing timestamp result = %+v", result) + } + var status, errorText string + if err := db.SQL.QueryRow(`SELECT status,error FROM relay_events WHERE id=?`, result.RelayEventID).Scan(&status, &errorText); err != nil { + t.Fatal(err) + } + if status != "ignored" || errorText != "notification posting time is missing or invalid" { + t.Fatalf("stored missing timestamp state=%q error=%q", status, errorText) + } +} func TestRetryResumesPreviouslyReceivedRelayEvent(t *testing.T) { ctx := context.Background() db := openRelayDB(t) @@ -359,16 +720,17 @@ func TestRetryResumesPreviouslyReceivedRelayEvent(t *testing.T) { service := NewService(db, paymentService) service.Now = func() time.Time { return receivedAt } sourceID := strings.Repeat("f", 64) - _, err := db.SQL.Exec(`INSERT INTO relay_events(id,device_id,source_event_id,package_name,posted_at,received_at,status) - VALUES('relay_existing',?,?,?,?,?,'received')`, deviceID, sourceID, observations.PaytmBusinessPackage, occurredAt.UnixMilli(), receivedAt.UnixMilli()) - if err != nil { - t.Fatal(err) - } body := marshalEvent(t, EventInput{ SchemaVersion: 1, EventID: sourceID, PackageName: observations.PaytmBusinessPackage, PostedAtMS: occurredAt.UnixMilli(), Title: "Payment Received on Paytm", Text: "₹100.37 Received from Rahul", + AmountHintPaise: 10037, }) + _, err := db.SQL.Exec(`INSERT INTO relay_events(id,device_id,source_event_id,package_name,posted_at,received_at,amount_hint_paise,title,text,status) + VALUES('relay_existing',?,?,?,?,?,10037,?,?,'received')`, deviceID, sourceID, observations.PaytmBusinessPackage, occurredAt.UnixMilli(), receivedAt.UnixMilli(), "Payment Received on Paytm", "₹100.37 Received from Rahul") + if err != nil { + t.Fatal(err) + } result, err := service.IngestSigned(ctx, signedAuth(t, priv, deviceID, receivedAt, body), body) if err != nil { t.Fatal(err) @@ -380,38 +742,54 @@ func TestRetryResumesPreviouslyReceivedRelayEvent(t *testing.T) { t.Fatal("retry did not resume the existing relay event") } } -func TestSignedKotakGoogleMessagesEventMatchesKotakPayment(t *testing.T) { +func TestSignedBlockedMessageAndEmailPackagesAreRejectedBeforeStorage(t *testing.T) { ctx := context.Background() db := openRelayDB(t) - createdAt := time.Date(2026, 9, 1, 5, 30, 0, 0, time.UTC) - insertProfile(t, db, "kotak", "kotak_sms", "merchant@kotak", true, createdAt) - paymentService, created := createPayment(t, db, createdAt, "kotak-1") - occurredAt := createdAt.Add(2 * time.Minute) - receivedAt := occurredAt.Add(time.Second) - paymentService.Now = func() time.Time { return receivedAt } - priv, deviceID := enrollTestDevice(t, db, createdAt.Add(-time.Minute)) - service := NewService(db, paymentService) - service.Now = func() time.Time { return receivedAt } + now := time.Date(2026, 9, 1, 5, 30, 0, 0, time.UTC) + priv, deviceID := enrollTestDevice(t, db, now.Add(-time.Minute)) + service := NewService(db, payments.NewService(db)) + service.Now = func() time.Time { return now } + for index, packageName := range []string{observations.GoogleMessagesPackage, observations.GmailPackage} { + body := marshalEvent(t, EventInput{ + SchemaVersion: 1, EventID: strings.Repeat(string(rune('1'+index)), 64), + PackageName: packageName, PostedAtMS: now.UnixMilli(), + Title: "Payment received", Text: "Received Rs. 100.37 from maya@okaxis", + }) + if _, err := service.IngestSigned(ctx, signedAuth(t, priv, deviceID, now, body), body); err == nil || !strings.Contains(err.Error(), "no longer accepted") { + t.Fatalf("package %s error=%v", packageName, err) + } + } + if countRows(t, db, "relay_events") != 0 { + t.Fatal("blocked notification packages must not be stored") + } +} + +func TestUntrustedGenericNotificationCannotAutoConfirm(t *testing.T) { + db := openRelayDB(t) + now := time.Date(2026, 9, 4, 7, 45, 0, 0, time.UTC) + insertProfile(t, db, "paytm", "paytm_notification", "merchant@paytm", true, now.Add(-time.Hour)) + paymentService, created := createPayment(t, db, now, "untrusted-generic") + priv, deviceID := enrollTestDevice(t, db, now.Add(-time.Hour)) + relayService := NewService(db, paymentService) + relayService.Now = func() time.Time { return now.Add(time.Minute) } body := marshalEvent(t, EventInput{ - SchemaVersion: 1, EventID: strings.Repeat("1", 64), - PackageName: observations.GoogleMessagesPackage, - PostedAtMS: occurredAt.UnixMilli(), - Title: "Kotak Mahindra Bank", - Text: "Kotak: Received Rs. 100.37 from maya@okaxis", + SchemaVersion: 1, EventID: strings.Repeat("e", 64), + PackageName: "com.example.wallet", PostedAtMS: now.Add(time.Minute).UnixMilli(), + Text: "Payment received ₹100.37 from Rahul", }) - result, err := service.IngestSigned(ctx, signedAuth(t, priv, deviceID, receivedAt, body), body) + result, err := relayService.IngestSigned(context.Background(), signedAuth(t, priv, deviceID, now.Add(time.Minute), body), body) if err != nil { t.Fatal(err) } - if result.Status != "matched" || result.PaymentID != created.Payment.ID || !result.Transitioned { - t.Fatalf("Kotak result = %+v", result) + if result.Status != "ignored" || result.PaymentID != "" || countRows(t, db, "payment_observations") != 0 { + t.Fatalf("untrusted generic result=%+v observations=%d", result, countRows(t, db, "payment_observations")) } - got, err := paymentService.Get(ctx, created.Payment.ID) + got, err := paymentService.Get(context.Background(), created.Payment.ID) if err != nil { t.Fatal(err) } - if got.Payment.Status != "paid" || got.Payment.PayerUPIID != "maya@okaxis" { - t.Fatalf("Kotak paid payment = %+v", got.Payment) + if got.Payment.Status != "pending" { + t.Fatalf("untrusted generic notification changed payment to %q", got.Payment.Status) } } @@ -425,7 +803,7 @@ func TestGenericWalletNotificationMatchesActiveProfilePayment(t *testing.T) { relayService.Now = func() time.Time { return now.Add(time.Minute) } body := marshalEvent(t, EventInput{ SchemaVersion: 1, EventID: strings.Repeat("a", 64), - PackageName: "com.example.wallet", PostedAtMS: now.Add(time.Minute).UnixMilli(), + PackageName: observations.GooglePayPackage, PostedAtMS: now.Add(time.Minute).UnixMilli(), Text: "Payment received ₹100.37 from Rahul", }) result, err := relayService.IngestSigned(context.Background(), signedAuth(t, priv, deviceID, now.Add(time.Minute), body), body) @@ -439,8 +817,8 @@ func TestGenericWalletNotificationMatchesActiveProfilePayment(t *testing.T) { if err != nil { t.Fatal(err) } - if got.Payment.Status != "paid" { - t.Fatalf("payment status=%s", got.Payment.Status) + if got.Payment.Status != "paid" || got.Payment.PayerName != "Rahul" { + t.Fatalf("payment=%+v", got.Payment) } } @@ -448,7 +826,7 @@ func TestGenericWalletNotificationUsesReservationProfileAfterActiveSwitch(t *tes db := openRelayDB(t) now := time.Date(2026, 9, 5, 8, 15, 0, 0, time.UTC) insertProfile(t, db, "old-profile", "paytm_notification", "old@upi", true, now.Add(-time.Hour)) - paymentService, created := createPayment(t, db, now, "generic-profile-switch") + paymentService, _ := createPayment(t, db, now, "generic-profile-switch") if _, err := db.SQL.Exec(`UPDATE collection_profiles SET active=0,updated_at=? WHERE id='old-profile'`, now.Add(20*time.Second).UnixMilli()); err != nil { t.Fatal(err) } @@ -461,87 +839,41 @@ func TestGenericWalletNotificationUsesReservationProfileAfterActiveSwitch(t *tes relayService.Now = func() time.Time { return received } body := marshalEvent(t, EventInput{ SchemaVersion: 1, EventID: strings.Repeat("d", 64), - PackageName: "com.example.wallet", PostedAtMS: occurred.UnixMilli(), + PackageName: observations.GooglePayPackage, PostedAtMS: occurred.UnixMilli(), Text: "Payment received ₹100.37 from Rahul", }) result, err := relayService.IngestSigned(context.Background(), signedAuth(t, priv, deviceID, received, body), body) if err != nil { t.Fatal(err) } - if result.Status != "matched" || result.PaymentID != created.Payment.ID || !result.Transitioned { + if result.Status != "matched" || result.PaymentID == "" || !result.Transitioned { t.Fatalf("generic delayed result=%+v", result) } - var profileID string - if err := db.SQL.QueryRow(`SELECT collection_profile_id FROM payment_observations WHERE matched_payment_id=?`, created.Payment.ID).Scan(&profileID); err != nil { + var profileID, matchResult string + if err := db.SQL.QueryRow(`SELECT collection_profile_id,match_result FROM payment_observations WHERE relay_event_id=?`, result.RelayEventID).Scan(&profileID, &matchResult); err != nil { t.Fatal(err) } - if profileID != "old-profile" { - t.Fatalf("observation profile=%q want old-profile", profileID) + if profileID != "old-profile" || matchResult != "matched" { + t.Fatalf("observation profile=%q result=%q", profileID, matchResult) } } -func TestGenericWalletNotificationIsAmbiguousAcrossProfileReservations(t *testing.T) { - db := openRelayDB(t) - now := time.Date(2026, 9, 5, 9, 0, 0, 0, time.UTC) - insertProfile(t, db, "profile-a", "paytm_notification", "a@upi", true, now.Add(-time.Hour)) - _, first := createPayment(t, db, now, "generic-ambiguous-a") - if _, err := db.SQL.Exec(`UPDATE collection_profiles SET active=0,updated_at=? WHERE id='profile-a'`, now.Add(10*time.Second).UnixMilli()); err != nil { - t.Fatal(err) - } - insertProfile(t, db, "profile-b", "paytm_notification", "b@upi", true, now.Add(10*time.Second)) - secondService, second := createPayment(t, db, now.Add(10*time.Second), "generic-ambiguous-b") - if first.Payment.PayableAmountPaise != second.Payment.PayableAmountPaise { - t.Fatalf("test requires overlapping amount reservations: %d != %d", first.Payment.PayableAmountPaise, second.Payment.PayableAmountPaise) - } - - priv, deviceID := enrollTestDevice(t, db, now.Add(-time.Hour)) - occurred := now.Add(time.Minute) - received := occurred.Add(time.Second) - relayService := NewService(db, secondService) - relayService.Now = func() time.Time { return received } - body := marshalEvent(t, EventInput{ - SchemaVersion: 1, EventID: strings.Repeat("e", 64), - PackageName: "com.example.wallet", PostedAtMS: occurred.UnixMilli(), - Text: "₹100.37 received from Rahul", - }) - result, err := relayService.IngestSigned(context.Background(), signedAuth(t, priv, deviceID, received, body), body) - if err != nil { - t.Fatal(err) - } - if result.Status != "ambiguous" || result.PaymentID != "" || result.Transitioned { - t.Fatalf("ambiguous generic result=%+v", result) - } - if countRows(t, db, "payment_observations") != 0 { - t.Fatal("ambiguous cross-profile notification must not claim an observation profile") - } - for _, paymentID := range []string{first.Payment.ID, second.Payment.ID} { - var status string - if err := db.SQL.QueryRow(`SELECT status FROM payments WHERE id=?`, paymentID).Scan(&status); err != nil { - t.Fatal(err) - } - if status != "pending" { - t.Fatalf("payment %s status=%s want pending", paymentID, status) - } - } -} - -func TestTwoEnabledPhonesCanRelaySameIncomingPaymentSafely(t *testing.T) { +func TestOneEnabledPhoneCanRelaySameIncomingPaymentSafely(t *testing.T) { db := openRelayDB(t) now := time.Date(2026, 9, 4, 8, 30, 0, 0, time.UTC) insertProfile(t, db, "paytm", "paytm_notification", "merchant@paytm", true, now.Add(-time.Hour)) - paymentService, created := createPayment(t, db, now, "two-phone-match") - firstPriv, firstID := enrollTestDevice(t, db, now.Add(-time.Hour)) - secondPriv, secondID := enrollTestDevice(t, db, now.Add(-30*time.Minute)) + paymentService, created := createPayment(t, db, now, "one-phone-match") + priv, deviceID := enrollTestDevice(t, db, now.Add(-time.Hour)) relayService := NewService(db, paymentService) occurred := now.Add(time.Minute) relayService.Now = func() time.Time { return occurred.Add(time.Second) } - first := marshalEvent(t, EventInput{SchemaVersion: 1, EventID: strings.Repeat("b", 64), PackageName: "com.example.wallet", PostedAtMS: occurred.UnixMilli(), Text: "₹100.37 received from Rahul"}) - one, err := relayService.IngestSigned(context.Background(), signedAuth(t, firstPriv, firstID, occurred.Add(time.Second), first), first) + first := marshalEvent(t, EventInput{SchemaVersion: 1, EventID: strings.Repeat("b", 64), PackageName: observations.PaytmBusinessPackage, PostedAtMS: occurred.UnixMilli(), Text: "₹100.37 received from Rahul"}) + one, err := relayService.IngestSigned(context.Background(), signedAuth(t, priv, deviceID, occurred.Add(time.Second), first), first) if err != nil { t.Fatal(err) } - second := marshalEvent(t, EventInput{SchemaVersion: 1, EventID: strings.Repeat("c", 64), PackageName: "com.example.wallet", PostedAtMS: occurred.Add(500 * time.Millisecond).UnixMilli(), Text: "₹100.37 received from Rahul"}) - two, err := relayService.IngestSigned(context.Background(), signedAuth(t, secondPriv, secondID, occurred.Add(2*time.Second), second), second) + second := marshalEvent(t, EventInput{SchemaVersion: 1, EventID: strings.Repeat("c", 64), PackageName: observations.PaytmBusinessPackage, PostedAtMS: occurred.Add(500 * time.Millisecond).UnixMilli(), Text: "₹100.37 received from Rahul"}) + two, err := relayService.IngestSigned(context.Background(), signedAuth(t, priv, deviceID, occurred.Add(2*time.Second), second), second) if err != nil { t.Fatal(err) } @@ -552,6 +884,6 @@ func TestTwoEnabledPhonesCanRelaySameIncomingPaymentSafely(t *testing.T) { t.Fatalf("second=%+v", two) } if countRows(t, db, "relay_events") != 2 || countRows(t, db, "payment_observations") != 2 { - t.Fatal("two-phone audit evidence was not retained") + t.Fatal("audit evidence was not retained") } } diff --git a/internal/v4/runtime/app_test.go b/internal/v4/runtime/app_test.go index 2c4f075..6acbdfa 100644 --- a/internal/v4/runtime/app_test.go +++ b/internal/v4/runtime/app_test.go @@ -21,6 +21,29 @@ import ( const testAdminPassword = "correct horse battery staple" +func TestRawEventRetentionDefaultsAndBounds(t *testing.T) { + cfg, err := (Config{DataDir: t.TempDir()}).normalized() + if err != nil { + t.Fatal(err) + } + if cfg.RawEventRetention != 7*24*time.Hour { + t.Fatalf("default raw event retention = %s", cfg.RawEventRetention) + } + for _, retention := range []time.Duration{59 * time.Minute, 31 * 24 * time.Hour} { + _, err := (Config{DataDir: t.TempDir(), RawEventRetention: retention}).normalized() + if err == nil { + t.Fatalf("raw event retention %s should be rejected", retention) + } + } + cfg, err = (Config{DataDir: t.TempDir(), RawEventRetention: 24 * time.Hour}).normalized() + if err != nil { + t.Fatal(err) + } + if cfg.RawEventRetention != 24*time.Hour { + t.Fatalf("explicit raw event retention = %s", cfg.RawEventRetention) + } +} + func newTestApp(t *testing.T, mutate func(*Config)) *App { t.Helper() cfg := Config{ diff --git a/internal/v4/runtime/config.go b/internal/v4/runtime/config.go index bb2c585..5edee69 100644 --- a/internal/v4/runtime/config.go +++ b/internal/v4/runtime/config.go @@ -20,9 +20,16 @@ type Config struct { BackupDir string BackupHourUTC int BackupRetention int + RawEventRetention time.Duration ExpiryInterval time.Duration } +const ( + defaultRawEventRetention = 7 * 24 * time.Hour + minRawEventRetention = 1 * time.Hour + maxRawEventRetention = 30 * 24 * time.Hour +) + func (c Config) normalized() (Config, error) { if strings.TrimSpace(c.DataDir) == "" { return Config{}, errors.New("PayGate v4 data directory is required") @@ -57,6 +64,12 @@ func (c Config) normalized() (Config, error) { if c.BackupRetention < 1 || c.BackupRetention > 365 { return Config{}, errors.New("backup retention must be between 1 and 365") } + if c.RawEventRetention == 0 { + c.RawEventRetention = defaultRawEventRetention + } + if c.RawEventRetention < minRawEventRetention || c.RawEventRetention > maxRawEventRetention { + return Config{}, errors.New("raw event retention must be between 1h and 30d") + } if c.ExpiryInterval <= 0 { c.ExpiryInterval = 30 * time.Second } diff --git a/internal/v4/runtime/workers.go b/internal/v4/runtime/workers.go index 98322f3..a803bb7 100644 --- a/internal/v4/runtime/workers.go +++ b/internal/v4/runtime/workers.go @@ -18,6 +18,7 @@ func (a *App) RunWorkers(ctx context.Context) { go a.Webhooks.Run(ctx) go a.expiryWorker(ctx) go a.backupWorker(ctx) + go a.rawEventWorker(ctx) } func (a *App) expiryWorker(ctx context.Context) { @@ -43,6 +44,41 @@ func (a *App) expiryWorker(ctx context.Context) { } } } +func (a *App) rawEventWorker(ctx context.Context) { + const ( + interval = time.Hour + batch = 500 + ) + redact := func() { + var total int64 + for { + before := time.Now().UTC().Add(-a.Config.RawEventRetention) + count, err := a.Relay.RedactRawEvents(ctx, before, batch) + if err != nil { + slog.Error("redact expired relay notification bodies", "error", err) + return + } + total += count + if count < batch { + break + } + } + if total > 0 { + slog.Info("redacted expired relay notification bodies", "count", total) + } + } + redact() + ticker := time.NewTicker(interval) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + redact() + } + } +} func (a *App) BackupNow(ctx context.Context, at time.Time) (string, error) { if a == nil || a.DB == nil { return "", fmt.Errorf("PayGate storage is unavailable") diff --git a/internal/v4/storage/db.go b/internal/v4/storage/db.go index d1f60b3..5b7ec69 100644 --- a/internal/v4/storage/db.go +++ b/internal/v4/storage/db.go @@ -3,8 +3,10 @@ package storage import ( "context" "database/sql" + "errors" "fmt" "net/url" + "os" "path/filepath" "strings" "time" @@ -17,6 +19,8 @@ const ( schemaVersion = 4 ) +var ErrBusy = errors.New("sqlite database busy") + type DB struct { SQL *sql.DB Path string @@ -40,6 +44,30 @@ func (tx *ImmediateTx) QueryRowContext(ctx context.Context, query string, args . return tx.conn.QueryRowContext(ctx, query, args...) } +type sqliteCodeError interface { + Code() int +} + +func isSQLiteBusy(err error) bool { + if errors.Is(err, ErrBusy) { + return true + } + var coded sqliteCodeError + if errors.As(err, &coded) { + baseCode := coded.Code() & 0xff + return baseCode == 5 || baseCode == 6 + } + message := strings.ToUpper(err.Error()) + return strings.Contains(message, "SQLITE_BUSY") || strings.Contains(message, "SQLITE_LOCKED") +} + +func wrapTransactionError(operation string, err error) error { + if isSQLiteBusy(err) { + return fmt.Errorf("%w: %s: %w", ErrBusy, operation, err) + } + return fmt.Errorf("%s: %w", operation, err) +} + // WithImmediateTx runs fn inside BEGIN IMMEDIATE on a dedicated pooled connection. // Use it for short payment-critical write transactions. Network calls must never // happen inside fn. Ordinary reads should use DB.SQL directly. @@ -51,7 +79,7 @@ func (db *DB) WithImmediateTx(ctx context.Context, fn func(*ImmediateTx) error) defer conn.Close() if _, err := conn.ExecContext(ctx, "BEGIN IMMEDIATE"); err != nil { - return fmt.Errorf("begin immediate transaction: %w", err) + return wrapTransactionError("begin immediate transaction", err) } done := false defer func() { @@ -61,15 +89,68 @@ func (db *DB) WithImmediateTx(ctx context.Context, fn func(*ImmediateTx) error) }() if err := fn(&ImmediateTx{conn: conn}); err != nil { + if isSQLiteBusy(err) { + return wrapTransactionError("immediate transaction callback", err) + } return err } if _, err := conn.ExecContext(ctx, "COMMIT"); err != nil { - return fmt.Errorf("commit immediate transaction: %w", err) + return wrapTransactionError("commit immediate transaction", err) } done = true return nil } +const databaseFileMode os.FileMode = 0o600 + +func prepareDatabaseFile(path string) error { + info, err := os.Lstat(path) + if errors.Is(err, os.ErrNotExist) { + file, createErr := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_EXCL, databaseFileMode) + if createErr != nil { + return fmt.Errorf("create sqlite database: %w", createErr) + } + if closeErr := file.Close(); closeErr != nil { + return fmt.Errorf("close sqlite database: %w", closeErr) + } + return nil + } + if err != nil { + return fmt.Errorf("inspect sqlite database: %w", err) + } + if info.Mode()&os.ModeSymlink != 0 || !info.Mode().IsRegular() { + return fmt.Errorf("sqlite database must be a regular file") + } + if info.Mode().Perm() != databaseFileMode { + if err := os.Chmod(path, databaseFileMode); err != nil { + return fmt.Errorf("harden sqlite database permissions: %w", err) + } + } + return nil +} + +func verifyDatabaseSidecars(path string) error { + for _, suffix := range []string{"-wal", "-shm", "-journal"} { + sidecar := path + suffix + info, err := os.Lstat(sidecar) + if errors.Is(err, os.ErrNotExist) { + continue + } + if err != nil { + return fmt.Errorf("inspect sqlite %s file: %w", suffix, err) + } + if info.Mode()&os.ModeSymlink != 0 || !info.Mode().IsRegular() { + return fmt.Errorf("sqlite %s file must be a regular file", suffix) + } + if info.Mode().Perm() != databaseFileMode { + if err := os.Chmod(sidecar, databaseFileMode); err != nil { + return fmt.Errorf("harden sqlite %s permissions: %w", suffix, err) + } + } + } + return nil +} + func Open(ctx context.Context, path string) (*DB, error) { if strings.TrimSpace(path) == "" { return nil, fmt.Errorf("sqlite path is required") @@ -78,6 +159,12 @@ func Open(ctx context.Context, path string) (*DB, error) { if err != nil { return nil, fmt.Errorf("resolve sqlite path: %w", err) } + if err := prepareDatabaseFile(abs); err != nil { + return nil, err + } + if err := verifyDatabaseSidecars(abs); err != nil { + return nil, err + } q := url.Values{} q.Add("_pragma", "journal_mode(WAL)") q.Add("_pragma", "synchronous(FULL)") @@ -100,6 +187,10 @@ func Open(ctx context.Context, path string) (*DB, error) { raw.Close() return nil, fmt.Errorf("ping sqlite: %w", err) } + if err := verifyDatabaseSidecars(abs); err != nil { + raw.Close() + return nil, err + } if err := db.verifyPragmas(ctx); err != nil { raw.Close() return nil, err @@ -112,6 +203,14 @@ func Open(ctx context.Context, path string) (*DB, error) { raw.Close() return nil, err } + if err := verifyDatabaseSidecars(abs); err != nil { + raw.Close() + return nil, err + } + if err := db.ensureRelayPayloadIntegrity(ctx); err != nil { + raw.Close() + return nil, err + } return db, nil } diff --git a/internal/v4/storage/db_test.go b/internal/v4/storage/db_test.go index 3835b60..9a256a6 100644 --- a/internal/v4/storage/db_test.go +++ b/internal/v4/storage/db_test.go @@ -100,32 +100,36 @@ func TestPaymentIdentityDoesNotUseNameOrExternalID(t *testing.T) { } } -func TestActiveAmountReservationIsUniqueButHistoryIsRetained(t *testing.T) { +func TestActiveAmountReservationIsGloballyUniqueButHistoryIsRetained(t *testing.T) { db := openTestDB(t) now := int64(1_788_200_000_000) insertProfile(t, db.SQL, "paytm", true, now) + insertProfile(t, db.SQL, "kotak", false, now) insertPayment(t, db.SQL, "pay_1", "Person A", "evt_123", 10037, now) insertPayment(t, db.SQL, "pay_2", "Person B", "evt_123", 10037, now+1) + if _, err := db.SQL.Exec(`UPDATE payments SET collection_profile_id='kotak' WHERE id='pay_2'`); err != nil { + t.Fatal(err) + } if _, err := db.SQL.Exec(`INSERT INTO amount_reservations(id,collection_profile_id,payable_amount_paise,payment_id,reserved_at,reserved_until,last_used_at) VALUES('res_1','paytm',10037,'pay_1',?,?,?)`, now, now+900_000, now); err != nil { t.Fatal(err) } if _, err := db.SQL.Exec(`INSERT INTO amount_reservations(id,collection_profile_id,payable_amount_paise,payment_id,reserved_at,reserved_until,last_used_at) - VALUES('res_2','paytm',10037,'pay_2',?,?,?)`, now+1, now+900_001, now+1); err == nil { - t.Fatal("expected duplicate active profile+amount reservation to fail") + VALUES('res_2','kotak',10037,'pay_2',?,?,?)`, now+1, now+900_001, now+1); err == nil { + t.Fatal("expected duplicate active payable amount to fail") } if _, err := db.SQL.Exec(`UPDATE amount_reservations SET released_at=? WHERE id='res_1'`, now+900_000); err != nil { t.Fatal(err) } if _, err := db.SQL.Exec(`INSERT INTO amount_reservations(id,collection_profile_id,payable_amount_paise,payment_id,reserved_at,reserved_until,last_used_at) - VALUES('res_2','paytm',10037,'pay_2',?,?,?)`, now+900_001, now+1_800_001, now+900_001); err != nil { + VALUES('res_2','kotak',10037,'pay_2',?,?,?)`, now+900_001, now+1_800_001, now+900_001); err != nil { t.Fatalf("reuse after release should succeed: %v", err) } var count int - if err := db.SQL.QueryRow(`SELECT COUNT(*) FROM amount_reservations WHERE collection_profile_id='paytm' AND payable_amount_paise=10037`).Scan(&count); err != nil { + if err := db.SQL.QueryRow(`SELECT COUNT(*) FROM amount_reservations WHERE payable_amount_paise=10037`).Scan(&count); err != nil { t.Fatal(err) } if count != 2 { @@ -334,6 +338,116 @@ func TestOrdinaryReadTransactionDoesNotAcquireWriterLock(t *testing.T) { t.Fatalf("ordinary read transaction blocked writer: %v", err) } } +func TestOpenMigratesV4CrossProfileDuplicateReservationsWithoutRewritingAmounts(t *testing.T) { + ctx := context.Background() + path := filepath.Join(t.TempDir(), "paygate-v4-duplicates.db") + raw, err := sql.Open("sqlite", "file:"+filepath.ToSlash(path)) + if err != nil { + t.Fatal(err) + } + if _, err := raw.ExecContext(ctx, `CREATE TABLE schema_migrations(version INTEGER PRIMARY KEY, applied_at INTEGER NOT NULL) STRICT;`); err != nil { + t.Fatal(err) + } + if _, err := raw.ExecContext(ctx, schemaV1); err != nil { + t.Fatal(err) + } + if _, err := raw.ExecContext(ctx, schemaV2); err != nil { + t.Fatal(err) + } + if _, err := raw.ExecContext(ctx, `INSERT INTO schema_migrations(version,applied_at) VALUES(1,1),(2,2)`); err != nil { + t.Fatal(err) + } + v4 := &DB{SQL: raw, Path: path} + if err := v4.applyV3(ctx); err != nil { + t.Fatal(err) + } + if err := v4.runMigrationTx(ctx, 4, applyV4); err != nil { + t.Fatal(err) + } + + now := int64(1_788_200_000_000) + if _, err := raw.ExecContext(ctx, `INSERT INTO collection_profiles(id,label,upi_id,parser,enabled,active,created_at,updated_at) VALUES + ('paytm','Paytm','paytm@upi','paytm_notification',1,1,?,?), + ('kotak','Kotak','kotak@upi','kotak_sms',1,0,?,?)`, now, now, now, now); err != nil { + t.Fatal(err) + } + insert := `INSERT INTO payments(id,name,metadata_json,requested_amount_paise,payable_amount_paise,adjustment_paise,collection_profile_id,upi_id_snapshot,status,created_at,expires_at,grace_until,reuse_after) + VALUES(?,?,'{}',10000,10037,37,?,?,'pending',?,?,?,?)` + if _, err := raw.ExecContext(ctx, insert, "pay_paytm", "Paytm payer", "paytm", "paytm@upi", now, now+300_000, now+600_000, now+900_000); err != nil { + t.Fatal(err) + } + if _, err := raw.ExecContext(ctx, insert, "pay_kotak", "Kotak payer", "kotak", "kotak@upi", now+1, now+300_001, now+600_001, now+900_001); err != nil { + t.Fatal(err) + } + reservation := `INSERT INTO amount_reservations(id,collection_profile_id,payable_amount_paise,payment_id,reserved_at,reserved_until,last_used_at) VALUES(?,?,?,?,?,?,?)` + if _, err := raw.ExecContext(ctx, reservation, "res_paytm", "paytm", 10037, "pay_paytm", now, now+900_000, now); err != nil { + t.Fatal(err) + } + if _, err := raw.ExecContext(ctx, reservation, "res_kotak", "kotak", 10037, "pay_kotak", now+1, now+900_001, now+1); err != nil { + t.Fatalf("v4 should allow cross-profile duplicate amount: %v", err) + } + if err := raw.Close(); err != nil { + t.Fatal(err) + } + + db, err := Open(ctx, path) + if err != nil { + t.Fatalf("upgrade with live v4 duplicates failed: %v", err) + } + defer db.Close() + var version int + if err := db.SQL.QueryRowContext(ctx, `SELECT MAX(version) FROM schema_migrations`).Scan(&version); err != nil { + t.Fatal(err) + } + if version != schemaVersion { + t.Fatalf("schema version=%d want=%d", version, schemaVersion) + } + rows, err := db.SQL.QueryContext(ctx, `SELECT payment_id,payable_amount_paise,global_unique_enforced FROM amount_reservations ORDER BY payment_id`) + if err != nil { + t.Fatal(err) + } + defer rows.Close() + seen := 0 + for rows.Next() { + var paymentID string + var amount int64 + var enforced int + if err := rows.Scan(&paymentID, &amount, &enforced); err != nil { + t.Fatal(err) + } + if amount != 10037 || enforced != 0 { + t.Fatalf("grandfathered reservation %s amount=%d enforced=%d", paymentID, amount, enforced) + } + seen++ + } + if err := rows.Err(); err != nil { + t.Fatal(err) + } + if seen != 2 { + t.Fatalf("grandfathered reservations=%d want=2", seen) + } + + if _, err := db.SQL.ExecContext(ctx, insert, "pay_blocked", "Blocked", "paytm", "paytm@upi", now+2, now+300_002, now+600_002, now+900_002); err != nil { + t.Fatal(err) + } + if _, err := db.SQL.ExecContext(ctx, reservation, "res_blocked", "paytm", 10037, "pay_blocked", now+2, now+900_002, now+2); err == nil { + t.Fatal("new reservation reused a grandfathered live amount") + } + if _, err := db.SQL.ExecContext(ctx, `UPDATE payments SET payable_amount_paise=10048,adjustment_paise=48 WHERE id='pay_blocked'`); err != nil { + t.Fatal(err) + } + if _, err := db.SQL.ExecContext(ctx, reservation, "res_new", "paytm", 10048, "pay_blocked", now+2, now+900_002, now+2); err != nil { + t.Fatalf("new globally unique reservation failed: %v", err) + } + var enforced int + if err := db.SQL.QueryRowContext(ctx, `SELECT global_unique_enforced FROM amount_reservations WHERE id='res_new'`).Scan(&enforced); err != nil { + t.Fatal(err) + } + if enforced != 1 { + t.Fatalf("new reservation enforcement=%d want=1", enforced) + } +} + func TestOpenMigratesV1DatabaseToRelayHealthSchema(t *testing.T) { ctx := context.Background() path := filepath.Join(t.TempDir(), "paygate-v1.db") @@ -470,7 +584,7 @@ func TestObservationSchemaSupportsCorroborationAndFutureSources(t *testing.T) { } } -func TestMultiRelayCompatibilityKeepsSchemaV4RollbackReadable(t *testing.T) { +func TestMultiRelayCompatibilityKeepsSchemaRollbackReadable(t *testing.T) { db, err := Open(context.Background(), filepath.Join(t.TempDir(), "paygate.db")) if err != nil { t.Fatal(err) @@ -481,8 +595,8 @@ func TestMultiRelayCompatibilityKeepsSchemaV4RollbackReadable(t *testing.T) { if err := db.SQL.QueryRow(`SELECT COALESCE(MAX(version),0) FROM schema_migrations`).Scan(&version); err != nil { t.Fatal(err) } - if version != 4 { - t.Fatalf("schema version=%d want=4 for rollback compatibility", version) + if version != schemaVersion { + t.Fatalf("schema version=%d want=%d for rollback compatibility", version, schemaVersion) } var indexes int @@ -493,3 +607,44 @@ func TestMultiRelayCompatibilityKeepsSchemaV4RollbackReadable(t *testing.T) { t.Fatalf("singleton relay index still present: %d", indexes) } } + +func TestRollbackV4SQLShapeRemainsOperational(t *testing.T) { + db := openTestDB(t) + now := int64(1_788_200_000_000) + + var version int + if err := db.SQL.QueryRow(`SELECT COALESCE(MAX(version),0) FROM schema_migrations`).Scan(&version); err != nil { + t.Fatal(err) + } + if version != 4 { + t.Fatalf("rollback schema version=%d want=4", version) + } + + insertProfile(t, db.SQL, "paytm", true, now) + insertProfile(t, db.SQL, "kotak", false, now) + insertPayment(t, db.SQL, "pay_v4_a", "A", "evt", 10037, now) + insertPayment(t, db.SQL, "pay_v4_b", "B", "evt", 10037, now+1) + if _, err := db.SQL.Exec(`UPDATE payments SET collection_profile_id='kotak',upi_id_snapshot='kotak@upi' WHERE id='pay_v4_b'`); err != nil { + t.Fatal(err) + } + + oldInsert := `INSERT INTO amount_reservations(id,collection_profile_id,payable_amount_paise,payment_id,reserved_at,reserved_until,last_used_at) VALUES(?,?,?,?,?,?,?)` + if _, err := db.SQL.Exec(oldInsert, "res_v4_a", "paytm", 10037, "pay_v4_a", now, now+900_000, now); err != nil { + t.Fatalf("old-v4 reservation insert failed: %v", err) + } + if _, err := db.SQL.Exec(oldInsert, "res_v4_b", "kotak", 10037, "pay_v4_b", now+1, now+900_001, now+1); err == nil { + t.Fatal("old-v4 SQL bypassed global active amount guard") + } + + if _, err := db.SQL.Exec(`INSERT INTO relay_devices(id,name,public_key_pem,enabled,enrolled_at,app_version,device_model,android_version) VALUES('v4-device','Phone','pem',1,?,'v4','model','android')`, now); err != nil { + t.Fatalf("old-v4 relay insert failed: %v", err) + } + var enabled int + var enrolled int64 + if err := db.SQL.QueryRow(`SELECT enabled,enrolled_at FROM relay_devices WHERE id='v4-device'`).Scan(&enabled, &enrolled); err != nil { + t.Fatalf("old-v4 relay read failed: %v", err) + } + if enabled != 1 || enrolled != now { + t.Fatalf("old-v4 relay read enabled=%d enrolled=%d", enabled, enrolled) + } +} diff --git a/internal/v4/storage/restore.go b/internal/v4/storage/restore.go new file mode 100644 index 0000000..a5bdb7d --- /dev/null +++ b/internal/v4/storage/restore.go @@ -0,0 +1,907 @@ +package storage + +import ( + "context" + "crypto/sha256" + "database/sql" + "encoding/hex" + "errors" + "fmt" + "io" + "net/url" + "os" + "path/filepath" + "strings" +) + +const restoreFileMode os.FileMode = 0o600 + +// RestoreReport contains non-sensitive counts from an isolated backup validation. +type RestoreReport struct { + SchemaVersion int + Payments int + RelayEvents int + WebhookDeliveries int + RelayDevices int +} + +// RestoreDrill copies a standalone backup into a private temporary directory, +// opens it through the production storage path, and verifies integrity and +// representative tables. It never opens or writes the live database. +func RestoreDrill(ctx context.Context, backupPath, livePath, expectedSHA256 string) (report RestoreReport, err error) { + if ctx == nil { + ctx = context.Background() + } + backupPath, sourceInfo, expected, err := validateRestoreSource(backupPath, livePath, expectedSHA256) + if err != nil { + return RestoreReport{}, err + } + + drillDir, err := os.MkdirTemp("", "paygate-restore-drill-") + if err != nil { + return RestoreReport{}, fmt.Errorf("create restore drill directory: %w", err) + } + defer func() { + if removeErr := os.RemoveAll(drillDir); err == nil && removeErr != nil { + err = fmt.Errorf("remove restore drill directory: %w", removeErr) + } + }() + drillPath := filepath.Join(drillDir, "paygate.db") + + if err := copyRestoreSource(backupPath, sourceInfo, drillPath, expected); err != nil { + return RestoreReport{}, err + } + if err := validateRestoreDatabase(ctx, drillPath); err != nil { + return RestoreReport{}, fmt.Errorf("validate isolated restore: %w", err) + } + db, err := Open(ctx, drillPath) + if err != nil { + return RestoreReport{}, fmt.Errorf("open isolated restore: %w", err) + } + report, err = inspectRestoredDatabase(ctx, db) + closeErr := db.Close() + if err != nil { + return RestoreReport{}, err + } + if closeErr != nil { + return RestoreReport{}, fmt.Errorf("close isolated restore: %w", closeErr) + } + return report, nil +} + +func validateRestoreSource(backupPath, livePath, expectedSHA256 string) (string, os.FileInfo, []byte, error) { + if strings.TrimSpace(backupPath) == "" { + return "", nil, nil, errors.New("restore backup path is required") + } + backupPath, err := filepath.Abs(strings.TrimSpace(backupPath)) + if err != nil { + return "", nil, nil, fmt.Errorf("resolve restore backup path: %w", err) + } + info, err := os.Lstat(backupPath) + if err != nil { + return "", nil, nil, fmt.Errorf("inspect restore backup: %w", err) + } + if info.Mode()&os.ModeSymlink != 0 || !info.Mode().IsRegular() { + return "", nil, nil, errors.New("restore backup must be a regular file, not a symlink") + } + if info.Size() <= 0 { + return "", nil, nil, errors.New("restore backup must be non-empty") + } + if info.Mode().Perm() != restoreFileMode { + return "", nil, nil, fmt.Errorf("restore backup permissions must be %04o, got %04o", restoreFileMode.Perm(), info.Mode().Perm()) + } + for _, suffix := range []string{"-wal", "-shm", "-journal"} { + if sidecarInfo, sidecarErr := os.Lstat(backupPath + suffix); sidecarErr == nil { + if sidecarInfo.Mode()&os.ModeSymlink != 0 || !sidecarInfo.Mode().IsRegular() { + return "", nil, nil, fmt.Errorf("restore backup sidecar %s is not a regular file", suffix) + } + return "", nil, nil, fmt.Errorf("restore backup has sidecar %s; use a completed standalone backup", suffix) + } else if !errors.Is(sidecarErr, os.ErrNotExist) { + return "", nil, nil, fmt.Errorf("inspect restore backup sidecar %s: %w", suffix, sidecarErr) + } + } + if strings.TrimSpace(livePath) != "" { + livePath, err = filepath.Abs(strings.TrimSpace(livePath)) + if err != nil { + return "", nil, nil, fmt.Errorf("resolve live database path: %w", err) + } + if filepath.Clean(livePath) == filepath.Clean(backupPath) { + return "", nil, nil, errors.New("restore backup cannot be the live database") + } + liveInfo, statErr := os.Stat(livePath) + if statErr == nil { + if os.SameFile(info, liveInfo) { + return "", nil, nil, errors.New("restore backup cannot be the live database") + } + } else if !errors.Is(statErr, os.ErrNotExist) { + return "", nil, nil, fmt.Errorf("inspect live database path: %w", statErr) + } + } + + var expected []byte + if value := strings.TrimSpace(expectedSHA256); value != "" { + expected, err = hex.DecodeString(value) + if err != nil || len(expected) != sha256.Size { + return "", nil, nil, errors.New("restore SHA-256 must be 64 hexadecimal characters") + } + } + return backupPath, info, expected, nil +} + +func copyRestoreSource(sourcePath string, expectedInfo os.FileInfo, destination string, expectedHash []byte) error { + source, err := os.Open(sourcePath) + if err != nil { + return fmt.Errorf("open restore backup: %w", err) + } + defer source.Close() + openedInfo, err := source.Stat() + if err != nil { + return fmt.Errorf("stat restore backup: %w", err) + } + if !os.SameFile(expectedInfo, openedInfo) || openedInfo.Mode()&os.ModeSymlink != 0 || !openedInfo.Mode().IsRegular() { + return errors.New("restore backup changed while it was being opened") + } + + destinationFile, err := os.OpenFile(destination, os.O_WRONLY|os.O_CREATE|os.O_EXCL, restoreFileMode) + if err != nil { + return fmt.Errorf("create isolated restore copy: %w", err) + } + hasher := sha256.New() + written, copyErr := io.Copy(io.MultiWriter(destinationFile, hasher), source) + if copyErr == nil { + copyErr = destinationFile.Sync() + } + closeErr := destinationFile.Close() + if copyErr != nil { + return fmt.Errorf("copy restore backup after %d bytes: %w", written, copyErr) + } + if closeErr != nil { + return fmt.Errorf("close isolated restore copy: %w", closeErr) + } + afterInfo, err := source.Stat() + if err != nil { + return fmt.Errorf("stat restore backup after copy: %w", err) + } + if !os.SameFile(expectedInfo, afterInfo) || afterInfo.Size() != expectedInfo.Size() || + !afterInfo.ModTime().Equal(expectedInfo.ModTime()) { + return errors.New("restore backup changed while it was being copied") + } + for _, suffix := range []string{"-wal", "-shm", "-journal"} { + if _, sidecarErr := os.Lstat(sourcePath + suffix); sidecarErr == nil { + return fmt.Errorf("restore backup sidecar %s appeared while it was being copied", suffix) + } else if !errors.Is(sidecarErr, os.ErrNotExist) { + return fmt.Errorf("inspect restore backup sidecar %s after copy: %w", suffix, sidecarErr) + } + } + if len(expectedHash) > 0 && !strings.EqualFold(hex.EncodeToString(hasher.Sum(nil)), hex.EncodeToString(expectedHash)) { + return errors.New("restore backup SHA-256 does not match the expected value") + } + return nil +} + +func validateRestoreDatabase(ctx context.Context, path string) error { + query := url.Values{} + query.Set("mode", "ro") + query.Set("_query_only", "true") + raw, err := sql.Open("sqlite", "file:"+filepath.ToSlash(path)+"?"+query.Encode()) + if err != nil { + return fmt.Errorf("open restore validation connection: %w", err) + } + defer raw.Close() + raw.SetMaxOpenConns(1) + if err := raw.PingContext(ctx); err != nil { + return fmt.Errorf("ping restore validation connection: %w", err) + } + var integrity string + if err := raw.QueryRowContext(ctx, `PRAGMA integrity_check`).Scan(&integrity); err != nil { + return fmt.Errorf("restore pre-migration integrity check: %w", err) + } + if strings.ToLower(strings.TrimSpace(integrity)) != "ok" { + return fmt.Errorf("restore pre-migration integrity check returned %q", integrity) + } + rows, err := raw.QueryContext(ctx, `SELECT version FROM schema_migrations ORDER BY version`) + if err != nil { + return fmt.Errorf("read restore schema migrations before migration: %w", err) + } + var versions []int + for rows.Next() { + var version int + if err := rows.Scan(&version); err != nil { + rows.Close() + return fmt.Errorf("scan restore schema migration before migration: %w", err) + } + versions = append(versions, version) + } + if err := rows.Err(); err != nil { + rows.Close() + return fmt.Errorf("iterate restore schema migrations before migration: %w", err) + } + if err := rows.Close(); err != nil { + return fmt.Errorf("close restore schema migrations before migration: %w", err) + } + if len(versions) == 0 { + return errors.New("restore backup has no schema migrations") + } + for index, version := range versions { + if version != index+1 { + return fmt.Errorf("restore schema migrations are not contiguous at version %d", version) + } + } + lastVersion := versions[len(versions)-1] + if lastVersion > 7 { + return fmt.Errorf("restore schema version %d is newer than supported transitional schema 7", lastVersion) + } + hasColumn := func(tableName, columnName string) (bool, error) { + rows, err := raw.QueryContext(ctx, fmt.Sprintf("PRAGMA table_info(%s)", tableName)) + if err != nil { + return false, err + } + defer rows.Close() + for rows.Next() { + var cid, notNull, pk int + var name, kind string + var defaultValue any + if err := rows.Scan(&cid, &name, &kind, ¬Null, &defaultValue, &pk); err != nil { + return false, err + } + if name == columnName { + return true, nil + } + } + return false, rows.Err() + } + hasGlobalMarker, err := hasColumn("amount_reservations", "global_unique_enforced") + if err != nil { + return fmt.Errorf("inspect restore global amount marker: %w", err) + } + hasEpochRequired, err := hasColumn("relay_devices", "epoch_required") + if err != nil { + return fmt.Errorf("inspect restore relay epoch requirement: %w", err) + } + if lastVersion >= 6 && !hasEpochRequired { + return errors.New("restore transitional schema is missing relay epoch requirement") + } + if lastVersion >= 7 && !hasGlobalMarker { + return errors.New("restore transitional schema is missing global amount marker") + } + type restoreExpectedColumn struct { + name string + kind string + pk bool + } + type restoreColumn struct { + name string + kind string + pk bool + notNull bool + defaultValue sql.NullString + } + requiredColumns := map[string][]restoreExpectedColumn{ + "schema_migrations": {{"version", "INTEGER", true}, {"applied_at", "INTEGER", false}}, + "collection_profiles": {{"id", "TEXT", true}, {"label", "TEXT", false}, {"upi_id", "TEXT", false}, {"payee_name", "TEXT", false}, {"parser", "TEXT", false}, {"enabled", "INTEGER", false}, {"active", "INTEGER", false}, {"created_at", "INTEGER", false}, {"updated_at", "INTEGER", false}}, + "payments": {{"id", "TEXT", true}, {"name", "TEXT", false}, {"external_id", "TEXT", false}, {"metadata_json", "TEXT", false}, {"requested_amount_paise", "INTEGER", false}, {"payable_amount_paise", "INTEGER", false}, {"adjustment_paise", "INTEGER", false}, {"currency", "TEXT", false}, {"collection_profile_id", "TEXT", false}, {"upi_id_snapshot", "TEXT", false}, {"payee_name_snapshot", "TEXT", false}, {"status", "TEXT", false}, {"created_at", "INTEGER", false}, {"expires_at", "INTEGER", false}, {"grace_until", "INTEGER", false}, {"reuse_after", "INTEGER", false}, {"paid_at", "INTEGER", false}, {"payer_name", "TEXT", false}, {"payer_upi_id", "TEXT", false}, {"internal_note", "TEXT", false}}, + "amount_reservations": {{"id", "TEXT", true}, {"collection_profile_id", "TEXT", false}, {"payable_amount_paise", "INTEGER", false}, {"payment_id", "TEXT", false}, {"reserved_at", "INTEGER", false}, {"reserved_until", "INTEGER", false}, {"released_at", "INTEGER", false}, {"last_used_at", "INTEGER", false}}, + "idempotency_keys": {{"scope", "TEXT", true}, {"key_hash", "BLOB", true}, {"request_hash", "BLOB", false}, {"payment_id", "TEXT", false}, {"created_at", "INTEGER", false}, {"expires_at", "INTEGER", false}}, + "relay_devices": {{"id", "TEXT", true}, {"name", "TEXT", false}, {"public_key_pem", "TEXT", false}, {"enabled", "INTEGER", false}, {"enrolled_at", "INTEGER", false}, {"last_seen_at", "INTEGER", false}, {"last_heartbeat_at", "INTEGER", false}, {"app_version", "TEXT", false}, {"device_model", "TEXT", false}, {"android_version", "TEXT", false}}, + "pairing_sessions": {{"id", "TEXT", true}, {"token_hash", "BLOB", false}, {"replace_existing", "INTEGER", false}, {"created_at", "INTEGER", false}, {"expires_at", "INTEGER", false}, {"consumed_at", "INTEGER", false}}, + "relay_events": {{"id", "TEXT", true}, {"device_id", "TEXT", false}, {"source_event_id", "TEXT", false}, {"package_name", "TEXT", false}, {"posted_at", "INTEGER", false}, {"received_at", "INTEGER", false}, {"amount_hint_paise", "INTEGER", false}, {"title", "TEXT", false}, {"text", "TEXT", false}, {"big_text", "TEXT", false}, {"status", "TEXT", false}, {"error", "TEXT", false}}, + "payment_observations": {{"id", "TEXT", true}, {"relay_event_id", "TEXT", false}, {"source", "TEXT", false}, {"collection_profile_id", "TEXT", false}, {"amount_paise", "INTEGER", false}, {"payer_name", "TEXT", false}, {"payer_upi_id", "TEXT", false}, {"occurred_at", "INTEGER", false}, {"occurred_at_source", "TEXT", false}, {"received_at", "INTEGER", false}, {"matched_payment_id", "TEXT", false}, {"match_result", "TEXT", false}}, + "payment_history": {{"id", "TEXT", true}, {"payment_id", "TEXT", false}, {"type", "TEXT", false}, {"actor", "TEXT", false}, {"summary", "TEXT", false}, {"changes_json", "TEXT", false}, {"created_at", "INTEGER", false}}, + "webhook_deliveries": {{"id", "TEXT", true}, {"event_type", "TEXT", false}, {"payment_id", "TEXT", false}, {"payload_json", "TEXT", false}, {"status", "TEXT", false}, {"attempts", "INTEGER", false}, {"next_attempt_at", "INTEGER", false}, {"last_http_status", "INTEGER", false}, {"last_error", "TEXT", false}, {"created_at", "INTEGER", false}, {"delivered_at", "INTEGER", false}}, + "api_keys": {{"id", "TEXT", true}, {"label", "TEXT", false}, {"secret_hash", "BLOB", false}, {"enabled", "INTEGER", false}, {"created_at", "INTEGER", false}, {"last_used_at", "INTEGER", false}}, + "admin_credentials": {{"singleton", "INTEGER", true}, {"password_hash", "TEXT", false}, {"updated_at", "INTEGER", false}}, + "admin_sessions": {{"token_hash", "BLOB", true}, {"created_at", "INTEGER", false}, {"expires_at", "INTEGER", false}, {"last_seen_at", "INTEGER", false}, {"revoked_at", "INTEGER", false}}, + "settings": {{"key", "TEXT", true}, {"value", "TEXT", false}, {"updated_at", "INTEGER", false}}, + } + if versions[len(versions)-1] >= 2 { + requiredColumns["relay_devices"] = append(requiredColumns["relay_devices"], + restoreExpectedColumn{"notification_access", "INTEGER", false}, + restoreExpectedColumn{"listener_connected", "INTEGER", false}, + restoreExpectedColumn{"battery_optimization_exempt", "INTEGER", false}, + restoreExpectedColumn{"power_save_mode", "INTEGER", false}, + restoreExpectedColumn{"background_restricted", "INTEGER", false}, + restoreExpectedColumn{"foreground_service", "INTEGER", false}, + restoreExpectedColumn{"pending_count", "INTEGER", false}, + restoreExpectedColumn{"failed_count", "INTEGER", false}, + restoreExpectedColumn{"last_successful_delivery_at", "INTEGER", false}, + restoreExpectedColumn{"last_client_error", "TEXT", false}) + } + if hasEpochRequired { + requiredColumns["relay_devices"] = append(requiredColumns["relay_devices"], + restoreExpectedColumn{"epoch_required", "INTEGER", false}) + } + if hasGlobalMarker { + requiredColumns["amount_reservations"] = append(requiredColumns["amount_reservations"], + restoreExpectedColumn{"global_unique_enforced", "INTEGER", false}) + } + requiredNotNull := map[string][]string{ + "schema_migrations": {"applied_at"}, + "collection_profiles": {"label", "upi_id", "parser", "enabled", "active", "created_at", "updated_at"}, + "payments": {"name", "metadata_json", "requested_amount_paise", "payable_amount_paise", "adjustment_paise", "currency", "collection_profile_id", "upi_id_snapshot", "status", "created_at", "expires_at", "grace_until", "reuse_after"}, + "amount_reservations": {"collection_profile_id", "payable_amount_paise", "payment_id", "reserved_at", "reserved_until", "last_used_at"}, + "idempotency_keys": {"scope", "key_hash", "request_hash", "payment_id", "created_at", "expires_at"}, + "relay_devices": {"public_key_pem", "enabled", "enrolled_at"}, + "pairing_sessions": {"token_hash", "replace_existing", "created_at", "expires_at"}, + "relay_events": {"device_id", "source_event_id", "package_name", "posted_at", "received_at", "status"}, + "payment_observations": {"relay_event_id", "source", "collection_profile_id", "amount_paise", "occurred_at", "occurred_at_source", "received_at", "match_result"}, + "payment_history": {"payment_id", "type", "actor", "summary", "changes_json", "created_at"}, + "webhook_deliveries": {"event_type", "payment_id", "payload_json", "status", "attempts", "created_at"}, + "api_keys": {"label", "secret_hash", "enabled", "created_at"}, + "admin_credentials": {"password_hash", "updated_at"}, + "admin_sessions": {"created_at", "expires_at"}, + "settings": {"value", "updated_at"}, + } + if hasEpochRequired { + requiredNotNull["relay_devices"] = append(requiredNotNull["relay_devices"], "epoch_required") + } + if hasGlobalMarker { + requiredNotNull["amount_reservations"] = append(requiredNotNull["amount_reservations"], "global_unique_enforced") + } + type restoreForeignKey struct { + table string + from string + to string + onUpdate string + onDelete string + } + requiredForeignKeys := map[string][]restoreForeignKey{ + "payments": { + {"collection_profiles", "collection_profile_id", "id", "RESTRICT", "RESTRICT"}, + }, + "amount_reservations": { + {"collection_profiles", "collection_profile_id", "id", "RESTRICT", "RESTRICT"}, + {"payments", "payment_id", "id", "RESTRICT", "RESTRICT"}, + }, + "idempotency_keys": { + {"payments", "payment_id", "id", "RESTRICT", "RESTRICT"}, + }, + "relay_events": { + {"relay_devices", "device_id", "id", "RESTRICT", "RESTRICT"}, + }, + "payment_observations": { + {"relay_events", "relay_event_id", "id", "RESTRICT", "RESTRICT"}, + {"collection_profiles", "collection_profile_id", "id", "RESTRICT", "RESTRICT"}, + {"payments", "matched_payment_id", "id", "RESTRICT", "SET NULL"}, + }, + "payment_history": { + {"payments", "payment_id", "id", "RESTRICT", "RESTRICT"}, + }, + "webhook_deliveries": { + {"payments", "payment_id", "id", "RESTRICT", "RESTRICT"}, + }, + } + requiredChecks := map[string]bool{ + "collection_profiles": true, "payments": true, "amount_reservations": true, + "idempotency_keys": true, "relay_devices": true, "pairing_sessions": true, + "relay_events": true, "payment_observations": true, "payment_history": true, + "webhook_deliveries": true, "api_keys": true, "admin_credentials": true, + "admin_sessions": true, + } + requiredCheckFragments := map[string][]string{ + "collection_profiles": { + "LENGTH(TRIM(LABEL)) BETWEEN 1 AND 120", + "LENGTH(TRIM(UPI_ID)) BETWEEN 3 AND 255", + "ENABLED IN (0,1)", + "ACTIVE IN (0,1)", + "ACTIVE = 0 OR ENABLED = 1", + }, + "payments": { + "LENGTH(TRIM(NAME)) BETWEEN 1 AND 120", + "REQUESTED_AMOUNT_PAISE > 0", + "REQUESTED_AMOUNT_PAISE % 100 = 0", + "PAYABLE_AMOUNT_PAISE > REQUESTED_AMOUNT_PAISE", + "PAYABLE_AMOUNT_PAISE % 100 BETWEEN 1 AND 99", + "PAYABLE_AMOUNT_PAISE = REQUESTED_AMOUNT_PAISE + ADJUSTMENT_PAISE", + "ADJUSTMENT_PAISE BETWEEN 1 AND 199", + "CURRENCY = 'INR'", + "JSON_VALID(METADATA_JSON)", + "STATUS IN ('PENDING','PAID','EXPIRED','CANCELLED')", + "CREATED_AT < EXPIRES_AT AND EXPIRES_AT < GRACE_UNTIL AND GRACE_UNTIL < REUSE_AFTER", + "STATUS = 'PAID' AND PAID_AT IS NOT NULL", + "STATUS != 'PAID' AND PAID_AT IS NULL", + }, + "amount_reservations": { + "PAYABLE_AMOUNT_PAISE > 0", + "PAYABLE_AMOUNT_PAISE % 100 BETWEEN 1 AND 99", + "RESERVED_UNTIL > RESERVED_AT", + "RELEASED_AT IS NULL OR RELEASED_AT >= RESERVED_AT", + }, + "idempotency_keys": {"EXPIRES_AT > CREATED_AT"}, + "relay_devices": {"ENABLED IN (0,1)"}, + "pairing_sessions": { + "REPLACE_EXISTING IN (0,1)", + "EXPIRES_AT > CREATED_AT", + "CONSUMED_AT IS NULL OR (CONSUMED_AT >= CREATED_AT AND CONSUMED_AT <= EXPIRES_AT)", + }, + "relay_events": { + "AMOUNT_HINT_PAISE IS NULL OR (AMOUNT_HINT_PAISE > 0 AND AMOUNT_HINT_PAISE % 100 BETWEEN 1 AND 99)", + "STATUS IN ('RECEIVED','PARSED','IGNORED','MATCHED','UNMATCHED','AMBIGUOUS','ERROR')", + }, + "payment_observations": { + "AMOUNT_PAISE > 0", + "AMOUNT_PAISE % 100 BETWEEN 1 AND 99", + "OCCURRED_AT_SOURCE IN ('NOTIFICATION_TEXT','NOTIFICATION_POSTED_AT','SERVER_RECEIVED_AT')", + }, + "payment_history": { + "ACTOR IN ('SYSTEM','ADMIN')", + "JSON_VALID(CHANGES_JSON)", + }, + "webhook_deliveries": { + "EVENT_TYPE IN ('PAYMENT.CREATED','PAYMENT.PAID','PAYMENT.EXPIRED','PAYMENT.CANCELLED','PAYMENT.UPDATED')", + "JSON_VALID(PAYLOAD_JSON)", + "STATUS IN ('PENDING','RETRY','DELIVERED','EXHAUSTED')", + "ATTEMPTS >= 0", + "STATUS = 'DELIVERED' AND DELIVERED_AT IS NOT NULL", + "STATUS != 'DELIVERED' AND DELIVERED_AT IS NULL", + }, + "api_keys": {"ENABLED IN (0,1)"}, + "admin_credentials": {"SINGLETON = 1"}, + "admin_sessions": {"EXPIRES_AT > CREATED_AT"}, + } + if hasEpochRequired { + requiredCheckFragments["relay_devices"] = append(requiredCheckFragments["relay_devices"], "EPOCH_REQUIRED IN (0,1)") + } + if hasGlobalMarker { + requiredCheckFragments["amount_reservations"] = append(requiredCheckFragments["amount_reservations"], "GLOBAL_UNIQUE_ENFORCED IN (0,1)") + } + if versions[len(versions)-1] >= 3 { + requiredCheckFragments["collection_profiles"] = append(requiredCheckFragments["collection_profiles"], + "PARSER IN ('PAYTM_NOTIFICATION','KOTAK_SMS','LEGACY')", + "PARSER != 'LEGACY' OR (ENABLED = 0 AND ACTIVE = 0)") + } else { + requiredCheckFragments["collection_profiles"] = append(requiredCheckFragments["collection_profiles"], + "PARSER IN ('PAYTM_NOTIFICATION','KOTAK_SMS')") + } + if versions[len(versions)-1] >= 4 { + requiredCheckFragments["payment_observations"] = append(requiredCheckFragments["payment_observations"], + "LENGTH(TRIM(SOURCE)) BETWEEN 1 AND 64", + "MATCH_RESULT IN ('MATCHED','CORROBORATED','UNMATCHED','AMBIGUOUS','IGNORED','ERROR')") + } else { + requiredCheckFragments["payment_observations"] = append(requiredCheckFragments["payment_observations"], + "SOURCE IN ('PAYTM_NOTIFICATION','KOTAK_SMS')", + "MATCH_RESULT IN ('MATCHED','UNMATCHED','AMBIGUOUS','IGNORED','ERROR')") + } + for table, columns := range requiredColumns { + var createSQL sql.NullString + if err := raw.QueryRowContext(ctx, `SELECT sql FROM sqlite_master WHERE type='table' AND name=?`, table).Scan(&createSQL); err != nil { + return fmt.Errorf("read restore table %s definition: %w", table, err) + } + if !createSQL.Valid || !strings.Contains(strings.ToUpper(createSQL.String), "STRICT") { + return fmt.Errorf("restore table %s is not a strict production table", table) + } + upperSQL := strings.ToUpper(createSQL.String) + if requiredChecks[table] && !strings.Contains(upperSQL, "CHECK") { + return fmt.Errorf("restore table %s is missing production checks", table) + } + for _, fragment := range requiredCheckFragments[table] { + if !strings.Contains(upperSQL, fragment) { + return fmt.Errorf("restore table %s is missing production check %s", table, fragment) + } + } + tableRows, err := raw.QueryContext(ctx, fmt.Sprintf("PRAGMA table_info(%s)", table)) + if err != nil { + return fmt.Errorf("inspect restore table %s columns: %w", table, err) + } + found := make(map[string]restoreColumn, len(columns)) + for tableRows.Next() { + var cid, notNull, primaryKey int + var name, columnType string + var defaultValue sql.NullString + if err := tableRows.Scan(&cid, &name, &columnType, ¬Null, &defaultValue, &primaryKey); err != nil { + tableRows.Close() + return fmt.Errorf("scan restore table %s columns: %w", table, err) + } + if _, duplicate := found[name]; duplicate { + tableRows.Close() + return fmt.Errorf("restore table %s has duplicate column %s", table, name) + } + if name == "payload_hash" && table == "relay_events" { + if !strings.EqualFold(strings.TrimSpace(columnType), "BLOB") { + tableRows.Close() + return fmt.Errorf("restore table %s payload_hash must be BLOB", table) + } + found[name] = restoreColumn{name: name, kind: "BLOB", notNull: notNull != 0, defaultValue: defaultValue} + continue + } + found[name] = restoreColumn{name: name, kind: strings.ToUpper(strings.TrimSpace(columnType)), pk: primaryKey > 0, notNull: notNull != 0, defaultValue: defaultValue} + } + if err := tableRows.Err(); err != nil { + tableRows.Close() + return fmt.Errorf("iterate restore table %s columns: %w", table, err) + } + if err := tableRows.Close(); err != nil { + return fmt.Errorf("close restore table %s columns: %w", table, err) + } + for name := range found { + known := false + for _, column := range columns { + if column.name == name { + known = true + break + } + } + if !known && !(table == "relay_events" && name == "payload_hash") { + return fmt.Errorf("restore table %s has unexpected column %s", table, name) + } + } + for _, name := range requiredNotNull[table] { + actual, ok := found[name] + if !ok || !actual.notNull { + return fmt.Errorf("restore table %s column %s must be NOT NULL", table, name) + } + } + expectedCount := len(columns) + if table == "relay_events" && found["payload_hash"].name != "" { + expectedCount++ + } + if len(found) != expectedCount { + return fmt.Errorf("restore table %s has %d columns, want %d", table, len(found), expectedCount) + } + for _, column := range columns { + actual, ok := found[column.name] + if !ok { + return fmt.Errorf("restore table %s is missing column %s", table, column.name) + } + if actual.kind != column.kind { + return fmt.Errorf("restore table %s column %s has type %s, want %s", table, column.name, actual.kind, column.kind) + } + if actual.pk != column.pk { + return fmt.Errorf("restore table %s column %s primary-key flag mismatch", table, column.name) + } + if column.name == "global_unique_enforced" && strings.Trim(strings.TrimSpace(actual.defaultValue.String), "'\"") != "1" { + return fmt.Errorf("restore table %s column %s must default to 1", table, column.name) + } + if column.name == "epoch_required" && strings.Trim(strings.TrimSpace(actual.defaultValue.String), "'\"") != "0" { + return fmt.Errorf("restore table %s column %s must default to 0", table, column.name) + } + } + if expected := requiredForeignKeys[table]; len(expected) > 0 { + fkRows, err := raw.QueryContext(ctx, fmt.Sprintf("PRAGMA foreign_key_list(%s)", table)) + if err != nil { + return fmt.Errorf("inspect restore table %s foreign keys: %w", table, err) + } + actual := make(map[string]struct{}, len(expected)) + for fkRows.Next() { + var id, seq int + var referencedTable, from, to, onUpdate, onDelete, match string + if err := fkRows.Scan(&id, &seq, &referencedTable, &from, &to, &onUpdate, &onDelete, &match); err != nil { + fkRows.Close() + return fmt.Errorf("scan restore table %s foreign keys: %w", table, err) + } + actual[strings.Join([]string{referencedTable, from, to, strings.ToUpper(onUpdate), strings.ToUpper(onDelete)}, "\x00")] = struct{}{} + } + if err := fkRows.Err(); err != nil { + fkRows.Close() + return fmt.Errorf("iterate restore table %s foreign keys: %w", table, err) + } + if err := fkRows.Close(); err != nil { + return fmt.Errorf("close restore table %s foreign keys: %w", table, err) + } + if len(actual) != len(expected) { + return fmt.Errorf("restore table %s has %d foreign keys, want %d", table, len(actual), len(expected)) + } + for _, foreignKey := range expected { + key := strings.Join([]string{foreignKey.table, foreignKey.from, foreignKey.to, foreignKey.onUpdate, foreignKey.onDelete}, "\x00") + if _, ok := actual[key]; !ok { + return fmt.Errorf("restore table %s is missing foreign key %s(%s) references %s(%s) ON UPDATE %s ON DELETE %s", table, table, foreignKey.from, foreignKey.table, foreignKey.to, foreignKey.onUpdate, foreignKey.onDelete) + } + } + } + } + type restoreIndex struct { + table string + unique bool + columns []string + fragments []string + } + requiredIndexes := map[string]restoreIndex{ + "uq_collection_profiles_one_active": {"collection_profiles", true, []string{"active"}, []string{"ACTIVE = 1", "WHERE"}}, + "idx_payments_external_id": {"payments", false, []string{"external_id"}, []string{"EXTERNAL_ID"}}, + "idx_payments_status_created": {"payments", false, []string{"status", "created_at"}, []string{"STATUS", "CREATED_AT"}}, + "idx_payments_profile_payable": {"payments", false, []string{"collection_profile_id", "payable_amount_paise"}, []string{"COLLECTION_PROFILE_ID", "PAYABLE_AMOUNT_PAISE"}}, + "idx_amount_reservations_history": {"amount_reservations", false, []string{"collection_profile_id", "payable_amount_paise", "reserved_at"}, []string{"COLLECTION_PROFILE_ID", "PAYABLE_AMOUNT_PAISE", "RESERVED_AT"}}, + "idx_amount_reservations_release": {"amount_reservations", false, []string{"released_at", "reserved_until"}, []string{"RELEASED_AT", "RESERVED_UNTIL"}}, + "idx_relay_events_received": {"relay_events", false, []string{"received_at"}, []string{"RECEIVED_AT"}}, + "idx_observations_amount_time": {"payment_observations", false, []string{"collection_profile_id", "amount_paise", "occurred_at"}, []string{"COLLECTION_PROFILE_ID", "AMOUNT_PAISE", "OCCURRED_AT"}}, + "idx_payment_history_payment": {"payment_history", false, []string{"payment_id", "created_at"}, []string{"PAYMENT_ID", "CREATED_AT"}}, + "idx_webhook_delivery_queue": {"webhook_deliveries", false, []string{"status", "next_attempt_at", "created_at"}, []string{"STATUS", "NEXT_ATTEMPT_AT", "CREATED_AT"}}, + "idx_admin_sessions_expiry": {"admin_sessions", false, []string{"expires_at"}, []string{"EXPIRES_AT"}}, + } + if lastVersion <= 4 { + requiredIndexes["uq_active_profile_payable"] = restoreIndex{ + table: "amount_reservations", unique: true, columns: []string{"collection_profile_id", "payable_amount_paise"}, + fragments: []string{"COLLECTION_PROFILE_ID", "PAYABLE_AMOUNT_PAISE", "RELEASED_AT", "IS NULL", "WHERE"}, + } + } + if lastVersion >= 4 { + requiredIndexes["idx_observations_payment"] = restoreIndex{ + table: "payment_observations", unique: false, columns: []string{"matched_payment_id", "occurred_at"}, + fragments: []string{"MATCHED_PAYMENT_ID", "OCCURRED_AT", "IS NOT NULL", "WHERE"}, + } + } + if hasGlobalMarker { + requiredIndexes["uq_active_payable"] = restoreIndex{ + table: "amount_reservations", unique: true, columns: []string{"payable_amount_paise"}, + fragments: []string{"PAYABLE_AMOUNT_PAISE", "RELEASED_AT", "IS NULL", "WHERE", "GLOBAL_UNIQUE_ENFORCED", "=1"}, + } + } else if lastVersion >= 5 { + requiredIndexes["uq_active_payable"] = restoreIndex{ + table: "amount_reservations", unique: true, columns: []string{"payable_amount_paise"}, + fragments: []string{"PAYABLE_AMOUNT_PAISE", "RELEASED_AT", "IS NULL", "WHERE"}, + } + } + readIndexColumns := func(indexName string) ([]string, error) { + quotedName := strings.ReplaceAll(indexName, "'", "''") + rows, err := raw.QueryContext(ctx, fmt.Sprintf("PRAGMA index_info('%s')", quotedName)) + if err != nil { + return nil, err + } + var columns []string + for rows.Next() { + var seq, cid int + var columnName sql.NullString + if err := rows.Scan(&seq, &cid, &columnName); err != nil { + rows.Close() + return nil, err + } + if !columnName.Valid { + rows.Close() + return nil, fmt.Errorf("index %s contains an expression", indexName) + } + columns = append(columns, columnName.String) + } + if err := rows.Err(); err != nil { + rows.Close() + return nil, err + } + if err := rows.Close(); err != nil { + return nil, err + } + return columns, nil + } + sameIndexColumns := func(actual, expected []string) bool { + if len(actual) != len(expected) { + return false + } + for index := range expected { + if !strings.EqualFold(actual[index], expected[index]) { + return false + } + } + return true + } + requiredUniqueColumns := map[string][]string{ + "relay_events": {"device_id", "source_event_id"}, + "payment_observations": {"relay_event_id"}, + "amount_reservations": {"payment_id"}, + "pairing_sessions": {"token_hash"}, + "api_keys": {"secret_hash"}, + } + for table, expectedColumns := range requiredUniqueColumns { + rows, err := raw.QueryContext(ctx, fmt.Sprintf("PRAGMA index_list(%s)", table)) + if err != nil { + return fmt.Errorf("inspect restore table %s indexes: %w", table, err) + } + var uniqueNames []string + for rows.Next() { + var seq, unique, partial int + var origin, indexName string + if err := rows.Scan(&seq, &indexName, &unique, &origin, &partial); err != nil { + rows.Close() + return fmt.Errorf("scan restore table %s indexes: %w", table, err) + } + if unique != 0 { + uniqueNames = append(uniqueNames, indexName) + } + } + if err := rows.Err(); err != nil { + rows.Close() + return fmt.Errorf("iterate restore table %s indexes: %w", table, err) + } + if err := rows.Close(); err != nil { + return fmt.Errorf("close restore table %s indexes: %w", table, err) + } + found := false + for _, indexName := range uniqueNames { + columns, err := readIndexColumns(indexName) + if err != nil { + return fmt.Errorf("inspect restore index %s: %w", indexName, err) + } + if sameIndexColumns(columns, expectedColumns) { + found = true + break + } + } + if !found { + return fmt.Errorf("restore table %s is missing unique index on (%s)", table, strings.Join(expectedColumns, ", ")) + } + } + relayIndexRows, err := raw.QueryContext(ctx, `PRAGMA index_list(relay_devices)`) + if err != nil { + return fmt.Errorf("inspect relay device indexes: %w", err) + } + var relayUniqueNames []string + for relayIndexRows.Next() { + var seq, unique, partial int + var origin, indexName string + if err := relayIndexRows.Scan(&seq, &indexName, &unique, &origin, &partial); err != nil { + relayIndexRows.Close() + return fmt.Errorf("scan relay device indexes: %w", err) + } + if unique != 0 { + relayUniqueNames = append(relayUniqueNames, indexName) + } + } + if err := relayIndexRows.Err(); err != nil { + relayIndexRows.Close() + return fmt.Errorf("iterate relay device indexes: %w", err) + } + if err := relayIndexRows.Close(); err != nil { + return fmt.Errorf("close relay device indexes: %w", err) + } + for _, indexName := range relayUniqueNames { + columns, err := readIndexColumns(indexName) + if err != nil { + return fmt.Errorf("inspect relay device index %s: %w", indexName, err) + } + if sameIndexColumns(columns, []string{"enabled"}) { + return errors.New("restore backup retains singleton relay_devices enabled uniqueness") + } + } + requiredIndexDefinitions := map[string]string{ + "uq_collection_profiles_one_active": "CREATE UNIQUE INDEX uq_collection_profiles_one_active ON collection_profiles(active) WHERE active = 1", + "idx_payments_external_id": "CREATE INDEX idx_payments_external_id ON payments(external_id)", + "idx_payments_status_created": "CREATE INDEX idx_payments_status_created ON payments(status, created_at DESC)", + "idx_payments_profile_payable": "CREATE INDEX idx_payments_profile_payable ON payments(collection_profile_id, payable_amount_paise)", + "idx_amount_reservations_history": "CREATE INDEX idx_amount_reservations_history ON amount_reservations(collection_profile_id, payable_amount_paise, reserved_at DESC)", + "idx_amount_reservations_release": "CREATE INDEX idx_amount_reservations_release ON amount_reservations(released_at, reserved_until)", + "idx_relay_events_received": "CREATE INDEX idx_relay_events_received ON relay_events(received_at DESC)", + "idx_observations_amount_time": "CREATE INDEX idx_observations_amount_time ON payment_observations(collection_profile_id, amount_paise, occurred_at)", + "idx_payment_history_payment": "CREATE INDEX idx_payment_history_payment ON payment_history(payment_id, created_at)", + "idx_webhook_delivery_queue": "CREATE INDEX idx_webhook_delivery_queue ON webhook_deliveries(status, next_attempt_at, created_at)", + "idx_admin_sessions_expiry": "CREATE INDEX idx_admin_sessions_expiry ON admin_sessions(expires_at)", + } + if lastVersion <= 4 { + requiredIndexDefinitions["uq_active_profile_payable"] = "CREATE UNIQUE INDEX uq_active_profile_payable ON amount_reservations(collection_profile_id, payable_amount_paise) WHERE released_at IS NULL" + } + if lastVersion >= 4 { + requiredIndexDefinitions["idx_observations_payment"] = "CREATE INDEX idx_observations_payment ON payment_observations(matched_payment_id, occurred_at) WHERE matched_payment_id IS NOT NULL" + } + if hasGlobalMarker { + requiredIndexDefinitions["uq_active_payable"] = "CREATE UNIQUE INDEX uq_active_payable ON amount_reservations(payable_amount_paise) WHERE released_at IS NULL AND global_unique_enforced=1" + } else if lastVersion >= 5 { + requiredIndexDefinitions["uq_active_payable"] = "CREATE UNIQUE INDEX uq_active_payable ON amount_reservations(payable_amount_paise) WHERE released_at IS NULL" + } + canonicalSQL := func(value string) string { + return strings.Join(strings.Fields(strings.ToUpper(value)), " ") + } + requiredTriggers := map[string][]string{} + if hasGlobalMarker { + requiredTriggers["trg_amount_reservations_global_unique_insert"] = []string{ + "BEFORE INSERT ON AMOUNT_RESERVATIONS", "NEW.GLOBAL_UNIQUE_ENFORCED<>1", + "NEW.RELEASED_AT IS NULL", "PAYABLE_AMOUNT_PAISE=NEW.PAYABLE_AMOUNT_PAISE", "RAISE(ABORT", + } + updateFragments := []string{ + "BEFORE UPDATE OF PAYABLE_AMOUNT_PAISE,RELEASED_AT,GLOBAL_UNIQUE_ENFORCED ON AMOUNT_RESERVATIONS", + "OLD.GLOBAL_UNIQUE_ENFORCED=1", "NEW.GLOBAL_UNIQUE_ENFORCED<>1", + "NEW.RELEASED_AT IS NULL", "ID<>NEW.ID", "RAISE(ABORT", + } + if lastVersion <= 4 { + updateFragments = append(updateFragments, + "NEW.GLOBAL_UNIQUE_ENFORCED NOT IN (0,1)", "NEW.GLOBAL_UNIQUE_ENFORCED=0") + } + requiredTriggers["trg_amount_reservations_global_unique_update"] = updateFragments + } + + for name, expected := range requiredIndexes { + var tableName, indexSQL string + if err := raw.QueryRowContext(ctx, `SELECT tbl_name,sql FROM sqlite_master WHERE type='index' AND name=?`, name).Scan(&tableName, &indexSQL); err != nil { + return fmt.Errorf("read restore index %s: %w", name, err) + } + if !strings.EqualFold(tableName, expected.table) { + return fmt.Errorf("restore index %s belongs to table %s, want %s", name, tableName, expected.table) + } + indexRows, err := raw.QueryContext(ctx, fmt.Sprintf("PRAGMA index_list(%s)", expected.table)) + if err != nil { + return fmt.Errorf("inspect restore index %s metadata: %w", name, err) + } + var metadataFound, metadataUnique bool + for indexRows.Next() { + var seq, unique, partial int + var indexName, origin string + if err := indexRows.Scan(&seq, &indexName, &unique, &origin, &partial); err != nil { + indexRows.Close() + return fmt.Errorf("scan restore index %s metadata: %w", name, err) + } + if indexName == name { + metadataFound = true + metadataUnique = unique != 0 + } + } + if err := indexRows.Err(); err != nil { + indexRows.Close() + return fmt.Errorf("iterate restore index %s metadata: %w", name, err) + } + if err := indexRows.Close(); err != nil { + return fmt.Errorf("close restore index %s metadata: %w", name, err) + } + if !metadataFound { + return fmt.Errorf("restore index %s is not listed by SQLite", name) + } + if expected.unique != metadataUnique { + return fmt.Errorf("restore index %s uniqueness mismatch", name) + } + columns, err := readIndexColumns(name) + if err != nil { + return fmt.Errorf("inspect restore index %s columns: %w", name, err) + } + if !sameIndexColumns(columns, expected.columns) { + return fmt.Errorf("restore index %s has columns (%s), want (%s)", name, strings.Join(columns, ", "), strings.Join(expected.columns, ", ")) + } + upperIndexSQL := strings.ToUpper(indexSQL) + if expected.unique != strings.Contains(upperIndexSQL, "CREATE UNIQUE INDEX") { + return fmt.Errorf("restore index %s uniqueness mismatch", name) + } + for _, fragment := range expected.fragments { + if !strings.Contains(upperIndexSQL, fragment) { + return fmt.Errorf("restore index %s is missing definition fragment %s", name, fragment) + } + } + if want := requiredIndexDefinitions[name]; want != "" && canonicalSQL(indexSQL) != canonicalSQL(want) { + return fmt.Errorf("restore index %s definition mismatch", name) + } + } + for name, fragments := range requiredTriggers { + var tableName, triggerSQL string + if err := raw.QueryRowContext(ctx, `SELECT tbl_name,sql FROM sqlite_master WHERE type='trigger' AND name=?`, name).Scan(&tableName, &triggerSQL); err != nil { + return fmt.Errorf("read restore trigger %s: %w", name, err) + } + if !strings.EqualFold(tableName, "amount_reservations") { + return fmt.Errorf("restore trigger %s belongs to table %s, want amount_reservations", name, tableName) + } + upperTriggerSQL := strings.ToUpper(triggerSQL) + for _, fragment := range fragments { + if !strings.Contains(upperTriggerSQL, fragment) { + return fmt.Errorf("restore trigger %s is missing definition fragment %s", name, fragment) + } + } + } + return nil +} + +func inspectRestoredDatabase(ctx context.Context, db *DB) (RestoreReport, error) { + var report RestoreReport + var integrity string + if err := db.SQL.QueryRowContext(ctx, `PRAGMA integrity_check`).Scan(&integrity); err != nil { + return RestoreReport{}, fmt.Errorf("restore integrity check: %w", err) + } + if strings.ToLower(strings.TrimSpace(integrity)) != "ok" { + return RestoreReport{}, fmt.Errorf("restore integrity check returned %q", integrity) + } + rows, err := db.SQL.QueryContext(ctx, `PRAGMA foreign_key_check`) + if err != nil { + return RestoreReport{}, fmt.Errorf("restore foreign-key check: %w", err) + } + defer rows.Close() + if rows.Next() { + return RestoreReport{}, errors.New("restore foreign-key check found violations") + } + if err := rows.Err(); err != nil { + return RestoreReport{}, fmt.Errorf("iterate restore foreign-key check: %w", err) + } + + if err := db.SQL.QueryRowContext(ctx, `SELECT COALESCE(MAX(version),0) FROM schema_migrations`).Scan(&report.SchemaVersion); err != nil { + return RestoreReport{}, fmt.Errorf("read restored schema version: %w", err) + } + for query, target := range map[string]*int{ + `SELECT COUNT(*) FROM payments`: &report.Payments, + `SELECT COUNT(*) FROM relay_events`: &report.RelayEvents, + `SELECT COUNT(*) FROM webhook_deliveries`: &report.WebhookDeliveries, + `SELECT COUNT(*) FROM relay_devices`: &report.RelayDevices, + } { + if err := db.SQL.QueryRowContext(ctx, query).Scan(target); err != nil { + return RestoreReport{}, fmt.Errorf("read restored table count: %w", err) + } + } + return report, nil +} diff --git a/internal/v4/storage/restore_test.go b/internal/v4/storage/restore_test.go new file mode 100644 index 0000000..126eef1 --- /dev/null +++ b/internal/v4/storage/restore_test.go @@ -0,0 +1,387 @@ +package storage + +import ( + "context" + "crypto/sha256" + "database/sql" + "encoding/hex" + "fmt" + "os" + "path/filepath" + "strings" + "testing" +) + +func TestRestoreDrillValidatesIsolatedBackup(t *testing.T) { + ctx := context.Background() + dir := t.TempDir() + livePath := filepath.Join(dir, "paygate.db") + live, err := Open(ctx, livePath) + if err != nil { + t.Fatal(err) + } + backupPath := filepath.Join(dir, "backup.db") + if err := live.BackupTo(ctx, backupPath); err != nil { + live.Close() + t.Fatal(err) + } + if err := live.Close(); err != nil { + t.Fatal(err) + } + + raw, err := os.ReadFile(backupPath) + if err != nil { + t.Fatal(err) + } + digest := sha256.Sum256(raw) + report, err := RestoreDrill(ctx, backupPath, livePath, hex.EncodeToString(digest[:])) + if err != nil { + t.Fatal(err) + } + if report.SchemaVersion != schemaVersion || report.Payments != 0 || report.RelayEvents != 0 || report.WebhookDeliveries != 0 || report.RelayDevices != 0 { + t.Fatalf("restore report = %+v", report) + } + if _, err := os.Stat(livePath); err != nil { + t.Fatalf("live database was not preserved: %v", err) + } +} + +func TestRestoreDrillRejectsLivePathAndWrongHash(t *testing.T) { + ctx := context.Background() + dir := t.TempDir() + livePath := filepath.Join(dir, "paygate.db") + live, err := Open(ctx, livePath) + if err != nil { + t.Fatal(err) + } + backupPath := filepath.Join(dir, "backup.db") + if err := live.BackupTo(ctx, backupPath); err != nil { + live.Close() + t.Fatal(err) + } + if err := live.Close(); err != nil { + t.Fatal(err) + } + + if _, err := RestoreDrill(ctx, backupPath, backupPath, ""); err == nil || !strings.Contains(err.Error(), "live database") { + t.Fatalf("same path error = %v", err) + } + wrongHash := strings.Repeat("0", sha256.Size*2) + if _, err := RestoreDrill(ctx, backupPath, livePath, wrongHash); err == nil || !strings.Contains(err.Error(), "SHA-256") { + t.Fatalf("wrong hash error = %v", err) + } +} + +func TestRestoreDrillRejectsTruncatedBackup(t *testing.T) { + ctx := context.Background() + dir := t.TempDir() + livePath := filepath.Join(dir, "live.db") + live, err := Open(ctx, livePath) + if err != nil { + t.Fatal(err) + } + backupPath := filepath.Join(dir, "backup.db") + if err := live.BackupTo(ctx, backupPath); err != nil { + live.Close() + t.Fatal(err) + } + if err := live.Close(); err != nil { + t.Fatal(err) + } + raw, err := os.ReadFile(backupPath) + if err != nil { + t.Fatal(err) + } + truncatedPath := filepath.Join(dir, "truncated.db") + if err := os.WriteFile(truncatedPath, raw[:len(raw)/2], 0o600); err != nil { + t.Fatal(err) + } + if _, err := RestoreDrill(ctx, truncatedPath, livePath, ""); err == nil { + t.Fatal("truncated backup was accepted") + } +} + +func TestRestoreDrillRejectsEmptyBackup(t *testing.T) { + ctx := context.Background() + dir := t.TempDir() + backupPath := filepath.Join(dir, "empty.db") + if err := os.WriteFile(backupPath, nil, 0o600); err != nil { + t.Fatal(err) + } + if _, err := RestoreDrill(ctx, backupPath, filepath.Join(dir, "live.db"), ""); err == nil || !strings.Contains(err.Error(), "non-empty") { + t.Fatalf("empty backup error = %v", err) + } +} + +func TestRestoreDrillRejectsHeaderOnlyDatabase(t *testing.T) { + ctx := context.Background() + dir := t.TempDir() + backupPath := filepath.Join(dir, "header-only.db") + if err := os.WriteFile(backupPath, []byte("SQLite format 3\x00"), 0o600); err != nil { + t.Fatal(err) + } + if _, err := RestoreDrill(ctx, backupPath, filepath.Join(dir, "live.db"), ""); err == nil { + t.Fatal("header-only backup was accepted") + } +} + +func TestRestoreDrillRejectsMissingReservationUniqueness(t *testing.T) { + ctx := context.Background() + dir := t.TempDir() + livePath := filepath.Join(dir, "live.db") + db, err := Open(ctx, livePath) + if err != nil { + t.Fatal(err) + } + if _, err := db.SQL.ExecContext(ctx, `DROP INDEX uq_active_payable`); err != nil { + db.Close() + t.Fatal(err) + } + backupPath := filepath.Join(dir, "backup.db") + if err := db.BackupTo(ctx, backupPath); err != nil { + db.Close() + t.Fatal(err) + } + if err := db.Close(); err != nil { + t.Fatal(err) + } + if _, err := RestoreDrill(ctx, backupPath, livePath, ""); err == nil || !strings.Contains(err.Error(), "uq_active_payable") { + t.Fatalf("missing uniqueness index error = %v", err) + } +} + +func TestRestoreDrillRejectsMissingGlobalAmountTrigger(t *testing.T) { + ctx := context.Background() + dir := t.TempDir() + livePath := filepath.Join(dir, "live.db") + db, err := Open(ctx, livePath) + if err != nil { + t.Fatal(err) + } + if _, err := db.SQL.ExecContext(ctx, `DROP TRIGGER trg_amount_reservations_global_unique_insert`); err != nil { + db.Close() + t.Fatal(err) + } + backupPath := filepath.Join(dir, "backup.db") + if err := db.BackupTo(ctx, backupPath); err != nil { + db.Close() + t.Fatal(err) + } + if err := db.Close(); err != nil { + t.Fatal(err) + } + if _, err := RestoreDrill(ctx, backupPath, livePath, ""); err == nil || !strings.Contains(err.Error(), "trg_amount_reservations_global_unique_insert") { + t.Fatalf("missing global amount trigger error = %v", err) + } +} + +func buildV4RestoreFixture(t *testing.T, path string) (*sql.DB, *DB) { + t.Helper() + ctx := context.Background() + raw, err := sql.Open("sqlite", "file:"+filepath.ToSlash(path)) + if err != nil { + t.Fatal(err) + } + if _, err := raw.ExecContext(ctx, `CREATE TABLE schema_migrations(version INTEGER PRIMARY KEY, applied_at INTEGER NOT NULL) STRICT;`); err != nil { + raw.Close() + t.Fatal(err) + } + if _, err := raw.ExecContext(ctx, schemaV1); err != nil { + raw.Close() + t.Fatal(err) + } + if _, err := raw.ExecContext(ctx, schemaV2); err != nil { + raw.Close() + t.Fatal(err) + } + if _, err := raw.ExecContext(ctx, `INSERT INTO schema_migrations(version,applied_at) VALUES(1,1),(2,2)`); err != nil { + raw.Close() + t.Fatal(err) + } + db := &DB{SQL: raw, Path: path} + if err := db.applyV3(ctx); err != nil { + raw.Close() + t.Fatal(err) + } + if err := db.runMigrationTx(ctx, 4, applyV4); err != nil { + raw.Close() + t.Fatal(err) + } + if err := db.ensureMultiRelayCompatibility(ctx); err != nil { + raw.Close() + t.Fatal(err) + } + if err := db.ensureRelayPayloadIntegrity(ctx); err != nil { + raw.Close() + t.Fatal(err) + } + return raw, db +} + +func finishRestoreFixture(t *testing.T, raw *sql.DB, path string) { + t.Helper() + if err := raw.Close(); err != nil { + t.Fatal(err) + } + if err := os.Chmod(path, restoreFileMode); err != nil { + t.Fatal(err) + } +} + +func installTransitionalRestoreState(ctx context.Context, db *DB, version int) error { + tx, err := db.SQL.BeginTx(ctx, nil) + if err != nil { + return err + } + defer tx.Rollback() + if _, err := tx.ExecContext(ctx, `ALTER TABLE amount_reservations ADD COLUMN global_unique_enforced INTEGER NOT NULL DEFAULT 1 CHECK(global_unique_enforced IN (0,1))`); err != nil { + return err + } + if version >= 6 { + if _, err := tx.ExecContext(ctx, `ALTER TABLE relay_devices ADD COLUMN epoch_required INTEGER NOT NULL DEFAULT 0 CHECK(epoch_required IN (0,1))`); err != nil { + return err + } + } + if err := installGlobalAmountUniqueness(ctx, tx); err != nil { + return err + } + for migration := 5; migration <= version; migration++ { + if _, err := tx.ExecContext(ctx, `INSERT INTO schema_migrations(version,applied_at) VALUES(?,?)`, migration, migration); err != nil { + return err + } + } + return tx.Commit() +} + +func installHistoricalMarkerAwareTransitionalState(ctx context.Context, db *DB, version int) error { + tx, err := db.SQL.BeginTx(ctx, nil) + if err != nil { + return err + } + defer tx.Rollback() + if _, err := tx.ExecContext(ctx, `ALTER TABLE amount_reservations ADD COLUMN global_unique_enforced INTEGER NOT NULL DEFAULT 1 CHECK(global_unique_enforced IN (0,1))`); err != nil { + return err + } + if version >= 6 { + if _, err := tx.ExecContext(ctx, `ALTER TABLE relay_devices ADD COLUMN epoch_required INTEGER NOT NULL DEFAULT 0 CHECK(epoch_required IN (0,1))`); err != nil { + return err + } + } + if _, err := tx.ExecContext(ctx, ` +DROP INDEX IF EXISTS uq_active_profile_payable; +CREATE UNIQUE INDEX uq_active_payable ON amount_reservations(payable_amount_paise) + WHERE released_at IS NULL AND global_unique_enforced=1; +CREATE TRIGGER trg_amount_reservations_global_unique_insert +BEFORE INSERT ON amount_reservations +WHEN NEW.global_unique_enforced<>1 OR ( + NEW.released_at IS NULL AND EXISTS ( + SELECT 1 FROM amount_reservations + WHERE released_at IS NULL AND payable_amount_paise=NEW.payable_amount_paise + ) +) +BEGIN + SELECT RAISE(ABORT, 'active payable amount already reserved'); +END; +CREATE TRIGGER trg_amount_reservations_global_unique_update +BEFORE UPDATE OF payable_amount_paise,released_at,global_unique_enforced ON amount_reservations +WHEN (OLD.global_unique_enforced=1 AND NEW.global_unique_enforced<>1) OR ( + NEW.released_at IS NULL + AND (NEW.global_unique_enforced=1 OR OLD.released_at IS NOT NULL OR NEW.payable_amount_paise<>OLD.payable_amount_paise) + AND EXISTS ( + SELECT 1 FROM amount_reservations + WHERE id<>NEW.id AND released_at IS NULL AND payable_amount_paise=NEW.payable_amount_paise + ) +) +BEGIN + SELECT RAISE(ABORT, 'active payable amount already reserved'); +END;`); err != nil { + return err + } + for migration := 5; migration <= version; migration++ { + if _, err := tx.ExecContext(ctx, `INSERT INTO schema_migrations(version,applied_at) VALUES(?,?)`, migration, migration); err != nil { + return err + } + } + return tx.Commit() +} + +func TestRestoreDrillAcceptsHistoricalMarkerAwareTransitionalMigrations(t *testing.T) { + for _, version := range []int{5, 7} { + t.Run(fmt.Sprintf("v%d", version), func(t *testing.T) { + ctx := context.Background() + dir := t.TempDir() + backupPath := filepath.Join(dir, "backup.db") + raw, db := buildV4RestoreFixture(t, backupPath) + if err := installHistoricalMarkerAwareTransitionalState(ctx, db, version); err != nil { + raw.Close() + t.Fatal(err) + } + finishRestoreFixture(t, raw, backupPath) + report, err := RestoreDrill(ctx, backupPath, filepath.Join(dir, "live.db"), "") + if err != nil { + t.Fatalf("restore historical marker-aware v%d backup: %v", version, err) + } + if report.SchemaVersion != schemaVersion { + t.Fatalf("restored schema=%d want=%d", report.SchemaVersion, schemaVersion) + } + }) + } +} + +func TestRestoreDrillAcceptsMarkerAwareInterruptedMigrations(t *testing.T) { + for _, version := range []int{5, 6, 7} { + t.Run(fmt.Sprintf("v%d", version), func(t *testing.T) { + ctx := context.Background() + dir := t.TempDir() + backupPath := filepath.Join(dir, "backup.db") + raw, db := buildV4RestoreFixture(t, backupPath) + if err := installTransitionalRestoreState(ctx, db, version); err != nil { + raw.Close() + t.Fatal(err) + } + finishRestoreFixture(t, raw, backupPath) + report, err := RestoreDrill(ctx, backupPath, filepath.Join(dir, "live.db"), "") + if err != nil { + t.Fatalf("restore marker-aware v%d backup: %v", version, err) + } + if report.SchemaVersion != schemaVersion { + t.Fatalf("restored schema=%d want=%d", report.SchemaVersion, schemaVersion) + } + }) + } +} + +func TestRestoreDrillAcceptsLegacyV5GlobalIndexWithoutMarker(t *testing.T) { + ctx := context.Background() + dir := t.TempDir() + backupPath := filepath.Join(dir, "backup.db") + raw, db := buildV4RestoreFixture(t, backupPath) + tx, err := raw.BeginTx(ctx, nil) + if err != nil { + raw.Close() + t.Fatal(err) + } + if _, err := tx.ExecContext(ctx, `CREATE UNIQUE INDEX uq_active_payable ON amount_reservations(payable_amount_paise) WHERE released_at IS NULL`); err != nil { + tx.Rollback() + raw.Close() + t.Fatal(err) + } + if _, err := tx.ExecContext(ctx, `INSERT INTO schema_migrations(version,applied_at) VALUES(5,5)`); err != nil { + tx.Rollback() + raw.Close() + t.Fatal(err) + } + if err := tx.Commit(); err != nil { + raw.Close() + t.Fatal(err) + } + _ = db + finishRestoreFixture(t, raw, backupPath) + report, err := RestoreDrill(ctx, backupPath, filepath.Join(dir, "live.db"), "") + if err != nil { + t.Fatalf("legacy v5 restore failed: %v", err) + } + if report.SchemaVersion != schemaVersion { + t.Fatalf("restored schema=%d want=%d", report.SchemaVersion, schemaVersion) + } +} diff --git a/internal/v4/storage/schema.go b/internal/v4/storage/schema.go index 0fba5a2..c816489 100644 --- a/internal/v4/storage/schema.go +++ b/internal/v4/storage/schema.go @@ -4,8 +4,14 @@ import ( "context" "database/sql" "fmt" + "strings" ) +type schemaQueryer interface { + QueryContext(context.Context, string, ...any) (*sql.Rows, error) + QueryRowContext(context.Context, string, ...any) *sql.Row +} + func (db *DB) migrate(ctx context.Context) error { if _, err := db.SQL.ExecContext(ctx, ` CREATE TABLE IF NOT EXISTS schema_migrations ( @@ -15,12 +21,21 @@ CREATE TABLE IF NOT EXISTS schema_migrations ( `); err != nil { return fmt.Errorf("create schema_migrations: %w", err) } - var current int - if err := db.SQL.QueryRowContext(ctx, `SELECT COALESCE(MAX(version), 0) FROM schema_migrations`).Scan(¤t); err != nil { - return fmt.Errorf("read schema version: %w", err) + versions, err := readSchemaVersions(ctx, db.SQL) + if err != nil { + return err } - if current > schemaVersion { - return fmt.Errorf("database schema %d is newer than supported %d", current, schemaVersion) + for index, version := range versions { + if version != index+1 { + return fmt.Errorf("schema migrations are not contiguous at version %d", version) + } + if version > 7 { + return fmt.Errorf("database schema %d is newer than supported transitional schema 7", version) + } + } + current := 0 + if len(versions) > 0 { + current = versions[len(versions)-1] } if current < 1 { if err := db.runMigrationTx(ctx, 1, applyV1); err != nil { @@ -44,7 +59,145 @@ CREATE TABLE IF NOT EXISTS schema_migrations ( if err := db.runMigrationTx(ctx, 4, applyV4); err != nil { return err } - current = 4 + } + return db.reconcileCompatibility(ctx) +} + +func readSchemaVersions(ctx context.Context, queryer schemaQueryer) ([]int, error) { + rows, err := queryer.QueryContext(ctx, `SELECT version FROM schema_migrations ORDER BY version`) + if err != nil { + return nil, fmt.Errorf("read schema migrations: %w", err) + } + defer rows.Close() + var versions []int + for rows.Next() { + var version int + if err := rows.Scan(&version); err != nil { + return nil, fmt.Errorf("scan schema migration: %w", err) + } + versions = append(versions, version) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("iterate schema migrations: %w", err) + } + return versions, nil +} + +func (db *DB) reconcileCompatibility(ctx context.Context) error { + return db.WithImmediateTx(ctx, func(tx *ImmediateTx) error { + versions, err := readSchemaVersions(ctx, tx) + if err != nil { + return err + } + for _, version := range versions { + if version > 7 { + return fmt.Errorf("database schema %d is newer than supported transitional schema 7", version) + } + } + if err := reconcileCompatibilityTx(ctx, tx); err != nil { + return err + } + if _, err := tx.ExecContext(ctx, `DELETE FROM schema_migrations WHERE version > 4`); err != nil { + return fmt.Errorf("remove transitional schema ledgers: %w", err) + } + return nil + }) +} + +func reconcileCompatibilityTx(ctx context.Context, tx schemaQueryer) error { + for _, statement := range []string{ + `DROP INDEX IF EXISTS uq_active_payable`, + `DROP TRIGGER IF EXISTS trg_amount_reservations_global_unique_insert`, + `DROP TRIGGER IF EXISTS trg_amount_reservations_global_unique_update`, + `DROP INDEX IF EXISTS uq_active_profile_payable`, + } { + if _, err := tx.(interface { + ExecContext(context.Context, string, ...any) (sql.Result, error) + }).ExecContext(ctx, statement); err != nil { + return fmt.Errorf("remove compatibility object: %w", err) + } + } + exec := tx.(interface { + ExecContext(context.Context, string, ...any) (sql.Result, error) + }) + hasMarker, err := tableColumnExists(ctx, tx, "amount_reservations", "global_unique_enforced") + if err != nil { + return err + } + if !hasMarker { + if _, err := exec.ExecContext(ctx, `ALTER TABLE amount_reservations ADD COLUMN global_unique_enforced INTEGER NOT NULL DEFAULT 1 CHECK(global_unique_enforced IN (0,1))`); err != nil { + return fmt.Errorf("add global amount enforcement marker: %w", err) + } + } else if err := validateCompatibilityColumn(ctx, tx, "amount_reservations", "global_unique_enforced", "1", "GLOBAL_UNIQUE_ENFORCED IN (0,1)"); err != nil { + return err + } + hasEpoch, err := tableColumnExists(ctx, tx, "relay_devices", "epoch_required") + if err != nil { + return err + } + if !hasEpoch { + if _, err := exec.ExecContext(ctx, `ALTER TABLE relay_devices ADD COLUMN epoch_required INTEGER NOT NULL DEFAULT 0 CHECK(epoch_required IN (0,1))`); err != nil { + return fmt.Errorf("add relay epoch requirement: %w", err) + } + } else if err := validateCompatibilityColumn(ctx, tx, "relay_devices", "epoch_required", "0", "EPOCH_REQUIRED IN (0,1)"); err != nil { + return err + } + if _, err := exec.ExecContext(ctx, `UPDATE amount_reservations SET global_unique_enforced = CASE + WHEN released_at IS NULL AND payable_amount_paise IN ( + SELECT payable_amount_paise FROM amount_reservations + WHERE released_at IS NULL GROUP BY payable_amount_paise HAVING COUNT(*) > 1 + ) THEN 0 ELSE 1 END`); err != nil { + return fmt.Errorf("grandfather pre-compatibility overlapping amount reservations: %w", err) + } + if _, err := exec.ExecContext(ctx, `CREATE UNIQUE INDEX uq_active_profile_payable + ON amount_reservations(collection_profile_id, payable_amount_paise) WHERE released_at IS NULL`); err != nil { + return fmt.Errorf("recreate profile amount uniqueness: %w", err) + } + if err := installGlobalAmountUniqueness(ctx, exec); err != nil { + return err + } + return nil +} + +func validateCompatibilityColumn(ctx context.Context, queryer schemaQueryer, tableName, columnName, wantDefault, checkFragment string) error { + rows, err := queryer.QueryContext(ctx, fmt.Sprintf("PRAGMA table_info(%s)", tableName)) + if err != nil { + return fmt.Errorf("inspect %s columns: %w", tableName, err) + } + defer rows.Close() + var found bool + var columnType string + var notNull int + var defaultValue sql.NullString + for rows.Next() { + var cid, primaryKey int + var name, typ string + var value sql.NullString + if err := rows.Scan(&cid, &name, &typ, ¬Null, &value, &primaryKey); err != nil { + return fmt.Errorf("scan %s columns: %w", tableName, err) + } + if name == columnName { + found = true + columnType = typ + defaultValue = value + } + } + if err := rows.Err(); err != nil { + return fmt.Errorf("iterate %s columns: %w", tableName, err) + } + if !found { + return fmt.Errorf("%s column %s is missing", tableName, columnName) + } + if !strings.EqualFold(strings.TrimSpace(columnType), "INTEGER") || notNull == 0 || + strings.Trim(strings.TrimSpace(defaultValue.String), "'\"") != wantDefault { + return fmt.Errorf("%s column %s has incompatible definition", tableName, columnName) + } + var createSQL string + if err := queryer.QueryRowContext(ctx, `SELECT sql FROM sqlite_master WHERE type='table' AND name=?`, tableName).Scan(&createSQL); err != nil { + return fmt.Errorf("read %s definition: %w", tableName, err) + } + if !strings.Contains(strings.ToUpper(createSQL), checkFragment) { + return fmt.Errorf("%s column %s is missing check constraint", tableName, columnName) } return nil } @@ -131,16 +284,88 @@ ALTER TABLE collection_profiles_v3 RENAME TO collection_profiles; CREATE UNIQUE INDEX uq_collection_profiles_one_active ON collection_profiles(active) WHERE active = 1; ` -// ensureMultiRelayCompatibility removes the historical singleton relay index -// without advancing schema_migrations. The change is backwards-compatible with -// the v4 runtime, so an image rollback can still open the database. +// ensureMultiRelayCompatibility removes the historical singleton relay index. func (db *DB) ensureMultiRelayCompatibility(ctx context.Context) error { - if _, err := db.SQL.ExecContext(ctx, `DROP INDEX IF EXISTS uq_relay_devices_one_enabled;`); err != nil { + if _, err := db.SQL.ExecContext(ctx, `DROP INDEX IF EXISTS uq_relay_devices_one_enabled`); err != nil { return fmt.Errorf("enable multi-relay compatibility: %w", err) } return nil } +func (db *DB) ensureRelayPayloadIntegrity(ctx context.Context) error { + return db.WithImmediateTx(ctx, func(tx *ImmediateTx) error { + rows, err := tx.QueryContext(ctx, `PRAGMA table_info(relay_events)`) + if err != nil { + return fmt.Errorf("inspect relay event columns: %w", err) + } + hasPayloadHash := false + for rows.Next() { + var cid, notNull, primaryKey int + var name, columnType string + var defaultValue sql.NullString + if err := rows.Scan(&cid, &name, &columnType, ¬Null, &defaultValue, &primaryKey); err != nil { + rows.Close() + return fmt.Errorf("scan relay event columns: %w", err) + } + if name == "payload_hash" { + hasPayloadHash = true + } + } + if err := rows.Err(); err != nil { + rows.Close() + return fmt.Errorf("iterate relay event columns: %w", err) + } + if err := rows.Close(); err != nil { + return fmt.Errorf("close relay event columns: %w", err) + } + if !hasPayloadHash { + if _, err := tx.ExecContext(ctx, `ALTER TABLE relay_events ADD COLUMN payload_hash BLOB`); err != nil { + return fmt.Errorf("add relay event payload hash: %w", err) + } + } + if _, err := tx.ExecContext(ctx, relayPayloadIntegritySQL); err != nil { + return fmt.Errorf("create payment reservation consistency triggers: %w", err) + } + return nil + }) +} + +const relayPayloadIntegritySQL = ` +CREATE TRIGGER IF NOT EXISTS amount_reservations_payment_consistency_insert +BEFORE INSERT ON amount_reservations +WHEN NOT EXISTS ( + SELECT 1 FROM payments + WHERE id = NEW.payment_id + AND collection_profile_id = NEW.collection_profile_id + AND payable_amount_paise = NEW.payable_amount_paise +) +BEGIN + SELECT RAISE(ABORT, 'amount reservation does not match payment'); +END; +CREATE TRIGGER IF NOT EXISTS amount_reservations_payment_consistency_update +BEFORE UPDATE OF payment_id,collection_profile_id,payable_amount_paise ON amount_reservations +WHEN NOT EXISTS ( + SELECT 1 FROM payments + WHERE id = NEW.payment_id + AND collection_profile_id = NEW.collection_profile_id + AND payable_amount_paise = NEW.payable_amount_paise +) +BEGIN + SELECT RAISE(ABORT, 'amount reservation does not match payment'); +END; +CREATE TRIGGER IF NOT EXISTS payments_reservation_consistency_update +BEFORE UPDATE OF collection_profile_id,payable_amount_paise ON payments +WHEN EXISTS ( + SELECT 1 FROM amount_reservations + WHERE payment_id = NEW.id + AND (collection_profile_id <> NEW.collection_profile_id + OR payable_amount_paise <> NEW.payable_amount_paise) +) +BEGIN + SELECT RAISE(ABORT, 'payment does not match amount reservation'); +END; +` + func applyV4(ctx context.Context, tx *sql.Tx) error { if _, err := tx.ExecContext(ctx, schemaV4); err != nil { return fmt.Errorf("apply schema v4: %w", err) @@ -148,6 +373,71 @@ func applyV4(ctx context.Context, tx *sql.Tx) error { return nil } +func tableColumnExists(ctx context.Context, queryer schemaQueryer, tableName, columnName string) (bool, error) { + rows, err := queryer.QueryContext(ctx, fmt.Sprintf("PRAGMA table_info(%s)", tableName)) + if err != nil { + return false, fmt.Errorf("inspect %s columns: %w", tableName, err) + } + defer rows.Close() + for rows.Next() { + var cid, notNull, pk int + var name, kind string + var defaultValue sql.NullString + if err := rows.Scan(&cid, &name, &kind, ¬Null, &defaultValue, &pk); err != nil { + return false, fmt.Errorf("scan %s columns: %w", tableName, err) + } + if name == columnName { + return true, nil + } + } + if err := rows.Err(); err != nil { + return false, fmt.Errorf("iterate %s columns: %w", tableName, err) + } + return false, nil +} + +func installGlobalAmountUniqueness(ctx context.Context, exec interface { + ExecContext(context.Context, string, ...any) (sql.Result, error) +}) error { + if _, err := exec.ExecContext(ctx, ` +DROP INDEX IF EXISTS uq_active_payable; +DROP TRIGGER IF EXISTS trg_amount_reservations_global_unique_insert; +DROP TRIGGER IF EXISTS trg_amount_reservations_global_unique_update; +CREATE UNIQUE INDEX uq_active_payable ON amount_reservations(payable_amount_paise) + WHERE released_at IS NULL AND global_unique_enforced=1; +CREATE TRIGGER trg_amount_reservations_global_unique_insert +BEFORE INSERT ON amount_reservations +WHEN NEW.global_unique_enforced<>1 OR ( + NEW.released_at IS NULL AND EXISTS ( + SELECT 1 FROM amount_reservations + WHERE released_at IS NULL AND payable_amount_paise=NEW.payable_amount_paise + ) +) +BEGIN + SELECT RAISE(ABORT, 'active payable amount already reserved'); +END; +CREATE TRIGGER trg_amount_reservations_global_unique_update +BEFORE UPDATE OF payable_amount_paise,released_at,global_unique_enforced ON amount_reservations +WHEN NEW.global_unique_enforced NOT IN (0,1) + OR (OLD.global_unique_enforced=1 AND NEW.global_unique_enforced<>1) + OR (NEW.global_unique_enforced=0 AND NEW.released_at IS NULL + AND (OLD.released_at IS NOT NULL OR NEW.payable_amount_paise<>OLD.payable_amount_paise)) + OR ( + NEW.released_at IS NULL AND NEW.global_unique_enforced=1 + AND EXISTS ( + SELECT 1 FROM amount_reservations + WHERE id<>NEW.id AND released_at IS NULL AND payable_amount_paise=NEW.payable_amount_paise + ) + ) +BEGIN + SELECT RAISE(ABORT, 'active payable amount already reserved'); +END; +`); err != nil { + return fmt.Errorf("install global amount uniqueness: %w", err) + } + return nil +} + const schemaV4 = ` CREATE TABLE payment_observations_v4 ( id TEXT PRIMARY KEY, @@ -179,6 +469,7 @@ func applyV2(ctx context.Context, tx *sql.Tx) error { } const schemaV2 = ` + ALTER TABLE relay_devices ADD COLUMN notification_access INTEGER CHECK(notification_access IS NULL OR notification_access IN (0,1)); ALTER TABLE relay_devices ADD COLUMN listener_connected INTEGER CHECK(listener_connected IS NULL OR listener_connected IN (0,1)); ALTER TABLE relay_devices ADD COLUMN battery_optimization_exempt INTEGER CHECK(battery_optimization_exempt IS NULL OR battery_optimization_exempt IN (0,1)); diff --git a/internal/v4/webhooks/service.go b/internal/v4/webhooks/service.go index 577b949..c41b7ff 100644 --- a/internal/v4/webhooks/service.go +++ b/internal/v4/webhooks/service.go @@ -8,11 +8,13 @@ import ( "errors" "fmt" "io" + "net" "net/http" "net/url" "strconv" "strings" "sync" + "syscall" "time" "github.com/Phloraxx/payment-api/internal/v4/storage" @@ -30,6 +32,9 @@ type Config struct { Endpoint string Secret string AllowInsecureHTTP bool + // allowPrivateNetwork is test-only and intentionally not configurable by + // runtime settings or environment variables. + allowPrivateNetwork bool } type Service struct { DB *storage.DB @@ -53,7 +58,7 @@ type Delivery struct { func NewService(db *storage.DB, cfg Config) *Service { return &Service{ - DB: db, config: cfg, HTTPClient: newHTTPClient(), Now: time.Now, + DB: db, config: cfg, HTTPClient: newHTTPClient(cfg.allowPrivateNetwork), Now: time.Now, MaxAttempts: defaultMaxAttempts, BatchSize: defaultBatchSize, Lease: defaultLease, wake: make(chan struct{}, 1), } @@ -264,13 +269,21 @@ func (s *Service) RetryOne(ctx context.Context, id string) error { if id == "" { return errors.New("webhook id is required") } - result, err := s.DB.SQL.ExecContext(ctx, `UPDATE webhook_deliveries - SET status='pending',attempts=0,next_attempt_at=?,last_http_status=NULL,last_error=NULL,delivered_at=NULL - WHERE id=? AND status IN ('retry','exhausted')`, s.now().UnixMilli(), id) + var rowsAffected int64 + err := s.DB.WithImmediateTx(ctx, func(tx *storage.ImmediateTx) error { + result, err := tx.ExecContext(ctx, `UPDATE webhook_deliveries + SET status='pending',attempts=0,next_attempt_at=?,last_http_status=NULL,last_error=NULL,delivered_at=NULL + WHERE id=? AND status IN ('retry','exhausted')`, s.now().UnixMilli(), id) + if err != nil { + return fmt.Errorf("retry webhook: %w", err) + } + rowsAffected, _ = result.RowsAffected() + return nil + }) if err != nil { - return fmt.Errorf("retry webhook: %w", err) + return err } - if rows, _ := result.RowsAffected(); rows != 1 { + if rowsAffected != 1 { return errors.New("webhook is not retryable") } s.Wake() @@ -297,6 +310,11 @@ func ValidateConfig(cfg Config) error { if u.Scheme != "https" && !(cfg.AllowInsecureHTTP && u.Scheme == "http") { return fmt.Errorf("%w: HTTPS endpoint is required", ErrInvalidConfig) } + if !cfg.allowPrivateNetwork { + if ip := net.ParseIP(u.Hostname()); ip != nil && restrictedIP(ip) { + return fmt.Errorf("%w: endpoint must not target a private or local address", ErrInvalidConfig) + } + } return nil } @@ -340,19 +358,70 @@ func (s *Service) now() time.Time { func (s *Service) client() *http.Client { if s.HTTPClient == nil { - return newHTTPClient() + cfg := s.ConfigSnapshot() + return newHTTPClient(cfg.allowPrivateNetwork) } return s.HTTPClient } -func newHTTPClient() *http.Client { + +func newHTTPClient(allowPrivateNetwork bool) *http.Client { return &http.Client{ Timeout: 10 * time.Second, + Transport: &http.Transport{ + Proxy: nil, + DialContext: restrictedDialer(allowPrivateNetwork), + }, CheckRedirect: func(_ *http.Request, _ []*http.Request) error { return http.ErrUseLastResponse }, } } +func restrictedDialer(allowPrivateNetwork bool) func(context.Context, string, string) (net.Conn, error) { + dialer := &net.Dialer{Timeout: 10 * time.Second, KeepAlive: 30 * time.Second} + if !allowPrivateNetwork { + dialer.ControlContext = func(_ context.Context, network, address string, _ syscall.RawConn) error { + if network != "tcp" && network != "tcp4" && network != "tcp6" { + return fmt.Errorf("unsupported webhook network %q", network) + } + host, _, err := net.SplitHostPort(address) + if err != nil { + return fmt.Errorf("parse webhook destination %q: %w", address, err) + } + ip := net.ParseIP(strings.Trim(host, "[]")) + if ip == nil { + return errors.New("webhook destination did not resolve to an IP address") + } + if restrictedIP(ip) { + return errors.New("webhook destination resolves to a private address") + } + return nil + } + } + return dialer.DialContext +} + +func restrictedIP(ip net.IP) bool { + if ip == nil { + return true + } + if v4 := ip.To4(); v4 != nil { + return v4[0] == 0 || + (v4[0] == 10) || + (v4[0] == 100 && v4[1] >= 64 && v4[1] <= 127) || + (v4[0] == 127) || + (v4[0] == 169 && v4[1] == 254) || + (v4[0] == 172 && v4[1] >= 16 && v4[1] <= 31) || + (v4[0] == 192 && (v4[1] == 0 || v4[1] == 168)) || + (v4[0] == 198 && (v4[1] == 18 || v4[1] == 19)) || + (v4[0] >= 224) || + (v4[0] == 255 && v4[1] == 255) + } + return ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || + ip.IsUnspecified() || ip.IsMulticast() || + (ip[0] == 0xfe && ip[1]&0xc0 == 0xc0) +} + func nullableStatus(status int) any { if status == 0 { return nil diff --git a/internal/v4/webhooks/service_test.go b/internal/v4/webhooks/service_test.go index 67c748b..a8d4614 100644 --- a/internal/v4/webhooks/service_test.go +++ b/internal/v4/webhooks/service_test.go @@ -6,6 +6,7 @@ import ( "encoding/json" "errors" "io" + "net" "net/http" "net/http/httptest" "path/filepath" @@ -58,7 +59,10 @@ func newWebhookFixture(t *testing.T) webhookFixture { } func newTestService(f webhookFixture, endpoint string) *Service { - s := NewService(f.db, Config{Endpoint: endpoint, Secret: testSecret, AllowInsecureHTTP: true}) + s := NewService(f.db, Config{ + Endpoint: endpoint, Secret: testSecret, AllowInsecureHTTP: true, + allowPrivateNetwork: true, + }) s.Now = func() time.Time { return *f.now } return s } @@ -243,8 +247,11 @@ func TestConfigurationValidation(t *testing.T) { cases := []Config{ {Endpoint: "http://example.com/hook", Secret: testSecret}, {Endpoint: "https://user:pass@example.com/hook", Secret: testSecret}, + {Endpoint: "http://127.0.0.1/hook", Secret: testSecret, AllowInsecureHTTP: true}, {Endpoint: "https://example.com/hook?token=x", Secret: testSecret}, {Endpoint: "https://example.com/hook", Secret: "short"}, + {Endpoint: "https://127.0.0.1/hook", Secret: testSecret}, + {Endpoint: "https://[::1]/hook", Secret: testSecret}, } for _, cfg := range cases { s := NewService(f.db, cfg) @@ -253,3 +260,25 @@ func TestConfigurationValidation(t *testing.T) { } } } + +func TestRestrictedDialerRejectsLocalDestination(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer listener.Close() + + conn, err := restrictedDialer(false)(context.Background(), "tcp", listener.Addr().String()) + if err == nil { + conn.Close() + t.Fatal("restricted dialer connected to loopback") + } +} + +func TestRestrictedIPRejectsSpecialRanges(t *testing.T) { + for _, value := range []string{"0.0.0.1", "100.64.0.1", "192.0.0.1", "198.18.0.1", "225.1.2.3", "239.255.255.255", "255.255.255.255", "fec0::1"} { + if !restrictedIP(net.ParseIP(value)) { + t.Fatalf("restrictedIP(%q) = false", value) + } + } +} diff --git a/internal/v4web/dist/index.html b/internal/v4web/dist/index.html index a8d6a03..965468f 100644 --- a/internal/v4web/dist/index.html +++ b/internal/v4web/dist/index.html @@ -8,8 +8,8 @@ PayGate · Payments control plane - - + +
diff --git a/web-v4/src/App.tsx b/web-v4/src/App.tsx index 2e43c25..6589ca3 100644 --- a/web-v4/src/App.tsx +++ b/web-v4/src/App.tsx @@ -10,7 +10,7 @@ import { Spinner, cx } from "./ui"; type Tab = "overview" | "payments" | "activity" | "settings"; const tabs: Array<{ id: Tab; label: string }> = [ - { id: "overview", label: "Overview" }, + { id: "overview", label: "Dashboard" }, { id: "payments", label: "Payments" }, { id: "activity", label: "Activity" }, { id: "settings", label: "Settings" }, @@ -52,13 +52,13 @@ export function App() {
Live - Secure admin session + PGAdmin
{tab === "overview" && } - {tab === "payments" && setOpenPaymentId(undefined)} />} + {tab === "payments" && setOpenPaymentId(undefined)} onOpenSettings={() => setTab("settings")} onOpenActivity={() => setTab("activity")} />} {tab === "activity" && } {tab === "settings" && setSession("out")} />}
diff --git a/web-v4/src/Brand.tsx b/web-v4/src/Brand.tsx index aebc7ef..a0966ef 100644 --- a/web-v4/src/Brand.tsx +++ b/web-v4/src/Brand.tsx @@ -1,3 +1,5 @@ +import { useId } from "react"; + export function Brand({ compact = false, subtitle }: { compact?: boolean; subtitle?: string }) { return
@@ -6,13 +8,16 @@ export function Brand({ compact = false, subtitle }: { compact?: boolean; subtit } export function PayGateMark({ className = "" }: { className?: string }) { + const uid = useId().replace(/:/g, ""); + const markGradient = `pg-mark-gradient-${uid}`; + const barGradient = `pg-bar-gradient-${uid}`; return ; } diff --git a/web-v4/src/PaymentsPage.tsx b/web-v4/src/PaymentsPage.tsx index b2649e2..4efdc15 100644 --- a/web-v4/src/PaymentsPage.tsx +++ b/web-v4/src/PaymentsPage.tsx @@ -1,65 +1,226 @@ -import { useCallback, useEffect, useMemo, useState } from "react"; -import { ApiError, editPayment, getPayment, getProfiles, listPayments, retryWebhook } from "./api"; -import type { Payment, PaymentDetail, PaymentStatus, Profile } from "./types"; -import { Badge, Empty, ErrorNotice, Modal, SectionHead, Spinner, dateTime, money } from "./ui"; - -const PAGE_SIZE = 50; -const statuses: Array<{ id: "" | PaymentStatus; label: string }> = [ - { id: "", label: "All statuses" }, { id: "pending", label: "Pending" }, { id: "paid", label: "Paid" }, - { id: "expired", label: "Expired" }, { id: "cancelled", label: "Cancelled" }, +import { useCallback, useEffect, useMemo, useRef, useState } from "react"; +import { ApiError, editPayment, getDevices, getOverview, getPayment, getProfiles, listPayments, retryWebhook } from "./api"; +import type { DeviceInfo, Overview, Payment, PaymentDetail, PaymentStatus, Profile } from "./types"; +import { Badge, Empty, ErrorNotice, Modal, SectionHead, Spinner, dateTime, money, relativeTime } from "./ui"; + +const PAGE_SIZE = 25; +const statusFilters: Array<{ id: "" | PaymentStatus; label: string }> = [ + { id: "", label: "All" }, + { id: "paid", label: "Paid" }, + { id: "pending", label: "Pending" }, + { id: "expired", label: "Expired" }, + { id: "cancelled", label: "Cancelled" }, ]; -export function PaymentsPage({ initialPaymentId, onInitialConsumed }: { initialPaymentId?: string; onInitialConsumed: () => void }) { +export function PaymentsPage({ + initialPaymentId, + onInitialConsumed, + onOpenSettings, + onOpenActivity, +}: { + initialPaymentId?: string; + onInitialConsumed: () => void; + onOpenSettings: () => void; + onOpenActivity: () => void; +}) { const [items, setItems] = useState([]); const [total, setTotal] = useState(0); const [profiles, setProfiles] = useState([]); + const [overview, setOverview] = useState(); + const [devices, setDevices] = useState([]); const [q, setQ] = useState(""); const [query, setQuery] = useState(""); - const [status, setStatus] = useState(""); + const [status, setStatus] = useState<"" | PaymentStatus>(""); const [profile, setProfile] = useState(""); const [offset, setOffset] = useState(0); const [loading, setLoading] = useState(true); + const [refreshing, setRefreshing] = useState(false); const [error, setError] = useState(""); + const [operationsError, setOperationsError] = useState(""); const [selected, setSelected] = useState(); - const load = useCallback(async () => { - setLoading(true); setError(""); + const paymentRequest = useRef(0); + const operationsRequest = useRef(0); + + const loadPayments = useCallback(async (silent = false) => { + const request = ++paymentRequest.current; + if (silent) setRefreshing(true); else setLoading(true); + setError(""); try { - const [result, profileItems] = await Promise.all([ - listPayments({ q: query, status, profile, limit: PAGE_SIZE, offset }), getProfiles(), - ]); - setItems(result.items); setTotal(result.total); setProfiles(profileItems); - } catch (e) { setError(e instanceof ApiError ? e.message : "Could not load payments."); } - finally { setLoading(false); } + const result = await listPayments({ q: query, status, profile, limit: PAGE_SIZE, offset }); + if (request !== paymentRequest.current) return; + setItems(result.items); + setTotal(result.total); + } catch (e) { + if (request !== paymentRequest.current) return; + setError(e instanceof ApiError ? e.message : "Could not load payments."); + } finally { + if (request === paymentRequest.current) { + setLoading(false); + setRefreshing(false); + } + } }, [query, status, profile, offset]); - useEffect(() => { void load(); }, [load]); - useEffect(() => { if (initialPaymentId) { setSelected(initialPaymentId); onInitialConsumed(); } }, [initialPaymentId, onInitialConsumed]); + + const loadOperations = useCallback(async () => { + const request = ++operationsRequest.current; + setOperationsError(""); + const [profileResult, overviewResult, deviceResult] = await Promise.allSettled([getProfiles(), getOverview(), getDevices()]); + if (request !== operationsRequest.current) return; + let failures = 0; + if (profileResult.status === "fulfilled") setProfiles(profileResult.value); else failures++; + if (overviewResult.status === "fulfilled") setOverview(overviewResult.value); else failures++; + if (deviceResult.status === "fulfilled") setDevices(deviceResult.value); else failures++; + if (failures) setOperationsError("Some operational status data could not be refreshed."); + }, []); + + const refreshAll = useCallback(async () => { + await Promise.all([loadPayments(true), loadOperations()]); + }, [loadPayments, loadOperations]); + + useEffect(() => { void loadPayments(); }, [loadPayments]); + useEffect(() => { void loadOperations(); }, [loadOperations]); + useEffect(() => { + const timer = window.setInterval(() => { + if (document.visibilityState === "visible") { + void loadPayments(true); + void loadOperations(); + } + }, 30_000); + return () => window.clearInterval(timer); + }, [loadPayments, loadOperations]); + useEffect(() => { + if (initialPaymentId) { + setSelected(initialPaymentId); + onInitialConsumed(); + } + }, [initialPaymentId, onInitialConsumed]); + const pages = Math.max(1, Math.ceil(total / PAGE_SIZE)); const page = Math.floor(offset / PAGE_SIZE) + 1; + const unhealthyDevices = useMemo(() => devices.filter((device) => device.enabled && !device.operational), [devices]); + const attentionCount = (overview?.unmatched_today ?? 0) + (overview?.expiring_soon ?? 0) + unhealthyDevices.length + (overview?.webhooks.exhausted ?? 0); + + return
+ void refreshAll()}>{refreshing ? "Refreshing…" : "Refresh"}} + /> - return <> - void load()}>Refresh} /> -
-
{ e.preventDefault(); setOffset(0); setQuery(q.trim()); }} className="search-box"> setQ(e.target.value)} placeholder="Search name, payment ID, event ID, payer…"/>
- - -
{error && } -
-
{total.toLocaleString("en-IN")} paymentsPage {page} of {pages}
- {loading && !items.length ?
Loading payments…
: !items.length ? :
- {items.map((payment) => setSelected(payment.id)}> - - - - - - - )} -
PaymentExact amountStatusCollectionCreated
{payment.name}{payment.external_id || payment.id}{money(payment.payable_amount_paise)}requested {money(payment.requested_amount_paise)}{profiles.find((p) => p.id === payment.collection_profile_id)?.label ?? payment.collection_profile_id}{payment.upi_id_snapshot}{dateTime(payment.created_at)}{payment.payer_name || "—"}
} -
{Math.min(offset + 1, total || 0)}–{Math.min(offset + PAGE_SIZE, total)} of {total}
+ {operationsError &&
{operationsError}
} + +
+ + + + +
- {selected && setSelected(undefined)} onChanged={() => void load()} />} - ; + +
+
+
+

Payment activity

Today · {(overview?.payments_today ?? 0).toLocaleString("en-IN")} created
+
{ event.preventDefault(); setOffset(0); setQuery(q.trim()); }}> + + setQ(event.target.value)} placeholder="Search registrant, event or payment ID" /> + {query && } +
+
+ +
+
+ {statusFilters.map((item) => )} +
+ +
+ + {loading && !items.length ?
Loading payments…
: !items.length ? :
+ + + {items.map((payment) => setSelected(payment.id)}> + + + + + + + + + )} +
RegistrantEventRequestedPayableStatusEvidenceTime
{payment.name || "Unnamed payment"}{payment.id}{payment.external_id || "—"}{money(payment.requested_amount_paise)}{money(payment.payable_amount_paise)}
+
} + + {total > PAGE_SIZE &&
+ {Math.min(offset + 1, total)}–{Math.min(offset + PAGE_SIZE, total)} of {total} +
Page {page} of {pages}
+
} +
+ + +
+ + {selected && setSelected(undefined)} onChanged={() => void refreshAll()} />} +
; +} + +function SummaryMetric({ value, label }: { value: string; label: string }) { + return
{value}{label}
; +} + +function AttentionRow({ count, label, detail, action, onClick }: { count: number; label: string; detail: string; action: string; onClick: () => void }) { + if (count <= 0) return null; + return ; +} + +function EvidenceLabel({ payment }: { payment: Payment }) { + if (!payment.evidence_package) return —; + return {paymentAppLabel(payment.evidence_package)}; +} + +function paymentAppLabel(packageName?: string): string { + switch (packageName) { + case "in.org.npci.upiapp": return "BHIM"; + case "com.google.android.apps.nbu.paisa.merchant": + case "com.google.android.apps.nbu.paisa.user": return "Google Pay"; + case "com.paytm.business": return "Paytm Business"; + case "net.one97.paytm": return "Paytm"; + case "com.phonepe.app": return "PhonePe"; + case "in.amazon.mShop.android.shopping": return "Amazon Pay"; + case "money.super.payments": return "super.money"; + case "com.kotak811mobilebankingapp.instantsavingsupiscanandpayrecharge": return "Kotak"; + default: return packageName ? "Payment app" : "—"; + } +} + +function timeOnly(value?: string | null): string { + if (!value) return "—"; + const date = new Date(value); + if (!Number.isFinite(date.getTime())) return "—"; + return new Intl.DateTimeFormat("en-IN", { hour: "numeric", minute: "2-digit" }).format(date); } function PaymentDrawer({ id, profiles, onClose, onChanged }: { id: string; profiles: Profile[]; onClose: () => void; onChanged: () => void }) { @@ -69,35 +230,69 @@ function PaymentDrawer({ id, profiles, onClose, onChanged }: { id: string; profi const [busy, setBusy] = useState(false); const load = useCallback(async () => { setError(""); - try { setDetail(await getPayment(id)); } catch (e) { setError(e instanceof ApiError ? e.message : "Could not load payment."); } + try { setDetail(await getPayment(id)); } + catch (e) { setError(e instanceof ApiError ? e.message : "Could not load payment."); } }, [id]); useEffect(() => { void load(); }, [load]); const payment = detail?.payment; - return - {!detail && !error &&
Loading payment…
} + + return + {!detail && !error &&
Loading payment…
} {error && } {payment && detail && <> -

Exact amount

{money(payment.payable_amount_paise)}{payment.id}
-
- - - p.id === payment.collection_profile_id)?.label ?? payment.collection_profile_id}/> - - - - - -
-
+
{payment.paid_at ? dateTime(payment.paid_at) : dateTime(payment.created_at)}
+
Payment ID{payment.id}
+ + + + + + + + + + + + + + + + + + + + + + + + + + + {payment.internal_note &&
Internal note

{payment.internal_note}

} -
Metadata
{JSON.stringify(payment.metadata ?? {}, null, 2)}
-

Timeline

{detail.history.length ? detail.history.map((item) =>
{item.summary}{item.actor} · {dateTime(item.created_at)}
) :

No history recorded.

}
-

Webhook deliveries

{detail.webhooks.length ? detail.webhooks.map((hook) =>
{hook.event_type}{hook.last_http_status ? `HTTP ${hook.last_http_status} · ` : ""}{hook.attempts} attempt{hook.attempts === 1 ? "" : "s"}{hook.last_error && {hook.last_error}}
{hook.status}{hook.status === "exhausted" && }
) :

No webhook deliveries for this payment.

}
+ +

Timeline

{detail.history.length ? detail.history.map((item) =>
{item.summary}{item.actor} · {dateTime(item.created_at)}
) :

No history recorded.

}
+ +
Webhook deliveries & metadata +
+

Webhook deliveries

{detail.webhooks.length ? detail.webhooks.map((hook) =>
{hook.event_type}{hook.last_http_status ? `HTTP ${hook.last_http_status} · ` : ""}{hook.attempts} attempt{hook.attempts === 1 ? "" : "s"}{hook.last_error && {hook.last_error}}
{hook.status}{hook.status === "exhausted" && }
) :

No webhook deliveries.

}
+
{JSON.stringify(payment.metadata ?? {}, null, 2)}
+
+
+ +
{editing && setEditing(false)} onSaved={async () => { setEditing(false); await load(); onChanged(); }} />} }
; } +function DetailSection({ title, children }: { title: string; children: React.ReactNode }) { + return

{title}

{children}
; +} +function DetailRow({ label, value, strong = false }: { label: string; value: string; strong?: boolean }) { + return
{label}{value}
; +} + function EditPaymentModal({ payment, onClose, onSaved }: { payment: Payment; onClose: () => void; onSaved: () => Promise }) { const [name, setName] = useState(payment.name); const [externalId, setExternalId] = useState(payment.external_id ?? ""); @@ -106,10 +301,13 @@ function EditPaymentModal({ payment, onClose, onSaved }: { payment: Payment; onC const [payerUPI, setPayerUPI] = useState(payment.payer_upi_id ?? ""); const [note, setNote] = useState(payment.internal_note ?? ""); const [metadata, setMetadata] = useState(JSON.stringify(payment.metadata ?? {}, null, 2)); - const [error, setError] = useState(""); const [busy, setBusy] = useState(false); + const [error, setError] = useState(""); + const [busy, setBusy] = useState(false); async function save() { - setError(""); let parsed: Record | null; - try { parsed = metadata.trim() ? JSON.parse(metadata) as Record : {}; } catch { setError("Metadata must be valid JSON."); return; } + setError(""); + let parsed: Record | null; + try { parsed = metadata.trim() ? JSON.parse(metadata) as Record : {}; } + catch { setError("Metadata must be valid JSON."); return; } setBusy(true); try { await editPayment(payment.id, { name: name.trim(), external_id: externalId.trim(), status, payer_name: payerName.trim(), payer_upi_id: payerUPI.trim(), internal_note: note.trim(), metadata: parsed }); @@ -118,19 +316,18 @@ function EditPaymentModal({ payment, onClose, onSaved }: { payment: Payment; onC finally { setBusy(false); } } return
- - + + -
-