diff --git a/internal/api/handlers/admin.go b/internal/api/handlers/admin.go index cf8d7bb3..4e10b0b8 100644 --- a/internal/api/handlers/admin.go +++ b/internal/api/handlers/admin.go @@ -2141,6 +2141,21 @@ func (h *AdminHandler) HandleUpdateSetting(w http.ResponseWriter, r *http.Reques "subtitle_ai.transcribe_quota_period must be day, week, or month") return } + case notifications.SettingApplePushDeliveryEnabled: + enabled, err := strconv.ParseBool(strings.TrimSpace(req.Value)) + if err != nil { + writeError(w, http.StatusBadRequest, "bad_request", "notifications.apple_push_delivery_enabled must be true or false") + return + } + req.Value = strconv.FormatBool(enabled) + case notifications.SettingPushRelayURL, notifications.SettingPushRelayDeploymentID, notifications.SettingPushRelayAPIKey: + // The registration flow persists the relay URL, deployment id, and API + // key together; a direct write to any of them desyncs the stored URL + // from the credentials the relay minted for it (and feeds an arbitrary + // id into the next rotation request). + writeError(w, http.StatusBadRequest, "bad_request", + key+" is managed by the push relay registration flow; use POST /admin/notifications/push/relay/register") + return case catalog.SearchSettingProvider: switch strings.TrimSpace(strings.ToLower(req.Value)) { case catalog.SearchProviderPostgres, catalog.SearchProviderMeilisearch: diff --git a/internal/api/handlers/admin_apple_push.go b/internal/api/handlers/admin_apple_push.go new file mode 100644 index 00000000..31affd00 --- /dev/null +++ b/internal/api/handlers/admin_apple_push.go @@ -0,0 +1,279 @@ +package handlers + +import ( + "bytes" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/url" + "strings" + "time" + + "github.com/Silo-Server/silo-server/internal/notifications" +) + +type AdminApplePushHandler struct { + system *notifications.System + settings ServerSettingsStore + client httpDoer +} + +type httpDoer interface { + Do(*http.Request) (*http.Response, error) +} + +func NewAdminApplePushHandler(system *notifications.System, settings ServerSettingsStore) *AdminApplePushHandler { + return &AdminApplePushHandler{ + system: system, + settings: settings, + client: &http.Client{Timeout: 10 * time.Second}, + } +} + +type adminApplePushTestRequest struct { + ProfileID string `json:"profile_id"` + ServerDeviceID string `json:"server_device_id"` +} + +type adminApplePushTestResponse struct { + AttemptID string `json:"attempt_id"` + PushDeviceID string `json:"push_device_id"` + ServerDeviceID string `json:"server_device_id"` + Outcome string `json:"outcome"` + RelayRequestID string `json:"relay_request_id,omitempty"` + UpstreamStatus *int `json:"upstream_status,omitempty"` + UpstreamReason string `json:"upstream_reason,omitempty"` + FailureMessage string `json:"failure_message,omitempty"` +} + +type adminPushRelayRegisterRequest struct { + RelayURL string `json:"relay_url"` +} + +type pushRelayRegisterRequest struct { + DeploymentID string `json:"deployment_id,omitempty"` +} + +type pushRelayRegisterResponse struct { + RequestID string `json:"request_id"` + DeploymentID string `json:"deployment_id"` + APIKey string `json:"api_key"` + KeyPrefix string `json:"key_prefix"` + APNsTopics []string `json:"apns_topics"` +} + +type pushRelayNestedError struct { + Error struct { + Code string `json:"code"` + Message string `json:"message"` + RequestID string `json:"request_id"` + } `json:"error"` +} + +type adminPushRelayRegisterResponse struct { + RelayURL string `json:"relay_url"` + DeploymentID string `json:"deployment_id"` + KeyPrefix string `json:"key_prefix"` + APIKeyConfigured bool `json:"api_key_configured"` + RelayRequestID string `json:"relay_request_id,omitempty"` + APNsTopics []string `json:"apns_topics,omitempty"` +} + +// HandleTest handles POST /admin/notifications/push/apple/test. +func (h *AdminApplePushHandler) HandleTest(w http.ResponseWriter, r *http.Request) { + if h == nil || h.system == nil { + writeError(w, http.StatusServiceUnavailable, "unavailable", "Apple push delivery is not available") + return + } + var req adminApplePushTestRequest + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeError(w, http.StatusBadRequest, "bad_request", "Invalid request body") + return + } + result, err := h.system.SendApplePushTest(r.Context(), req.ProfileID, req.ServerDeviceID) + if err != nil { + switch { + case errors.Is(err, notifications.ErrPushDeliveryInvalid): + writeError(w, http.StatusBadRequest, "bad_request", err.Error()) + case errors.Is(err, notifications.ErrPushDeliveryNotFound): + writeError(w, http.StatusNotFound, "not_found", "Apple push device not found") + case errors.Is(err, notifications.ErrPushDeliveryUnavailable): + writeError(w, http.StatusServiceUnavailable, "unavailable", "Apple push delivery is not available") + default: + writeError(w, http.StatusInternalServerError, "internal_error", "Failed to send Apple push test") + } + return + } + writeJSON(w, http.StatusOK, adminApplePushTestResponse{ + AttemptID: result.AttemptID, + PushDeviceID: result.PushDeviceID, + ServerDeviceID: result.ServerDeviceID, + Outcome: result.Outcome, + RelayRequestID: result.RelayRequestID, + UpstreamStatus: result.UpstreamStatus, + UpstreamReason: result.UpstreamReason, + FailureMessage: result.FailureMessage, + }) +} + +// HandleRegisterRelay handles POST /admin/notifications/push/relay/register. +func (h *AdminApplePushHandler) HandleRegisterRelay(w http.ResponseWriter, r *http.Request) { + if h == nil || h.settings == nil { + writeError(w, http.StatusServiceUnavailable, "unavailable", "Settings store is not available") + return + } + var req adminPushRelayRegisterRequest + dec := json.NewDecoder(r.Body) + dec.DisallowUnknownFields() + if err := dec.Decode(&req); err != nil { + writeError(w, http.StatusBadRequest, "bad_request", "Invalid request body") + return + } + relayURL, err := normalizePushRelayURL(req.RelayURL) + if err != nil { + writeError(w, http.StatusBadRequest, "bad_request", err.Error()) + return + } + + existingDeploymentID, err := h.settings.Get(r.Context(), notifications.SettingPushRelayDeploymentID) + if err != nil { + writeError(w, http.StatusInternalServerError, "settings_error", "Failed to load relay deployment id") + return + } + relayResp, err := h.registerWithRelay(r, relayURL, strings.TrimSpace(existingDeploymentID)) + if err != nil { + status, code, message := mapRelayRegistrationError(err) + writeError(w, status, code, message) + return + } + if strings.TrimSpace(relayResp.DeploymentID) == "" || strings.TrimSpace(relayResp.APIKey) == "" { + writeError(w, http.StatusBadGateway, "relay_bad_response", "Push relay returned an incomplete registration response") + return + } + + if err := h.settings.Set(r.Context(), notifications.SettingPushRelayURL, relayURL); err != nil { + writeError(w, http.StatusInternalServerError, "settings_error", "Failed to save push relay URL") + return + } + if err := h.settings.Set(r.Context(), notifications.SettingPushRelayDeploymentID, relayResp.DeploymentID); err != nil { + writeError(w, http.StatusInternalServerError, "settings_error", "Failed to save push relay deployment id") + return + } + if err := h.settings.Set(r.Context(), notifications.SettingPushRelayAPIKey, relayResp.APIKey); err != nil { + writeError(w, http.StatusInternalServerError, "settings_error", "Failed to save push relay API key") + return + } + if h.system != nil && h.system.Settings != nil { + h.system.Settings.Invalidate( + notifications.SettingPushRelayURL, + notifications.SettingPushRelayDeploymentID, + notifications.SettingPushRelayAPIKey, + ) + } + writeJSON(w, http.StatusOK, adminPushRelayRegisterResponse{ + RelayURL: relayURL, + DeploymentID: relayResp.DeploymentID, + KeyPrefix: relayResp.KeyPrefix, + APIKeyConfigured: true, + RelayRequestID: relayResp.RequestID, + APNsTopics: relayResp.APNsTopics, + }) +} + +func (h *AdminApplePushHandler) registerWithRelay(r *http.Request, relayURL, deploymentID string) (*pushRelayRegisterResponse, error) { + body, err := json.Marshal(pushRelayRegisterRequest{ + DeploymentID: deploymentID, + }) + if err != nil { + return nil, err + } + httpReq, err := http.NewRequestWithContext(r.Context(), http.MethodPost, relayURL+"/v1/deployments/register", bytes.NewReader(body)) + if err != nil { + return nil, err + } + httpReq.Header.Set("Content-Type", "application/json") + httpReq.Header.Set("User-Agent", "Silo-Server/PushRelayRegistration") + + resp, err := h.client.Do(httpReq) + if err != nil { + return nil, relayRegistrationError{status: http.StatusBadGateway, code: "relay_unreachable", message: "Push relay could not be reached"} + } + defer func() { _ = resp.Body.Close() }() + + data, _ := io.ReadAll(io.LimitReader(resp.Body, 16<<10)) + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + var parsed pushRelayNestedError + _ = json.Unmarshal(data, &parsed) + code := parsed.Error.Code + if code == "" { + code = fmt.Sprintf("relay_http_%d", resp.StatusCode) + } + message := parsed.Error.Message + if message == "" { + message = http.StatusText(resp.StatusCode) + } + return nil, relayRegistrationError{status: resp.StatusCode, code: code, message: message} + } + var parsed pushRelayRegisterResponse + if err := json.Unmarshal(data, &parsed); err != nil { + return nil, relayRegistrationError{status: http.StatusBadGateway, code: "relay_bad_response", message: "Push relay returned invalid JSON"} + } + return &parsed, nil +} + +type relayRegistrationError struct { + status int + code string + message string +} + +func (e relayRegistrationError) Error() string { + return e.code +} + +func mapRelayRegistrationError(err error) (int, string, string) { + var relayErr relayRegistrationError + if !errors.As(err, &relayErr) { + return http.StatusInternalServerError, "internal_error", "Failed to register push relay" + } + switch relayErr.status { + case http.StatusForbidden: + return http.StatusUnprocessableEntity, "relay_deployment_rejected", "Push relay rejected this deployment" + case http.StatusTooManyRequests: + return http.StatusTooManyRequests, "relay_rate_limited", relayErr.message + case http.StatusServiceUnavailable: + return http.StatusServiceUnavailable, "relay_unavailable", relayErr.message + default: + if relayErr.status >= 500 { + return http.StatusBadGateway, "relay_error", relayErr.message + } + return http.StatusBadGateway, relayErr.code, relayErr.message + } +} + +// normalizePushRelayURLValue trims a relay URL and enforces https with a +// host. Empty input stays empty so callers choose their own default. +func normalizePushRelayURLValue(raw string) (string, error) { + value := strings.TrimRight(strings.TrimSpace(raw), "/") + if value == "" { + return "", nil + } + parsed, err := url.Parse(value) + if err != nil || parsed.Scheme != "https" || parsed.Host == "" { + return "", errors.New("relay_url must be an https URL") + } + return value, nil +} + +func normalizePushRelayURL(raw string) (string, error) { + value, err := normalizePushRelayURLValue(raw) + if err != nil { + return "", err + } + if value == "" { + return notifications.DefaultPushRelayURL, nil + } + return value, nil +} diff --git a/internal/api/handlers/admin_apple_push_test.go b/internal/api/handlers/admin_apple_push_test.go new file mode 100644 index 00000000..bd97dc72 --- /dev/null +++ b/internal/api/handlers/admin_apple_push_test.go @@ -0,0 +1,123 @@ +package handlers + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/Silo-Server/silo-server/internal/notifications" +) + +func TestAdminApplePushHandlerUnavailableWithoutSystem(t *testing.T) { + h := NewAdminApplePushHandler(nil, nil) + rec := httptest.NewRecorder() + h.HandleTest(rec, httptest.NewRequest(http.MethodPost, "/admin/notifications/push/apple/test", strings.NewReader(`{}`))) + if rec.Code != http.StatusServiceUnavailable { + t.Fatalf("HandleTest without system = %d, want 503", rec.Code) + } +} + +func TestAdminApplePushHandlerRejectsInvalidJSON(t *testing.T) { + h := NewAdminApplePushHandler(¬ifications.System{}, nil) + rec := httptest.NewRecorder() + h.HandleTest(rec, httptest.NewRequest(http.MethodPost, "/admin/notifications/push/apple/test", strings.NewReader(`{`))) + if rec.Code != http.StatusBadRequest { + t.Fatalf("HandleTest invalid JSON = %d, want 400", rec.Code) + } +} + +func TestAdminApplePushHandlerRegistersRelayAndStoresKey(t *testing.T) { + settings := &fakeServerSettingsStore{values: map[string]string{ + notifications.SettingPushRelayDeploymentID: "01EXISTING", + }} + var relayReq map[string]string + relay := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/v1/deployments/register" { + t.Fatalf("relay path = %s", r.URL.Path) + } + if err := json.NewDecoder(r.Body).Decode(&relayReq); err != nil { + t.Fatalf("decode relay request: %v", err) + } + writeJSON(w, http.StatusOK, pushRelayRegisterResponse{ + RequestID: "relay-request", + DeploymentID: "01RETURNED", + APIKey: "rk_live_raw-key", + KeyPrefix: "rk_live_raw", + APNsTopics: []string{"org.siloserver.silo"}, + }) + })) + t.Cleanup(relay.Close) + + h := NewAdminApplePushHandler(¬ifications.System{Settings: notifications.NewSettings(settings)}, settings) + h.client = relay.Client() + + rec := httptest.NewRecorder() + h.HandleRegisterRelay(rec, httptest.NewRequest(http.MethodPost, "/admin/notifications/push/relay/register", strings.NewReader(`{ + "relay_url":"`+relay.URL+`" + }`))) + + if rec.Code != http.StatusBadRequest { + t.Fatalf("http relay URL status = %d, want 400 because HTTPS is required", rec.Code) + } + + httpsRelay := httptest.NewTLSServer(relay.Config.Handler) + t.Cleanup(httpsRelay.Close) + h.client = httpsRelay.Client() + rec = httptest.NewRecorder() + h.HandleRegisterRelay(rec, httptest.NewRequest(http.MethodPost, "/admin/notifications/push/relay/register", strings.NewReader(`{ + "relay_url":"`+httpsRelay.URL+`" + }`))) + + if rec.Code != http.StatusOK { + t.Fatalf("status = %d (%s), want 200", rec.Code, rec.Body.String()) + } + if relayReq["deployment_id"] != "01EXISTING" || len(relayReq) != 1 { + t.Fatalf("relay request = %+v", relayReq) + } + if settings.values[notifications.SettingPushRelayURL] != httpsRelay.URL { + t.Fatalf("stored relay URL = %q", settings.values[notifications.SettingPushRelayURL]) + } + if settings.values[notifications.SettingPushRelayDeploymentID] != "01RETURNED" { + t.Fatalf("stored deployment id = %q", settings.values[notifications.SettingPushRelayDeploymentID]) + } + if settings.values[notifications.SettingPushRelayAPIKey] != "rk_live_raw-key" { + t.Fatalf("stored api key = %q", settings.values[notifications.SettingPushRelayAPIKey]) + } + + var body map[string]any + if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil { + t.Fatalf("decode response: %v", err) + } + if _, ok := body["api_key"]; ok { + t.Fatal("response leaked raw api_key") + } + if body["api_key_configured"] != true || body["deployment_id"] != "01RETURNED" { + t.Fatalf("response = %+v", body) + } +} + +func TestAdminApplePushHandlerMapsRelayRateLimit(t *testing.T) { + settings := &fakeServerSettingsStore{values: map[string]string{}} + relay := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + writeJSON(w, http.StatusTooManyRequests, map[string]any{ + "error": map[string]string{"code": "rate_limited", "message": "too many deployment registrations from this network"}, + }) + })) + t.Cleanup(relay.Close) + + h := NewAdminApplePushHandler(¬ifications.System{}, settings) + h.client = relay.Client() + rec := httptest.NewRecorder() + h.HandleRegisterRelay(rec, httptest.NewRequest(http.MethodPost, "/admin/notifications/push/relay/register", strings.NewReader(`{ + "relay_url":"`+relay.URL+`" + }`))) + + if rec.Code != http.StatusTooManyRequests { + t.Fatalf("status = %d (%s), want 429", rec.Code, rec.Body.String()) + } + if settings.values[notifications.SettingPushRelayAPIKey] != "" { + t.Fatal("api key was stored after relay rejected registration") + } +} diff --git a/internal/api/handlers/notifications.go b/internal/api/handlers/notifications.go index 7a4d351c..f062f22d 100644 --- a/internal/api/handlers/notifications.go +++ b/internal/api/handlers/notifications.go @@ -43,6 +43,8 @@ type notificationSyncResponse struct { UnreadCount int `json:"unread_count"` } +type notificationApplePushDisplayResponse = notifications.NotificationDisplay + type unreadCountResponse struct { Count int `json:"count"` } @@ -150,6 +152,29 @@ func (h *NotificationsHandler) HandleGet(w http.ResponseWriter, r *http.Request) writeJSON(w, http.StatusOK, h.system.PayloadForRow(r.Context(), *row)) } +// HandleApplePushDisplay handles GET /notifications/push/apple/display/{delivery_id}. +// It returns only compact display metadata for notification-service-extension +// enrichment, scoped to the active profile. +func (h *NotificationsHandler) HandleApplePushDisplay(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Cache-Control", "no-store") + profileID := apimw.GetProfileID(r.Context()) + id := chi.URLParam(r, "delivery_id") + if id == "" { + writeError(w, http.StatusNotFound, "not_found", "Notification not found") + return + } + row, err := h.system.Deliveries.GetByID(r.Context(), profileID, id) + if err != nil { + writeError(w, http.StatusInternalServerError, "internal_error", "Failed to load notification") + return + } + if row == nil { + writeError(w, http.StatusNotFound, "not_found", "Notification not found") + return + } + writeJSON(w, http.StatusOK, notificationApplePushDisplayResponse(notifications.BuildNotificationDisplay(*row))) +} + // HandleUnreadCount handles GET /notifications/unread-count. func (h *NotificationsHandler) HandleUnreadCount(w http.ResponseWriter, r *http.Request) { profileID := apimw.GetProfileID(r.Context()) @@ -318,9 +343,7 @@ type capabilityWebhooks struct { } // HandleCapability handles GET /notifications/capability. Clients render -// setup UI from this response instead of introspecting admin settings. Push -// channels report unavailable until they ship -// (docs/superpowers/plans/notifications/02-03). +// setup UI from this response instead of introspecting admin settings. func (h *NotificationsHandler) HandleCapability(w http.ResponseWriter, r *http.Request) { webhooks := capabilityWebhooks{Available: false, MaxPerProfile: 0, SupportedTypes: []string{}} if h.system.Webhooks != nil && h.system.Settings.WebhooksEnabled(r.Context()) { @@ -364,9 +387,21 @@ func (h *NotificationsHandler) HandleCapability(w http.ResponseWriter, r *http.R DigestHour: h.system.Settings.DiscordDigestHour(r.Context()), } } + applePush := capabilityPush{Available: false, Provider: "off", SupportedModes: []string{"in_app_only"}} + // Like web push above, availability requires both the wiring (cipher/store) + // and the admin delivery toggle: Available must mean "setup will actually + // deliver", not "the server could store a token". + if h.system.PushDevices != nil && h.system.PushDevices.Available() && + h.system.Settings.ApplePushDeliveryEnabled(r.Context()) { + applePush = capabilityPush{ + Available: true, + Provider: notifications.PushProviderSiloRelay, + SupportedModes: []string{notifications.PushModePrivatePush, notifications.PushModeInAppOnly}, + } + } writeJSON(w, http.StatusOK, capabilityResponse{ InApp: capabilityInApp{Enabled: h.system.Settings.UIEnabled(r.Context())}, - ApplePush: capabilityPush{Available: false, Provider: "off", SupportedModes: []string{"in_app_only"}}, + ApplePush: applePush, AndroidPush: capabilityPush{Available: false, Provider: "off", SupportedModes: []string{"in_app_only"}}, WebPush: webPush, Webhooks: webhooks, diff --git a/internal/api/handlers/notifications_apple_display_test.go b/internal/api/handlers/notifications_apple_display_test.go new file mode 100644 index 00000000..2e94a188 --- /dev/null +++ b/internal/api/handlers/notifications_apple_display_test.go @@ -0,0 +1,134 @@ +package handlers + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "os" + "strings" + "testing" + + apimw "github.com/Silo-Server/silo-server/internal/api/middleware" + "github.com/Silo-Server/silo-server/internal/notifications" + "github.com/go-chi/chi/v5" + "github.com/jackc/pgx/v5/pgxpool" +) + +func TestHandleApplePushDisplayDB(t *testing.T) { + dsn := os.Getenv("SILO_TEST_DATABASE_URL") + if dsn == "" { + t.Skip("set SILO_TEST_DATABASE_URL to run DB-backed notification display handler test") + } + + ctx := context.Background() + config, err := pgxpool.ParseConfig(dsn) + if err != nil { + t.Fatalf("parse db config: %v", err) + } + config.MaxConns = 1 + pool, err := pgxpool.NewWithConfig(ctx, config) + if err != nil { + t.Fatalf("connect db: %v", err) + } + t.Cleanup(pool.Close) + + if _, err := pool.Exec(ctx, ` + CREATE TEMP TABLE notification_deliveries ( + id text PRIMARY KEY, + release_event_id text, + user_id integer NOT NULL, + profile_id text NOT NULL, + library_id integer, + series_id text, + episode_id text, + type text NOT NULL, + reason_flags jsonb NOT NULL, + status text NOT NULL DEFAULT 'delivered', + read_at timestamptz, + delivered_at timestamptz, + created_at timestamptz NOT NULL DEFAULT now() + ); + CREATE TEMP TABLE episodes ( + content_id text PRIMARY KEY, + title text, + season_number integer, + episode_number integer, + overview text + ); + CREATE TEMP TABLE media_items ( + content_id text PRIMARY KEY, + title text, + poster_path text, + poster_thumbhash text, + poster_source_path text, + type text, + year integer, + overview text, + genres text[], + content_rating text, + rating_imdb double precision, + rating_tmdb double precision, + imdb_id text, + tmdb_id text, + tvdb_id text + ); + INSERT INTO media_items (content_id, title, type, genres) + VALUES ('series-1', 'Severance', 'series', ARRAY[]::text[]); + INSERT INTO episodes (content_id, title, season_number, episode_number) + VALUES ('episode-1', 'Hello, Ms. Cobel', 2, 1); + INSERT INTO notification_deliveries ( + id, release_event_id, user_id, profile_id, library_id, series_id, episode_id, + type, reason_flags, status + ) VALUES ( + 'delivery-1', 'event-1', 42, 'profile-1', 7, 'series-1', 'episode-1', + 'episode.available', '{"favorite":true}', 'delivered' + ); + `); err != nil { + t.Fatalf("seed temp tables: %v", err) + } + + handler := NewNotificationsHandler(¬ifications.System{ + Deliveries: notifications.NewDeliveryRepository(pool), + }, nil) + router := chi.NewRouter() + router.Get("/notifications/push/apple/display/{delivery_id}", handler.HandleApplePushDisplay) + + req := httptest.NewRequest(http.MethodGet, "/notifications/push/apple/display/delivery-1", nil) + req = req.WithContext(apimw.SetProfileID(req.Context(), "profile-1")) + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, body = %s", rr.Code, rr.Body.String()) + } + if got := rr.Header().Get("Cache-Control"); got != "no-store" { + t.Fatalf("Cache-Control = %q", got) + } + var response notifications.NotificationDisplay + if err := json.NewDecoder(rr.Body).Decode(&response); err != nil { + t.Fatalf("decode response: %v", err) + } + if response.DeliveryID != "delivery-1" || + response.Title != "The latest episode of Severance S02E01 just dropped!" || + response.Body != "Hello, Ms. Cobel" || + response.ThreadID != "series:series-1" || + response.Category != "episode_available" || + response.URL != "/item/episode-1" { + t.Fatalf("response = %+v", response) + } + + req = httptest.NewRequest(http.MethodGet, "/notifications/push/apple/display/delivery-1", nil) + req = req.WithContext(apimw.SetProfileID(req.Context(), "other-profile")) + rr = httptest.NewRecorder() + router.ServeHTTP(rr, req) + if rr.Code != http.StatusNotFound { + t.Fatalf("cross-profile status = %d, body = %s", rr.Code, rr.Body.String()) + } + if got := rr.Header().Get("Cache-Control"); got != "no-store" { + t.Fatalf("cross-profile Cache-Control = %q", got) + } + if strings.Contains(rr.Body.String(), "Severance") { + t.Fatalf("cross-profile response leaked display metadata: %s", rr.Body.String()) + } +} diff --git a/internal/api/handlers/notifications_push_devices.go b/internal/api/handlers/notifications_push_devices.go new file mode 100644 index 00000000..8780646f --- /dev/null +++ b/internal/api/handlers/notifications_push_devices.go @@ -0,0 +1,75 @@ +package handlers + +import ( + "encoding/json" + "errors" + "net/http" + + apimw "github.com/Silo-Server/silo-server/internal/api/middleware" + "github.com/Silo-Server/silo-server/internal/notifications" +) + +type applePushRegisterRequest struct { + DeviceID string `json:"device_id"` + APNsToken string `json:"apns_token"` + APNsEnvironment string `json:"apns_environment"` + APNsTopic string `json:"apns_topic"` + PushMode string `json:"push_mode"` +} + +type applePushRegisterResponse struct { + ID string `json:"id"` + ServerDeviceID string `json:"server_device_id"` + Enabled bool `json:"enabled"` + PushMode string `json:"push_mode"` +} + +func (h *NotificationsHandler) pushDevices() *notifications.PushDeviceService { + if h == nil || h.system == nil { + return nil + } + return h.system.PushDevices +} + +// HandleRegisterApplePushDevice handles POST /devices/push/apple. +func (h *NotificationsHandler) HandleRegisterApplePushDevice(w http.ResponseWriter, r *http.Request) { + service := h.pushDevices() + if service == nil || !service.Available() { + writeError(w, http.StatusServiceUnavailable, "unavailable", "Apple push registration is not available") + return + } + + var req applePushRegisterRequest + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeError(w, http.StatusBadRequest, "bad_request", "Invalid request body") + return + } + + device, err := service.RegisterApple(r.Context(), apimw.GetUserID(r.Context()), apimw.GetProfileID(r.Context()), notifications.ApplePushRegistrationInput{ + DeviceID: req.DeviceID, + APNsToken: req.APNsToken, + APNsEnvironment: req.APNsEnvironment, + APNsTopic: req.APNsTopic, + PushMode: req.PushMode, + }) + if err != nil { + switch { + case errors.Is(err, notifications.ErrPushDeviceInvalid): + writeError(w, http.StatusBadRequest, "bad_request", err.Error()) + case errors.Is(err, notifications.ErrPushDeviceUnsupported): + writeError(w, http.StatusUnprocessableEntity, "unsupported_push_device", err.Error()) + case errors.Is(err, notifications.ErrPushDeviceUnavailable): + writeError(w, http.StatusServiceUnavailable, "unavailable", "Apple push registration is not available") + default: + writeError(w, http.StatusInternalServerError, "internal_error", "Failed to register Apple push device") + } + return + } + + writeJSON(w, http.StatusOK, applePushRegisterResponse{ + ID: device.ID, + ServerDeviceID: device.ServerDeviceID, + Enabled: device.Enabled, + PushMode: device.PushMode, + }) +} diff --git a/internal/api/handlers/notifications_push_devices_test.go b/internal/api/handlers/notifications_push_devices_test.go new file mode 100644 index 00000000..1cac7070 --- /dev/null +++ b/internal/api/handlers/notifications_push_devices_test.go @@ -0,0 +1,152 @@ +package handlers + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + + apimw "github.com/Silo-Server/silo-server/internal/api/middleware" + "github.com/Silo-Server/silo-server/internal/auth" + "github.com/Silo-Server/silo-server/internal/notifications" + "github.com/Silo-Server/silo-server/internal/secret" +) + +type handlerPushStore struct { + got notifications.ApplePushDeviceRegistration + calls int + device *notifications.PushDevice + err error +} + +func (f *handlerPushStore) UpsertApple(ctx context.Context, registration notifications.ApplePushDeviceRegistration, cipher *secret.Cipher) (*notifications.PushDevice, error) { + f.calls++ + f.got = registration + if f.err != nil { + return nil, f.err + } + if f.device != nil { + return f.device, nil + } + return ¬ifications.PushDevice{ + ID: "push-row", + ServerDeviceID: "server-device", + Enabled: true, + PushMode: registration.PushMode, + }, nil +} + +func handlerPushCipher(t *testing.T) *secret.Cipher { + t.Helper() + cipher, err := secret.New([]byte("01234567890123456789012345678901")) + if err != nil { + t.Fatalf("new cipher: %v", err) + } + return cipher +} + +func newApplePushRequest(body string) *http.Request { + req := httptest.NewRequest(http.MethodPost, "/api/v1/devices/push/apple", strings.NewReader(body)) + ctx := apimw.SetClaims(req.Context(), &auth.Claims{UserID: 42, Role: "user", TokenType: auth.TokenTypeAccess}) + ctx = apimw.SetProfileID(ctx, "profile-1") + return req.WithContext(ctx) +} + +func TestHandleRegisterApplePushDevice(t *testing.T) { + store := &handlerPushStore{} + handler := NewNotificationsHandler(¬ifications.System{ + PushDevices: notifications.NewPushDeviceService(store, handlerPushCipher(t)), + }, nil) + + body := `{ + "device_id":"local-device", + "apns_token":"` + strings.Repeat("a", 64) + `", + "apns_environment":"production", + "apns_topic":"org.siloserver.silo", + "push_mode":"private_push" + }` + rr := httptest.NewRecorder() + handler.HandleRegisterApplePushDevice(rr, newApplePushRequest(body)) + + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, body = %s", rr.Code, rr.Body.String()) + } + var response applePushRegisterResponse + if err := json.NewDecoder(rr.Body).Decode(&response); err != nil { + t.Fatalf("decode response: %v", err) + } + if response.ID != "push-row" || response.ServerDeviceID != "server-device" || !response.Enabled || response.PushMode != notifications.PushModePrivatePush { + t.Fatalf("unexpected response: %+v", response) + } + if store.calls != 1 { + t.Fatalf("store calls = %d, want 1", store.calls) + } + if store.got.UserID != 42 || store.got.ProfileID != "profile-1" || store.got.DeviceID != "local-device" { + t.Fatalf("unexpected stored registration: %+v", store.got) + } +} + +func TestHandleRegisterApplePushDeviceErrors(t *testing.T) { + tests := []struct { + name string + system *notifications.System + body string + wantStatus int + }{ + { + name: "service unavailable", + system: ¬ifications.System{}, + body: `{}`, + wantStatus: http.StatusServiceUnavailable, + }, + { + name: "invalid json", + system: ¬ifications.System{ + PushDevices: notifications.NewPushDeviceService(&handlerPushStore{}, handlerPushCipher(t)), + }, + body: `{`, + wantStatus: http.StatusBadRequest, + }, + { + name: "invalid field", + system: ¬ifications.System{ + PushDevices: notifications.NewPushDeviceService(&handlerPushStore{}, handlerPushCipher(t)), + }, + body: `{ + "device_id":"local-device", + "apns_token":"abcd", + "apns_environment":"production", + "apns_topic":"org.siloserver.silo", + "push_mode":"private_push" + }`, + wantStatus: http.StatusBadRequest, + }, + { + name: "unsupported topic", + system: ¬ifications.System{ + PushDevices: notifications.NewPushDeviceService(&handlerPushStore{}, handlerPushCipher(t)), + }, + body: `{ + "device_id":"local-device", + "apns_token":"` + strings.Repeat("a", 64) + `", + "apns_environment":"production", + "apns_topic":"com.example.app", + "push_mode":"private_push" + }`, + wantStatus: http.StatusUnprocessableEntity, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + handler := NewNotificationsHandler(tt.system, nil) + rr := httptest.NewRecorder() + handler.HandleRegisterApplePushDevice(rr, newApplePushRequest(tt.body)) + if rr.Code != tt.wantStatus { + t.Fatalf("status = %d, want %d, body = %s", rr.Code, tt.wantStatus, rr.Body.String()) + } + }) + } +} diff --git a/internal/api/router.go b/internal/api/router.go index 3127d6d5..06a1cda3 100644 --- a/internal/api/router.go +++ b/internal/api/router.go @@ -1626,6 +1626,7 @@ func NewRouter(deps Dependencies) chi.Router { } notificationsHandler := handlers.NewNotificationsHandler(deps.Notifications, deps.EventsHub) r.With(apimw.RequireProfile).Post("/events/ws-ticket", notificationsHandler.HandleMintWSTicket) + r.With(apimw.RequireProfile).Post("/devices/push/apple", notificationsHandler.HandleRegisterApplePushDevice) // Discord DM channel: the linked identity and mode hang off // the login account, not a profile, so these stay outside // the RequireProfile subrouter below (static paths coexist @@ -1644,6 +1645,7 @@ func NewRouter(deps Dependencies) chi.Router { r.Get("/capability", notificationsHandler.HandleCapability) r.Get("/preferences", notificationsHandler.HandleGetPreferences) r.Put("/preferences", notificationsHandler.HandleUpdatePreferences) + r.Get("/push/apple/display/{delivery_id}", notificationsHandler.HandleApplePushDisplay) r.Get("/email-preferences", notificationsHandler.HandleGetEmailPreferences) r.Put("/email-preferences", notificationsHandler.HandleUpdateEmailPreferences) r.Put("/email-preferences/address", notificationsHandler.HandleRequestEmailAddress) @@ -2303,6 +2305,15 @@ func NewRouter(deps Dependencies) chi.Router { if discordNotificationsHandler != nil { r.Post("/notifications/discord/test", discordNotificationsHandler.HandleAdminTest) } + if deps.Notifications != nil || settingsRepo != nil { + applePushHandler := handlers.NewAdminApplePushHandler(deps.Notifications, settingsRepo) + if deps.Notifications != nil { + r.Post("/notifications/push/apple/test", applePushHandler.HandleTest) + } + if settingsRepo != nil { + r.Post("/notifications/push/relay/register", applePushHandler.HandleRegisterRelay) + } + } if deps.Notifications != nil && deps.Notifications.ServerChannels != nil { serverChannelsHandler := handlers.NewAdminServerChannelsHandler(deps.Notifications) r.Route("/notifications/server-channels", func(r chi.Router) { diff --git a/internal/catalog/encrypted_settings_repo.go b/internal/catalog/encrypted_settings_repo.go index 7bf40ebe..aaad3581 100644 --- a/internal/catalog/encrypted_settings_repo.go +++ b/internal/catalog/encrypted_settings_repo.go @@ -108,6 +108,9 @@ var SensitiveSettingKeys = map[string]bool{ // value by the notifications system; clients receive the public half via // the capability endpoint, never from the settings store). "notifications.web_push.vapid_keypair": true, + + // Silo push relay bearer credential for APNs/FCM delivery. + "notifications.push_relay_api_key": true, } // EncryptedSettingsRepo decorates a raw settings store, transparently diff --git a/internal/catalog/encrypted_settings_repo_test.go b/internal/catalog/encrypted_settings_repo_test.go index 44788228..8f9dfe9b 100644 --- a/internal/catalog/encrypted_settings_repo_test.go +++ b/internal/catalog/encrypted_settings_repo_test.go @@ -158,6 +158,7 @@ func TestSensitiveSettingKeys_Audited(t *testing.T) { "discord.client_secret", "discord.bot_token", "notifications.web_push.vapid_keypair", + "notifications.push_relay_api_key", } for _, k := range mustHave { if !SensitiveSettingKeys[k] { diff --git a/internal/notifications/channel_dispatcher.go b/internal/notifications/channel_dispatcher.go new file mode 100644 index 00000000..db01dcc4 --- /dev/null +++ b/internal/notifications/channel_dispatcher.go @@ -0,0 +1,129 @@ +package notifications + +import ( + "context" + "log/slog" + "sync" + "time" +) + +// channelDispatcher is the shared dispatch core behind the per-channel +// Dispatcher implementations (webhooks, web push, Apple push). dispatch never +// blocks the fanout loop on destination I/O: it hands the delivery ID to a +// bounded worker pool that claims the delivery's pending outbox attempts and +// sends them. A full queue simply drops the hand-off — the durable pending +// rows are picked up by the retry sweep, so delivery is delayed, never lost. +type channelDispatcher[A any] struct { + channel string // log label, e.g. "web push" + queue chan string + logger *slog.Logger + claimPending func(ctx context.Context, deliveryID string) ([]A, error) + process func(ctx context.Context, attempt A) + + // Optional integrated retry sweep: when claimDue is set, run also drains + // due retries and recovers stale pending rows, skipping ticks while + // enabled reports the channel off. Webhooks instead run the standalone + // WebhookRetryWorker, whose lifecycle the system manages separately. + enabled func(ctx context.Context) bool + claimDue func(ctx context.Context, limit int) ([]A, error) + claimLimit int +} + +// dispatch queues the delivery's attempts for immediate send without blocking. +func (d *channelDispatcher[A]) dispatch(deliveryID string) { + select { + case d.queue <- deliveryID: + default: + d.logger.Warn(d.channel+" dispatch queue full; deferring to retry worker", + "delivery_id", deliveryID) + } +} + +// run consumes the dispatch queue with a bounded worker pool (plus the retry +// sweep when configured) until ctx is canceled. One slow destination cannot +// block other deliveries. +func (d *channelDispatcher[A]) run(ctx context.Context) { + var wg sync.WaitGroup + for range webhookDispatchWorkers { + wg.Add(1) + go func() { + defer wg.Done() + for { + select { + case <-ctx.Done(): + return + case deliveryID := <-d.queue: + d.processDelivery(ctx, deliveryID) + } + } + }() + } + if d.claimDue != nil { + wg.Add(1) + go func() { + defer wg.Done() + runRetrySweep(ctx, d.channel, d.logger, d.enabled, d.claimDue, d.claimLimit, d.process) + }() + } + wg.Wait() +} + +func (d *channelDispatcher[A]) processDelivery(ctx context.Context, deliveryID string) { + attempts, err := d.claimPending(ctx, deliveryID) + if err != nil { + if ctx.Err() == nil { + d.logger.Warn(d.channel+" attempt claim failed", "delivery_id", deliveryID, "error", err) + } + return + } + for _, attempt := range attempts { + if ctx.Err() != nil { + return + } + d.process(ctx, attempt) + } +} + +// runRetrySweep polls for due attempts until ctx is canceled, draining every +// due batch each tick. Shared by the integrated dispatcher sweeps and the +// standalone webhook retry worker. +func runRetrySweep[A any]( + ctx context.Context, + channel string, + logger *slog.Logger, + enabled func(context.Context) bool, + claim func(context.Context, int) ([]A, error), + limit int, + process func(context.Context, A), +) { + ticker := time.NewTicker(webhookRetryInterval) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + } + if !enabled(ctx) { + continue + } + for { + attempts, err := claim(ctx, limit) + if err != nil { + if ctx.Err() == nil { + logger.Warn(channel+" retry claim failed", "error", err) + } + break + } + if len(attempts) == 0 { + break + } + for _, attempt := range attempts { + if ctx.Err() != nil { + return + } + process(ctx, attempt) + } + } + } +} diff --git a/internal/notifications/display.go b/internal/notifications/display.go new file mode 100644 index 00000000..dcdb2779 --- /dev/null +++ b/internal/notifications/display.go @@ -0,0 +1,139 @@ +package notifications + +import ( + "fmt" + "strings" +) + +// NotificationDisplay is the compact, user-facing display metadata clients use +// when they need a native notification title/body without fetching a full inbox +// page. +type NotificationDisplay struct { + DeliveryID string `json:"delivery_id"` + Title string `json:"title"` + Body string `json:"body,omitempty"` + ThreadID string `json:"thread_id,omitempty"` + Category string `json:"category"` + URL string `json:"url"` +} + +// BuildNotificationDisplay renders a delivery into native-notification display +// metadata. Keep this in sync with inbox/websocket semantics by only deriving +// from DeliveryRow. +func BuildNotificationDisplay(row DeliveryRow) NotificationDisplay { + display := NotificationDisplay{ + DeliveryID: row.ID, + Title: genericNotificationTitle, + Category: "notification", + URL: "/notifications", + } + switch row.Type { + case DeliveryTypeEpisodeAvailable: + display.Category = "episode_available" + display.Title = episodeDisplayTitle(row) + display.Body = episodeDisplayBody(row) + if row.SeriesID != nil && *row.SeriesID != "" { + display.ThreadID = "series:" + *row.SeriesID + } + if row.EpisodeID != nil && *row.EpisodeID != "" { + display.URL = "/item/" + *row.EpisodeID + } + case DeliveryTypeRequestFulfilled: + display.Category = "request_fulfilled" + display.Title = "Your request is now available" + if row.SeriesTitle != "" { + display.Title = row.SeriesTitle + " is now available" + } + display.Body = "Your media request has arrived in the library." + flags := parseRequestFlags(row.ReasonFlags) + if flags.RequestID != "" { + display.ThreadID = "request:" + flags.RequestID + } else if row.SeriesID != nil && *row.SeriesID != "" { + display.ThreadID = "item:" + *row.SeriesID + } + if row.SeriesID != nil && *row.SeriesID != "" { + display.URL = "/item/" + *row.SeriesID + } + case DeliveryTypeRequestApproved: + flags := parseRequestFlags(row.ReasonFlags) + display.Category = "request_approved" + display.Title = "Your request was approved" + if flags.Title != "" { + display.Title = flags.Title + " was approved" + } + display.Body = "Your media request was approved." + display.ThreadID = requestThreadID(flags) + case DeliveryTypeRequestDeclined: + flags := parseRequestFlags(row.ReasonFlags) + display.Category = "request_declined" + display.Title = "Your request was declined" + if flags.Title != "" { + display.Title = flags.Title + " was declined" + } + display.Body = "Your media request was declined." + if flags.Reason != "" { + // Admin-typed free text with no upstream length cap; keep native + // notification bodies bounded. + display.Body = "Reason: " + truncateDisplayText(flags.Reason, displayBodyMaxLen) + } + display.ThreadID = requestThreadID(flags) + case DeliveryTypeWebhookAutoDisabled: + display.Category = "webhook_auto_disabled" + display.Title = "A webhook stopped working" + display.Body = "Open notification settings to fix it." + display.ThreadID = "settings:notifications" + display.URL = "/settings/notifications" + default: + if row.Type != "" { + display.Category = strings.ReplaceAll(row.Type, ".", "_") + } + } + return display +} + +func episodeDisplayTitle(row DeliveryRow) string { + code := episodeDisplayCode(row) + switch { + case row.SeriesTitle != "" && code != "": + return fmt.Sprintf("The latest episode of %s %s just dropped!", row.SeriesTitle, code) + case row.SeriesTitle != "": + return "The latest episode of " + row.SeriesTitle + " just dropped!" + case code != "": + return "New episode " + code + " available" + default: + return "New episode available" + } +} + +func episodeDisplayBody(row DeliveryRow) string { + if row.EpisodeTitle != "" { + return row.EpisodeTitle + } + return episodeDisplayCode(row) +} + +func episodeDisplayCode(row DeliveryRow) string { + if row.SeasonNumber != nil && row.EpisodeNumber != nil { + return fmt.Sprintf("S%02dE%02d", *row.SeasonNumber, *row.EpisodeNumber) + } + return "" +} + +// displayBodyMaxLen bounds free-text notification bodies; matches the varchar +// caps used elsewhere in the delivery pipeline. +const displayBodyMaxLen = 240 + +func truncateDisplayText(s string, max int) string { + runes := []rune(s) + if len(runes) <= max { + return s + } + return strings.TrimSpace(string(runes[:max-1])) + "…" +} + +func requestThreadID(flags RequestFlags) string { + if flags.RequestID == "" { + return "" + } + return "request:" + flags.RequestID +} diff --git a/internal/notifications/fanout_worker.go b/internal/notifications/fanout_worker.go index 75eed316..6760b63a 100644 --- a/internal/notifications/fanout_worker.go +++ b/internal/notifications/fanout_worker.go @@ -40,6 +40,7 @@ type FanoutWorker struct { webhooks *WebhookRepository rateLimiter *profileRateLimiter webPush *WebPushRepository + pushDevices *PushDeviceRepository } // SetWebhookOutbox wires durable webhook attempt enqueueing into the fanout @@ -55,6 +56,12 @@ func (w *FanoutWorker) SetWebPushOutbox(webPush *WebPushRepository) { w.webPush = webPush } +// SetPushOutbox wires durable Apple push attempt enqueueing into the fanout +// transaction. +func (w *FanoutWorker) SetPushOutbox(pushDevices *PushDeviceRepository) { + w.pushDevices = pushDevices +} + // NewFanoutWorker creates a FanoutWorker. func NewFanoutWorker( pool *pgxpool.Pool, @@ -299,6 +306,9 @@ func (w *FanoutWorker) fanOutEvent(ctx context.Context, tx pgx.Tx, event Release if err := w.enqueueWebPushOutbox(ctx, tx, inserted); err != nil { return nil, 0, err } + if err := w.enqueuePushOutbox(ctx, tx, inserted); err != nil { + return nil, 0, err + } notifiedProfiles := make([]string, 0, len(inserted)) for _, row := range inserted { @@ -344,6 +354,31 @@ type pendingDelivery struct { flags ReasonFlags } +// enqueuePushOutbox inserts pending Apple push attempt rows for each newly +// inserted delivery and enabled private-push device. +func (w *FanoutWorker) enqueuePushOutbox(ctx context.Context, tx pgx.Tx, inserted []InsertedDelivery) error { + if w.pushDevices == nil || len(inserted) == 0 || !w.settings.ApplePushDeliveryEnabled(ctx) { + return nil + } + profileSet := make(map[string]struct{}, len(inserted)) + profileIDs := make([]string, 0, len(inserted)) + for _, row := range inserted { + if _, ok := profileSet[row.ProfileID]; !ok { + profileSet[row.ProfileID] = struct{}{} + profileIDs = append(profileIDs, row.ProfileID) + } + } + devicesByProfile, err := w.pushDevices.ListEnabledAppleByProfiles(ctx, tx, profileIDs) + if err != nil { + return err + } + attempts := make([]PushDeliveryAttempt, 0, len(inserted)) + for _, row := range inserted { + attempts = append(attempts, newPushDeliveryAttempts(row.ID, devicesByProfile[row.ProfileID])...) + } + return w.pushDevices.EnqueuePushAttempts(ctx, tx, attempts) +} + // enqueueWebPushOutbox inserts `pending` web push attempt rows for each newly // inserted delivery, inside the fanout transaction. Subscriptions have no // per-reason filters: profile preferences already gated delivery creation, diff --git a/internal/notifications/operational_dispatch.go b/internal/notifications/operational_dispatch.go index 4a5ab164..9edfe275 100644 --- a/internal/notifications/operational_dispatch.go +++ b/internal/notifications/operational_dispatch.go @@ -18,8 +18,8 @@ type OperationalDispatch struct { } // DispatchOperational durably creates one operational delivery. The inbox row -// and the per-target webhook / web push outbox rows commit in a single -// transaction — a crash afterwards delays channel sends instead of dropping +// and the per-target webhook / web push / Apple push outbox rows commit in a +// single transaction — a crash afterwards delays channel sends instead of dropping // them, because the retry workers recover pending outbox rows — then realtime // and channel dispatch run post-commit. Returns nil when the delivery deduped // away (the partial unique indexes make operational notices idempotent). @@ -79,6 +79,16 @@ func (s *System) DispatchOperational(ctx context.Context, delivery Delivery, opt return nil, err } } + if s.pushDeviceRepo != nil && s.Settings.ApplePushDeliveryEnabled(ctx) { + devicesByProfile, err := s.pushDeviceRepo.ListEnabledAppleByProfiles(ctx, tx, []string{delivery.ProfileID}) + if err != nil { + return nil, err + } + attempts := newPushDeliveryAttempts(row.ID, devicesByProfile[delivery.ProfileID]) + if err := s.pushDeviceRepo.EnqueuePushAttempts(ctx, tx, attempts); err != nil { + return nil, err + } + } if err := tx.Commit(ctx); err != nil { return nil, fmt.Errorf("commit operational dispatch: %w", err) } diff --git a/internal/notifications/operational_dispatch_test.go b/internal/notifications/operational_dispatch_test.go new file mode 100644 index 00000000..43e09741 --- /dev/null +++ b/internal/notifications/operational_dispatch_test.go @@ -0,0 +1,157 @@ +package notifications + +import ( + "context" + "io" + "log/slog" + "os" + "testing" + + "github.com/jackc/pgx/v5/pgxpool" +) + +func TestDispatchOperationalEnqueuesApplePushAttempts(t *testing.T) { + dsn := os.Getenv("SILO_TEST_DATABASE_URL") + if dsn == "" { + t.Skip("set SILO_TEST_DATABASE_URL to run DB-backed operational push dispatch test") + } + + ctx := context.Background() + config, err := pgxpool.ParseConfig(dsn) + if err != nil { + t.Fatalf("parse db config: %v", err) + } + config.MaxConns = 1 + pool, err := pgxpool.NewWithConfig(ctx, config) + if err != nil { + t.Fatalf("connect db: %v", err) + } + t.Cleanup(pool.Close) + + if _, err := pool.Exec(ctx, ` + CREATE TEMP TABLE notification_deliveries ( + id text PRIMARY KEY, + release_event_id text, + user_id integer NOT NULL, + profile_id text NOT NULL, + library_id integer, + series_id text, + episode_id text, + type text NOT NULL, + reason_flags jsonb NOT NULL DEFAULT '{}'::jsonb, + status text NOT NULL DEFAULT 'delivered', + read_at timestamptz, + delivered_at timestamptz, + created_at timestamptz NOT NULL DEFAULT now() + ) ON COMMIT PRESERVE ROWS; + + CREATE TEMP TABLE push_devices ( + id text PRIMARY KEY, + user_id integer NOT NULL, + profile_id text NOT NULL, + device_id varchar(128) NOT NULL, + platform text NOT NULL, + provider text NOT NULL, + apns_environment text, + apns_topic text, + apns_token_ciphertext text, + apns_token_hash text, + server_device_id text NOT NULL, + push_mode text NOT NULL DEFAULT 'private_push', + enabled boolean NOT NULL DEFAULT true, + last_seen_at timestamptz, + last_success_at timestamptz, + last_failure_at timestamptz, + last_failure_code text, + created_at timestamptz NOT NULL DEFAULT now(), + updated_at timestamptz NOT NULL DEFAULT now() + ) ON COMMIT PRESERVE ROWS; + + CREATE TEMP TABLE push_delivery_attempts ( + id text PRIMARY KEY, + notification_delivery_id text, + push_device_id text NOT NULL, + trigger_type text NOT NULL, + provider text NOT NULL, + platform text NOT NULL, + attempt_number integer NOT NULL DEFAULT 0, + attempted_at timestamptz, + next_retry_at timestamptz, + outcome text NOT NULL DEFAULT 'pending', + relay_request_id text, + upstream_status integer, + upstream_reason text, + failure_message text, + created_at timestamptz NOT NULL DEFAULT now(), + updated_at timestamptz NOT NULL DEFAULT now(), + UNIQUE (notification_delivery_id, push_device_id, trigger_type) + ) ON COMMIT PRESERVE ROWS; + `); err != nil { + t.Fatalf("create temp notification push tables: %v", err) + } + + if _, err := pool.Exec(ctx, ` + INSERT INTO push_devices + (id, user_id, profile_id, device_id, platform, provider, apns_environment, apns_topic, + apns_token_ciphertext, apns_token_hash, server_device_id, push_mode, enabled) + VALUES + ('device-private', 42, 'profile-1', 'local-private', 'apple', 'silo_relay', 'sandbox', + 'org.siloserver.silo', 'ciphertext', 'hash-1', 'server-private', 'private_push', true), + ('device-in-app', 42, 'profile-1', 'local-in-app', 'apple', 'silo_relay', 'sandbox', + 'org.siloserver.silo', 'ciphertext', 'hash-2', 'server-in-app', 'in_app_only', true), + ('device-disabled', 42, 'profile-1', 'local-disabled', 'apple', 'silo_relay', 'sandbox', + 'org.siloserver.silo', 'ciphertext', 'hash-3', 'server-disabled', 'private_push', false), + ('device-other-profile', 42, 'profile-2', 'local-other', 'apple', 'silo_relay', 'sandbox', + 'org.siloserver.silo', 'ciphertext', 'hash-4', 'server-other', 'private_push', true) + `); err != nil { + t.Fatalf("seed push devices: %v", err) + } + + system := &System{ + pool: pool, + Settings: NewSettings(mapSettingReader{SettingApplePushDeliveryEnabled: "true"}), + Deliveries: NewDeliveryRepository(pool), + pushDeviceRepo: NewPushDeviceRepository(pool), + dispatcher: NewMultiDispatcher(), + logger: slog.New(slog.NewTextHandler(io.Discard, nil)), + } + + inserted, err := system.DispatchOperational(ctx, Delivery{ + ID: "delivery-request-1", + UserID: 42, + ProfileID: "profile-1", + Type: DeliveryTypeRequestFulfilled, + ReasonFlags: []byte(`{}`), + }, OperationalDispatch{}) + if err != nil { + t.Fatalf("dispatch operational: %v", err) + } + if inserted == nil || inserted.ID != "delivery-request-1" { + t.Fatalf("inserted = %+v", inserted) + } + + var deliveryID, pushDeviceID, triggerType, provider, platform, outcome string + if err := pool.QueryRow(ctx, ` + SELECT notification_delivery_id, push_device_id, trigger_type, provider, platform, outcome + FROM push_delivery_attempts + `).Scan(&deliveryID, &pushDeviceID, &triggerType, &provider, &platform, &outcome); err != nil { + t.Fatalf("query push attempt: %v", err) + } + if deliveryID != inserted.ID || + pushDeviceID != "device-private" || + triggerType != PushTriggerDelivery || + provider != PushProviderSiloRelay || + platform != PushPlatformApple || + outcome != PushOutcomePending { + t.Fatalf("unexpected push attempt: delivery=%q device=%q trigger=%q provider=%q platform=%q outcome=%q", + deliveryID, pushDeviceID, triggerType, provider, platform, outcome) + } + + var count int + if err := pool.QueryRow(ctx, `SELECT count(*) FROM push_delivery_attempts`).Scan(&count); err != nil { + t.Fatalf("count push attempts: %v", err) + } + if count != 1 { + t.Fatalf("push attempt count = %d, want 1", count) + } +} diff --git a/internal/notifications/push_delivery.go b/internal/notifications/push_delivery.go new file mode 100644 index 00000000..7955baa1 --- /dev/null +++ b/internal/notifications/push_delivery.go @@ -0,0 +1,365 @@ +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 +} diff --git a/internal/notifications/push_devices.go b/internal/notifications/push_devices.go new file mode 100644 index 00000000..d9595caf --- /dev/null +++ b/internal/notifications/push_devices.go @@ -0,0 +1,393 @@ +package notifications + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "errors" + "fmt" + "regexp" + "strings" + "time" + + "github.com/Silo-Server/silo-server/internal/secret" + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgxpool" + "github.com/oklog/ulid/v2" +) + +const ( + PushPlatformApple = "apple" + PushProviderSiloRelay = "silo_relay" + PushModeOff = "off" + PushModeInAppOnly = "in_app_only" + PushModePrivatePush = "private_push" + APNsEnvironmentProd = "production" + APNsEnvironmentSandbox = "sandbox" + ApplePushTopicSilo = "org.siloserver.silo" +) + +var ( + ErrPushDeviceUnavailable = errors.New("push device registration unavailable") + ErrPushDeviceInvalid = errors.New("invalid push device registration") + ErrPushDeviceUnsupported = errors.New("unsupported push device registration") + + apnsTokenHexPattern = regexp.MustCompile(`^[0-9a-f]+$`) +) + +// PushDevice represents one profile-scoped notification endpoint. +type PushDevice struct { + ID string + UserID int + ProfileID string + DeviceID string + Platform string + Provider string + APNsEnvironment string + APNsTopic string + APNsTokenCiphertext string + APNsTokenHash string + ServerDeviceID string + PushMode string + Enabled bool + LastSeenAt *time.Time + LastSuccessAt *time.Time + LastFailureAt *time.Time + LastFailureCode *string + CreatedAt time.Time + UpdatedAt time.Time +} + +type ApplePushRegistrationInput struct { + DeviceID string + APNsToken string + APNsEnvironment string + APNsTopic string + PushMode string +} + +type ApplePushDeviceRegistration struct { + UserID int + ProfileID string + DeviceID string + APNsToken string + APNsEnvironment string + APNsTopic string + PushMode string +} + +type PushDeviceStore interface { + UpsertApple(ctx context.Context, registration ApplePushDeviceRegistration, cipher *secret.Cipher) (*PushDevice, error) +} + +type PushDeviceRepository struct { + pool *pgxpool.Pool +} + +func NewPushDeviceRepository(pool *pgxpool.Pool) *PushDeviceRepository { + return &PushDeviceRepository{pool: pool} +} + +type PushDeviceService struct { + store PushDeviceStore + cipher *secret.Cipher +} + +func NewPushDeviceService(store PushDeviceStore, cipher *secret.Cipher) *PushDeviceService { + return &PushDeviceService{store: store, cipher: cipher} +} + +func (s *PushDeviceService) Available() bool { + return s != nil && s.store != nil && s.cipher != nil +} + +func (s *PushDeviceService) RegisterApple(ctx context.Context, userID int, profileID string, input ApplePushRegistrationInput) (*PushDevice, error) { + if !s.Available() { + return nil, ErrPushDeviceUnavailable + } + if userID <= 0 { + return nil, fmt.Errorf("%w: user_id is required", ErrPushDeviceInvalid) + } + profileID = strings.TrimSpace(profileID) + if profileID == "" { + return nil, fmt.Errorf("%w: profile_id is required", ErrPushDeviceInvalid) + } + registration, err := normalizeApplePushRegistration(input) + if err != nil { + return nil, err + } + registration.UserID = userID + registration.ProfileID = profileID + return s.store.UpsertApple(ctx, registration, s.cipher) +} + +func normalizeApplePushRegistration(input ApplePushRegistrationInput) (ApplePushDeviceRegistration, error) { + deviceID := strings.TrimSpace(input.DeviceID) + if deviceID == "" { + return ApplePushDeviceRegistration{}, fmt.Errorf("%w: device_id is required", ErrPushDeviceInvalid) + } + if len(deviceID) > 128 { + return ApplePushDeviceRegistration{}, fmt.Errorf("%w: device_id is too long", ErrPushDeviceInvalid) + } + + token := strings.ToLower(strings.TrimSpace(input.APNsToken)) + if token == "" { + return ApplePushDeviceRegistration{}, fmt.Errorf("%w: apns_token is required", ErrPushDeviceInvalid) + } + if len(token) < 64 || len(token) > 256 || !apnsTokenHexPattern.MatchString(token) { + return ApplePushDeviceRegistration{}, fmt.Errorf("%w: apns_token must be hex encoded", ErrPushDeviceInvalid) + } + + environment := strings.ToLower(strings.TrimSpace(input.APNsEnvironment)) + if environment != APNsEnvironmentProd && environment != APNsEnvironmentSandbox { + return ApplePushDeviceRegistration{}, fmt.Errorf("%w: apns_environment must be production or sandbox", ErrPushDeviceInvalid) + } + + topic := strings.TrimSpace(input.APNsTopic) + if topic == "" { + return ApplePushDeviceRegistration{}, fmt.Errorf("%w: apns_topic is required", ErrPushDeviceInvalid) + } + if topic != ApplePushTopicSilo { + return ApplePushDeviceRegistration{}, fmt.Errorf("%w: apns_topic is not supported", ErrPushDeviceUnsupported) + } + + pushMode := strings.TrimSpace(input.PushMode) + if pushMode == "" { + pushMode = PushModePrivatePush + } + switch pushMode { + case PushModeOff, PushModeInAppOnly, PushModePrivatePush: + default: + return ApplePushDeviceRegistration{}, fmt.Errorf("%w: push_mode is not supported", ErrPushDeviceUnsupported) + } + + return ApplePushDeviceRegistration{ + DeviceID: deviceID, + APNsToken: token, + APNsEnvironment: environment, + APNsTopic: topic, + PushMode: pushMode, + }, nil +} + +func apnsTokenHash(token string) string { + sum := sha256.Sum256([]byte(strings.ToLower(strings.TrimSpace(token)))) + return hex.EncodeToString(sum[:]) +} + +func pushDeviceAPNsTokenAAD(id string) string { + return secret.RowAAD("push_devices", "apns_token", id) +} + +const pushDeviceColumns = ` + id, + user_id, + profile_id, + device_id, + platform, + provider, + apns_environment, + apns_topic, + apns_token_ciphertext, + apns_token_hash, + server_device_id, + push_mode, + enabled, + last_seen_at, + last_success_at, + last_failure_at, + last_failure_code, + created_at, + updated_at` + +func (r *PushDeviceRepository) UpsertApple(ctx context.Context, registration ApplePushDeviceRegistration, cipher *secret.Cipher) (*PushDevice, error) { + if r == nil || r.pool == nil || cipher == nil { + return nil, ErrPushDeviceUnavailable + } + + tx, err := r.pool.Begin(ctx) + if err != nil { + return nil, fmt.Errorf("begin push device upsert: %w", err) + } + defer func() { _ = tx.Rollback(ctx) }() + + // A device install registers for the profile it is currently signed into. + // Purge the same install's registrations under other profiles (attempts + // cascade with them) so a profile switch on a shared device doesn't leave + // the previous profile's notifications flowing to it. + if _, err := tx.Exec(ctx, + `DELETE FROM push_devices WHERE device_id = $1 AND platform = $2 AND profile_id <> $3`, + registration.DeviceID, PushPlatformApple, registration.ProfileID); err != nil { + return nil, fmt.Errorf("purge reassigned push device: %w", err) + } + + device, err := r.selectAppleForUpdate(ctx, tx, registration.ProfileID, registration.DeviceID) + if err != nil { + return nil, err + } + + if device == nil { + device, err = r.insertApple(ctx, tx, registration, cipher) + if err != nil && !errors.Is(err, pgx.ErrNoRows) { + return nil, err + } + if device != nil { + if err := tx.Commit(ctx); err != nil { + return nil, fmt.Errorf("commit push device insert: %w", err) + } + return device, nil + } + + device, err = r.selectAppleForUpdate(ctx, tx, registration.ProfileID, registration.DeviceID) + if err != nil { + return nil, err + } + if device == nil { + return nil, fmt.Errorf("push device upsert conflict row missing") + } + } + + device, err = r.updateApple(ctx, tx, registration, cipher, device) + if err != nil { + return nil, err + } + if err := tx.Commit(ctx); err != nil { + return nil, fmt.Errorf("commit push device update: %w", err) + } + return device, nil +} + +// DeleteAllForProfile removes push registrations for a deleted profile. +func (r *PushDeviceRepository) DeleteAllForProfile(ctx context.Context, profileID string) error { + if r == nil || r.pool == nil { + return nil + } + _, err := r.pool.Exec(ctx, `DELETE FROM push_devices WHERE profile_id = $1`, profileID) + return err +} + +func (r *PushDeviceRepository) selectAppleForUpdate(ctx context.Context, tx pgx.Tx, profileID, deviceID string) (*PushDevice, error) { + row := tx.QueryRow(ctx, `SELECT `+pushDeviceColumns+` FROM push_devices WHERE profile_id = $1 AND device_id = $2 AND platform = $3 FOR UPDATE`, profileID, deviceID, PushPlatformApple) + device, err := scanPushDevice(row) + if errors.Is(err, pgx.ErrNoRows) { + return nil, nil + } + if err != nil { + return nil, fmt.Errorf("select push device: %w", err) + } + return device, nil +} + +func (r *PushDeviceRepository) insertApple(ctx context.Context, tx pgx.Tx, registration ApplePushDeviceRegistration, cipher *secret.Cipher) (*PushDevice, error) { + id := ulid.Make().String() + serverDeviceID := ulid.Make().String() + ciphertext, err := cipher.Encrypt(registration.APNsToken, pushDeviceAPNsTokenAAD(id)) + if err != nil { + return nil, fmt.Errorf("encrypt apns token: %w", err) + } + + row := tx.QueryRow(ctx, ` + INSERT INTO push_devices ( + id, + user_id, + profile_id, + device_id, + platform, + provider, + apns_environment, + apns_topic, + apns_token_ciphertext, + apns_token_hash, + server_device_id, + push_mode, + enabled, + last_seen_at + ) + VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, true, now()) + ON CONFLICT (profile_id, device_id, platform) DO NOTHING + RETURNING `+pushDeviceColumns, + id, + registration.UserID, + registration.ProfileID, + registration.DeviceID, + PushPlatformApple, + PushProviderSiloRelay, + registration.APNsEnvironment, + registration.APNsTopic, + ciphertext, + apnsTokenHash(registration.APNsToken), + serverDeviceID, + registration.PushMode, + ) + device, err := scanPushDevice(row) + if err != nil { + return nil, err + } + return device, nil +} + +func (r *PushDeviceRepository) updateApple(ctx context.Context, tx pgx.Tx, registration ApplePushDeviceRegistration, cipher *secret.Cipher, existing *PushDevice) (*PushDevice, error) { + ciphertext, err := cipher.Encrypt(registration.APNsToken, pushDeviceAPNsTokenAAD(existing.ID)) + if err != nil { + return nil, fmt.Errorf("encrypt apns token: %w", err) + } + + row := tx.QueryRow(ctx, ` + UPDATE push_devices + SET user_id = $1, + provider = $2, + apns_environment = $3, + apns_topic = $4, + apns_token_ciphertext = $5, + apns_token_hash = $6, + push_mode = $7, + enabled = true, + last_seen_at = now(), + last_failure_at = NULL, + last_failure_code = NULL, + updated_at = now() + WHERE id = $8 + RETURNING `+pushDeviceColumns, + registration.UserID, + PushProviderSiloRelay, + registration.APNsEnvironment, + registration.APNsTopic, + ciphertext, + apnsTokenHash(registration.APNsToken), + registration.PushMode, + existing.ID, + ) + device, err := scanPushDevice(row) + if err != nil { + return nil, fmt.Errorf("update push device: %w", err) + } + return device, nil +} + +func scanPushDevice(row pgx.Row) (*PushDevice, error) { + var device PushDevice + if err := row.Scan( + &device.ID, + &device.UserID, + &device.ProfileID, + &device.DeviceID, + &device.Platform, + &device.Provider, + &device.APNsEnvironment, + &device.APNsTopic, + &device.APNsTokenCiphertext, + &device.APNsTokenHash, + &device.ServerDeviceID, + &device.PushMode, + &device.Enabled, + &device.LastSeenAt, + &device.LastSuccessAt, + &device.LastFailureAt, + &device.LastFailureCode, + &device.CreatedAt, + &device.UpdatedAt, + ); err != nil { + return nil, err + } + return &device, nil +} diff --git a/internal/notifications/push_devices_test.go b/internal/notifications/push_devices_test.go new file mode 100644 index 00000000..754fed8a --- /dev/null +++ b/internal/notifications/push_devices_test.go @@ -0,0 +1,340 @@ +package notifications + +import ( + "context" + "errors" + "os" + "strings" + "testing" + + "github.com/Silo-Server/silo-server/internal/secret" + "github.com/jackc/pgx/v5/pgxpool" +) + +type fakePushDeviceStore struct { + got ApplePushDeviceRegistration + calls int + device *PushDevice + err error +} + +func (f *fakePushDeviceStore) UpsertApple(ctx context.Context, registration ApplePushDeviceRegistration, cipher *secret.Cipher) (*PushDevice, error) { + f.calls++ + f.got = registration + if f.err != nil { + return nil, f.err + } + if f.device != nil { + return f.device, nil + } + return &PushDevice{ + ID: "device-row", + ServerDeviceID: "server-device", + Enabled: true, + PushMode: registration.PushMode, + }, nil +} + +func testPushCipher(t *testing.T) *secret.Cipher { + t.Helper() + cipher, err := secret.New([]byte("01234567890123456789012345678901")) + if err != nil { + t.Fatalf("new cipher: %v", err) + } + return cipher +} + +func validApplePushInput() ApplePushRegistrationInput { + return ApplePushRegistrationInput{ + DeviceID: "local-device", + APNsToken: strings.Repeat("a", 64), + APNsEnvironment: APNsEnvironmentProd, + APNsTopic: ApplePushTopicSilo, + PushMode: PushModePrivatePush, + } +} + +func TestPushDeviceServiceRegisterAppleNormalizesAndStores(t *testing.T) { + store := &fakePushDeviceStore{} + service := NewPushDeviceService(store, testPushCipher(t)) + + input := validApplePushInput() + input.DeviceID = " local-device " + input.APNsToken = strings.ToUpper(input.APNsToken) + input.APNsEnvironment = "Production" + input.PushMode = "" + + device, err := service.RegisterApple(context.Background(), 42, " profile-1 ", input) + if err != nil { + t.Fatalf("register apple: %v", err) + } + if device.ServerDeviceID != "server-device" || !device.Enabled || device.PushMode != PushModePrivatePush { + t.Fatalf("unexpected device response: %+v", device) + } + if store.calls != 1 { + t.Fatalf("store calls = %d, want 1", store.calls) + } + if store.got.UserID != 42 || store.got.ProfileID != "profile-1" { + t.Fatalf("unexpected owner scope: %+v", store.got) + } + if store.got.DeviceID != "local-device" { + t.Fatalf("device_id = %q", store.got.DeviceID) + } + if store.got.APNsToken != strings.Repeat("a", 64) { + t.Fatalf("apns token was not canonicalized") + } + if store.got.APNsEnvironment != APNsEnvironmentProd { + t.Fatalf("environment = %q", store.got.APNsEnvironment) + } + if store.got.PushMode != PushModePrivatePush { + t.Fatalf("push mode = %q", store.got.PushMode) + } +} + +func TestPushDeviceServiceRegisterAppleRejectsInvalidOrUnsupported(t *testing.T) { + tests := []struct { + name string + mutate func(*ApplePushRegistrationInput) + wantErr error + }{ + { + name: "missing device id", + mutate: func(input *ApplePushRegistrationInput) { input.DeviceID = " " }, + wantErr: ErrPushDeviceInvalid, + }, + { + name: "malformed token", + mutate: func(input *ApplePushRegistrationInput) { input.APNsToken = strings.Repeat("z", 64) }, + wantErr: ErrPushDeviceInvalid, + }, + { + name: "short token", + mutate: func(input *ApplePushRegistrationInput) { input.APNsToken = "abcd" }, + wantErr: ErrPushDeviceInvalid, + }, + { + name: "invalid environment", + mutate: func(input *ApplePushRegistrationInput) { input.APNsEnvironment = "development" }, + wantErr: ErrPushDeviceInvalid, + }, + { + name: "unsupported topic", + mutate: func(input *ApplePushRegistrationInput) { input.APNsTopic = "com.example.app" }, + wantErr: ErrPushDeviceUnsupported, + }, + { + name: "unsupported push mode", + mutate: func(input *ApplePushRegistrationInput) { input.PushMode = "custom_apns" }, + wantErr: ErrPushDeviceUnsupported, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + input := validApplePushInput() + tt.mutate(&input) + store := &fakePushDeviceStore{} + service := NewPushDeviceService(store, testPushCipher(t)) + + _, err := service.RegisterApple(context.Background(), 42, "profile-1", input) + if !errors.Is(err, tt.wantErr) { + t.Fatalf("error = %v, want %v", err, tt.wantErr) + } + if store.calls != 0 { + t.Fatalf("store was called for rejected input") + } + }) + } +} + +func TestPushDeviceAPNsTokenEncryptionUsesRowAAD(t *testing.T) { + cipher := testPushCipher(t) + token := strings.Repeat("a", 64) + ciphertext, err := cipher.Encrypt(token, pushDeviceAPNsTokenAAD("row-1")) + if err != nil { + t.Fatalf("encrypt token: %v", err) + } + if ciphertext == token || strings.Contains(ciphertext, token) { + t.Fatalf("ciphertext exposes token: %q", ciphertext) + } + plaintext, err := cipher.Decrypt(ciphertext, pushDeviceAPNsTokenAAD("row-1")) + if err != nil { + t.Fatalf("decrypt token: %v", err) + } + if plaintext != token { + t.Fatalf("plaintext = %q, want token", plaintext) + } + if _, err := cipher.Decrypt(ciphertext, pushDeviceAPNsTokenAAD("row-2")); err == nil { + t.Fatalf("decrypt with wrong row AAD succeeded") + } + if got, want := apnsTokenHash(strings.ToUpper(token)), apnsTokenHash(token); got != want { + t.Fatalf("hash should be token-case agnostic: %q != %q", got, want) + } +} + +// newPushDeviceTestRepo connects to SILO_TEST_DATABASE_URL (skipping when +// unset) and shadows push_devices with a session-local temp table, pinning the +// pool to one connection so every query sees it. +func newPushDeviceTestRepo(t *testing.T) (*PushDeviceRepository, *pgxpool.Pool) { + t.Helper() + dsn := os.Getenv("SILO_TEST_DATABASE_URL") + if dsn == "" { + t.Skip("set SILO_TEST_DATABASE_URL to run DB-backed push device repository test") + } + + ctx := context.Background() + config, err := pgxpool.ParseConfig(dsn) + if err != nil { + t.Fatalf("parse db config: %v", err) + } + config.MaxConns = 1 + pool, err := pgxpool.NewWithConfig(ctx, config) + if err != nil { + t.Fatalf("connect db: %v", err) + } + t.Cleanup(pool.Close) + + if _, err := pool.Exec(ctx, ` + CREATE TEMP TABLE push_devices ( + id text PRIMARY KEY, + user_id integer NOT NULL, + profile_id text NOT NULL, + device_id varchar(128) NOT NULL, + platform text NOT NULL, + provider text NOT NULL, + apns_environment text, + apns_topic text, + apns_token_ciphertext text, + apns_token_hash text, + server_device_id text NOT NULL, + push_mode text NOT NULL DEFAULT 'private_push', + enabled boolean NOT NULL DEFAULT true, + last_seen_at timestamptz, + last_success_at timestamptz, + last_failure_at timestamptz, + last_failure_code text, + created_at timestamptz NOT NULL DEFAULT now(), + updated_at timestamptz NOT NULL DEFAULT now(), + CONSTRAINT push_devices_profile_device_platform_key UNIQUE (profile_id, device_id, platform), + CONSTRAINT push_devices_server_device_id_key UNIQUE (server_device_id) + ) ON COMMIT PRESERVE ROWS`); err != nil { + t.Fatalf("create temp push_devices table: %v", err) + } + return NewPushDeviceRepository(pool), pool +} + +func TestPushDeviceRepositoryUpsertApplePreservesStableIDs(t *testing.T) { + ctx := context.Background() + repo, pool := newPushDeviceTestRepo(t) + cipher := testPushCipher(t) + registration := ApplePushDeviceRegistration{ + UserID: 42, + ProfileID: "profile-1", + DeviceID: "local-device", + APNsToken: strings.Repeat("a", 64), + APNsEnvironment: APNsEnvironmentProd, + APNsTopic: ApplePushTopicSilo, + PushMode: PushModePrivatePush, + } + + first, err := repo.UpsertApple(ctx, registration, cipher) + if err != nil { + t.Fatalf("first upsert: %v", err) + } + registration.APNsToken = strings.Repeat("b", 64) + registration.PushMode = PushModeInAppOnly + second, err := repo.UpsertApple(ctx, registration, cipher) + if err != nil { + t.Fatalf("second upsert: %v", err) + } + + if second.ID != first.ID { + t.Fatalf("row id changed on token rotation: %q != %q", second.ID, first.ID) + } + if second.ServerDeviceID != first.ServerDeviceID { + t.Fatalf("server device id changed on token rotation: %q != %q", second.ServerDeviceID, first.ServerDeviceID) + } + if second.APNsTokenHash != apnsTokenHash(registration.APNsToken) { + t.Fatalf("token hash = %q, want rotated hash", second.APNsTokenHash) + } + if first.APNsTokenHash == second.APNsTokenHash { + t.Fatalf("token hash did not change after rotation") + } + plaintext, err := cipher.Decrypt(second.APNsTokenCiphertext, pushDeviceAPNsTokenAAD(second.ID)) + if err != nil { + t.Fatalf("decrypt rotated token: %v", err) + } + if plaintext != registration.APNsToken { + t.Fatalf("rotated plaintext = %q", plaintext) + } + if !second.Enabled || second.PushMode != PushModeInAppOnly { + t.Fatalf("upsert did not re-enable/update mode: %+v", second) + } + + var count int + if err := pool.QueryRow(ctx, `SELECT count(*) FROM push_devices`).Scan(&count); err != nil { + t.Fatalf("count rows: %v", err) + } + if count != 1 { + t.Fatalf("row count = %d, want 1", count) + } +} + +func TestPushDeviceRepositoryUpsertApplePurgesOtherProfiles(t *testing.T) { + ctx := context.Background() + repo, pool := newPushDeviceTestRepo(t) + cipher := testPushCipher(t) + + registration := ApplePushDeviceRegistration{ + UserID: 42, + ProfileID: "profile-parent", + DeviceID: "shared-phone", + APNsToken: strings.Repeat("a", 64), + APNsEnvironment: APNsEnvironmentProd, + APNsTopic: ApplePushTopicSilo, + PushMode: PushModePrivatePush, + } + if _, err := repo.UpsertApple(ctx, registration, cipher); err != nil { + t.Fatalf("register under first profile: %v", err) + } + + // A different install on another profile must be untouched by the purge. + other := registration + other.ProfileID = "profile-parent" + other.DeviceID = "other-phone" + if _, err := repo.UpsertApple(ctx, other, cipher); err != nil { + t.Fatalf("register unrelated device: %v", err) + } + + // The shared phone switches profiles and re-registers: the old profile's + // row for that install must be gone, not left enabled. + registration.ProfileID = "profile-kid" + device, err := repo.UpsertApple(ctx, registration, cipher) + if err != nil { + t.Fatalf("register under second profile: %v", err) + } + if device.ProfileID != "profile-kid" { + t.Fatalf("profile = %q, want profile-kid", device.ProfileID) + } + + rows, err := pool.Query(ctx, `SELECT profile_id, device_id FROM push_devices ORDER BY profile_id`) + if err != nil { + t.Fatalf("list rows: %v", err) + } + defer rows.Close() + got := map[string]string{} + for rows.Next() { + var profileID, deviceID string + if err := rows.Scan(&profileID, &deviceID); err != nil { + t.Fatalf("scan row: %v", err) + } + got[profileID] = deviceID + } + if err := rows.Err(); err != nil { + t.Fatalf("rows: %v", err) + } + want := map[string]string{"profile-kid": "shared-phone", "profile-parent": "other-phone"} + if len(got) != len(want) || got["profile-kid"] != want["profile-kid"] || got["profile-parent"] != want["profile-parent"] { + t.Fatalf("rows after reassignment = %v, want %v", got, want) + } +} diff --git a/internal/notifications/push_sender.go b/internal/notifications/push_sender.go new file mode 100644 index 00000000..77cd0449 --- /dev/null +++ b/internal/notifications/push_sender.go @@ -0,0 +1,343 @@ +package notifications + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "log/slog" + "net/http" + "strings" + "time" + + "github.com/Silo-Server/silo-server/internal/secret" +) + +var pushRetrySchedule = []time.Duration{ + 0, + 30 * time.Second, + 2 * time.Minute, + 10 * time.Minute, + 30 * time.Minute, +} + +const ( + pushMaxAttempts = 5 + pushDispatchQueue = 512 + pushRetryClaimLimit = 100 + relayAppleSendPath = "/v1/apple/send" +) + +func pushRetryDelay(completedAttempt int) (time.Duration, bool) { + if completedAttempt < 1 || completedAttempt >= pushMaxAttempts { + return 0, false + } + return pushRetrySchedule[completedAttempt] - pushRetrySchedule[completedAttempt-1], true +} + +type pushRelayAppleRequest struct { + Token string `json:"token"` + Environment string `json:"environment"` + Topic string `json:"topic"` + Mode string `json:"mode"` + ServerDeviceID string `json:"server_device_id"` + DeliveryID string `json:"delivery_id"` + CollapseID *string `json:"collapse_id,omitempty"` +} + +type pushRelayAppleResponse struct { + RequestID string `json:"request_id"` + APNsID string `json:"apns_id"` + Status string `json:"status"` +} + +type pushRelayErrorResponse struct { + Error struct { + Code string `json:"code"` + Message string `json:"message"` + RequestID string `json:"request_id"` + } `json:"error"` +} + +type pushSendResult struct { + OK bool + HTTPStatus int + RetryAfter time.Duration + RelayRequestID string + UpstreamReason string + Message string + TerminalDevice bool +} + +type pushSender struct { + devices *PushDeviceRepository + deliveries *DeliveryRepository + cipher *secret.Cipher + settings *Settings + client *http.Client + logger *slog.Logger +} + +func newPushSender(devices *PushDeviceRepository, deliveries *DeliveryRepository, cipher *secret.Cipher, settings *Settings) *pushSender { + return &pushSender{ + devices: devices, + deliveries: deliveries, + cipher: cipher, + settings: settings, + client: newWebhookHTTPClient(nil), + logger: slog.Default().With("component", "notifications.apple_push"), + } +} + +func (s *pushSender) processAttempt(ctx context.Context, attempt PushDeliveryAttempt) *PushDeliveryAttempt { + device, err := s.devices.getPushDeviceByID(ctx, attempt.PushDeviceID) + if err != nil || device == nil { + if err == nil { + return s.finalize(ctx, attempt, PushOutcomeFailed, "", "push device deleted", nil, "", nil) + } + if ctx.Err() == nil { + s.logger.Warn("push device lookup failed", "attempt_id", attempt.ID, "error", err) + } + return nil + } + if !device.Enabled || device.PushMode != PushModePrivatePush || !s.settings.ApplePushDeliveryEnabled(ctx) { + return s.finalize(ctx, attempt, PushOutcomeFailed, "delivery_disabled", "apple push delivery disabled", nil, "", nil) + } + if attempt.NotificationDeliveryID != nil { + row, err := s.deliveries.GetRowByID(ctx, *attempt.NotificationDeliveryID) + if err != nil { + if ctx.Err() == nil { + s.logger.Warn("push delivery lookup failed", + "attempt_id", attempt.ID, + "delivery_id", *attempt.NotificationDeliveryID, + "error", err) + } + return nil + } + if row == nil { + return s.finalize(ctx, attempt, PushOutcomeFailed, "delivery_missing", "delivery row missing", nil, "", nil) + } + if row.ProfileID != device.ProfileID { + return s.finalize(ctx, attempt, PushOutcomeFailed, "device_reassigned", "push device reassigned", nil, "", nil) + } + } + + token, err := s.cipher.Decrypt(device.APNsTokenCiphertext, pushDeviceAPNsTokenAAD(device.ID)) + if err != nil { + return s.finalize(ctx, attempt, PushOutcomeFailed, "decrypt_failed", "APNs token decrypt failed", nil, "", nil) + } + result := s.send(ctx, attempt, device, token) + attemptNumber := attempt.AttemptNumber + 1 + if result.OK { + updated, _ := s.devices.FinalizePushAttempt(ctx, attempt.ID, PushOutcomeDelivered, attemptNumber, + result.RelayRequestID, &result.HTTPStatus, result.UpstreamReason, "", nil) + _ = s.devices.RecordPushSuccess(ctx, device.ID) + return updated + } + + statusPtr := (*int)(nil) + if result.HTTPStatus > 0 { + statusPtr = &result.HTTPStatus + } + code := result.UpstreamReason + if code == "" { + code = result.Message + } + if code == "" { + code = "push_delivery_failed" + } + _ = s.devices.RecordPushFailure(ctx, device.ID, code, result.TerminalDevice) + + delay, more := pushRetryDelay(attemptNumber) + if result.RetryAfter > 0 { + delay = result.RetryAfter + } + if more && !result.TerminalDevice && retryableHTTPStatus(result.HTTPStatus) { + nextRetry := time.Now().Add(delay) + return s.finalize(ctx, attempt, PushOutcomeRetrying, result.UpstreamReason, result.Message, statusPtr, result.RelayRequestID, &nextRetry) + } + return s.finalize(ctx, attempt, PushOutcomeFailed, result.UpstreamReason, result.Message, statusPtr, result.RelayRequestID, nil) +} + +func (s *pushSender) finalize(ctx context.Context, attempt PushDeliveryAttempt, outcome string, reason, message string, statusPtr *int, relayRequestID string, nextRetryAt *time.Time) *PushDeliveryAttempt { + attemptNumber := attempt.AttemptNumber + 1 + updated, err := s.devices.FinalizePushAttempt(ctx, attempt.ID, outcome, attemptNumber, relayRequestID, statusPtr, reason, message, nextRetryAt) + if err != nil && ctx.Err() == nil { + s.logger.Warn("finalize push attempt failed", "attempt_id", attempt.ID, "error", err) + } + return updated +} + +func (s *pushSender) send(ctx context.Context, attempt PushDeliveryAttempt, device *PushDevice, token string) pushSendResult { + apiKey := s.settings.PushRelayAPIKey(ctx) + if apiKey == "" { + return pushSendResult{Message: "push relay API key not configured", UpstreamReason: "relay_api_key_missing"} + } + relayURL := s.settings.PushRelayURL(ctx) + deliveryID := attempt.ID + if attempt.NotificationDeliveryID != nil { + deliveryID = *attempt.NotificationDeliveryID + } + collapseID := deliveryID + body, err := json.Marshal(pushRelayAppleRequest{ + Token: token, + Environment: device.APNsEnvironment, + Topic: device.APNsTopic, + Mode: "private_alert", + ServerDeviceID: device.ServerDeviceID, + DeliveryID: deliveryID, + CollapseID: &collapseID, + }) + if err != nil { + return pushSendResult{Message: "relay payload build failed", UpstreamReason: "payload_build_failed"} + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, relayURL+relayAppleSendPath, bytes.NewReader(body)) + if err != nil { + return pushSendResult{Message: "invalid push relay URL", UpstreamReason: "invalid_relay_url"} + } + if req.URL.Scheme != schemeHTTPS { + return pushSendResult{Message: "invalid push relay URL", UpstreamReason: "invalid_relay_url"} + } + req.Header.Set("Authorization", "Bearer "+apiKey) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("User-Agent", "Silo-Push/1.0") + req.Header.Set("Idempotency-Key", fmt.Sprintf("%s:%d", attempt.ID, attempt.AttemptNumber+1)) + + resp, err := s.client.Do(req) + if err != nil { + return pushSendResult{Message: classifyWebhookError(err), UpstreamReason: "relay_unreachable"} + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode >= 200 && resp.StatusCode < 300 { + var parsed pushRelayAppleResponse + _ = json.NewDecoder(io.LimitReader(resp.Body, 16<<10)).Decode(&parsed) + if parsed.RequestID == "" { + parsed.RequestID = resp.Header.Get("X-Request-ID") + } + return pushSendResult{ + OK: true, + HTTPStatus: resp.StatusCode, + RelayRequestID: parsed.RequestID, + } + } + + var parsed pushRelayErrorResponse + data, _ := io.ReadAll(io.LimitReader(resp.Body, 16<<10)) + _ = json.Unmarshal(data, &parsed) + code := parsed.Error.Code + if code == "" { + code = fmt.Sprintf("http_%d", resp.StatusCode) + } + message := parsed.Error.Message + if message == "" { + message = http.StatusText(resp.StatusCode) + } + if message == "" { + message = fmt.Sprintf("HTTP %d", resp.StatusCode) + } + return pushSendResult{ + HTTPStatus: resp.StatusCode, + RetryAfter: parseRetryAfter(resp.Header.Get("Retry-After"), time.Now()), + RelayRequestID: parsed.Error.RequestID, + UpstreamReason: code, + Message: strings.TrimSpace(message), + TerminalDevice: resp.StatusCode == http.StatusUnprocessableEntity && code == "apns_rejected", + } +} + +// PushDispatcher implements the channel Dispatcher interface for Apple push +// on top of the shared channelDispatcher core, with the retry/recovery sweep +// integrated. +type PushDispatcher struct { + core channelDispatcher[PushDeliveryAttempt] +} + +func newPushDispatcher(sender *pushSender) *PushDispatcher { + return &PushDispatcher{core: channelDispatcher[PushDeliveryAttempt]{ + channel: "apple push", + queue: make(chan string, pushDispatchQueue), + logger: slog.Default().With("component", "notifications.apple_push.dispatch"), + claimPending: sender.devices.ClaimPendingPushForDelivery, + process: func(ctx context.Context, attempt PushDeliveryAttempt) { + sender.processAttempt(ctx, attempt) + }, + enabled: sender.settings.ApplePushDeliveryEnabled, + claimDue: sender.devices.ClaimDuePushAttempts, + claimLimit: pushRetryClaimLimit, + }} +} + +// Dispatch queues the delivery's Apple push attempts for immediate send. +func (d *PushDispatcher) Dispatch(_ context.Context, delivery DeliveryRow) error { + if d == nil { + return nil + } + d.core.dispatch(delivery.ID) + return nil +} + +// Run consumes the dispatch queue and the retry/recovery sweep until ctx is +// canceled. +func (d *PushDispatcher) Run(ctx context.Context) { + d.core.run(ctx) +} + +type ApplePushTestResult struct { + AttemptID string + PushDeviceID string + ServerDeviceID string + Outcome string + RelayRequestID string + UpstreamStatus *int + UpstreamReason string + FailureMessage string +} + +func (s *System) SendApplePushTest(ctx context.Context, profileID, serverDeviceID string) (*ApplePushTestResult, error) { + if s == nil || s.pushDeviceRepo == nil || s.pushSender == nil { + return nil, ErrPushDeliveryUnavailable + } + if !s.Settings.ApplePushDeliveryEnabled(ctx) { + return nil, ErrPushDeliveryUnavailable + } + if s.Settings.PushRelayAPIKey(ctx) == "" { + return nil, ErrPushDeliveryUnavailable + } + attempt, device, err := s.pushDeviceRepo.EnqueueAppleTestAttempt(ctx, profileID, serverDeviceID) + if err != nil { + return nil, err + } + claimed, err := s.pushDeviceRepo.ClaimPushAttemptByID(ctx, attempt.ID) + if err != nil { + return nil, err + } + if len(claimed) != 1 { + return nil, fmt.Errorf("push test attempt was not claimable") + } + updated := s.pushSender.processAttempt(ctx, claimed[0]) + if updated == nil { + updated, _ = s.pushDeviceRepo.GetPushAttempt(ctx, attempt.ID) + } + if updated == nil { + return nil, fmt.Errorf("push test attempt disappeared") + } + result := &ApplePushTestResult{ + AttemptID: updated.ID, + PushDeviceID: updated.PushDeviceID, + ServerDeviceID: device.ServerDeviceID, + Outcome: updated.Outcome, + UpstreamStatus: updated.UpstreamStatus, + } + if updated.RelayRequestID != nil { + result.RelayRequestID = *updated.RelayRequestID + } + if updated.UpstreamReason != nil { + result.UpstreamReason = *updated.UpstreamReason + } + if updated.FailureMessage != nil { + result.FailureMessage = *updated.FailureMessage + } + return result, nil +} diff --git a/internal/notifications/push_sender_test.go b/internal/notifications/push_sender_test.go new file mode 100644 index 00000000..11ff4382 --- /dev/null +++ b/internal/notifications/push_sender_test.go @@ -0,0 +1,169 @@ +package notifications + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" +) + +func TestApplePushDeliverySettings(t *testing.T) { + ctx := context.Background() + settings := NewSettings(nil) + if settings.ApplePushDeliveryEnabled(ctx) { + t.Fatal("ApplePushDeliveryEnabled must default to false") + } + if got := settings.PushRelayURL(ctx); got != DefaultPushRelayURL { + t.Fatalf("PushRelayURL default = %q, want %q", got, DefaultPushRelayURL) + } + + settings = NewSettings(mapSettingReader{ + SettingApplePushDeliveryEnabled: "true", + SettingPushRelayURL: "https://push.example.test/", + SettingPushRelayAPIKey: " relay-key ", + }) + if !settings.ApplePushDeliveryEnabled(ctx) { + t.Fatal("ApplePushDeliveryEnabled = false with setting on") + } + if got := settings.PushRelayURL(ctx); got != "https://push.example.test" { + t.Fatalf("PushRelayURL = %q", got) + } + if got := settings.PushRelayAPIKey(ctx); got != "relay-key" { + t.Fatalf("PushRelayAPIKey was not trimmed") + } +} + +func TestPushSenderSendBuildsRelayRequest(t *testing.T) { + token := strings.Repeat("a", 64) + var got struct { + auth string + idempotencyKey string + body pushRelayAppleRequest + } + server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + got.auth = r.Header.Get("Authorization") + got.idempotencyKey = r.Header.Get("Idempotency-Key") + if r.URL.Path != relayAppleSendPath { + t.Fatalf("path = %q, want %q", r.URL.Path, relayAppleSendPath) + } + if err := json.NewDecoder(r.Body).Decode(&got.body); err != nil { + t.Fatalf("decode request body: %v", err) + } + _ = json.NewEncoder(w).Encode(pushRelayAppleResponse{ + RequestID: "relay-request-1", + APNsID: "apns-1", + Status: "accepted", + }) + })) + defer server.Close() + + settings := NewSettings(mapSettingReader{ + SettingPushRelayURL: server.URL, + SettingPushRelayAPIKey: "relay-key", + }) + sender := newPushSender(nil, nil, nil, settings) + sender.client = server.Client() + + deliveryID := "delivery-1" + result := sender.send(context.Background(), PushDeliveryAttempt{ + ID: "attempt-1", + NotificationDeliveryID: &deliveryID, + AttemptNumber: 1, + }, &PushDevice{ + APNsEnvironment: APNsEnvironmentSandbox, + APNsTopic: ApplePushTopicSilo, + ServerDeviceID: "server-device-1", + }, token) + + if !result.OK || result.RelayRequestID != "relay-request-1" { + t.Fatalf("result = %+v", result) + } + if got.auth != "Bearer relay-key" { + t.Fatalf("Authorization = %q", got.auth) + } + if got.idempotencyKey != "attempt-1:2" { + t.Fatalf("Idempotency-Key = %q", got.idempotencyKey) + } + if got.body.Token != token || got.body.Mode != "private_alert" || got.body.DeliveryID != deliveryID { + t.Fatalf("relay body = %+v", got.body) + } + if got.body.CollapseID == nil || *got.body.CollapseID != deliveryID { + t.Fatalf("collapse_id = %+v, want delivery id", got.body.CollapseID) + } +} + +func TestPushSenderSendMapsRelayTerminalAPNsRejection(t *testing.T) { + server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusUnprocessableEntity) + _ = json.NewEncoder(w).Encode(pushRelayErrorResponse{ + Error: struct { + Code string `json:"code"` + Message string `json:"message"` + RequestID string `json:"request_id"` + }{ + Code: "apns_rejected", + Message: "APNs rejected the notification: BadDeviceToken", + RequestID: "relay-request-2", + }, + }) + })) + defer server.Close() + + settings := NewSettings(mapSettingReader{ + SettingPushRelayURL: server.URL, + SettingPushRelayAPIKey: "relay-key", + }) + sender := newPushSender(nil, nil, nil, settings) + sender.client = server.Client() + + result := sender.send(context.Background(), PushDeliveryAttempt{ID: "attempt-1"}, &PushDevice{ + APNsEnvironment: APNsEnvironmentSandbox, + APNsTopic: ApplePushTopicSilo, + ServerDeviceID: "server-device-1", + }, strings.Repeat("a", 64)) + + if result.OK || !result.TerminalDevice || result.HTTPStatus != http.StatusUnprocessableEntity { + t.Fatalf("terminal result = %+v", result) + } + if result.UpstreamReason != "apns_rejected" || result.RelayRequestID != "relay-request-2" { + t.Fatalf("terminal diagnostic = %+v", result) + } +} + +func TestPushSenderSendMapsRelayRetryAfter(t *testing.T) { + server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Retry-After", "30") + w.WriteHeader(http.StatusTooManyRequests) + _ = json.NewEncoder(w).Encode(pushRelayErrorResponse{ + Error: struct { + Code string `json:"code"` + Message string `json:"message"` + RequestID string `json:"request_id"` + }{ + Code: "upstream_rate_limited", + Message: "APNs upstream rate limited the request", + RequestID: "relay-request-3", + }, + }) + })) + defer server.Close() + + settings := NewSettings(mapSettingReader{ + SettingPushRelayURL: server.URL, + SettingPushRelayAPIKey: "relay-key", + }) + sender := newPushSender(nil, nil, nil, settings) + sender.client = server.Client() + + result := sender.send(context.Background(), PushDeliveryAttempt{ID: "attempt-1"}, &PushDevice{ + APNsEnvironment: APNsEnvironmentSandbox, + APNsTopic: ApplePushTopicSilo, + ServerDeviceID: "server-device-1", + }, strings.Repeat("a", 64)) + + if result.OK || result.TerminalDevice || result.RetryAfter == 0 { + t.Fatalf("retryable result = %+v", result) + } +} diff --git a/internal/notifications/settings.go b/internal/notifications/settings.go index 657f3f95..ed5b5de2 100644 --- a/internal/notifications/settings.go +++ b/internal/notifications/settings.go @@ -51,6 +51,11 @@ const ( SettingServerChannelsBatchSeconds = "notifications.server_channels.batch_seconds" SettingServerChannelMentionRequesters = "notifications.server_channels.mention_requesters" + SettingApplePushDeliveryEnabled = "notifications.apple_push_delivery_enabled" + SettingPushRelayURL = "notifications.push_relay_url" + SettingPushRelayDeploymentID = "notifications.push_relay_deployment_id" + SettingPushRelayAPIKey = "notifications.push_relay_api_key" + // Discord application credentials live under the discord.* namespace // (admin-configured, alongside email.smtp_*). The secret and bot token // are registered in catalog.SensitiveSettingKeys and encrypted at rest. @@ -78,6 +83,9 @@ const ( // already passed, and the window is what keeps those visible. defaultServerChannelsBatchSeconds = 300 minServerChannelsBatchSeconds = 120 + // DefaultPushRelayURL is the public Silo relay origin used when no + // notifications.push_relay_url override is stored. + DefaultPushRelayURL = "https://push.siloserver.org" settingsCacheTTL = 15 * time.Second ) @@ -347,6 +355,35 @@ func (s *Settings) ServerChannelsBatchWindow(ctx context.Context) time.Duration defaultServerChannelsBatchSeconds, minServerChannelsBatchSeconds, 3600)) * time.Second } +// ApplePushDeliveryEnabled gates relay sends and the capability endpoint's +// apple_push availability, mirroring how web push advertises itself. The +// device registration endpoint stays available independently so clients that +// already hold tokens keep them fresh across admin toggles. +func (s *Settings) ApplePushDeliveryEnabled(ctx context.Context) bool { + return s.boolSetting(ctx, SettingApplePushDeliveryEnabled, false) +} + +// PushRelayURL is the public Silo relay origin. +func (s *Settings) PushRelayURL(ctx context.Context) string { + value := strings.TrimRight(strings.TrimSpace(s.raw(ctx, SettingPushRelayURL)), "/") + if value == "" { + return DefaultPushRelayURL + } + return value +} + +// PushRelayAPIKey is the bearer credential for the Silo relay. +func (s *Settings) PushRelayAPIKey(ctx context.Context) string { + return strings.TrimSpace(s.raw(ctx, SettingPushRelayAPIKey)) +} + +// PushRelayDeploymentID is the relay account identifier returned during +// self-registration. The dispatcher does not need it; registration uses it for +// subsequent key rotations. +func (s *Settings) PushRelayDeploymentID(ctx context.Context) string { + return strings.TrimSpace(s.raw(ctx, SettingPushRelayDeploymentID)) +} + // DiscordClientID is the Discord application's OAuth2 client ID. func (s *Settings) DiscordClientID(ctx context.Context) string { return strings.TrimSpace(s.raw(ctx, SettingDiscordClientID)) diff --git a/internal/notifications/system.go b/internal/notifications/system.go index 10f4c202..5eb4678d 100644 --- a/internal/notifications/system.go +++ b/internal/notifications/system.go @@ -54,6 +54,9 @@ type System struct { // WebPush is nil when the settings store is not writable (VAPID keys // could not be provisioned). WebPush *WebPushService + // PushDevices is nil when no at-rest cipher is configured; APNs tokens are + // credentials and must not be stored in plaintext. + PushDevices *PushDeviceService // EmailPrefs is nil when no mail sender was provided. EmailPrefs *EmailPrefsRepository // DiscordPrefs holds Discord DM link + mode state; the channel only @@ -73,6 +76,9 @@ type System struct { webhookRetry *WebhookRetryWorker webPushRepo *WebPushRepository webPushDispatcher *WebPushDispatcher + pushDeviceRepo *PushDeviceRepository + pushDispatcher *PushDispatcher + pushSender *pushSender // serverChannelWorker sweeps release_events into admin broadcast posts; // nil without the at-rest cipher. serverChannelWorker *serverChannelWorker @@ -119,6 +125,10 @@ func NewSystem( var webhookDispatcher *WebhookDispatcher var webhookRetry *WebhookRetryWorker var sender *webhookSender + var pushDeviceService *PushDeviceService + var pushDeviceRepo *PushDeviceRepository + var pushSenderInst *pushSender + var pushDispatcher *PushDispatcher if cipher != nil { webhookRepo = NewWebhookRepository(pool) sender = newWebhookSender(webhookRepo, deliveries, cipher, settings) @@ -126,6 +136,11 @@ func NewSystem( webhookDispatcher = newWebhookDispatcher(sender) webhookRetry = newWebhookRetryWorker(sender) dispatchers = append(dispatchers, webhookDispatcher) + pushDeviceRepo = NewPushDeviceRepository(pool) + pushDeviceService = NewPushDeviceService(pushDeviceRepo, cipher) + pushSenderInst = newPushSender(pushDeviceRepo, deliveries, cipher, settings) + pushDispatcher = newPushDispatcher(pushSenderInst) + dispatchers = append(dispatchers, pushDispatcher) } // Admin server channels (broadcast destinations) share the cipher @@ -187,6 +202,9 @@ func NewSystem( if webPushRepo != nil { fanout.SetWebPushOutbox(webPushRepo) } + if pushDeviceRepo != nil { + fanout.SetPushOutbox(pushDeviceRepo) + } detector := NewAvailabilityDetector(releases, settings) detector.SetFanoutNudge(func() { fanout.Nudge() @@ -207,6 +225,7 @@ func NewSystem( Webhooks: webhookService, ServerChannels: serverChannelService, WebPush: webPushService, + PushDevices: pushDeviceService, EmailPrefs: emailPrefs, DiscordPrefs: discordPrefs, mailSender: mailSender, @@ -218,6 +237,9 @@ func NewSystem( webhookRetry: webhookRetry, webPushRepo: webPushRepo, webPushDispatcher: webPushDispatcher, + pushDeviceRepo: pushDeviceRepo, + pushDispatcher: pushDispatcher, + pushSender: pushSenderInst, serverChannelWorker: serverChannelSweep, dispatcher: multiDispatcher, pool: pool, @@ -354,6 +376,13 @@ func (s *System) Start(ctx context.Context) { } }() } + if s.pushDispatcher != nil { + s.wg.Add(1) + go func() { + defer s.wg.Done() + s.pushDispatcher.Run(ctx) + }() + } } // Wait blocks until the background loops exit (after their context is @@ -390,6 +419,11 @@ func (s *System) PurgeProfile(ctx context.Context, profileID string) error { return fmt.Errorf("purge web push subscriptions: %w", err) } } + if s.pushDeviceRepo != nil { + if err := s.pushDeviceRepo.DeleteAllForProfile(ctx, profileID); err != nil { + return fmt.Errorf("purge push devices: %w", err) + } + } if s.EmailPrefs != nil { if err := s.EmailPrefs.DeleteForProfile(ctx, profileID); err != nil { return fmt.Errorf("purge email prefs: %w", err) diff --git a/internal/notifications/webhook_dispatcher.go b/internal/notifications/webhook_dispatcher.go index ddf434a8..c28fd701 100644 --- a/internal/notifications/webhook_dispatcher.go +++ b/internal/notifications/webhook_dispatcher.go @@ -3,7 +3,6 @@ package notifications import ( "context" "log/slog" - "sync" "time" ) @@ -15,23 +14,20 @@ const ( ) // WebhookDispatcher implements the channel Dispatcher interface for outbound -// webhooks. Dispatch never blocks the fanout loop on destination HTTP: it -// hands the delivery ID to a bounded worker pool that claims the delivery's -// pending outbox attempts and sends them. A full queue simply drops the -// hand-off — the durable `pending` rows are picked up by the retry worker's -// outbox recovery sweep, so delivery is delayed, never lost. +// webhooks on top of the shared channelDispatcher core. Retry/recovery runs in +// the standalone WebhookRetryWorker. type WebhookDispatcher struct { - sender *webhookSender - queue chan string - logger *slog.Logger + core channelDispatcher[DeliveryAttempt] } func newWebhookDispatcher(sender *webhookSender) *WebhookDispatcher { - return &WebhookDispatcher{ - sender: sender, - queue: make(chan string, webhookDispatchQueue), - logger: slog.Default().With("component", "notifications.webhooks.dispatch"), - } + return &WebhookDispatcher{core: channelDispatcher[DeliveryAttempt]{ + channel: "webhook", + queue: make(chan string, webhookDispatchQueue), + logger: slog.Default().With("component", "notifications.webhooks.dispatch"), + claimPending: sender.webhooks.ClaimPendingForDelivery, + process: sender.processAttempt, + }} } // Dispatch queues the delivery's webhook attempts for immediate send. @@ -44,50 +40,14 @@ func (d *WebhookDispatcher) Dispatch(_ context.Context, delivery DeliveryRow) er // webhook, or a broken webhook would loop forever. return nil } - select { - case d.queue <- delivery.ID: - default: - d.logger.Warn("webhook dispatch queue full; deferring to retry worker", - "delivery_id", delivery.ID) - } + d.core.dispatch(delivery.ID) return nil } // Run consumes the dispatch queue with a bounded worker pool until ctx is // canceled. One slow destination cannot block other deliveries. func (d *WebhookDispatcher) Run(ctx context.Context) { - var wg sync.WaitGroup - for range webhookDispatchWorkers { - wg.Add(1) - go func() { - defer wg.Done() - for { - select { - case <-ctx.Done(): - return - case deliveryID := <-d.queue: - d.processDelivery(ctx, deliveryID) - } - } - }() - } - wg.Wait() -} - -func (d *WebhookDispatcher) processDelivery(ctx context.Context, deliveryID string) { - attempts, err := d.sender.webhooks.ClaimPendingForDelivery(ctx, deliveryID) - if err != nil { - if ctx.Err() == nil { - d.logger.Warn("webhook attempt claim failed", "delivery_id", deliveryID, "error", err) - } - return - } - for _, attempt := range attempts { - if ctx.Err() != nil { - return - } - d.sender.processAttempt(ctx, attempt) - } + d.core.run(ctx) } // WebhookRetryWorker drains due retries and recovers stale pending outbox @@ -107,34 +67,7 @@ func newWebhookRetryWorker(sender *webhookSender) *WebhookRetryWorker { // Run polls for due attempts until ctx is canceled. func (w *WebhookRetryWorker) Run(ctx context.Context) { - ticker := time.NewTicker(webhookRetryInterval) - defer ticker.Stop() - for { - select { - case <-ctx.Done(): - return - case <-ticker.C: - } - if !w.sender.settings.WebhooksEnabled(ctx) { - continue - } - for { - attempts, err := w.sender.webhooks.ClaimDue(ctx, webhookRetryClaimLimit) - if err != nil { - if ctx.Err() == nil { - w.logger.Warn("webhook retry claim failed", "error", err) - } - break - } - if len(attempts) == 0 { - break - } - for _, attempt := range attempts { - if ctx.Err() != nil { - return - } - w.sender.processAttempt(ctx, attempt) - } - } - } + runRetrySweep(ctx, "webhook", w.logger, + w.sender.settings.WebhooksEnabled, w.sender.webhooks.ClaimDue, + webhookRetryClaimLimit, w.sender.processAttempt) } diff --git a/internal/notifications/webpush_logic_test.go b/internal/notifications/webpush_logic_test.go index ffb078c8..27960b9d 100644 --- a/internal/notifications/webpush_logic_test.go +++ b/internal/notifications/webpush_logic_test.go @@ -18,10 +18,10 @@ func TestBuildWebPushPayload(t *testing.T) { if err := json.Unmarshal(raw, &payload); err != nil { t.Fatal(err) } - if payload.Title != "New episode of Severance" { + if payload.Title != "The latest episode of Severance S02E01 just dropped!" { t.Fatalf("unexpected title %q", payload.Title) } - if payload.Body != "S2E1 — Hello, Ms. Cobel" { + if payload.Body != "Hello, Ms. Cobel" { t.Fatalf("unexpected body %q", payload.Body) } if payload.URL != "/item/episode-456" { @@ -83,6 +83,41 @@ func TestBuildWebPushPayload(t *testing.T) { }) } +func TestBuildNotificationDisplay(t *testing.T) { + t.Run("episode available", func(t *testing.T) { + display := BuildNotificationDisplay(webhookTestRow()) + if display.DeliveryID != "01DELIVERY" || + display.Title != "The latest episode of Severance S02E01 just dropped!" || + display.Body != "Hello, Ms. Cobel" || + display.ThreadID != "series:series-123" || + display.Category != "episode_available" || + display.URL != "/item/episode-456" { + t.Fatalf("display = %+v", display) + } + }) + + t.Run("request declined", func(t *testing.T) { + row := requestDeclinedTestRow() + display := BuildNotificationDisplay(row) + if display.Title != "Dune was declined" || + display.Body != "Reason: Already available in 4K" || + display.ThreadID != "request:01REQ" || + display.Category != "request_declined" || + display.URL != "/notifications" { + t.Fatalf("display = %+v", display) + } + }) + + t.Run("unknown type", func(t *testing.T) { + display := BuildNotificationDisplay(DeliveryRow{Delivery: Delivery{ID: "01X", Type: "future.type"}}) + if display.Title != genericNotificationTitle || + display.Category != "future_type" || + display.URL != "/notifications" { + t.Fatalf("display = %+v", display) + } + }) +} + func TestWebPushRetrySchedule(t *testing.T) { total := time.Duration(0) for attempt := 1; attempt < webPushMaxAttempts; attempt++ { diff --git a/internal/notifications/webpush_sender.go b/internal/notifications/webpush_sender.go index 359d29e6..af8dc1cf 100644 --- a/internal/notifications/webpush_sender.go +++ b/internal/notifications/webpush_sender.go @@ -7,7 +7,6 @@ import ( "io" "log/slog" "net/http" - "sync" "time" webpush "github.com/SherClockHolmes/webpush-go" @@ -51,70 +50,19 @@ type webPushPayload struct { // buildWebPushPayload renders a delivery for the service worker. func buildWebPushPayload(row DeliveryRow, posterURL string) ([]byte, error) { + display := BuildNotificationDisplay(row) payload := webPushPayload{ - Title: "Silo", - URL: "/notifications", + Title: display.Title, + Body: display.Body, + URL: display.URL, Tag: row.ID, DeliveryID: row.ID, } switch row.Type { case DeliveryTypeEpisodeAvailable: - if row.SeriesTitle != "" { - payload.Title = "New episode of " + row.SeriesTitle - } else { - payload.Title = "New episode available" - } - var code string - if row.SeasonNumber != nil && row.EpisodeNumber != nil { - code = fmt.Sprintf("S%dE%d", *row.SeasonNumber, *row.EpisodeNumber) - } - switch { - case code != "" && row.EpisodeTitle != "": - payload.Body = code + " — " + row.EpisodeTitle - case code != "": - payload.Body = code - default: - payload.Body = row.EpisodeTitle - } - if row.EpisodeID != nil { - payload.URL = "/item/" + *row.EpisodeID - } payload.Icon = posterURL case DeliveryTypeRequestFulfilled: - if row.SeriesTitle != "" { - payload.Title = row.SeriesTitle + " is now available" - } else { - payload.Title = "Your request is now available" - } - payload.Body = "Your media request has arrived in the library." - if row.SeriesID != nil { - payload.URL = "/item/" + *row.SeriesID - } payload.Icon = posterURL - case DeliveryTypeRequestApproved: - flags := parseRequestFlags(row.ReasonFlags) - payload.Title = "Your request was approved" - if flags.Title != "" { - payload.Title = flags.Title + " was approved" - } - payload.Body = "Your media request was approved." - case DeliveryTypeRequestDeclined: - flags := parseRequestFlags(row.ReasonFlags) - payload.Title = "Your request was declined" - if flags.Title != "" { - payload.Title = flags.Title + " was declined" - } - payload.Body = "Your media request was declined." - if flags.Reason != "" { - payload.Body = "Reason: " + flags.Reason - } - case DeliveryTypeWebhookAutoDisabled: - payload.Title = "A webhook stopped working" - payload.Body = "Open notification settings to fix it." - payload.URL = "/settings/notifications" - default: - // Unknown types render generically; the inbox has the details. - payload.Title = genericNotificationTitle } return json.Marshal(payload) } @@ -280,21 +228,23 @@ func (s *webPushSender) send(ctx context.Context, sub *WebPushSubscription, mess return resp.StatusCode, retryAfter, nil } -// WebPushDispatcher implements the channel Dispatcher interface: it hands -// delivery IDs to a bounded worker pool that claims and sends the pending -// outbox attempts. A full queue defers to the retry worker's recovery sweep. +// WebPushDispatcher implements the channel Dispatcher interface on top of the +// shared channelDispatcher core, with the retry/recovery sweep integrated. type WebPushDispatcher struct { - sender *webPushSender - queue chan string - logger *slog.Logger + core channelDispatcher[DeliveryAttempt] } func newWebPushDispatcher(sender *webPushSender) *WebPushDispatcher { - return &WebPushDispatcher{ - sender: sender, - queue: make(chan string, webhookDispatchQueue), - logger: slog.Default().With("component", "notifications.webpush.dispatch"), - } + return &WebPushDispatcher{core: channelDispatcher[DeliveryAttempt]{ + channel: "web push", + queue: make(chan string, webhookDispatchQueue), + logger: slog.Default().With("component", "notifications.webpush.dispatch"), + claimPending: sender.subscriptions.ClaimPendingForDelivery, + process: sender.processAttempt, + enabled: sender.settings.WebPushEnabled, + claimDue: sender.subscriptions.ClaimDue, + claimLimit: webhookRetryClaimLimit, + }} } // Dispatch queues the delivery's web push attempts for immediate send. @@ -302,83 +252,12 @@ func (d *WebPushDispatcher) Dispatch(_ context.Context, delivery DeliveryRow) er if d == nil { return nil } - select { - case d.queue <- delivery.ID: - default: - d.logger.Warn("web push dispatch queue full; deferring to retry worker", - "delivery_id", delivery.ID) - } + d.core.dispatch(delivery.ID) return nil } // Run consumes the dispatch queue and the retry/recovery sweep until ctx is // canceled. func (d *WebPushDispatcher) Run(ctx context.Context) { - var wg sync.WaitGroup - for range webhookDispatchWorkers { - wg.Add(1) - go func() { - defer wg.Done() - for { - select { - case <-ctx.Done(): - return - case deliveryID := <-d.queue: - d.processDelivery(ctx, deliveryID) - } - } - }() - } - - wg.Add(1) - go func() { - defer wg.Done() - ticker := time.NewTicker(webhookRetryInterval) - defer ticker.Stop() - for { - select { - case <-ctx.Done(): - return - case <-ticker.C: - } - if !d.sender.settings.WebPushEnabled(ctx) { - continue - } - for { - attempts, err := d.sender.subscriptions.ClaimDue(ctx, webhookRetryClaimLimit) - if err != nil { - if ctx.Err() == nil { - d.logger.Warn("web push retry claim failed", "error", err) - } - break - } - if len(attempts) == 0 { - break - } - for _, attempt := range attempts { - if ctx.Err() != nil { - return - } - d.sender.processAttempt(ctx, attempt) - } - } - } - }() - wg.Wait() -} - -func (d *WebPushDispatcher) processDelivery(ctx context.Context, deliveryID string) { - attempts, err := d.sender.subscriptions.ClaimPendingForDelivery(ctx, deliveryID) - if err != nil { - if ctx.Err() == nil { - d.logger.Warn("web push attempt claim failed", "delivery_id", deliveryID, "error", err) - } - return - } - for _, attempt := range attempts { - if ctx.Err() != nil { - return - } - d.sender.processAttempt(ctx, attempt) - } + d.core.run(ctx) } diff --git a/migrations/push_devices_test.go b/migrations/push_devices_test.go new file mode 100644 index 00000000..0b27ec33 --- /dev/null +++ b/migrations/push_devices_test.go @@ -0,0 +1,60 @@ +package migrations + +import ( + "strings" + "testing" +) + +func TestPushDevicesMigrationContract(t *testing.T) { + migrationBytes, err := FS.ReadFile("sql/20260701143000_push_devices.sql") + if err != nil { + t.Fatalf("read migration: %v", err) + } + migration := string(migrationBytes) + for _, want := range []string{ + "CREATE TABLE public.push_devices", + "CONSTRAINT push_devices_profile_device_platform_key UNIQUE (profile_id, device_id, platform)", + "CONSTRAINT push_devices_server_device_id_key UNIQUE (server_device_id)", + "CONSTRAINT push_devices_platform_check CHECK (platform IN ('apple'))", + "CONSTRAINT push_devices_provider_check CHECK (provider IN ('silo_relay'))", + "CONSTRAINT push_devices_apns_environment_check CHECK (apns_environment IN ('production', 'sandbox'))", + "CONSTRAINT push_devices_push_mode_check CHECK (push_mode IN ('off', 'in_app_only', 'private_push'))", + "CONSTRAINT push_devices_apple_fields_check CHECK", + "CREATE INDEX push_devices_profile_enabled_idx ON public.push_devices (profile_id) WHERE enabled", + } { + if !strings.Contains(migration, want) { + t.Fatalf("migration missing %q", want) + } + } + if strings.Contains(migration, "REFERENCES profiles") { + t.Fatal("push_devices migration must not add a profile foreign key") + } +} + +func TestPushDeliveryAttemptsMigrationContract(t *testing.T) { + migrationBytes, err := FS.ReadFile("sql/20260701170000_push_delivery_attempts.sql") + if err != nil { + t.Fatalf("read migration: %v", err) + } + migration := string(migrationBytes) + for _, want := range []string{ + "CREATE TABLE public.push_delivery_attempts", + "notification_delivery_id text REFERENCES public.notification_deliveries(id) ON DELETE CASCADE", + "push_device_id text NOT NULL REFERENCES public.push_devices(id) ON DELETE CASCADE", + "CONSTRAINT push_delivery_attempts_trigger_check CHECK (trigger_type IN ('delivery', 'test'))", + "CONSTRAINT push_delivery_attempts_delivery_required_check CHECK", + "CONSTRAINT push_delivery_attempts_provider_check CHECK (provider IN ('silo_relay'))", + "CONSTRAINT push_delivery_attempts_platform_check CHECK (platform IN ('apple'))", + "CONSTRAINT push_delivery_attempts_outcome_check CHECK (outcome IN ('pending', 'delivered', 'retrying', 'failed'))", + "CREATE UNIQUE INDEX push_delivery_attempts_delivery_unique", + "CREATE INDEX push_delivery_attempts_retry_idx", + "CREATE INDEX push_delivery_attempts_device_history_idx", + } { + if !strings.Contains(migration, want) { + t.Fatalf("migration missing %q", want) + } + } + if strings.Contains(migration, "REFERENCES profiles") { + t.Fatal("push delivery attempts migration must not add a profile foreign key") + } +} diff --git a/migrations/sql/20260701143000_push_devices.sql b/migrations/sql/20260701143000_push_devices.sql new file mode 100644 index 00000000..4a4cd79a --- /dev/null +++ b/migrations/sql/20260701143000_push_devices.sql @@ -0,0 +1,46 @@ +-- +goose Up +-- +goose StatementBegin +CREATE TABLE public.push_devices ( + id text PRIMARY KEY, + user_id integer NOT NULL, + profile_id text NOT NULL, + device_id varchar(128) NOT NULL, + platform text NOT NULL, + provider text NOT NULL, + apns_environment text, + apns_topic text, + apns_token_ciphertext text, + apns_token_hash text, + server_device_id text NOT NULL, + push_mode text NOT NULL DEFAULT 'private_push', + enabled boolean NOT NULL DEFAULT true, + last_seen_at timestamptz, + last_success_at timestamptz, + last_failure_at timestamptz, + last_failure_code text, + created_at timestamptz NOT NULL DEFAULT now(), + updated_at timestamptz NOT NULL DEFAULT now(), + CONSTRAINT push_devices_profile_device_platform_key UNIQUE (profile_id, device_id, platform), + CONSTRAINT push_devices_server_device_id_key UNIQUE (server_device_id), + CONSTRAINT push_devices_platform_check CHECK (platform IN ('apple')), + CONSTRAINT push_devices_provider_check CHECK (provider IN ('silo_relay')), + CONSTRAINT push_devices_apns_environment_check CHECK (apns_environment IN ('production', 'sandbox')), + CONSTRAINT push_devices_push_mode_check CHECK (push_mode IN ('off', 'in_app_only', 'private_push')), + CONSTRAINT push_devices_apple_fields_check CHECK ( + platform = 'apple' + AND apns_environment IS NOT NULL + AND apns_topic IS NOT NULL + AND apns_token_ciphertext IS NOT NULL + AND apns_token_hash IS NOT NULL + ) +); + +CREATE INDEX push_devices_profile_idx ON public.push_devices (profile_id); +CREATE INDEX push_devices_profile_enabled_idx ON public.push_devices (profile_id) WHERE enabled; +CREATE INDEX push_devices_apns_token_hash_idx ON public.push_devices (apns_token_hash); +-- +goose StatementEnd + +-- +goose Down +-- +goose StatementBegin +DROP TABLE IF EXISTS public.push_devices; +-- +goose StatementEnd diff --git a/migrations/sql/20260701170000_push_delivery_attempts.sql b/migrations/sql/20260701170000_push_delivery_attempts.sql new file mode 100644 index 00000000..1f815ec8 --- /dev/null +++ b/migrations/sql/20260701170000_push_delivery_attempts.sql @@ -0,0 +1,45 @@ +-- +goose Up +-- +goose StatementBegin +CREATE TABLE public.push_delivery_attempts ( + id text PRIMARY KEY, + notification_delivery_id text REFERENCES public.notification_deliveries(id) ON DELETE CASCADE, + push_device_id text NOT NULL REFERENCES public.push_devices(id) ON DELETE CASCADE, + trigger_type text NOT NULL DEFAULT 'delivery', + provider text NOT NULL, + platform text NOT NULL, + attempt_number integer NOT NULL DEFAULT 0, + attempted_at timestamptz NOT NULL DEFAULT now(), + next_retry_at timestamptz, + outcome text NOT NULL DEFAULT 'pending', + relay_request_id text, + upstream_status integer, + upstream_reason varchar(256), + failure_message varchar(256), + created_at timestamptz NOT NULL DEFAULT now(), + updated_at timestamptz NOT NULL DEFAULT now(), + CONSTRAINT push_delivery_attempts_trigger_check CHECK (trigger_type IN ('delivery', 'test')), + CONSTRAINT push_delivery_attempts_delivery_required_check CHECK ( + (trigger_type = 'delivery' AND notification_delivery_id IS NOT NULL) + OR trigger_type = 'test' + ), + CONSTRAINT push_delivery_attempts_provider_check CHECK (provider IN ('silo_relay')), + CONSTRAINT push_delivery_attempts_platform_check CHECK (platform IN ('apple')), + CONSTRAINT push_delivery_attempts_outcome_check CHECK (outcome IN ('pending', 'delivered', 'retrying', 'failed')) +); + +CREATE UNIQUE INDEX push_delivery_attempts_delivery_unique + ON public.push_delivery_attempts (push_device_id, notification_delivery_id, attempt_number) + WHERE notification_delivery_id IS NOT NULL; +CREATE INDEX push_delivery_attempts_delivery_idx + ON public.push_delivery_attempts (notification_delivery_id) + WHERE notification_delivery_id IS NOT NULL; +CREATE INDEX push_delivery_attempts_retry_idx + ON public.push_delivery_attempts (outcome, next_retry_at); +CREATE INDEX push_delivery_attempts_device_history_idx + ON public.push_delivery_attempts (push_device_id, attempted_at DESC); +-- +goose StatementEnd + +-- +goose Down +-- +goose StatementBegin +DROP TABLE IF EXISTS public.push_delivery_attempts; +-- +goose StatementEnd diff --git a/web/src/lib/adminSettingsSearch.ts b/web/src/lib/adminSettingsSearch.ts index 14fc4b85..8497c456 100644 --- a/web/src/lib/adminSettingsSearch.ts +++ b/web/src/lib/adminSettingsSearch.ts @@ -324,10 +324,17 @@ export const ADMIN_SETTINGS_GROUPS: AdminSettingsSearchGroup[] = [ id: "notifications", label: "Notifications", description: - "Server notification channels, release events, Discord, web push, and webhooks.", + "Server notification channels, release events, Silo Push Relay, Discord, web push, and webhooks.", keywords: [ "release events", "new episode", + "silo push relay", + "mobile push", + "apple push", + "android push", + "apns", + "push relay", + "privacy disclosure", "discord", "browser push", "web push", @@ -342,6 +349,11 @@ export const ADMIN_SETTINGS_GROUPS: AdminSettingsSearchGroup[] = [ "Delivery Channels", "In-App", "Web Push", + "Silo Push Relay", + "Relay URL", + "Deployment ID", + "Register Relay", + "Privacy Disclosure", "Email", "Allow Per-Episode Email", "Digest Hour", diff --git a/web/src/pages/admin-settings/AdminSettingsLayout.test.tsx b/web/src/pages/admin-settings/AdminSettingsLayout.test.tsx index d640ea6f..0937851a 100644 --- a/web/src/pages/admin-settings/AdminSettingsLayout.test.tsx +++ b/web/src/pages/admin-settings/AdminSettingsLayout.test.tsx @@ -1,3 +1,4 @@ +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { fireEvent, render, screen } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { renderToStaticMarkup } from "react-dom/server"; @@ -13,10 +14,26 @@ vi.mock("@/hooks/useSettingsForm", () => ({ })); function renderLayout(search = "") { + const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + return renderToStaticMarkup( - - - , + + + + + , + ); +} + +function renderInteractiveLayout(search = "") { + const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + + return render( + + + + + , ); } @@ -71,11 +88,7 @@ describe("AdminSettingsLayout", () => { }); it("filters admin settings sections from the search box", async () => { - render( - - - , - ); + renderInteractiveLayout(); await userEvent.type(screen.getByRole("searchbox", { name: "Search settings" }), "redis"); @@ -85,11 +98,7 @@ describe("AdminSettingsLayout", () => { }); it("matches individual admin setting labels", async () => { - render( - - - , - ); + renderInteractiveLayout(); await userEvent.type( screen.getByRole("searchbox", { name: "Search settings" }), @@ -101,11 +110,7 @@ describe("AdminSettingsLayout", () => { }); it("focuses admin settings search with Cmd+K", () => { - render( - - - , - ); + renderInteractiveLayout(); const searchBox = screen.getByRole("searchbox", { name: "Search settings" }); fireEvent.keyDown(document, { key: "k", metaKey: true }); diff --git a/web/src/pages/admin-settings/NotificationsAdminSettings.test.tsx b/web/src/pages/admin-settings/NotificationsAdminSettings.test.tsx new file mode 100644 index 00000000..c1486651 --- /dev/null +++ b/web/src/pages/admin-settings/NotificationsAdminSettings.test.tsx @@ -0,0 +1,115 @@ +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { renderToStaticMarkup } from "react-dom/server"; +import { MemoryRouter } from "react-router"; +import { describe, expect, it, vi } from "vitest"; + +import NotificationsAdminSettings from "./NotificationsAdminSettings"; + +const useSettingsFormMock = vi.fn(); + +vi.mock("@/hooks/useSettingsForm", () => ({ + useSettingsForm: (...args: unknown[]) => useSettingsFormMock(...args), +})); + +vi.mock("@/hooks/queries/admin/serverNotificationChannels", () => ({ + useServerNotificationChannels: () => ({ data: [] }), +})); + +function makeForm() { + return { + isLoading: false, + getValue: (key: string) => { + switch (key) { + case "notifications.release_events_enabled": + case "notifications.fanout_enabled": + case "notifications.ui_enabled": + case "notifications.web_push_enabled": + case "notifications.apple_push_delivery_enabled": + return "true"; + case "notifications.push_relay_url": + return "https://push.siloserver.org"; + case "notifications.push_relay_deployment_id": + return "01DEPLOYMENT"; + default: + return ""; + } + }, + setValue: vi.fn(), + dirtyCount: 0, + dirtyKeys: [], + isDirty: vi.fn(() => false), + save: vi.fn(), + discard: vi.fn(), + isSaving: false, + restartRequired: false, + sensitiveConfigured: ["notifications.push_relay_api_key"], + sensitiveManagedByEnv: [], + buildConnectionCheckRequest: vi.fn(), + }; +} + +function renderPage() { + const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + + return ( + + + + + + ); +} + +function renderStaticPage() { + return renderToStaticMarkup(renderPage()); +} + +describe("NotificationsAdminSettings", () => { + it("registers Silo Push Relay settings with the shared settings form", () => { + useSettingsFormMock.mockReturnValue(makeForm()); + + renderStaticPage(); + + expect(useSettingsFormMock).toHaveBeenCalledWith({ + keys: expect.arrayContaining([ + "notifications.apple_push_delivery_enabled", + "notifications.push_relay_deployment_id", + ]), + }); + const firstCall = useSettingsFormMock.mock.calls[0]; + if (!firstCall) { + throw new Error("useSettingsForm was not called"); + } + const [options] = firstCall as [{ keys: string[] }]; + expect(options.keys).not.toContain("notifications.push_relay_api_key"); + // The relay URL is persisted only via the registration endpoint, never + // through the settings form. + expect(options.keys).not.toContain("notifications.push_relay_url"); + }); + + it("shows the Silo Push Relay channel status", async () => { + useSettingsFormMock.mockReturnValue(makeForm()); + + render(renderPage()); + + expect(screen.getByText("Silo Push Relay")).toBeInTheDocument(); + expect(screen.getByText(/Mobile push delivery through Silo's relay/)).toBeInTheDocument(); + expect(screen.getByText(/Android support will use the same relay/)).toBeInTheDocument(); + expect(screen.getByText("Relay configured")).toBeInTheDocument(); + + await userEvent.click(screen.getByRole("button", { name: /Silo Push Relay/ })); + + expect(screen.getByText("Privacy disclosure")).toBeInTheDocument(); + expect(screen.getByText(/content-free request to Silo's push relay/)).toBeInTheDocument(); + expect(screen.getByText(/does not receive notification titles/)).toBeInTheDocument(); + expect(screen.getByText(/fetches private content directly/)).toBeInTheDocument(); + expect(screen.getByText("Deployment ID")).toBeInTheDocument(); + expect(screen.getByText("Rotate relay key")).toBeInTheDocument(); + expect(screen.queryByText("Relay API Key")).not.toBeInTheDocument(); + expect(screen.queryByText("Smoke Test Profile ID")).not.toBeInTheDocument(); + expect(screen.queryByText("Server Device ID")).not.toBeInTheDocument(); + expect(screen.queryByText("Send test push")).not.toBeInTheDocument(); + }); +}); diff --git a/web/src/pages/admin-settings/NotificationsAdminSettings.tsx b/web/src/pages/admin-settings/NotificationsAdminSettings.tsx index 1ed8d1d6..888d7ece 100644 --- a/web/src/pages/admin-settings/NotificationsAdminSettings.tsx +++ b/web/src/pages/admin-settings/NotificationsAdminSettings.tsx @@ -1,4 +1,5 @@ import { useId, useMemo, useState } from "react"; +import { useQueryClient } from "@tanstack/react-query"; import { Bell, BookOpen, @@ -9,10 +10,12 @@ import { Copy, ExternalLink, Inbox, + KeyRound, Loader2, Mail, Megaphone, MonitorSmartphone, + RadioTower, Rss, Send, TriangleAlert, @@ -26,6 +29,7 @@ import { api } from "@/api/client"; import { Button } from "@/components/ui/button"; import { Skeleton } from "@/components/ui/skeleton"; import { Switch } from "@/components/ui/switch"; +import { adminKeys } from "@/hooks/queries/keys"; import { useServerNotificationChannels } from "@/hooks/queries/admin/serverNotificationChannels"; import { useSettingsForm } from "@/hooks/useSettingsForm"; import { cn } from "@/lib/utils"; @@ -40,6 +44,11 @@ const KEYS = [ "notifications.ui_enabled", "notifications.webhooks_enabled", "notifications.web_push_enabled", + "notifications.apple_push_delivery_enabled", + // notifications.push_relay_url is intentionally absent: the server rejects + // direct writes; the relay URL is persisted by the registration endpoint + // together with the deployment id and API key. + "notifications.push_relay_deployment_id", "notifications.fanout.settle_seconds", "notifications.fanout.max_series_burst", "notifications.fanout.max_event_age_hours", @@ -71,6 +80,17 @@ interface DiscordTestResult { message?: string; } +interface AppleRelayRegisterResult { + relay_url: string; + deployment_id: string; + key_prefix: string; + api_key_configured: boolean; + relay_request_id?: string; + apns_topics?: string[]; +} + +const DEFAULT_PUSH_RELAY_URL = "https://push.siloserver.org"; + /** * Invite link for adding the bot to a Discord server. Membership alone is * enough to DM, so no permissions are requested. @@ -407,9 +427,120 @@ function TestDiscordRow({ unsaved }: { unsaved: boolean }) { ); } +function RegisterRelayRow({ + relayURL, + deploymentID, + urlEdited, + onRegistered, +}: { + relayURL: string; + deploymentID: string; + urlEdited: boolean; + onRegistered: () => void; +}) { + const queryClient = useQueryClient(); + const [pending, setPending] = useState(false); + const [result, setResult] = useState(null); + + const configured = deploymentID.trim() !== ""; + + const registerRelay = async () => { + if (pending) return; + setPending(true); + setResult(null); + try { + const response = await api( + "/admin/notifications/push/relay/register", + { + method: "POST", + body: JSON.stringify({ + relay_url: relayURL, + }), + }, + ); + setResult(response); + await Promise.all([ + queryClient.invalidateQueries({ queryKey: adminKeys.serverSettings() }), + queryClient.invalidateQueries({ + queryKey: [...adminKeys.serverSettings(), "sensitive-status"] as const, + }), + ]); + onRegistered(); + toast.success("Push relay registered"); + } catch (error) { + toast.error(error instanceof Error ? error.message : "Relay registration failed"); + } finally { + setPending(false); + } + }; + + return ( +
+ {}} + disabled + /> +
+ +
+ {urlEdited && ( +
+ The relay URL change is applied when you register; credentials are stored immediately. +
+ )} + {result && ( +
+ Registered {result.deployment_id} + {result.key_prefix ? ` — key ${result.key_prefix}` : ""} + {result.relay_request_id ? ` — relay ${result.relay_request_id}` : ""} +
+ )} +
+ ); +} + +function ApplePushPrivacyDisclosure() { + return ( +
+
Privacy disclosure
+
+

+ If you enable push notifications, your Silo Server sends a content-free request to Silo's + push relay so Silo can deliver notifications through Apple Push Notification service. +

+

+ The relay does not receive notification titles, message bodies, media names, user names, + profile names, or your server URL. It does process technical metadata needed to deliver + and operate the service, including an opaque deployment identifier, push delivery timing, + request status, app topic, the IP address your self-hosted Silo Server uses to contact the + relay, and a hashed device push token. Apple may also process standard APNs delivery + metadata. +

+

+ Push notifications are generic; the app fetches private content directly from your Silo + Server after receiving the push. +

+
+
+ ); +} + export default function NotificationsAdminSettings() { const form = useSettingsForm({ keys: useMemo(() => KEYS, []) }); const { data: serverChannels } = useServerNotificationChannels(); + // Local draft for the relay URL; null means "show the saved value". + const [pushRelayURLDraft, setPushRelayURLDraft] = useState(null); if (form.isLoading) { return ( @@ -444,10 +575,20 @@ export default function NotificationsAdminSettings() { const webPushOn = isOn("notifications.web_push_enabled"); const emailOn = isOn("notifications.email_enabled"); const serverChannelsOn = isOn("notifications.server_channels_enabled"); - // Discord and personal webhooks are opt-in (default off). + // Mobile push, Discord, and personal webhooks are opt-in (default off). + const applePushOn = form.getValue("notifications.apple_push_delivery_enabled") === "true"; const discordOn = form.getValue("notifications.discord_enabled") === "true"; const webhooksOn = form.getValue("notifications.webhooks_enabled") === "true"; + // The relay URL is not part of the settings form: the server only persists + // it through the registration endpoint, alongside the credentials it mints. + const savedPushRelayURL = form.getValue("notifications.push_relay_url") || DEFAULT_PUSH_RELAY_URL; + const pushRelayURL = pushRelayURLDraft ?? savedPushRelayURL; + const pushRelayURLEdited = pushRelayURL !== savedPushRelayURL; + const pushRelayDeploymentID = form.getValue("notifications.push_relay_deployment_id"); + const pushRelayAPIKeyReady = form.sensitiveConfigured.includes( + "notifications.push_relay_api_key", + ); const allowPrivate = form.getValue("notifications.webhooks.allow_private_destinations") === "true"; // The test endpoint reads the SAVED credentials; testing with unsaved @@ -463,7 +604,15 @@ export default function NotificationsAdminSettings() { (key) => form.getValue(key) !== "" || form.sensitiveConfigured.includes(key), ); - const channelStates = [uiOn, webPushOn, emailOn, discordOn, webhooksOn, serverChannelsOn]; + const channelStates = [ + uiOn, + webPushOn, + applePushOn, + emailOn, + discordOn, + webhooksOn, + serverChannelsOn, + ]; const enabledChannelCount = channelStates.filter(Boolean).length; const failingServerChannels = (serverChannels ?? []).filter( @@ -565,6 +714,38 @@ export default function NotificationsAdminSettings() { onEnabledChange={setToggle("notifications.web_push_enabled")} /> + Relay configured + ) : ( + Relay registration required + ) + } + > +
+ + setPushRelayURLDraft(v)} + /> + setPushRelayURLDraft(null)} + /> +
+
+