Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
31 changes: 20 additions & 11 deletions functions/handlers.js
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,9 @@ const { getMessaging } = require('firebase-admin/messaging');
const { FirestoreRateLimiter, ValkeyRateLimiter } = require('./rate-limiter');

const MAX_NOTIFICATIONS_PER_DAY = parseInt(process.env.MAX_NOTIFICATIONS_PER_DAY || '500');
// Widget pushes get a separate, lower daily cap: Apple's WidgetKit push budget is
// small, so a widget token needs far fewer sends than a normal device.
const MAX_WIDGET_PUSHES_PER_DAY = parseInt(process.env.MAX_WIDGET_PUSHES_PER_DAY || '100');
const REGION = (process.env.REGION || 'us-central1').toLowerCase();

const usingCloudFunctions = process.env.FUNCTION_TARGET !== undefined;
Expand All @@ -13,20 +16,25 @@ const messaging = getMessaging();
const logging = new Logging();
const debug = process.env.DEBUG === 'true';

// Use Valkey rate limiter if Valkey config is available, otherwise use Firestore
let rateLimiter;
// Use Valkey rate limiter if Valkey config is available, otherwise use Firestore.
const useValkey = process.env.VALKEY_HOST && process.env.VALKEY_PORT;
if (useValkey) {
rateLimiter = new ValkeyRateLimiter(
MAX_NOTIFICATIONS_PER_DAY,
debug,
process.env.VALKEY_HOST,
parseInt(process.env.VALKEY_PORT, 10),
);
} else {
rateLimiter = new FirestoreRateLimiter(MAX_NOTIFICATIONS_PER_DAY, debug);
function buildRateLimiter(maxPerDay) {
if (useValkey) {
return new ValkeyRateLimiter(
maxPerDay,
debug,
process.env.VALKEY_HOST,
parseInt(process.env.VALKEY_PORT, 10),
);
}
return new FirestoreRateLimiter(maxPerDay, debug);
}

const rateLimiter = buildRateLimiter(MAX_NOTIFICATIONS_PER_DAY);
// Separate instance so widgets have their own daily cap; widget tokens are a
// different key namespace from FCM tokens, so their counters never collide.
const widgetRateLimiter = buildRateLimiter(MAX_WIDGET_PUSHES_PER_DAY);

async function handleCheckRateLimits(req, res) {
const { push_token: token } = req.body;
if (!token) {
Expand Down Expand Up @@ -348,3 +356,4 @@ function buildLogMetadata(req) {

exports.handleRequest = handleRequest;
exports.handleCheckRateLimits = handleCheckRateLimits;
exports.widgetRateLimiter = widgetRateLimiter;
15 changes: 12 additions & 3 deletions functions/index.js
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ initializeApp();

const android = require('./android');
const legacy = require('./legacy');
const widgetPush = require('./widget-push');

const region = (functions.config().app && functions.config().app.region) || 'us-central1';
const regionalFunctions = functions.region(region).runWith({ timeoutSeconds: 10 });
Expand All @@ -22,9 +23,17 @@ exports.androidV1 = regionalFunctions.https.onRequest(async (req, res) =>
handleRequest(req, res, android.createPayload),
);

exports.sendPushNotification = regionalFunctions.https.onRequest(async (req, res) =>
handleRequest(req, res, legacy.createPayload),
);
exports.sendPushNotification = functions
.region(region)
.runWith({ timeoutSeconds: 10, secrets: widgetPush.SECRETS })
.https.onRequest(async (req, res) => {
// Widget push_subscription payloads target a WidgetKit push token, which FCM
// can't route — send those straight to APNs instead of through messaging.send.
if (req.body && req.body.push_subscription) {
return widgetPush.sendWidgetPush(req, res);
}
return handleRequest(req, res, legacy.createPayload);
});
Comment thread
hariharanjagan marked this conversation as resolved.

exports.checkRateLimits = regionalFunctions.https.onRequest(async (req, res) =>
handleCheckRateLimits(req, res),
Expand Down
235 changes: 235 additions & 0 deletions functions/test/widget-push.test.js
Original file line number Diff line number Diff line change
@@ -0,0 +1,235 @@
'use strict';

const crypto = require('crypto');
const { createMockRequest, createMockResponse } = require('./utils/mock-factories');

jest.mock('node:http2', () => ({ connect: jest.fn() }));
const http2 = require('node:http2');

const mockRateLimiter = {
recordAttempt: jest.fn(),
recordSuccess: jest.fn(),
recordError: jest.fn(),
};
jest.mock('../handlers', () => ({ widgetRateLimiter: mockRateLimiter }));

const widgetPush = require('../widget-push');

const WIDGET_TOKEN = '80f7d67347204c7dda85d331a95ec31c1e3c62b9173836ada8ed9abf';

// A real P-256 key so ES256 signing succeeds; its value is irrelevant to the mock.
const TEST_P8 = crypto
.generateKeyPairSync('ec', { namedCurve: 'P-256' })
.privateKey.export({ type: 'pkcs8', format: 'pem' });

// Makes http2.connect return a client whose request replays the given APNs
// responses in order (one per connect call, so we can exercise the fallback).
function mockApns(responses, onRequest) {
let index = 0;
http2.connect.mockImplementation(() => {
const response = responses[Math.min(index, responses.length - 1)];
index += 1;
const handlers = {};
const request = {
on: jest.fn((event, cb) => {
handlers[event] = cb;
return request;
}),
setEncoding: jest.fn(),
end: jest.fn(() => {
process.nextTick(() => {
handlers.response?.({ ':status': response.status, 'apns-id': response.apnsId });
if (response.body) handlers.data?.(response.body);
handlers.end?.();
});
}),
};
return {
on: jest.fn(),
request: jest.fn((headers) => {
onRequest?.(headers);
return request;
}),
close: jest.fn(),
};
});
}

function widgetRequest(overrides = {}) {
return createMockRequest({
body: {
push_subscription: { subscription_id: 'ios-widget-sensors', target: 'sensors' },
push_token: WIDGET_TOKEN,
registration_info: { app_id: 'io.test.HomeAssistant' },
...overrides,
},
});
}

describe('widget-push', () => {
beforeEach(() => {
jest.clearAllMocks();
process.env.APNS_KEY_P8 = TEST_P8;
process.env.APNS_KEY_ID = 'KEY1234567';
process.env.APNS_TEAM_ID = 'TEAM123456';
mockRateLimiter.recordAttempt.mockResolvedValue({ isRateLimited: false, rateLimits: {} });
mockRateLimiter.recordSuccess.mockResolvedValue({});
mockRateLimiter.recordError.mockResolvedValue({});
});

it('returns 201 and echoes the apns-id on a successful send', async () => {
mockApns([{ status: 200, apnsId: 'apns-success' }]);
const res = createMockResponse();

await widgetPush.sendWidgetPush(widgetRequest(), res);

expect(res.status).toHaveBeenCalledWith(201);
expect(res.send).toHaveBeenCalledWith(
expect.objectContaining({
target: WIDGET_TOKEN,
messageId: 'apns-success',
pushType: 'widgets',
}),
);
});

it('rate-limits per token and returns 429 without reaching APNs', async () => {
mockRateLimiter.recordAttempt.mockResolvedValueOnce({
isRateLimited: true,
rateLimits: { successful: 500 },
});
const res = createMockResponse();

await widgetPush.sendWidgetPush(widgetRequest(), res);

expect(res.status).toHaveBeenCalledWith(429);
expect(http2.connect).not.toHaveBeenCalled();
});

it('records the attempt outcome with the rate limiter on success', async () => {
mockApns([{ status: 200, apnsId: 'x' }]);

await widgetPush.sendWidgetPush(widgetRequest(), createMockResponse());

expect(mockRateLimiter.recordAttempt).toHaveBeenCalledWith(WIDGET_TOKEN);
expect(mockRateLimiter.recordSuccess).toHaveBeenCalledWith(WIDGET_TOKEN);
});

it('sends the widgets push type, widget topic and device path', async () => {
let headers;
mockApns([{ status: 200, apnsId: 'x' }], (h) => {
headers = h;
});

await widgetPush.sendWidgetPush(widgetRequest(), createMockResponse());

expect(headers['apns-push-type']).toBe('widgets');
expect(headers['apns-topic']).toBe('io.test.HomeAssistant.push-type.widgets');
expect(headers[':path']).toBe(`/3/device/${WIDGET_TOKEN}`);
expect(headers.authorization).toMatch(/^bearer /);
});

it('falls back to the sandbox host on BadDeviceToken', async () => {
mockApns([
{ status: 400, body: '{"reason":"BadDeviceToken"}' },
{ status: 200, apnsId: 'sandbox-ok' },
]);
const res = createMockResponse();

await widgetPush.sendWidgetPush(widgetRequest(), res);

expect(http2.connect).toHaveBeenCalledTimes(2);
expect(http2.connect).toHaveBeenNthCalledWith(1, 'https://api.push.apple.com');
expect(http2.connect).toHaveBeenNthCalledWith(2, 'https://api.sandbox.push.apple.com');
expect(res.status).toHaveBeenCalledWith(201);
});

it('returns 403 when no token is sent', async () => {
const res = createMockResponse();
await widgetPush.sendWidgetPush(widgetRequest({ push_token: null }), res);
expect(res.status).toHaveBeenCalledWith(403);
});

it('returns 400 when registration_info.app_id is missing', async () => {
const res = createMockResponse();
await widgetPush.sendWidgetPush(widgetRequest({ registration_info: {} }), res);
expect(res.status).toHaveBeenCalledWith(400);
});

it('rejects a non-hex push token without reaching APNs', async () => {
const res = createMockResponse();
await widgetPush.sendWidgetPush(widgetRequest({ push_token: `${WIDGET_TOKEN}/../evil` }), res);
expect(res.status).toHaveBeenCalledWith(400);
expect(http2.connect).not.toHaveBeenCalled();
});

it('rejects an app id with illegal characters without reaching APNs', async () => {
const res = createMockResponse();
await widgetPush.sendWidgetPush(
widgetRequest({ registration_info: { app_id: 'io.test.HomeAssistant\r\nx-evil: 1' } }),
res,
);
expect(res.status).toHaveBeenCalledWith(400);
expect(http2.connect).not.toHaveBeenCalled();
});

it('closes the HTTP/2 client and rejects on a session error', async () => {
const request = { on: jest.fn().mockReturnThis(), setEncoding: jest.fn(), end: jest.fn() };
const client = { on: jest.fn(), request: jest.fn(() => request), close: jest.fn() };
client.on.mockImplementation((event, cb) => {
if (event === 'error') process.nextTick(() => cb(new Error('boom')));
return client;
});
http2.connect.mockReturnValue(client);
const res = createMockResponse();

await widgetPush.sendWidgetPush(widgetRequest(), res);

expect(client.close).toHaveBeenCalled();
expect(res.status).toHaveBeenCalledWith(502);
});

it('closes the HTTP/2 client and rejects on a stream error', async () => {
const handlers = {};
const request = {
on: jest.fn((event, cb) => {
handlers[event] = cb;
return request;
}),
setEncoding: jest.fn(),
end: jest.fn(() => process.nextTick(() => handlers.error?.(new Error('stream boom')))),
};
const client = { on: jest.fn(), request: jest.fn(() => request), close: jest.fn() };
http2.connect.mockReturnValue(client);
const res = createMockResponse();

await widgetPush.sendWidgetPush(widgetRequest(), res);

expect(client.close).toHaveBeenCalled();
expect(res.status).toHaveBeenCalledWith(502);
});

it('returns 500 when the rate-limit check throws', async () => {
mockRateLimiter.recordAttempt.mockRejectedValueOnce(new Error('firestore down'));
const res = createMockResponse();

await widgetPush.sendWidgetPush(widgetRequest(), res);

expect(res.status).toHaveBeenCalledWith(500);
expect(http2.connect).not.toHaveBeenCalled();
});

it('returns 500 when APNs credentials are not configured', async () => {
delete process.env.APNS_KEY_P8;
const res = createMockResponse();
await widgetPush.sendWidgetPush(widgetRequest(), res);
expect(res.status).toHaveBeenCalledWith(500);
});

it('propagates a non-fallback APNs rejection status', async () => {
mockApns([{ status: 410, body: '{"reason":"Unregistered"}' }]);
const res = createMockResponse();
await widgetPush.sendWidgetPush(widgetRequest(), res);
expect(res.status).toHaveBeenCalledWith(410);
});
});
Loading