From a3dcf6d3f990cd43e7c693f6fb95d3d04e81daa9 Mon Sep 17 00:00:00 2001 From: Phloraxx Date: Tue, 1 Sep 2026 06:00:08 +0000 Subject: [PATCH] feat(v4): expose relay pairing events and heartbeat --- internal/v4/httpapi/relay.go | 152 ++++++++++++++++++++++++++++++ internal/v4/httpapi/relay_test.go | 146 ++++++++++++++++++++++++++++ internal/v4/operator/service.go | 5 +- internal/v4/relay/heartbeat.go | 91 ++++++++++++++++++ internal/v4/relay/pairing.go | 83 ++++++++++++---- internal/v4/storage/db.go | 2 +- internal/v4/storage/db_test.go | 53 +++++++++++ internal/v4/storage/schema.go | 28 ++++++ 8 files changed, 538 insertions(+), 22 deletions(-) create mode 100644 internal/v4/httpapi/relay.go create mode 100644 internal/v4/httpapi/relay_test.go create mode 100644 internal/v4/relay/heartbeat.go diff --git a/internal/v4/httpapi/relay.go b/internal/v4/httpapi/relay.go new file mode 100644 index 0000000..f490c4d --- /dev/null +++ b/internal/v4/httpapi/relay.go @@ -0,0 +1,152 @@ +package httpapi + +import ( + "errors" + "io" + "net/http" + "strings" + + "github.com/Phloraxx/payment-api/internal/v4/relay" +) + +const relayHTTPBodyLimit = 64 << 10 + +const ( + relayDeviceHeader = "X-PayGate-Relay-Device" + relayTimeHeader = "X-PayGate-Relay-Time" + relaySignatureHeader = "X-PayGate-Relay-Signature" +) + +type RelayHandler struct { + Relay *relay.Service + mux *http.ServeMux +} + +func NewRelayHandler(service *relay.Service) *RelayHandler { + h := &RelayHandler{Relay: service, mux: http.NewServeMux()} + h.mux.HandleFunc("POST /api/v4/relay/pair", h.pair) + h.mux.HandleFunc("POST "+relay.EventPath, h.event) + h.mux.HandleFunc("POST "+relay.HeartbeatPath, h.heartbeat) + return h +} +func (h *RelayHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Cache-Control", "no-store") + w.Header().Set("X-Content-Type-Options", "nosniff") + if h == nil || h.Relay == nil || h.mux == nil { + writeError(w, http.StatusServiceUnavailable, "service_unavailable", "PayGate relay is unavailable") + return + } + if r.URL.RawQuery != "" { + writeError(w, http.StatusBadRequest, "invalid_request", "Relay endpoints do not accept query parameters") + return + } + h.mux.ServeHTTP(w, r) +} + +type relayPairRequest struct { + Token string `json:"token"` + Name string `json:"name"` + PublicKeyPEM string `json:"public_key_pem"` + AppVersion string `json:"app_version,omitempty"` + DeviceModel string `json:"device_model,omitempty"` + AndroidVersion string `json:"android_version,omitempty"` +} + +func (h *RelayHandler) pair(w http.ResponseWriter, r *http.Request) { + if !isJSON(r.Header.Get("Content-Type")) { + writeError(w, http.StatusUnsupportedMediaType, "invalid_content_type", "Content-Type must be application/json") + return + } + var input relayPairRequest + if err := decodeStrictJSON(w, r, &input); err != nil { + writeError(w, http.StatusBadRequest, "invalid_request", err.Error()) + return + } + result, err := h.Relay.PairDevice(r.Context(), relay.PairDeviceInput{ + Token: input.Token, Name: input.Name, PublicKeyPEM: input.PublicKeyPEM, + AppVersion: input.AppVersion, DeviceModel: input.DeviceModel, AndroidVersion: input.AndroidVersion, + }) + if err != nil { + writeRelayPairError(w, err) + return + } + writeJSON(w, http.StatusOK, map[string]any{ + "device_id": result.DeviceID, "enabled": result.Enabled, + "replaced_device_id": emptyToNil(result.ReplacedDeviceID), + }) +} + +func (h *RelayHandler) event(w http.ResponseWriter, r *http.Request) { + raw, ok := readRelayBody(w, r) + if !ok { + return + } + result, err := h.Relay.IngestSigned(r.Context(), relayAuth(r, relay.EventPath), raw) + if err != nil { + writeRelayError(w, err) + return + } + writeJSON(w, http.StatusOK, result) +} + +func (h *RelayHandler) heartbeat(w http.ResponseWriter, r *http.Request) { + raw, ok := readRelayBody(w, r) + if !ok { + return + } + result, err := h.Relay.HeartbeatSigned(r.Context(), relayAuth(r, relay.HeartbeatPath), raw) + if err != nil { + writeRelayError(w, err) + return + } + writeJSON(w, http.StatusOK, result) +} + +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, + } +} + +func readRelayBody(w http.ResponseWriter, r *http.Request) ([]byte, bool) { + if !isJSON(r.Header.Get("Content-Type")) { + writeError(w, http.StatusUnsupportedMediaType, "invalid_content_type", "Content-Type must be application/json") + return nil, false + } + r.Body = http.MaxBytesReader(w, r.Body, relayHTTPBodyLimit) + raw, err := io.ReadAll(r.Body) + if err != nil { + writeError(w, http.StatusBadRequest, "invalid_request", "Relay body is too large or unreadable") + return nil, false + } + return raw, true +} +func writeRelayPairError(w http.ResponseWriter, err error) { + 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: + writeError(w, http.StatusInternalServerError, "internal_error", "Could not connect PayGate phone") + } +} + +func writeRelayError(w http.ResponseWriter, err error) { + var relayErr *relay.Error + if errors.As(err, &relayErr) { + writeError(w, relayErr.HTTPStatus, strings.ToLower(relayErr.Code), relayErr.Message) + return + } + 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 new file mode 100644 index 0000000..93fa1e4 --- /dev/null +++ b/internal/v4/httpapi/relay_test.go @@ -0,0 +1,146 @@ +package httpapi + +import ( + "context" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/sha256" + "crypto/x509" + "encoding/base64" + "encoding/hex" + "encoding/json" + "encoding/pem" + "net/http" + "net/http/httptest" + "path/filepath" + "strconv" + "strings" + "testing" + "time" + + "github.com/Phloraxx/payment-api/internal/v4/payments" + "github.com/Phloraxx/payment-api/internal/v4/relay" + "github.com/Phloraxx/payment-api/internal/v4/storage" +) + +type relayHTTPFixture struct { + db *storage.DB + service *relay.Service + handler *RelayHandler + private *ecdsa.PrivateKey + device string + now time.Time +} + +func newRelayHTTPFixture(t *testing.T) relayHTTPFixture { + t.Helper() + db, err := storage.Open(context.Background(), filepath.Join(t.TempDir(), "paygate.db")) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + paymentService := payments.NewService(db) + service := relay.NewService(db, paymentService) + now := time.Date(2026, 9, 1, 6, 30, 0, 0, time.UTC) + service.Now = func() time.Time { return now } + + private, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatal(err) + } + der, err := x509.MarshalPKIXPublicKey(&private.PublicKey) + if err != nil { + t.Fatal(err) + } + sum := sha256.Sum256(der) + deviceID := hex.EncodeToString(sum[:]) + return relayHTTPFixture{db: db, service: service, handler: NewRelayHandler(service), private: private, device: deviceID, now: now} +} +func (f relayHTTPFixture) publicKeyPEM(t *testing.T) string { + t.Helper() + der, err := x509.MarshalPKIXPublicKey(&f.private.PublicKey) + if err != nil { + t.Fatal(err) + } + return string(pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: der})) +} + +func pairRelayHTTP(t *testing.T, f relayHTTPFixture) { + t.Helper() + session, err := f.service.CreatePairing(context.Background(), false) + if err != nil { + t.Fatal(err) + } + body, _ := json.Marshal(map[string]any{ + "token": session.Token, "name": "Edge 60 Stylus", "public_key_pem": f.publicKeyPEM(t), + "app_version": "0.5.0", "device_model": "motorola edge 60 stylus", "android_version": "16", + }) + req := httptest.NewRequest(http.MethodPost, "/api/v4/relay/pair", strings.NewReader(string(body))) + req.Header.Set("Content-Type", "application/json") + rr := httptest.NewRecorder() + f.handler.ServeHTTP(rr, req) + if rr.Code != http.StatusOK || !strings.Contains(rr.Body.String(), f.device) { + t.Fatalf("pair status=%d body=%s", rr.Code, rr.Body.String()) + } +} +func signedRelayRequest(t *testing.T, f relayHTTPFixture, path string, body []byte) *http.Request { + t.Helper() + timestamp := strconv.FormatInt(f.now.UnixMilli(), 10) + canonical := relay.CanonicalRequest(http.MethodPost, path, timestamp, body) + digest := sha256.Sum256([]byte(canonical)) + signature, err := ecdsa.SignASN1(rand.Reader, f.private, digest[:]) + if err != nil { + t.Fatal(err) + } + req := httptest.NewRequest(http.MethodPost, path, strings.NewReader(string(body))) + req.Header.Set("Content-Type", "application/json") + req.Header.Set(relayDeviceHeader, f.device) + req.Header.Set(relayTimeHeader, timestamp) + req.Header.Set(relaySignatureHeader, base64.StdEncoding.EncodeToString(signature)) + return req +} + +func TestRelayPairHeartbeatAndHealthPersistence(t *testing.T) { + f := newRelayHTTPFixture(t) + 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)) + 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) + } +} +func TestRelaySignedEventAndSignatureFailure(t *testing.T) { + f := newRelayHTTPFixture(t) + 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)) + 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.Header.Set(relaySignatureHeader, base64.StdEncoding.EncodeToString([]byte("bad"))) + rr = httptest.NewRecorder() + f.handler.ServeHTTP(rr, bad) + if rr.Code != http.StatusUnauthorized { + t.Fatalf("bad signature status=%d body=%s", rr.Code, rr.Body.String()) + } +} + +func TestRelayRejectsQueryParameters(t *testing.T) { + f := newRelayHTTPFixture(t) + req := httptest.NewRequest(http.MethodPost, "/api/v4/relay/pair?token=leak", strings.NewReader(`{}`)) + req.Header.Set("Content-Type", "application/json") + rr := httptest.NewRecorder() + f.handler.ServeHTTP(rr, req) + if rr.Code != http.StatusBadRequest { + t.Fatalf("query status=%d body=%s", rr.Code, rr.Body.String()) + } +} diff --git a/internal/v4/operator/service.go b/internal/v4/operator/service.go index 74578ce..4f97405 100644 --- a/internal/v4/operator/service.go +++ b/internal/v4/operator/service.go @@ -187,7 +187,7 @@ func (s *Service) loadRelay(ctx context.Context, out *RelaySummary) error { var name string var lastSeen sql.NullInt64 var appVersion sql.NullString - err := s.DB.SQL.QueryRowContext(ctx, `SELECT COALESCE(name,''),last_seen_at,app_version FROM relay_devices WHERE enabled=1 LIMIT 1`). + err := s.DB.SQL.QueryRowContext(ctx, `SELECT COALESCE(name,''),COALESCE(last_heartbeat_at,last_seen_at),app_version FROM relay_devices WHERE enabled=1 LIMIT 1`). Scan(&name, &lastSeen, &appVersion) if errors.Is(err, sql.ErrNoRows) { return nil @@ -195,12 +195,13 @@ func (s *Service) loadRelay(ctx context.Context, out *RelaySummary) error { if err != nil { return fmt.Errorf("read relay summary: %w", err) } - out.Connected = true out.Name = name out.AppVersion = appVersion.String if lastSeen.Valid { value := time.UnixMilli(lastSeen.Int64).UTC() out.LastSeenAt = &value + age := s.now().Sub(value) + out.Connected = age >= -5*time.Minute && age <= time.Hour } return nil } diff --git a/internal/v4/relay/heartbeat.go b/internal/v4/relay/heartbeat.go new file mode 100644 index 0000000..4a62940 --- /dev/null +++ b/internal/v4/relay/heartbeat.go @@ -0,0 +1,91 @@ +package relay + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "strings" + "time" +) + +const HeartbeatPath = "/api/v4/relay/heartbeat" + +type HeartbeatInput struct { + SchemaVersion int `json:"schema_version"` + AppVersion string `json:"app_version,omitempty"` + AndroidVersion string `json:"android_version,omitempty"` + DeviceModel string `json:"device_model,omitempty"` + NotificationAccess bool `json:"notification_access"` + ListenerConnected bool `json:"listener_connected"` + BatteryOptimizationExempt bool `json:"battery_optimization_exempt"` + PowerSaveMode bool `json:"power_save_mode"` + BackgroundRestricted bool `json:"background_restricted"` + ForegroundService bool `json:"foreground_service"` + PendingCount int `json:"pending_count"` + FailedCount int `json:"failed_count"` + LastSuccessfulDeliveryMS int64 `json:"last_successful_delivery_at_ms,omitempty"` + LastClientError string `json:"last_client_error,omitempty"` +} +type HeartbeatResult struct { + ReceivedAt time.Time `json:"received_at"` +} + +func (s *Service) HeartbeatSigned(ctx context.Context, auth RequestAuth, rawBody []byte) (HeartbeatResult, error) { + if s == nil || s.DB == nil || s.DB.SQL == nil { + return HeartbeatResult{}, errors.New("relay storage is required") + } + if len(rawBody) == 0 || len(rawBody) > maxRawBodyBytes { + return HeartbeatResult{}, relayError("RELAY_HEARTBEAT_TOO_LARGE", "relay heartbeat body is empty or too large", 400) + } + if strings.ToUpper(strings.TrimSpace(auth.Method)) != "POST" || auth.Path != HeartbeatPath { + return HeartbeatResult{}, relayError("INVALID_RELAY_ENDPOINT", "relay signature is not for the v4 heartbeat endpoint", 401) + } + nowFn := s.Now + if nowFn == nil { + nowFn = time.Now + } + now := nowFn().UTC() + device, err := verifyRequest(ctx, s.DB, auth, rawBody, now) + if err != nil { + return HeartbeatResult{}, err + } + var input HeartbeatInput + if err := json.Unmarshal(rawBody, &input); err != nil { + return HeartbeatResult{}, relayError("INVALID_RELAY_HEARTBEAT", "relay heartbeat is not valid JSON", 400) + } + if input.SchemaVersion != SchemaVersion { + return HeartbeatResult{}, relayError("UNSUPPORTED_RELAY_SCHEMA", "schema_version must be 1", 400) + } + input.AppVersion = trimRunes(input.AppVersion, 64) + input.AndroidVersion = trimRunes(input.AndroidVersion, 64) + input.DeviceModel = trimRunes(input.DeviceModel, 255) + input.LastClientError = trimRunes(input.LastClientError, 512) + if input.PendingCount < 0 || input.PendingCount > 1_000_000 || input.FailedCount < 0 || input.FailedCount > 1_000_000 { + return HeartbeatResult{}, relayError("INVALID_RELAY_HEARTBEAT", "queue counts are outside the allowed range", 400) + } + var delivered any + if input.LastSuccessfulDeliveryMS > 0 { + t := time.UnixMilli(input.LastSuccessfulDeliveryMS).UTC() + if t.After(now.Add(5*time.Minute)) || t.Year() < 2020 { + return HeartbeatResult{}, relayError("INVALID_RELAY_HEARTBEAT", "last successful delivery time is invalid", 400) + } + 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) + if err != nil { + return HeartbeatResult{}, fmt.Errorf("persist relay heartbeat: %w", err) + } + if rows, _ := result.RowsAffected(); rows != 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 511919b..4378d41 100644 --- a/internal/v4/relay/pairing.go +++ b/internal/v4/relay/pairing.go @@ -48,14 +48,25 @@ type PairDeviceResult struct { } type DeviceInfo struct { - ID string `json:"id"` - Name string `json:"name"` - Enabled bool `json:"enabled"` - EnrolledAt time.Time `json:"enrolled_at"` - LastSeenAt *time.Time `json:"last_seen_at,omitempty"` - AppVersion string `json:"app_version,omitempty"` - DeviceModel string `json:"device_model,omitempty"` - AndroidVersion string `json:"android_version,omitempty"` + ID string `json:"id"` + Name string `json:"name"` + Enabled bool `json:"enabled"` + EnrolledAt time.Time `json:"enrolled_at"` + LastSeenAt *time.Time `json:"last_seen_at,omitempty"` + LastHeartbeatAt *time.Time `json:"last_heartbeat_at,omitempty"` + AppVersion string `json:"app_version,omitempty"` + DeviceModel string `json:"device_model,omitempty"` + AndroidVersion string `json:"android_version,omitempty"` + NotificationAccess *bool `json:"notification_access,omitempty"` + ListenerConnected *bool `json:"listener_connected,omitempty"` + BatteryOptimizationExempt *bool `json:"battery_optimization_exempt,omitempty"` + PowerSaveMode *bool `json:"power_save_mode,omitempty"` + BackgroundRestricted *bool `json:"background_restricted,omitempty"` + ForegroundService *bool `json:"foreground_service,omitempty"` + PendingCount *int `json:"pending_count,omitempty"` + FailedCount *int `json:"failed_count,omitempty"` + LastSuccessfulDeliveryAt *time.Time `json:"last_successful_delivery_at,omitempty"` + LastClientError string `json:"last_client_error,omitempty"` } func (s *Service) CreatePairing(ctx context.Context, replaceExisting bool) (PairingSession, error) { @@ -254,12 +265,19 @@ func (s *Service) ActiveDevice(ctx context.Context) (*DeviceInfo, error) { return nil, errors.New("relay storage is required") } var info DeviceInfo - var lastSeen sql.NullInt64 - var appVersion, model, androidVersion sql.NullString + var lastSeen, lastHeartbeat, lastDelivered sql.NullInt64 + var appVersion, model, androidVersion, lastError sql.NullString + var notificationAccess, listenerConnected, batteryExempt, powerSave, backgroundRestricted, foregroundService sql.NullInt64 + var pendingCount, failedCount sql.NullInt64 var enrolledAt int64 var enabled int - err := s.DB.SQL.QueryRowContext(ctx, `SELECT id,COALESCE(name,''),enabled,enrolled_at,last_seen_at,app_version,device_model,android_version - FROM relay_devices WHERE enabled=1 LIMIT 1`).Scan(&info.ID, &info.Name, &enabled, &enrolledAt, &lastSeen, &appVersion, &model, &androidVersion) + err := s.DB.SQL.QueryRowContext(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 + FROM relay_devices WHERE enabled=1 LIMIT 1`).Scan( + &info.ID, &info.Name, &enabled, &enrolledAt, &lastSeen, &lastHeartbeat, &appVersion, &model, &androidVersion, + ¬ificationAccess, &listenerConnected, &batteryExempt, &powerSave, &backgroundRestricted, &foregroundService, + &pendingCount, &failedCount, &lastDelivered, &lastError) if errors.Is(err, sql.ErrNoRows) { return nil, nil } @@ -268,12 +286,39 @@ func (s *Service) ActiveDevice(ctx context.Context) (*DeviceInfo, error) { } info.Enabled = enabled == 1 info.EnrolledAt = time.UnixMilli(enrolledAt).UTC() - info.AppVersion = appVersion.String - info.DeviceModel = model.String - info.AndroidVersion = androidVersion.String - if lastSeen.Valid { - value := time.UnixMilli(lastSeen.Int64).UTC() - info.LastSeenAt = &value - } + info.AppVersion, info.DeviceModel, info.AndroidVersion, info.LastClientError = appVersion.String, model.String, androidVersion.String, lastError.String + info.LastSeenAt = nullableTimePointer(lastSeen) + info.LastHeartbeatAt = nullableTimePointer(lastHeartbeat) + info.LastSuccessfulDeliveryAt = nullableTimePointer(lastDelivered) + info.NotificationAccess = nullableBoolPointer(notificationAccess) + info.ListenerConnected = nullableBoolPointer(listenerConnected) + info.BatteryOptimizationExempt = nullableBoolPointer(batteryExempt) + info.PowerSaveMode = nullableBoolPointer(powerSave) + info.BackgroundRestricted = nullableBoolPointer(backgroundRestricted) + info.ForegroundService = nullableBoolPointer(foregroundService) + info.PendingCount = nullableIntPointer(pendingCount) + info.FailedCount = nullableIntPointer(failedCount) return &info, nil } + +func nullableTimePointer(value sql.NullInt64) *time.Time { + if !value.Valid { + return nil + } + t := time.UnixMilli(value.Int64).UTC() + return &t +} +func nullableBoolPointer(value sql.NullInt64) *bool { + if !value.Valid { + return nil + } + v := value.Int64 == 1 + return &v +} +func nullableIntPointer(value sql.NullInt64) *int { + if !value.Valid { + return nil + } + v := int(value.Int64) + return &v +} diff --git a/internal/v4/storage/db.go b/internal/v4/storage/db.go index 7b7bde8..b797537 100644 --- a/internal/v4/storage/db.go +++ b/internal/v4/storage/db.go @@ -14,7 +14,7 @@ import ( const ( defaultBusyTimeoutMS = 5000 - schemaVersion = 1 + schemaVersion = 2 ) type DB struct { diff --git a/internal/v4/storage/db_test.go b/internal/v4/storage/db_test.go index e7f575d..318f82d 100644 --- a/internal/v4/storage/db_test.go +++ b/internal/v4/storage/db_test.go @@ -333,3 +333,56 @@ func TestOrdinaryReadTransactionDoesNotAcquireWriterLock(t *testing.T) { t.Fatalf("ordinary read transaction blocked writer: %v", err) } } +func TestOpenMigratesV1DatabaseToRelayHealthSchema(t *testing.T) { + ctx := context.Background() + path := filepath.Join(t.TempDir(), "paygate-v1.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, `INSERT INTO schema_migrations(version,applied_at) VALUES(1,1)`); err != nil { + t.Fatal(err) + } + if err := raw.Close(); err != nil { + t.Fatal(err) + } + + db, err := Open(ctx, path) + if err != nil { + t.Fatal(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 != 2 { + t.Fatalf("schema version=%d want=2", version) + } + rows, err := db.SQL.QueryContext(ctx, `PRAGMA table_info(relay_devices)`) + if err != nil { + t.Fatal(err) + } + defer rows.Close() + found := false + for rows.Next() { + var cid, notnull, pk int + var name, typ string + var defaultValue any + if err := rows.Scan(&cid, &name, &typ, ¬null, &defaultValue, &pk); err != nil { + t.Fatal(err) + } + if name == "notification_access" { + found = true + } + } + if !found { + t.Fatal("relay health columns were not added") + } +} diff --git a/internal/v4/storage/schema.go b/internal/v4/storage/schema.go index 299d09f..400a90e 100644 --- a/internal/v4/storage/schema.go +++ b/internal/v4/storage/schema.go @@ -37,12 +37,40 @@ CREATE TABLE IF NOT EXISTS schema_migrations ( return fmt.Errorf("record schema v1: %w", err) } } + if current < 2 { + if err := applyV2(ctx, tx); err != nil { + return err + } + if _, err := tx.ExecContext(ctx, `INSERT INTO schema_migrations(version, applied_at) VALUES(2, unixepoch('subsec') * 1000)`); err != nil { + return fmt.Errorf("record schema v2: %w", err) + } + } if err := tx.Commit(); err != nil { return fmt.Errorf("commit schema migration: %w", err) } return nil } +func applyV2(ctx context.Context, tx *sql.Tx) error { + if _, err := tx.ExecContext(ctx, schemaV2); err != nil { + return fmt.Errorf("apply schema v2: %w", err) + } + return nil +} + +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)); +ALTER TABLE relay_devices ADD COLUMN power_save_mode INTEGER CHECK(power_save_mode IS NULL OR power_save_mode IN (0,1)); +ALTER TABLE relay_devices ADD COLUMN background_restricted INTEGER CHECK(background_restricted IS NULL OR background_restricted IN (0,1)); +ALTER TABLE relay_devices ADD COLUMN foreground_service INTEGER CHECK(foreground_service IS NULL OR foreground_service IN (0,1)); +ALTER TABLE relay_devices ADD COLUMN pending_count INTEGER CHECK(pending_count IS NULL OR pending_count >= 0); +ALTER TABLE relay_devices ADD COLUMN failed_count INTEGER CHECK(failed_count IS NULL OR failed_count >= 0); +ALTER TABLE relay_devices ADD COLUMN last_successful_delivery_at INTEGER; +ALTER TABLE relay_devices ADD COLUMN last_client_error TEXT; +` + func applyV1(ctx context.Context, tx *sql.Tx) error { if _, err := tx.ExecContext(ctx, schemaV1); err != nil { return fmt.Errorf("apply schema v1: %w", err)