fix(devices): re-approving a known device rotates its token instead of 500

App reinstalls that preserve data (Android Studio installDebug over an
existing install) re-register with the same device_identifier, but the
devices row from the previous install still exists — device_identifier
is UNIQUE, so ApproveDevice's blind INSERT failed with a unique
violation and returned 500 'failed to create device' (reproduced via
curl: second approve with the same identifier = instant 500; the ~98s
in the original report was app-side retry/polling, not server wait).

ApproveDevice is now idempotent: look the device up by identifier
first; a row owned by the approving user gets its auth token rotated
via UpdateDeviceAuthToken (row id unchanged, so synced highlights/
bookmarks/progress anchored to it stay valid; fresh install = fresh
credentials, old token invalidated); a row owned by another user gets
409; unknown identifiers INSERT as before, with the 23505 race falling
through to the rotate path. DB failures are logged (they were silent).

Also guard the in-memory pendingRegistrations map with a mutex —
register/approve/reject/status/list all touch it from HTTP goroutines,
and a racing write is a Go runtime fatal, not an error. The approver's
credential publication and the status poller's approved-branch snapshot
now run under the lock so the token can never be read half-written.

Regression test: TestApproveDeviceReapprovalRotatesToken — register →
approve → re-register same identifier → approve (must be 200) → token
rotated, exactly one devices row, row carries the new token.
This commit is contained in:
John O'Keefe
2026-09-17 23:09:02 -04:00
parent b4c956aed4
commit 1cd8557b58
2 changed files with 204 additions and 20 deletions
+82
View File
@@ -252,6 +252,88 @@ func TestListPendingRegistrations(t *testing.T) {
assert.NotNil(t, pending, "Pending registrations should not be nil") assert.NotNil(t, pending, "Pending registrations should not be nil")
} }
// Regression test for the "approve returns 500 after app reinstall" bug:
// an app reinstall that preserves data (e.g. Android Studio installDebug
// over an existing install) re-registers with the SAME device_identifier,
// and the blind INSERT in ApproveDevice hit the UNIQUE(device_identifier)
// constraint. Re-approval must be idempotent: same row (id unchanged),
// rotated token, old token invalidated.
func TestApproveDeviceReapprovalRotatesToken(t *testing.T) {
setup := setupTestServer(t)
userID := getTestUserID(t, setup.DB)
identifier := fmt.Sprintf("repro-reinstall-%s", uuid.New().String())
registerAndApprove := func() string {
regRequest := map[string]interface{}{
"device_name": "Reinstall Device",
"device_type": "mobile",
"device_identifier": identifier,
}
regBody, _ := json.Marshal(regRequest)
req := httptest.NewRequest("POST", "/api/devices/register", bytes.NewReader(regBody))
req.Header.Set("Content-Type", "application/json")
rec := httptest.NewRecorder()
setup.Server.Config.Handler.ServeHTTP(rec, req)
require.Equal(t, http.StatusCreated, rec.Code, "Should initiate registration")
var regResponse map[string]interface{}
json.Unmarshal(rec.Body.Bytes(), &regResponse)
registrationID, ok := regResponse["registration_id"].(string)
require.True(t, ok, "Should have registration_id")
req = httptest.NewRequest("GET", fmt.Sprintf("/api/devices/approve/%s", registrationID), nil)
req.Header.Set("Authorization", "Bearer "+setup.Token)
rec = httptest.NewRecorder()
setup.Server.Config.Handler.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code,
"Approve must succeed even when the identifier already exists (was the reinstall 500)")
// The status response is single-use and carries the credentials.
statusBody, _ := json.Marshal(map[string]string{"registration_id": registrationID})
req = httptest.NewRequest("POST", "/api/devices/register/status", bytes.NewReader(statusBody))
req.Header.Set("Content-Type", "application/json")
rec = httptest.NewRecorder()
setup.Server.Config.Handler.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
var statusResponse map[string]interface{}
json.Unmarshal(rec.Body.Bytes(), &statusResponse)
require.Equal(t, "approved", statusResponse["status"])
token, ok := statusResponse["auth_token"].(string)
require.True(t, ok, "Should have auth_token")
require.NotEmpty(t, token)
return token
}
token1 := registerAndApprove()
token2 := registerAndApprove()
assert.NotEqual(t, token1, token2, "Re-approval must rotate the auth token (fresh install = fresh credentials)")
devices, err := setup.DB.ListDevicesByUser(context.Background(),
pgtype.UUID{Bytes: [16]byte(userID), Valid: true})
require.NoError(t, err)
rows := 0
for _, d := range devices {
if d.DeviceIdentifier == identifier {
rows++
}
}
assert.Equal(t, 1, rows, "Re-approval must reuse the existing devices row, not duplicate it")
// The rotated-in token must be the live one.
var live *database.Devices
for i := range devices {
if devices[i].DeviceIdentifier == identifier {
live = &devices[i]
}
}
require.NotNil(t, live)
assert.Equal(t, token2, live.AuthToken, "The devices row must carry the newly rotated token")
}
func TestApproveDeviceRegistration(t *testing.T) { func TestApproveDeviceRegistration(t *testing.T) {
setup := setupTestServer(t) setup := setupTestServer(t)
+114 -12
View File
@@ -6,11 +6,16 @@ import (
"crypto/rand" "crypto/rand"
"encoding/base64" "encoding/base64"
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"log/slog"
"net/http" "net/http"
"sync"
"time" "time"
"github.com/google/uuid" "github.com/google/uuid"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgconn"
"github.com/jackc/pgx/v5/pgtype" "github.com/jackc/pgx/v5/pgtype"
"github.com/labstack/echo/v5" "github.com/labstack/echo/v5"
"github.com/skip2/go-qrcode" "github.com/skip2/go-qrcode"
@@ -106,7 +111,14 @@ type PendingRegistration struct {
SyncEndpoints map[string]string SyncEndpoints map[string]string
} }
var pendingRegistrations = make(map[string]*PendingRegistration) // pendingRegistrations holds in-flight (unapproved) device registrations.
// HTTP handlers touch it from multiple goroutines — every access must hold
// pendingMu (Go maps are not safe for concurrent use; a racing write is a
// runtime fatal, not an error).
var (
pendingRegistrations = make(map[string]*PendingRegistration)
pendingMu sync.Mutex
)
func (h *DeviceHandler) InitiateRegistration(c *echo.Context) error { func (h *DeviceHandler) InitiateRegistration(c *echo.Context) error {
req := DeviceRegistrationRequest{} req := DeviceRegistrationRequest{}
@@ -130,7 +142,9 @@ func (h *DeviceHandler) InitiateRegistration(c *echo.Context) error {
CreatedAt: time.Now(), CreatedAt: time.Now(),
} }
pendingMu.Lock()
pendingRegistrations[registrationID] = registration pendingRegistrations[registrationID] = registration
pendingMu.Unlock()
authURL := fmt.Sprintf("%s/devices/approve/%s", h.cfg.BaseURL, registrationID) authURL := fmt.Sprintf("%s/devices/approve/%s", h.cfg.BaseURL, registrationID)
@@ -167,25 +181,32 @@ func (h *DeviceHandler) CheckRegistrationStatus(c *echo.Context) error {
return c.JSON(http.StatusBadRequest, map[string]string{"error": "invalid request format"}) return c.JSON(http.StatusBadRequest, map[string]string{"error": "invalid request format"})
} }
pendingMu.Lock()
registration, exists := pendingRegistrations[req.RegistrationID] registration, exists := pendingRegistrations[req.RegistrationID]
if !exists { if !exists {
pendingMu.Unlock()
return c.JSON(http.StatusNotFound, map[string]string{"error": "registration not found"}) return c.JSON(http.StatusNotFound, map[string]string{"error": "registration not found"})
} }
if time.Now().After(registration.ExpiresAt) { if time.Now().After(registration.ExpiresAt) {
delete(pendingRegistrations, req.RegistrationID) delete(pendingRegistrations, req.RegistrationID)
pendingMu.Unlock()
return c.JSON(http.StatusGone, map[string]string{"error": "registration expired"}) return c.JSON(http.StatusGone, map[string]string{"error": "registration expired"})
} }
if registration.Approved { if registration.Approved {
delete(pendingRegistrations, req.RegistrationID) delete(pendingRegistrations, req.RegistrationID)
// Snapshot under the lock: the approver wrote these fields, and
return c.JSON(http.StatusOK, DeviceAuthStatusResponse{ // the row is gone from the map — no other reader/writer remains.
resp := DeviceAuthStatusResponse{
Status: "approved", Status: "approved",
AuthToken: registration.AuthToken, AuthToken: registration.AuthToken,
DeviceID: registration.DeviceID, DeviceID: registration.DeviceID,
SyncEndpoints: registration.SyncEndpoints, SyncEndpoints: registration.SyncEndpoints,
}) }
pendingMu.Unlock()
return c.JSON(http.StatusOK, resp)
} }
return c.JSON(http.StatusOK, DeviceAuthStatusResponse{ return c.JSON(http.StatusOK, DeviceAuthStatusResponse{
@@ -552,17 +573,21 @@ func (h *DeviceHandler) ApproveDevice(c *echo.Context) error {
registrationID := c.Param("registration_id") registrationID := c.Param("registration_id")
pendingMu.Lock()
registration, exists := pendingRegistrations[registrationID] registration, exists := pendingRegistrations[registrationID]
if !exists { if !exists {
pendingMu.Unlock()
return c.JSON(http.StatusNotFound, map[string]string{"error": "registration not found"}) return c.JSON(http.StatusNotFound, map[string]string{"error": "registration not found"})
} }
if time.Now().After(registration.ExpiresAt) { if time.Now().After(registration.ExpiresAt) {
delete(pendingRegistrations, registrationID) delete(pendingRegistrations, registrationID)
pendingMu.Unlock()
return c.JSON(http.StatusGone, map[string]string{"error": "registration expired"}) return c.JSON(http.StatusGone, map[string]string{"error": "registration expired"})
} }
if registration.Approved { if registration.Approved {
pendingMu.Unlock()
return c.JSON(http.StatusOK, map[string]interface{}{ return c.JSON(http.StatusOK, map[string]interface{}{
"message": "device already approved", "message": "device already approved",
"device_name": registration.DeviceName, "device_name": registration.DeviceName,
@@ -570,6 +595,7 @@ func (h *DeviceHandler) ApproveDevice(c *echo.Context) error {
"approved": true, "approved": true,
}) })
} }
pendingMu.Unlock()
authToken, err := generateDeviceToken() authToken, err := generateDeviceToken()
if err != nil { if err != nil {
@@ -577,23 +603,87 @@ func (h *DeviceHandler) ApproveDevice(c *echo.Context) error {
} }
pgUserID := pgtype.UUID{Bytes: [16]byte(userUUID), Valid: true} pgUserID := pgtype.UUID{Bytes: [16]byte(userUUID), Valid: true}
syncEnabled := pgtype.Bool{Bool: true, Valid: true}
autoSync := pgtype.Bool{Bool: true, Valid: true}
syncFreq := pgtype.Int4{Int32: 5, Valid: true}
device, err := h.db.CreateDevice(c.Request().Context(), database.CreateDeviceParams{ // Re-approval: app reinstalls that preserve data (e.g. Android Studio
// installDebug over an existing install) re-register with the SAME
// device_identifier, but the devices row from the previous install
// still exists — device_identifier is UNIQUE, so a blind INSERT fails
// with a unique violation (the historical "approve returns 500 after
// reinstall" bug). Look the device up first and rotate its token
// instead; a fresh install MUST get fresh credentials, so the old
// token is invalidated either way.
rotateExisting := func(existing database.Devices) (database.Devices, bool, error) {
// The identifier is globally unique; a row owned by another user
// means cross-account identifier reuse — refuse it.
if !existing.UserID.Valid || existing.UserID.Bytes != pgUserID.Bytes {
return existing, false, nil
}
updated, err := h.db.UpdateDeviceAuthToken(c.Request().Context(),
database.UpdateDeviceAuthTokenParams{
ID: pgtype.UUID{Bytes: existing.ID.Bytes, Valid: true},
AuthToken: authToken,
})
return updated, true, err
}
existing, err := h.db.GetDeviceByIdentifier(c.Request().Context(),
registration.DeviceIdentifier)
var device database.Devices
switch {
case err == nil:
// Known device: rotate the token, keep the row (id unchanged —
// highlights/bookmarks/progress anchored to it stay valid).
var ok bool
device, ok, err = rotateExisting(existing)
if err != nil {
slog.Error("device re-approval failed",
"registration_id", registrationID, "error", err)
return c.JSON(http.StatusInternalServerError,
map[string]string{"error": "failed to approve device"})
}
if !ok {
return c.JSON(http.StatusConflict,
map[string]string{"error": "device identifier already registered to another user"})
}
case errors.Is(err, pgx.ErrNoRows):
device, err = h.db.CreateDevice(c.Request().Context(), database.CreateDeviceParams{
UserID: pgUserID, UserID: pgUserID,
DeviceName: registration.DeviceName, DeviceName: registration.DeviceName,
DeviceType: registration.DeviceType, DeviceType: registration.DeviceType,
DeviceIdentifier: registration.DeviceIdentifier, DeviceIdentifier: registration.DeviceIdentifier,
AuthToken: authToken, AuthToken: authToken,
SyncEnabled: syncEnabled, SyncEnabled: pgtype.Bool{Bool: true, Valid: true},
AutoSync: autoSync, AutoSync: pgtype.Bool{Bool: true, Valid: true},
SyncFrequencyMinutes: syncFreq, SyncFrequencyMinutes: pgtype.Int4{Int32: 5, Valid: true},
DeviceMetadata: []byte("{}"), DeviceMetadata: []byte("{}"),
}) })
if err != nil { if err != nil {
return c.JSON(http.StatusInternalServerError, map[string]string{"error": "failed to create device"}) // Concurrent approve racing us onto the unique index: fall
// through to the rotate path for the winner's row.
var pgErr *pgconn.PgError
if asErr := errors.As(err, &pgErr); asErr && pgErr.Code == "23505" {
if winner, lerr := h.db.GetDeviceByIdentifier(
c.Request().Context(), registration.DeviceIdentifier); lerr == nil {
var ok bool
device, ok, err = rotateExisting(winner)
if err == nil && !ok {
return c.JSON(http.StatusConflict,
map[string]string{"error": "device identifier already registered to another user"})
}
}
}
}
if err != nil {
slog.Error("device creation failed",
"registration_id", registrationID, "error", err)
return c.JSON(http.StatusInternalServerError,
map[string]string{"error": "failed to create device"})
}
default:
slog.Error("device lookup failed",
"registration_id", registrationID, "error", err)
return c.JSON(http.StatusInternalServerError,
map[string]string{"error": "failed to look up device"})
} }
syncEndpoints := map[string]string{} syncEndpoints := map[string]string{}
@@ -607,11 +697,16 @@ func (h *DeviceHandler) ApproveDevice(c *echo.Context) error {
syncEndpoints["library"] = fmt.Sprintf("%s/api/sync/kobo/library", h.cfg.BaseURL) syncEndpoints["library"] = fmt.Sprintf("%s/api/sync/kobo/library", h.cfg.BaseURL)
} }
// Publish the credentials under the map lock: the status poller
// snapshots these fields only after Approved flips, so the token can
// never be read half-written.
pendingMu.Lock()
registration.UserID = userUUID registration.UserID = userUUID
registration.Approved = true registration.Approved = true
registration.AuthToken = authToken registration.AuthToken = authToken
registration.DeviceID = device.ID.Bytes registration.DeviceID = device.ID.Bytes
registration.SyncEndpoints = syncEndpoints registration.SyncEndpoints = syncEndpoints
pendingMu.Unlock()
return c.JSON(http.StatusOK, map[string]interface{}{ return c.JSON(http.StatusOK, map[string]interface{}{
"message": "device approved successfully", "message": "device approved successfully",
@@ -625,12 +720,15 @@ func (h *DeviceHandler) ApproveDevice(c *echo.Context) error {
func (h *DeviceHandler) RejectDevice(c *echo.Context) error { func (h *DeviceHandler) RejectDevice(c *echo.Context) error {
registrationID := c.Param("registration_id") registrationID := c.Param("registration_id")
pendingMu.Lock()
_, exists := pendingRegistrations[registrationID] _, exists := pendingRegistrations[registrationID]
if !exists { if !exists {
pendingMu.Unlock()
return c.JSON(http.StatusNotFound, map[string]string{"error": "registration not found"}) return c.JSON(http.StatusNotFound, map[string]string{"error": "registration not found"})
} }
delete(pendingRegistrations, registrationID) delete(pendingRegistrations, registrationID)
pendingMu.Unlock()
return c.JSON(http.StatusOK, map[string]string{ return c.JSON(http.StatusOK, map[string]string{
"message": "device registration rejected", "message": "device registration rejected",
@@ -639,6 +737,7 @@ func (h *DeviceHandler) RejectDevice(c *echo.Context) error {
func (h *DeviceHandler) GetPendingRegistrationsData(c *echo.Context) ([]map[string]interface{}, error) { func (h *DeviceHandler) GetPendingRegistrationsData(c *echo.Context) ([]map[string]interface{}, error) {
registrations := []map[string]interface{}{} registrations := []map[string]interface{}{}
pendingMu.Lock()
for _, reg := range pendingRegistrations { for _, reg := range pendingRegistrations {
if reg.UserID == (uuid.UUID{}) { if reg.UserID == (uuid.UUID{}) {
registrations = append(registrations, map[string]interface{}{ registrations = append(registrations, map[string]interface{}{
@@ -652,6 +751,7 @@ func (h *DeviceHandler) GetPendingRegistrationsData(c *echo.Context) ([]map[stri
}) })
} }
} }
pendingMu.Unlock()
return registrations, nil return registrations, nil
} }
@@ -664,6 +764,7 @@ func (h *DeviceHandler) ListPendingRegistrations(c *echo.Context) error {
} }
registrations := []map[string]interface{}{} registrations := []map[string]interface{}{}
pendingMu.Lock()
for _, reg := range pendingRegistrations { for _, reg := range pendingRegistrations {
if reg.UserID == userUUID || reg.UserID == (uuid.UUID{}) { if reg.UserID == userUUID || reg.UserID == (uuid.UUID{}) {
registrations = append(registrations, map[string]interface{}{ registrations = append(registrations, map[string]interface{}{
@@ -677,6 +778,7 @@ func (h *DeviceHandler) ListPendingRegistrations(c *echo.Context) error {
}) })
} }
} }
pendingMu.Unlock()
return c.JSON(http.StatusOK, map[string]interface{}{ return c.JSON(http.StatusOK, map[string]interface{}{
"registrations": registrations, "registrations": registrations,