diff --git a/apps/relay/README.md b/apps/relay/README.md index 16126123..aaeed90f 100644 --- a/apps/relay/README.md +++ b/apps/relay/README.md @@ -48,7 +48,7 @@ added in the dashboard survives later deploys. | `APNS_PRIVATE_KEY` | yes | Contents of the `.p8` | | `APNS_ENV` | no | Environment tried first: `production` (default) or `sandbox`. A `BadDeviceToken` is retried once on the other, so TestFlight and development-signed builds both work | -Rate limits (`PER_IP` 60/min, `PER_TOKEN` 10/min) are Cloudflare Rate Limiting +Rate limits (`PER_IP` 60/min, `PER_TOKEN` 10/min, `PER_TOKEN_URGENT` 30/min for urgent pushes) are Cloudflare Rate Limiting bindings in `wrangler.jsonc`, keyed on the `CF-Connecting-IP` the edge sets. Workers logs stay off, so nothing about a request is kept. diff --git a/apps/relay/src/index.test.ts b/apps/relay/src/index.test.ts index 628d0472..1ec44b92 100644 --- a/apps/relay/src/index.test.ts +++ b/apps/relay/src/index.test.ts @@ -107,6 +107,37 @@ describe('POST /push', () => { expect(sent).toHaveLength(0); }); + test('an exhausted normal bucket still lets an urgent push through', async () => { + const perToken = limiter(false); + const urgent = limiter(); + const env = fakeEnv(pem, { PER_TOKEN: perToken, PER_TOKEN_URGENT: urgent }); + const { sent, fetchImpl } = apnsStub(ok); + expect((await push(env, fetchImpl, { ...cleartext, level: 'normal' })).status).toBe(429); + expect((await push(env, fetchImpl, { ...cleartext, level: 'urgent' })).status).toBe(200); + expect(urgent.keys).toEqual([TOKEN]); + expect(sent).toHaveLength(1); + }); + + test('an exhausted urgent bucket returns 429 with Retry-After', async () => { + const urgent = limiter(false); + const { sent, fetchImpl } = apnsStub(ok); + const res = await push(fakeEnv(pem, { PER_TOKEN_URGENT: urgent }), fetchImpl, { ...cleartext, level: 'urgent' }); + expect(res.status).toBe(429); + expect(Number(res.headers.get('retry-after'))).toBeGreaterThan(0); + expect(sent).toHaveLength(0); + }); + + test('an urgent push falls back to the shared bucket when the urgent binding is missing', async () => { + const { sent, fetchImpl } = apnsStub(ok); + const env = fakeEnv(pem, { + PER_TOKEN: limiter(false), + PER_TOKEN_URGENT: undefined as unknown as Env['PER_TOKEN_URGENT'], + }); + const res = await push(env, fetchImpl, { ...cleartext, level: 'urgent' }); + expect(res.status).toBe(429); + expect(sent).toHaveLength(0); + }); + test('rejects malformed and oversized bodies', async () => { const { fetchImpl } = apnsStub(ok); expect((await push(fakeEnv(pem), fetchImpl, '{')).status).toBe(400); diff --git a/apps/relay/src/index.ts b/apps/relay/src/index.ts index c1609a2c..0f56722a 100644 --- a/apps/relay/src/index.ts +++ b/apps/relay/src/index.ts @@ -33,6 +33,9 @@ function required(env: Env, name: keyof Env): string { return value; } +// The binding doesn't expose when its window resets, so this is a fixed, short hint. +const RETRY_AFTER_SECONDS = 5; + export function createApp(fetchImpl: typeof fetch = fetch) { // One per isolate, so the signed JWT is reused across requests as Apple asks. let relay: Relay | null = null; @@ -91,7 +94,11 @@ export function createApp(fetchImpl: typeof fetch = fetch) { } const req = parsed.data; - if (!(await c.env.PER_TOKEN.limit({ key: req.token })).success) return c.json({ error: 'rate_limited' }, 429); + // A deploy that hasn't attached the urgent binding falls back to the shared bucket. + const bucket = (req.level === 'urgent' && c.env.PER_TOKEN_URGENT) || c.env.PER_TOKEN; + if (!(await bucket.limit({ key: req.token })).success) { + return c.json({ error: 'rate_limited' }, 429, { 'Retry-After': String(RETRY_AFTER_SECONDS) }); + } const { primary, other, topic } = relayFor(c.env); let result: Awaited>; diff --git a/apps/relay/src/testing.ts b/apps/relay/src/testing.ts index 7f2e2268..5a9a3724 100644 --- a/apps/relay/src/testing.ts +++ b/apps/relay/src/testing.ts @@ -35,6 +35,7 @@ export function fakeEnv(pem: string, overrides: Partial = {}): Env { APNS_PRIVATE_KEY: pem, PER_IP: limiter(), PER_TOKEN: limiter(), + PER_TOKEN_URGENT: limiter(), ...overrides, }; } diff --git a/apps/relay/worker-configuration.d.ts b/apps/relay/worker-configuration.d.ts index 92e78f4f..5fe0ab74 100644 --- a/apps/relay/worker-configuration.d.ts +++ b/apps/relay/worker-configuration.d.ts @@ -1,9 +1,10 @@ /* eslint-disable */ -// Generated by Wrangler by running `wrangler types` (hash: ef768959876914efdecdaf5bcc3adfe1) +// Generated by Wrangler by running `wrangler types` (hash: 107a7050598bb003809013600d3b2d17) // Runtime types generated with workerd@1.20260926.1 2026-09-26 interface __BaseEnv_Env { PER_IP: RateLimit; PER_TOKEN: RateLimit; + PER_TOKEN_URGENT: RateLimit; } declare namespace Cloudflare { interface GlobalProps { diff --git a/apps/relay/wrangler.jsonc b/apps/relay/wrangler.jsonc index 892f37fd..56b474ea 100644 --- a/apps/relay/wrangler.jsonc +++ b/apps/relay/wrangler.jsonc @@ -21,6 +21,13 @@ "name": "PER_TOKEN", "namespace_id": "2002", "simple": { "limit": 10, "period": 60 } + }, + // Urgent pushes (prompts, questions) get their own bucket so a burst of + // routine ones can't lock them out. + { + "name": "PER_TOKEN_URGENT", + "namespace_id": "2003", + "simple": { "limit": 30, "period": 60 } } ] } diff --git a/apps/tether-notify/deliver_test.go b/apps/tether-notify/deliver_test.go new file mode 100644 index 00000000..63b7b374 --- /dev/null +++ b/apps/tether-notify/deliver_test.go @@ -0,0 +1,154 @@ +package main + +import ( + "net/http" + "net/http/httptest" + "sync/atomic" + "testing" + "time" +) + +func scripted(t *testing.T, replies ...func(http.ResponseWriter)) (string, *atomic.Int32) { + t.Helper() + var calls atomic.Int32 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + i := int(calls.Add(1)) - 1 + if i >= len(replies) { + i = len(replies) - 1 + } + replies[i](w) + })) + t.Cleanup(srv.Close) + return srv.URL, &calls +} + +func status(code int, retryAfter string) func(http.ResponseWriter) { + return func(w http.ResponseWriter) { + if retryAfter != "" { + w.Header().Set("Retry-After", retryAfter) + } + w.WriteHeader(code) + } +} + +// recordingRetrier runs on a fake clock that only moves when it sleeps. +func recordingRetrier() (*retrier, *[]time.Duration, *time.Time) { + var waits []time.Duration + clock := time.Now() + r := newRetrier(func() time.Time { return clock }, func(d time.Duration) { + waits = append(waits, d) + clock = clock.Add(d) + }) + return r, &waits, &clock +} + +func TestDeliverUrgentRetriesRateLimitHonoringRetryAfter(t *testing.T) { + url, calls := scripted(t, status(429, "2"), status(429, "30"), status(200, "")) + r, waits, _ := recordingRetrier() + if got := deliver(http.DefaultClient, url, relayRequest{}, r); got != delivered { + t.Fatalf("got %v, want delivered", got) + } + if calls.Load() != 3 { + t.Fatalf("calls = %d, want 3", calls.Load()) + } + if len(*waits) != 2 || (*waits)[0] != 2*time.Second || (*waits)[1] != maxRetryWait { + t.Fatalf("waits = %v, want [2s %v]", *waits, maxRetryWait) + } +} + +func TestDeliverUrgentRetries503WithBackoffWhenNoRetryAfter(t *testing.T) { + url, _ := scripted(t, status(503, ""), status(200, "")) + r, waits, _ := recordingRetrier() + if got := deliver(http.DefaultClient, url, relayRequest{}, r); got != delivered { + t.Fatalf("got %v, want delivered", got) + } + if len(*waits) != 1 || (*waits)[0] != defaultBackoff { + t.Fatalf("waits = %v", *waits) + } +} + +func TestDeliverUrgentGivesUpAfterMaxRetries(t *testing.T) { + url, calls := scripted(t, status(429, "1")) + r, _, _ := recordingRetrier() + if got := deliver(http.DefaultClient, url, relayRequest{}, r); got != failed { + t.Fatalf("got %v, want failed", got) + } + if calls.Load() != maxRetries+1 { + t.Fatalf("calls = %d, want %d", calls.Load(), maxRetries+1) + } +} + +func TestDeliverStopsWhenTimeBudgetIsSpent(t *testing.T) { + url, calls := scripted(t, status(429, "5")) + r, waits, clock := recordingRetrier() + *clock = r.deadline.Add(-6 * time.Second) + if got := deliver(http.DefaultClient, url, relayRequest{}, r); got != failed { + t.Fatalf("got %v, want failed", got) + } + if calls.Load() != 2 || len(*waits) != 1 { + t.Fatalf("calls = %d waits = %v, want one retry then stop", calls.Load(), *waits) + } +} + +func TestDeliverRetriesTransportErrors(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {})) + url := srv.URL + srv.Close() + r, waits, _ := recordingRetrier() + if got := deliver(http.DefaultClient, url, relayRequest{}, r); got != failed { + t.Fatalf("got %v, want failed", got) + } + if len(*waits) != maxRetries { + t.Fatalf("waits = %v, want %d", *waits, maxRetries) + } +} + +func TestDeliverWithoutRetrierMakesOneAttempt(t *testing.T) { + url, calls := scripted(t, status(429, "1"), status(200, "")) + if got := deliver(http.DefaultClient, url, relayRequest{}, nil); got != failed { + t.Fatalf("got %v, want failed", got) + } + if calls.Load() != 1 { + t.Fatalf("calls = %d, want 1", calls.Load()) + } +} + +func TestDeliverDoesNotRetryClientErrors(t *testing.T) { + url, calls := scripted(t, status(400, ""), status(200, "")) + r, waits, _ := recordingRetrier() + if got := deliver(http.DefaultClient, url, relayRequest{}, r); got != failed { + t.Fatalf("got %v, want failed", got) + } + if calls.Load() != 1 || len(*waits) != 0 { + t.Fatalf("calls = %d waits = %v", calls.Load(), *waits) + } +} + +func TestDeliverGoneIsNotRetried(t *testing.T) { + url, calls := scripted(t, status(410, "")) + r, _, _ := recordingRetrier() + if got := deliver(http.DefaultClient, url, relayRequest{}, r); got != gone || calls.Load() != 1 { + t.Fatalf("got %v after %d calls", got, calls.Load()) + } +} + +func TestDeliverStalledRelayIsCutOffAtTheDeadline(t *testing.T) { + done := make(chan struct{}) + srv := httptest.NewServer(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) { + select { + case <-r.Context().Done(): + case <-done: + case <-time.After(30 * time.Second): + } + })) + t.Cleanup(srv.Close) + t.Cleanup(func() { close(done) }) + r := &retrier{sleep: time.Sleep, now: time.Now, deadline: time.Now().Add(300 * time.Millisecond)} + start := time.Now() + if got := deliver(&http.Client{Timeout: 5 * time.Second}, srv.URL, relayRequest{}, r); got != failed { + t.Fatalf("got %v, want failed", got) + } + if elapsed := time.Since(start); elapsed > 5*time.Second { + t.Fatalf("took %v; the deadline should cut a stalled request off", elapsed) + } +} diff --git a/apps/tether-notify/main.go b/apps/tether-notify/main.go index 6125dbd4..7491f232 100644 --- a/apps/tether-notify/main.go +++ b/apps/tether-notify/main.go @@ -2,6 +2,7 @@ package main import ( "bytes" + "context" "crypto/sha256" "encoding/hex" "encoding/json" @@ -10,6 +11,7 @@ import ( "fmt" "net/http" "os" + "strconv" "strings" "time" ) @@ -215,6 +217,11 @@ func sendPush(content PushContent, collapse string, dryRun bool) error { url := strings.TrimRight(relayURL(), "/") + "/push" client := &http.Client{Timeout: 5 * time.Second} + var retry *retrier + if content.Level == levelUrgent && !dryRun { + retry = newRetrier(time.Now, time.Sleep) + } + var sent int for _, device := range devices { ciphertext, err := encryptPushContent(device.SecretKey, content) @@ -229,7 +236,7 @@ func sendPush(content PushContent, collapse string, dryRun bool) error { sent++ continue } - switch deliver(client, url, req) { + switch deliver(client, url, req, retry) { case delivered: sent++ case gone: @@ -253,22 +260,97 @@ const ( failed ) -func deliver(client *http.Client, url string, req relayRequest) deliveryResult { +const ( + maxRetries = 3 + maxRetryWait = 5 * time.Second + retryBudget = 12 * time.Second + defaultBackoff = time.Second +) + +// retrier bounds how long urgent pushes may stall the calling hook. The deadline +// covers request time as well as waits, and is shared across devices so it bounds +// one sendPush call: a relay that accepts and then stalls can't stretch a hold. +type retrier struct { + sleep func(time.Duration) + now func() time.Time + deadline time.Time +} + +func newRetrier(now func() time.Time, sleep func(time.Duration)) *retrier { + return &retrier{sleep: sleep, now: now, deadline: now().Add(retryBudget)} +} + +// next reports how long to wait before retry attempt (0-based), or false once +// the retry count or the time budget is spent. +func (r *retrier) next(attempt int, retryAfter time.Duration) (time.Duration, bool) { + if attempt >= maxRetries { + return 0, false + } + wait := retryAfter + if wait <= 0 { + wait = defaultBackoff << attempt + } + if wait > maxRetryWait { + wait = maxRetryWait + } + if r.now().Add(wait).After(r.deadline) { + return 0, false + } + return wait, true +} + +func parseRetryAfter(v string) time.Duration { + secs, err := strconv.Atoi(strings.TrimSpace(v)) + if err != nil || secs <= 0 { + return 0 + } + return time.Duration(secs) * time.Second +} + +// deliver posts one request. A nil retry means a single attempt; otherwise +// 429, 503 and transport errors are retried within the retrier's bounds. +func deliver(client *http.Client, url string, req relayRequest, retry *retrier) deliveryResult { body, _ := json.Marshal(req) - resp, err := client.Post(url, "application/json", bytes.NewReader(body)) + ctx := context.Background() + if retry != nil { + var cancel context.CancelFunc + ctx, cancel = context.WithDeadline(ctx, retry.deadline) + defer cancel() + } + for attempt := 0; ; attempt++ { + result, transient, retryAfter := post(ctx, client, url, body) + if !transient || retry == nil { + return result + } + wait, ok := retry.next(attempt, retryAfter) + if !ok { + return result + } + retry.sleep(wait) + } +} + +func post(ctx context.Context, client *http.Client, url string, body []byte) (result deliveryResult, transient bool, retryAfter time.Duration) { + httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body)) + if err != nil { + return failed, false, 0 + } + httpReq.Header.Set("Content-Type", "application/json") + resp, err := client.Do(httpReq) if err != nil { fmt.Fprintf(os.Stderr, "relay post failed: %v\n", err) - return failed + return failed, true, 0 } defer resp.Body.Close() switch { case resp.StatusCode == http.StatusGone: - return gone + return gone, false, 0 case resp.StatusCode >= 200 && resp.StatusCode < 300: - return delivered + return delivered, false, 0 default: fmt.Fprintf(os.Stderr, "relay returned %d\n", resp.StatusCode) - return failed + busy := resp.StatusCode == http.StatusTooManyRequests || resp.StatusCode == http.StatusServiceUnavailable + return failed, busy, parseRetryAfter(resp.Header.Get("Retry-After")) } }