Files
silo-server/internal/notifications/push_delivery.go
T
cf0db385f3 Add Apple push notifications support (#255)
* Add push notifications support

* fix(notifications): address push notification review findings

- Gate the capability endpoint's apple_push availability on the admin
  delivery toggle, matching web push: Available now means setup will
  actually deliver.
- Reject direct admin writes to push_relay_deployment_id/api_key; the
  relay issues them as a pair during registration and a lone write
  desyncs them (and poisons the next rotation request).
- Purge a device's registrations under other profiles when it
  re-registers, so a profile switch on a shared device stops the old
  profile's pushes (attempts cascade); adds a DB-backed test.
- Extract the shared channelDispatcher core + retry sweep and rebuild
  the webhook/web push/Apple push dispatchers on it instead of keeping
  three copies of the worker-pool/retry loop.
- Deduplicate relay URL validation (admin setting + register flow) and
  the push outbox attempt-building loops behind shared helpers.
- Cap free-text decline reasons in notification display bodies.
- Fix TestHandleApplePushDisplayDB expectations to match the shared
  display copy (test previously failed under SILO_TEST_DATABASE_URL).

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>

* fix(notifications): route push relay URL writes through registration only

Direct writes to notifications.push_relay_url via the admin settings
endpoint bypassed the relay registration flow, letting the stored URL
drift out of sync with the deployment id / API key pair the relay
minted for it. Reject the URL alongside the deployment id and API key
in the settings handler; POST /admin/notifications/push/relay/register
remains the only path that persists all three together.

The admin UI's Relay URL field now edits local draft state and is
applied by the Register/Rotate action instead of the settings save.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>

---------

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
2026-07-01 17:25:16 -04:00

366 lines
12 KiB
Go

package notifications
import (
"context"
"errors"
"fmt"
"strings"
"time"
"github.com/jackc/pgx/v5"
"github.com/oklog/ulid/v2"
)
const (
PushTriggerDelivery = "delivery"
PushTriggerTest = "test"
PushOutcomePending = "pending"
PushOutcomeDelivered = "delivered"
PushOutcomeRetrying = "retrying"
PushOutcomeFailed = "failed"
)
var (
ErrPushDeliveryUnavailable = errors.New("apple push delivery unavailable")
ErrPushDeliveryInvalid = errors.New("invalid apple push delivery request")
ErrPushDeliveryNotFound = errors.New("apple push device not found")
)
// PushDeliveryAttempt is one row in the APNs relay outbox/retry log.
type PushDeliveryAttempt struct {
ID string
NotificationDeliveryID *string
PushDeviceID string
TriggerType string
Provider string
Platform string
AttemptNumber int
AttemptedAt time.Time
NextRetryAt *time.Time
Outcome string
RelayRequestID *string
UpstreamStatus *int
UpstreamReason *string
FailureMessage *string
CreatedAt time.Time
UpdatedAt time.Time
}
// newPushDeliveryAttempts builds the pending outbox rows fanning one delivery
// out to its profile's eligible devices. Shared by the fanout worker and
// operational dispatch so both outbox paths enqueue identical rows.
func newPushDeliveryAttempts(deliveryID string, devices []PushDevice) []PushDeliveryAttempt {
attempts := make([]PushDeliveryAttempt, 0, len(devices))
for _, device := range devices {
attempts = append(attempts, PushDeliveryAttempt{
ID: ulid.Make().String(),
NotificationDeliveryID: &deliveryID,
PushDeviceID: device.ID,
TriggerType: PushTriggerDelivery,
})
}
return attempts
}
const pushAttemptReturning = `
RETURNING id, notification_delivery_id, push_device_id, trigger_type, provider, platform,
attempt_number, attempted_at, next_retry_at, outcome, relay_request_id,
upstream_status, upstream_reason, failure_message, created_at, updated_at`
func scanPushDeliveryAttempts(rows pgx.Rows) ([]PushDeliveryAttempt, error) {
defer rows.Close()
attempts := make([]PushDeliveryAttempt, 0, 8)
for rows.Next() {
var attempt PushDeliveryAttempt
if err := rows.Scan(
&attempt.ID,
&attempt.NotificationDeliveryID,
&attempt.PushDeviceID,
&attempt.TriggerType,
&attempt.Provider,
&attempt.Platform,
&attempt.AttemptNumber,
&attempt.AttemptedAt,
&attempt.NextRetryAt,
&attempt.Outcome,
&attempt.RelayRequestID,
&attempt.UpstreamStatus,
&attempt.UpstreamReason,
&attempt.FailureMessage,
&attempt.CreatedAt,
&attempt.UpdatedAt,
); err != nil {
return nil, fmt.Errorf("scan push attempt: %w", err)
}
attempts = append(attempts, attempt)
}
return attempts, rows.Err()
}
// ListEnabledAppleByProfiles loads delivery-eligible APNs devices keyed by profile.
func (r *PushDeviceRepository) ListEnabledAppleByProfiles(ctx context.Context, tx pgx.Tx, profileIDs []string) (map[string][]PushDevice, error) {
out := make(map[string][]PushDevice, len(profileIDs))
if len(profileIDs) == 0 {
return out, nil
}
rows, err := tx.Query(ctx, `SELECT `+pushDeviceColumns+`
FROM push_devices
WHERE profile_id = ANY($1)
AND platform = $2
AND provider = $3
AND push_mode = $4
AND enabled`,
profileIDs, PushPlatformApple, PushProviderSiloRelay, PushModePrivatePush)
if err != nil {
return nil, fmt.Errorf("list enabled push devices: %w", err)
}
defer rows.Close()
for rows.Next() {
device, err := scanPushDevice(rows)
if err != nil {
return nil, fmt.Errorf("scan enabled push device: %w", err)
}
out[device.ProfileID] = append(out[device.ProfileID], *device)
}
return out, rows.Err()
}
// EnqueuePushAttempts inserts pending APNs relay attempts in the fanout transaction.
func (r *PushDeviceRepository) EnqueuePushAttempts(ctx context.Context, tx pgx.Tx, attempts []PushDeliveryAttempt) error {
if len(attempts) == 0 {
return nil
}
var sb strings.Builder
sb.WriteString(`
INSERT INTO push_delivery_attempts
(id, notification_delivery_id, push_device_id, trigger_type, provider, platform, attempt_number, outcome)
VALUES `)
args := make([]any, 0, len(attempts)*8)
for i, attempt := range attempts {
if i > 0 {
sb.WriteString(", ")
}
base := len(args)
sb.WriteString(fmt.Sprintf("($%d,$%d,$%d,$%d,$%d,$%d,$%d,$%d)",
base+1, base+2, base+3, base+4, base+5, base+6, base+7, base+8))
args = append(args,
attempt.ID,
attempt.NotificationDeliveryID,
attempt.PushDeviceID,
defaultString(attempt.TriggerType, PushTriggerDelivery),
PushProviderSiloRelay,
PushPlatformApple,
0,
PushOutcomePending,
)
}
sb.WriteString(" ON CONFLICT DO NOTHING")
if _, err := tx.Exec(ctx, sb.String(), args...); err != nil {
return fmt.Errorf("enqueue push attempts: %w", err)
}
return nil
}
// EnqueueAppleTestAttempt creates a pending diagnostic attempt for one enabled device.
func (r *PushDeviceRepository) EnqueueAppleTestAttempt(ctx context.Context, profileID, serverDeviceID string) (*PushDeliveryAttempt, *PushDevice, error) {
if r == nil || r.pool == nil {
return nil, nil, ErrPushDeliveryUnavailable
}
profileID = strings.TrimSpace(profileID)
serverDeviceID = strings.TrimSpace(serverDeviceID)
if profileID == "" {
return nil, nil, fmt.Errorf("%w: profile_id is required", ErrPushDeliveryInvalid)
}
tx, err := r.pool.Begin(ctx)
if err != nil {
return nil, nil, fmt.Errorf("begin push test attempt: %w", err)
}
defer func() { _ = tx.Rollback(ctx) }()
query := `SELECT ` + pushDeviceColumns + `
FROM push_devices
WHERE profile_id = $1
AND platform = $2
AND provider = $3
AND push_mode = $4
AND enabled`
args := []any{profileID, PushPlatformApple, PushProviderSiloRelay, PushModePrivatePush}
if serverDeviceID != "" {
args = append(args, serverDeviceID)
query += fmt.Sprintf(" AND server_device_id = $%d", len(args))
}
query += ` ORDER BY last_seen_at DESC NULLS LAST, created_at DESC LIMIT 1 FOR UPDATE`
device, err := scanPushDevice(tx.QueryRow(ctx, query, args...))
if errors.Is(err, pgx.ErrNoRows) {
return nil, nil, ErrPushDeliveryNotFound
}
if err != nil {
return nil, nil, fmt.Errorf("select push test device: %w", err)
}
attemptID := ulid.Make().String()
rows, err := tx.Query(ctx, `
INSERT INTO push_delivery_attempts
(id, notification_delivery_id, push_device_id, trigger_type, provider, platform, attempt_number, outcome)
VALUES ($1, NULL, $2, $3, $4, $5, 0, $6)`+pushAttemptReturning,
attemptID, device.ID, PushTriggerTest, PushProviderSiloRelay, PushPlatformApple, PushOutcomePending)
if err != nil {
return nil, nil, fmt.Errorf("insert push test attempt: %w", err)
}
attempts, err := scanPushDeliveryAttempts(rows)
if err != nil {
return nil, nil, err
}
if len(attempts) != 1 {
return nil, nil, fmt.Errorf("insert push test attempt returned %d rows", len(attempts))
}
if err := tx.Commit(ctx); err != nil {
return nil, nil, fmt.Errorf("commit push test attempt: %w", err)
}
return &attempts[0], device, nil
}
func (r *PushDeviceRepository) GetPushAttempt(ctx context.Context, id string) (*PushDeliveryAttempt, error) {
rows, err := r.pool.Query(ctx, `SELECT * FROM (
SELECT id, notification_delivery_id, push_device_id, trigger_type, provider, platform,
attempt_number, attempted_at, next_retry_at, outcome, relay_request_id,
upstream_status, upstream_reason, failure_message, created_at, updated_at
FROM push_delivery_attempts
WHERE id = $1
) attempt`, id)
if err != nil {
return nil, fmt.Errorf("get push attempt: %w", err)
}
attempts, err := scanPushDeliveryAttempts(rows)
if err != nil {
return nil, err
}
if len(attempts) == 0 {
return nil, nil
}
return &attempts[0], nil
}
func (r *PushDeviceRepository) getPushDeviceByID(ctx context.Context, id string) (*PushDevice, error) {
device, err := scanPushDevice(r.pool.QueryRow(ctx,
`SELECT `+pushDeviceColumns+` FROM push_devices WHERE id = $1`, id))
if errors.Is(err, pgx.ErrNoRows) {
return nil, nil
}
if err != nil {
return nil, fmt.Errorf("get push device: %w", err)
}
return device, nil
}
func (r *PushDeviceRepository) ClaimPendingPushForDelivery(ctx context.Context, deliveryID string) ([]PushDeliveryAttempt, error) {
return r.claimPushAttempts(ctx, `
UPDATE push_delivery_attempts SET outcome = 'retrying', next_retry_at = now() + $2, updated_at = now()
WHERE id IN (
SELECT id FROM push_delivery_attempts
WHERE notification_delivery_id = $1 AND outcome = 'pending'
FOR UPDATE SKIP LOCKED
)`+pushAttemptReturning,
deliveryID, webhookClaimLease)
}
func (r *PushDeviceRepository) ClaimPushAttemptByID(ctx context.Context, attemptID string) ([]PushDeliveryAttempt, error) {
return r.claimPushAttempts(ctx, `
UPDATE push_delivery_attempts SET outcome = 'retrying', next_retry_at = now() + $2, updated_at = now()
WHERE id IN (
SELECT id FROM push_delivery_attempts
WHERE id = $1 AND outcome = 'pending'
FOR UPDATE SKIP LOCKED
)`+pushAttemptReturning,
attemptID, webhookClaimLease)
}
func (r *PushDeviceRepository) ClaimDuePushAttempts(ctx context.Context, limit int) ([]PushDeliveryAttempt, error) {
return r.claimPushAttempts(ctx, `
UPDATE push_delivery_attempts SET outcome = 'retrying', next_retry_at = now() + $2, updated_at = now()
WHERE id IN (
SELECT id FROM push_delivery_attempts
WHERE (outcome = 'retrying' AND next_retry_at <= now())
OR (outcome = 'pending' AND attempted_at <= now() - interval '60 seconds')
ORDER BY next_retry_at NULLS FIRST
LIMIT $1
FOR UPDATE SKIP LOCKED
)`+pushAttemptReturning,
limit, webhookClaimLease)
}
func (r *PushDeviceRepository) claimPushAttempts(ctx context.Context, query string, args ...any) ([]PushDeliveryAttempt, error) {
rows, err := r.pool.Query(ctx, query, args...)
if err != nil {
return nil, fmt.Errorf("claim push attempts: %w", err)
}
return scanPushDeliveryAttempts(rows)
}
func (r *PushDeviceRepository) FinalizePushAttempt(ctx context.Context, attemptID, outcome string, attemptNumber int, relayRequestID string, upstreamStatus *int, upstreamReason, failureMessage string, nextRetryAt *time.Time) (*PushDeliveryAttempt, error) {
var relayRequestIDPtr *string
if relayRequestID != "" {
relayRequestIDPtr = &relayRequestID
}
var upstreamReasonPtr *string
if upstreamReason != "" {
upstreamReasonPtr = &upstreamReason
}
var failureMessagePtr *string
if failureMessage != "" {
failureMessagePtr = &failureMessage
}
rows, err := r.pool.Query(ctx, `
UPDATE push_delivery_attempts
SET outcome = $2,
attempt_number = $3,
attempted_at = now(),
next_retry_at = $4,
relay_request_id = $5,
upstream_status = $6,
upstream_reason = left($7, 256),
failure_message = left($8, 256),
updated_at = now()
WHERE id = $1`+pushAttemptReturning,
attemptID, outcome, attemptNumber, nextRetryAt, relayRequestIDPtr, upstreamStatus, upstreamReasonPtr, failureMessagePtr)
if err != nil {
return nil, fmt.Errorf("finalize push attempt: %w", err)
}
attempts, err := scanPushDeliveryAttempts(rows)
if err != nil {
return nil, err
}
if len(attempts) == 0 {
return nil, nil
}
return &attempts[0], nil
}
func (r *PushDeviceRepository) RecordPushSuccess(ctx context.Context, deviceID string) error {
_, err := r.pool.Exec(ctx, `
UPDATE push_devices
SET last_success_at = now(), last_failure_at = NULL, last_failure_code = NULL, updated_at = now()
WHERE id = $1`, deviceID)
return err
}
func (r *PushDeviceRepository) RecordPushFailure(ctx context.Context, deviceID, code string, disable bool) error {
_, err := r.pool.Exec(ctx, `
UPDATE push_devices
SET last_failure_at = now(),
last_failure_code = left($2, 128),
enabled = CASE WHEN $3 THEN false ELSE enabled END,
updated_at = now()
WHERE id = $1`, deviceID, code, disable)
return err
}
func defaultString(value, fallback string) string {
if value == "" {
return fallback
}
return value
}