diff --git a/.github/labeller.yaml b/.github/labeller.yaml index abbf293ab..3f981f6b6 100644 --- a/.github/labeller.yaml +++ b/.github/labeller.yaml @@ -28,6 +28,9 @@ f:media_proxy: f:messages: - changed-files: - any-glob-to-any-file: fluxer_messages/**/* +f:push: + - changed-files: + - any-glob-to-any-file: fluxer_push/**/* f:snowflakes: - changed-files: - any-glob-to-any-file: fluxer_snowflakes/**/* diff --git a/.github/workflows/build-push.yaml b/.github/workflows/build-push.yaml new file mode 100644 index 000000000..61927dbcd --- /dev/null +++ b/.github/workflows/build-push.yaml @@ -0,0 +1,36 @@ +# SPDX-License-Identifier: AGPL-3.0-or-later +name: build push + +on: + workflow_dispatch: + inputs: + build-version: + description: "Explicit Fluxer CalVer build version (YYYY.MDD.MICRO, UTC HHMMSS without leading zeroes) to use instead of automatic UTC clock allocation" + type: string + required: false + default: "" + +permissions: + actions: read + contents: write + packages: write + +jobs: + approve: + name: approve build release + permissions: {} + runs-on: ubuntu-24.04 + environment: builds + timeout-minutes: 5 + steps: + - name: approved + run: echo "Build release approved." + + image: + needs: approve + uses: ./.github/workflows/_build-image.yaml + secrets: inherit + with: + image: fluxer-push + dockerfile: fluxer_push/Dockerfile + build-version: ${{ inputs['build-version'] }} diff --git a/Cargo.lock b/Cargo.lock index 5847af673..89f4ec1f9 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1607,6 +1607,7 @@ dependencies = [ "ff", "generic-array", "group", + "hkdf", "pem-rfc7468", "pkcs8", "rand_core 0.6.4", @@ -1881,6 +1882,31 @@ dependencies = [ "url", ] +[[package]] +name = "fluxer-push" +version = "0.1.0" +dependencies = [ + "anyhow", + "axum", + "base64 0.23.1", + "clap", + "fluxer-svc", + "futures", + "hmac 0.13.0", + "p256", + "rand 0.10.2", + "reqwest", + "ring", + "serde", + "serde_json", + "sha2 0.11.0", + "thiserror", + "tokio", + "tracing", + "tracing-subscriber", + "url", +] + [[package]] name = "fluxer-snowflakes" version = "0.1.0" @@ -2308,6 +2334,15 @@ version = "0.4.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70" +[[package]] +name = "hkdf" +version = "0.12.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7b5f8eb2ad728638ea2c7d47a21db23b7b58a72ed6a38256b8a1849f15fbbdf7" +dependencies = [ + "hmac 0.12.1", +] + [[package]] name = "hmac" version = "0.12.1" @@ -3893,6 +3928,7 @@ dependencies = [ "futures-channel", "futures-core", "futures-util", + "h2", "http 1.5.0", "http-body 1.1.0", "http-body-util", diff --git a/Cargo.toml b/Cargo.toml index 40ceceed3..f40a20ddc 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -7,6 +7,7 @@ members = [ "fluxer_gifs", "fluxer_svc", "fluxer_messages", + "fluxer_push", "fluxer_snowflakes", "tools/ci", "tools/dev", diff --git a/deploy/self-hosting/.env.example b/deploy/self-hosting/.env.example index d3ce2ec41..97b6f09d9 100644 --- a/deploy/self-hosting/.env.example +++ b/deploy/self-hosting/.env.example @@ -147,6 +147,12 @@ FLUXER_VAPID_PRIVATE_KEY=CHANGE_ME #FLUXER_PASSKEY_ADDITIONAL_ALLOWED_ORIGINS=https://chat.example.com #FLUXER_PASSKEY_ADDITIONAL_ALLOWED_ORIGINS=http://chat.example.com:19080 +# Notification jobs the push container holds at once, 1 to 1000000. +#FLUXER_PUSH_SERVICE_QUEUE_CAPACITY=10000 +# Provider requests the push container sends at once, 1 to 65536. +#FLUXER_PUSH_SERVICE_SEND_CONCURRENCY=256 + + # Optional media policies, both off by default. See the operator docs. # # CORS limits which web origins may read media. A request with no Origin is @@ -245,6 +251,7 @@ FLUXER_DISCOVERY_ENABLED=true #FLUXER_GATEWAY_MEMORY_LIMIT=1gb #FLUXER_GATEWAY_MEMORY_RESERVATION=384mb #FLUXER_MEDIA_PROXY_MEMORY_LIMIT=512mb +#FLUXER_PUSH_MEMORY_LIMIT=256mb #FLUXER_STATIC_PROXY_MEMORY_LIMIT=256mb #FLUXER_APP_PROXY_MEMORY_LIMIT=256mb #FLUXER_SNOWFLAKES_MEMORY_LIMIT=128mb diff --git a/deploy/self-hosting/docker-compose.yml b/deploy/self-hosting/docker-compose.yml index 495ad96ce..01a043ad9 100644 --- a/deploy/self-hosting/docker-compose.yml +++ b/deploy/self-hosting/docker-compose.yml @@ -482,6 +482,30 @@ services: seaweedfs-init: {condition: service_completed_successfully} nats: {condition: service_healthy} + push: + <<: *fluxer-service + image: ${FLUXER_REGISTRY:-ghcr.io/${FLUXER_REGISTRY_OWNER:-fluxerapp}}/fluxer-push:${FLUXER_IMAGE_TAG:-v1} + deploy: + resources: + limits: + memory: ${FLUXER_PUSH_MEMORY_LIMIT:-256mb} + environment: + <<: *fluxer-env + FLUXER_PUSH_SERVICE_HOST: 0.0.0.0 + FLUXER_PUSH_SERVICE_PORT: "8126" + FLUXER_PUSH_SERVICE_QUEUE_CAPACITY: "${FLUXER_PUSH_SERVICE_QUEUE_CAPACITY:-}" + FLUXER_PUSH_SERVICE_SEND_CONCURRENCY: "${FLUXER_PUSH_SERVICE_SEND_CONCURRENCY:-}" + healthcheck: + test: ["CMD", "/usr/local/bin/fluxer-push", "healthcheck"] + interval: 10s + timeout: 5s + retries: 30 + start_period: 60s + start_interval: 1s + depends_on: + nats: {condition: service_healthy} + api: {condition: service_healthy} + static-proxy: <<: *fluxer-service image: ${FLUXER_REGISTRY:-ghcr.io/${FLUXER_REGISTRY_OWNER:-fluxerapp}}/fluxer-static:${FLUXER_IMAGE_TAG:-v1} diff --git a/fluxer_admin/openapi-admin.json b/fluxer_admin/openapi-admin.json index 2294cacba..16ffb88d7 100644 --- a/fluxer_admin/openapi-admin.json +++ b/fluxer_admin/openapi-admin.json @@ -10525,6 +10525,7 @@ "gateway_rollout": {"$ref": "#/components/schemas/GatewayRolloutConfigResponse"}, "voice_noise_suppression": {"$ref": "#/components/schemas/VoiceNoiseSuppressionConfigResponse"}, "screen_share_delivery": {"$ref": "#/components/schemas/ScreenShareDeliveryConfigResponse"}, + "push_service_delivery": {"$ref": "#/components/schemas/PushServiceDeliveryConfigResponse"}, "experiment_delivery": {"$ref": "#/components/schemas/ExperimentDeliveryConfigResponse"}, "registration": { "type": "object", @@ -10953,6 +10954,7 @@ "gateway_rollout", "voice_noise_suppression", "screen_share_delivery", + "push_service_delivery", "experiment_delivery", "registration", "self_hosted", @@ -11091,6 +11093,10 @@ "nullable": true, "allOf": [{"$ref": "#/components/schemas/ScreenShareDeliveryConfigUpdateRequest"}] }, + "push_service_delivery": { + "nullable": true, + "allOf": [{"$ref": "#/components/schemas/PushServiceDeliveryConfigUpdateRequest"}] + }, "experiment_delivery": { "nullable": true, "allOf": [{"$ref": "#/components/schemas/ExperimentDeliveryConfigUpdateRequest"}] @@ -15184,6 +15190,24 @@ "poll_jitter_percent": {"type": "integer", "minimum": 0, "maximum": 50} } }, + "PushServiceDeliveryConfigUpdateRequest": { + "type": "object", + "properties": { + "enabled": {"type": "boolean"}, + "rollout_basis_points": {"type": "integer", "minimum": 0, "maximum": 10000}, + "rollout_salt": {"type": "string", "minLength": 1, "maxLength": 64, "pattern": "^[\\x20-\\x7e]+$"}, + "included_user_ids": { + "maxItems": 1000, + "type": "array", + "items": {"type": "string", "pattern": "^\\d{1,20}$"} + }, + "excluded_user_ids": { + "maxItems": 1000, + "type": "array", + "items": {"type": "string", "pattern": "^\\d{1,20}$"} + } + } + }, "ScreenShareDeliveryConfigUpdateRequest": { "type": "object", "properties": { @@ -15267,6 +15291,42 @@ "required": ["poll_interval_seconds", "poll_jitter_percent"], "additionalProperties": false }, + "PushServiceDeliveryConfigResponse": { + "type": "object", + "properties": { + "enabled": {"default": false, "type": "boolean"}, + "config_version": {"default": 0, "type": "integer", "minimum": 0, "maximum": 9007199254740991}, + "rollout_basis_points": {"default": 0, "type": "integer", "minimum": 0, "maximum": 10000}, + "rollout_salt": { + "default": "push-service-delivery-v1", + "type": "string", + "minLength": 1, + "maxLength": 64, + "pattern": "^[\\x20-\\x7e]+$" + }, + "included_user_ids": { + "default": [], + "maxItems": 1000, + "type": "array", + "items": {"type": "string", "pattern": "^\\d{1,20}$"} + }, + "excluded_user_ids": { + "default": [], + "maxItems": 1000, + "type": "array", + "items": {"type": "string", "pattern": "^\\d{1,20}$"} + } + }, + "required": [ + "enabled", + "config_version", + "rollout_basis_points", + "rollout_salt", + "included_user_ids", + "excluded_user_ids" + ], + "additionalProperties": false + }, "ScreenShareDeliveryConfigResponse": { "type": "object", "properties": { diff --git a/fluxer_admin/src/api/types/instance_config.rs b/fluxer_admin/src/api/types/instance_config.rs index f0d6a4140..66b73b864 100644 --- a/fluxer_admin/src/api/types/instance_config.rs +++ b/fluxer_admin/src/api/types/instance_config.rs @@ -25,6 +25,8 @@ pub struct InstanceConfigResponse { #[serde(default)] pub screen_share_delivery: ScreenShareDeliveryConfigResponse, #[serde(default)] + pub push_service_delivery: PushServiceDeliveryConfigResponse, + #[serde(default)] pub experiment_delivery: ExperimentDeliveryConfigResponse, } @@ -449,6 +451,7 @@ impl VoiceE2eeScope { } pub const EXPERIMENT_MAX_TARGETED_USERS: usize = 1_000; +pub const PUSH_SERVICE_DELIVERY_DEFAULT_SALT: &str = "push-service-delivery-v1"; pub const SCREEN_SHARE_DELIVERY_DEFAULT_SALT: &str = "screen-share-delivery-v1"; pub const VOICE_NS_MAX_GUILD_OVERRIDES: usize = 200; @@ -578,6 +581,44 @@ pub struct ScreenShareDeliveryConfigUpdateRequest { pub excluded_user_ids: Option>, } +#[derive(Clone, Debug, Deserialize, Serialize)] +#[serde(default)] +pub struct PushServiceDeliveryConfigResponse { + pub enabled: bool, + pub config_version: u64, + pub rollout_basis_points: u32, + pub rollout_salt: String, + pub included_user_ids: Vec, + pub excluded_user_ids: Vec, +} + +impl Default for PushServiceDeliveryConfigResponse { + fn default() -> Self { + Self { + enabled: false, + config_version: 0, + rollout_basis_points: 0, + rollout_salt: PUSH_SERVICE_DELIVERY_DEFAULT_SALT.to_owned(), + included_user_ids: Vec::new(), + excluded_user_ids: Vec::new(), + } + } +} + +#[derive(Clone, Debug, Default, Serialize)] +pub struct PushServiceDeliveryConfigUpdateRequest { + #[serde(skip_serializing_if = "Option::is_none")] + pub enabled: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub rollout_basis_points: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub rollout_salt: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub included_user_ids: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub excluded_user_ids: Option>, +} + #[derive(Clone, Debug, Deserialize, Serialize)] #[serde(default)] pub struct ExperimentDeliveryConfigResponse { @@ -696,6 +737,8 @@ pub struct InstanceConfigUpdateRequest { #[serde(skip_serializing_if = "Option::is_none")] pub screen_share_delivery: Option, #[serde(skip_serializing_if = "Option::is_none")] + pub push_service_delivery: Option, + #[serde(skip_serializing_if = "Option::is_none")] pub experiment_delivery: Option, } diff --git a/fluxer_admin/src/routes/reports.rs b/fluxer_admin/src/routes/reports.rs index 9ba227915..6d3ee95cf 100644 --- a/fluxer_admin/src/routes/reports.rs +++ b/fluxer_admin/src/routes/reports.rs @@ -80,7 +80,7 @@ async fn reports_list( return reports_error_page( config, &auth.0, - "That page is out of range. The reports search returns at most the first 10000 reports, so narrow the filters and start again.", + "That page is out of range. The reports search returns at most the first 10000 reports. Narrow the filters and start again.", ); } let search_query = query.q.as_deref().and_then(clean_string); diff --git a/fluxer_admin/src/routes/system_actions.rs b/fluxer_admin/src/routes/system_actions.rs index 87b8e4f69..5654d9403 100644 --- a/fluxer_admin/src/routes/system_actions.rs +++ b/fluxer_admin/src/routes/system_actions.rs @@ -18,9 +18,10 @@ use crate::{ InstancePolicyUpdateRequest, InstanceRegistrationConfigUpdateRequest, InstanceServicesUpdateRequest, InstanceYoutubeIntegrationUpdateRequest, LimitConfigUpdateRequest, LimitRule, LimitRuleFilters, NoiseSuppressionBackend, - PremiumMode, RegistrationMode, ScreenShareDeliveryConfigUpdateRequest, - SsoConfigUpdateRequest, VOICE_NS_MAX_GUILD_OVERRIDES, VoiceE2eeScope, - VoiceNoiseSuppressionConfigUpdateRequest, VoiceNoiseSuppressionGuildOverride, + PremiumMode, PushServiceDeliveryConfigUpdateRequest, RegistrationMode, + ScreenShareDeliveryConfigUpdateRequest, SsoConfigUpdateRequest, + VOICE_NS_MAX_GUILD_OVERRIDES, VoiceE2eeScope, VoiceNoiseSuppressionConfigUpdateRequest, + VoiceNoiseSuppressionGuildOverride, }, }, config::AdminConfig, @@ -211,6 +212,10 @@ pub async fn instance_config_post( Ok(update) => instance_config_result(client.update_instance_config(&update).await), Err(message) => FlashData::error(message), }, + "update_push_service_delivery" => match build_push_service_delivery_update(&form) { + Ok(update) => instance_config_result(client.update_instance_config(&update).await), + Err(message) => FlashData::error(message), + }, "update_experiment_delivery" => match build_experiment_delivery_update(&form) { Ok(update) => instance_config_result(client.update_instance_config(&update).await), Err(message) => FlashData::error(message), @@ -496,6 +501,21 @@ fn parse_experiment_rollout_salt( Ok(Some(salt.to_owned())) } +fn parse_push_service_delivery_rollout_salt( + form: &MultiValueForm, + key: &str, +) -> Result, String> { + let salt = parse_experiment_rollout_salt(form, key)?; + if let Some(value) = salt.as_deref() + && !value + .bytes() + .all(|byte| byte.is_ascii_graphic() || byte == b' ') + { + return Err("Rollout salt must use printable ASCII".to_owned()); + } + Ok(salt) +} + fn is_experiment_snowflake(value: &str) -> bool { !value.is_empty() && value.len() <= EXPERIMENT_MAX_SNOWFLAKE_LENGTH @@ -665,6 +685,38 @@ fn build_screen_share_delivery_update( }) } +fn build_push_service_delivery_update( + form: &MultiValueForm, +) -> Result { + Ok(InstanceConfigUpdateRequest { + push_service_delivery: Some(PushServiceDeliveryConfigUpdateRequest { + enabled: Some(form.bool_value("push_service_delivery_enabled")), + rollout_basis_points: parse_form_number( + form, + "push_service_delivery_rollout_basis_points", + "Rollout basis points", + 0, + EXPERIMENT_ROLLOUT_BASIS_POINTS_MAX, + )?, + rollout_salt: parse_push_service_delivery_rollout_salt( + form, + "push_service_delivery_rollout_salt", + )?, + included_user_ids: Some(parse_experiment_user_ids( + form.first("push_service_delivery_included_user_ids") + .unwrap_or_default(), + "Included user IDs", + )?), + excluded_user_ids: Some(parse_experiment_user_ids( + form.first("push_service_delivery_excluded_user_ids") + .unwrap_or_default(), + "Excluded user IDs", + )?), + }), + ..Default::default() + }) +} + fn build_experiment_delivery_update( form: &MultiValueForm, ) -> Result { diff --git a/fluxer_admin/src/templates/pages/instance_config.rs b/fluxer_admin/src/templates/pages/instance_config.rs index d2e56dfd5..9b26c08e4 100644 --- a/fluxer_admin/src/templates/pages/instance_config.rs +++ b/fluxer_admin/src/templates/pages/instance_config.rs @@ -5,10 +5,10 @@ use crate::{ AppPublicConfigResponse, EXPERIMENT_MAX_TARGETED_USERS, ExperimentDeliveryConfigResponse, GatewayRolloutConfigResponse, InstanceConfigResponse, InstanceIntegrationsResponse, InstanceMediaResponse, InstancePolicyResponse, InstanceRegistrationResponse, - LimitConfigResponse, NoiseSuppressionBackend, PendingRegistrationResponse, - RegistrationUrlResponse, SCREEN_SHARE_DELIVERY_DEFAULT_SALT, - ScreenShareDeliveryConfigResponse, SsoConfigResponse, VOICE_NS_MAX_GUILD_OVERRIDES, - VoiceNoiseSuppressionConfigResponse, + LimitConfigResponse, NoiseSuppressionBackend, PUSH_SERVICE_DELIVERY_DEFAULT_SALT, + PendingRegistrationResponse, PushServiceDeliveryConfigResponse, RegistrationUrlResponse, + SCREEN_SHARE_DELIVERY_DEFAULT_SALT, ScreenShareDeliveryConfigResponse, SsoConfigResponse, + VOICE_NS_MAX_GUILD_OVERRIDES, VoiceNoiseSuppressionConfigResponse, }, config::AdminConfig, middleware::auth::AuthContext, @@ -150,6 +150,7 @@ pub fn instance_config_page( (gateway_rollout_section(base, csrf_token, &instance_config.gateway_rollout)) (voice_noise_suppression_section(base, csrf_token, &instance_config.voice_noise_suppression)) (screen_share_delivery_section(base, csrf_token, &instance_config.screen_share_delivery)) + (push_service_delivery_section(base, csrf_token, &instance_config.push_service_delivery)) (experiment_delivery_section(base, csrf_token, &instance_config.experiment_delivery)) @if let Some(limit_config) = limit_config { (limit_config_section(base, limit_config)) @@ -1112,7 +1113,7 @@ fn voice_noise_suppression_section( p class="text-xs text-neutral-500" { "One snowflake per line, or comma separated. These users are targeted \ regardless of the percentage above. IDs must contain 1 to 20 decimal \ - digits. Invalid entries prevent the save; blank entries and duplicate \ + digits. Invalid entries prevent the save. Blank entries and duplicate \ IDs are ignored." } } @@ -1131,7 +1132,7 @@ fn voice_noise_suppression_section( )) p class="text-xs text-neutral-500" { "Same format. Exclusion wins over both the always-on list and the \ - percentage, so this is the per-user kill switch." + percentage. This is the per-user kill switch." } } @@ -1257,7 +1258,7 @@ fn screen_share_delivery_section( p class="text-xs text-neutral-500" { "One snowflake per line, or comma separated. These users are targeted \ regardless of the percentage above. IDs must contain 1 to 20 decimal \ - digits. Invalid entries prevent the save; blank entries and duplicate \ + digits. Invalid entries prevent the save. Blank entries and duplicate \ IDs are ignored." } } @@ -1276,7 +1277,7 @@ fn screen_share_delivery_section( )) p class="text-xs text-neutral-500" { "Same format. Exclusion wins over both the always-on list and the \ - percentage, so this is the per-user kill switch." + percentage. This is the per-user kill switch." } } @@ -1289,6 +1290,115 @@ fn screen_share_delivery_section( ) } +fn push_service_delivery_section( + base: &str, + csrf_token: &str, + push_service_delivery: &PushServiceDeliveryConfigResponse, +) -> Markup { + let status = if push_service_delivery.enabled { + ("Live", BadgeVariant::Success) + } else { + ("Inert", BadgeVariant::Default) + }; + let included_user_ids = push_service_delivery.included_user_ids.join("\n"); + let excluded_user_ids = push_service_delivery.excluded_user_ids.join("\n"); + section_card_with_description( + "Push Service Delivery", + "Routes push notification delivery for the selected accounts through the push service. \ + Accounts the rollout does not select keep the current path.", + html! { + form method="post" action={(base) "/instance-config?action=update_push_service_delivery"} { + (csrf_input(csrf_token)) + div class="space-y-6" { + div class="flex flex-wrap items-center gap-2" { + h3 class="text-sm font-semibold text-neutral-900" { "Master switch" } + (badge(status.0, status.1)) + span class="text-xs text-neutral-500" { + "Config version " (push_service_delivery.config_version) + } + } + (checkbox( + "push_service_delivery_enabled", + "true", + "Hand push notifications to the push service", + push_service_delivery.enabled, + true, + )) + p class="text-xs text-neutral-500" { + "Off is the safe state. With this unchecked every notification keeps the \ + current delivery path, so the rollout and targeting fields below have no \ + effect at all." + } + + h3 class="text-sm font-semibold text-neutral-900" { "Rollout" } + (number_field( + "push_service_delivery_rollout_basis_points", + "Rollout (basis points)", + &push_service_delivery.rollout_basis_points.to_string(), + Some(0), Some(10000), "1", + Some("Share of users bucketed into the canary, in basis points: 0 is nobody, 100 is 1%, 10000 is everybody."), + )) + div class="flex flex-col gap-2" { + (text_input( + "push_service_delivery_rollout_salt", + "Rollout Salt", + &push_service_delivery.rollout_salt, + PUSH_SERVICE_DELIVERY_DEFAULT_SALT, + )) + p class="text-xs text-neutral-500" { + "Seeds the bucketing hash. Changing it reshuffles which users fall \ + inside the percentage above. Leave it alone to keep the current \ + cohort stable." + } + } + div class="flex flex-col gap-2" { + (textarea_input( + "push_service_delivery_included_user_ids", + "Always-on User IDs", + "1500000000000000001\n1500000000000000002", + &included_user_ids, + 4, + false, + )) + (entry_count_hint( + push_service_delivery.included_user_ids.len(), + EXPERIMENT_MAX_TARGETED_USERS, + )) + p class="text-xs text-neutral-500" { + "One snowflake per line, or comma separated. These users are targeted \ + regardless of the percentage above. IDs must contain 1 to 20 decimal \ + digits. Invalid entries prevent the save. Blank entries and duplicate \ + IDs are ignored." + } + } + div class="flex flex-col gap-2" { + (textarea_input( + "push_service_delivery_excluded_user_ids", + "Never-on User IDs", + "1500000000000000003\n1500000000000000004", + &excluded_user_ids, + 4, + false, + )) + (entry_count_hint( + push_service_delivery.excluded_user_ids.len(), + EXPERIMENT_MAX_TARGETED_USERS, + )) + p class="text-xs text-neutral-500" { + "Same format. Exclusion wins over both the always-on list and the \ + percentage. This is the per-user kill switch." + } + } + + (form_actions(html! { + (submit_button("Save Push Service Delivery Configuration")) + })) + } + } + }, + ) +} + fn experiment_delivery_section( base: &str, csrf_token: &str, diff --git a/fluxer_admin/tests/api_deserialization.rs b/fluxer_admin/tests/api_deserialization.rs index f80361fbf..57094af88 100644 --- a/fluxer_admin/tests/api_deserialization.rs +++ b/fluxer_admin/tests/api_deserialization.rs @@ -418,6 +418,14 @@ fn deserialize_instance_config_response_with_unknown_keys() { "future_delivery_knob": 9, "excluded_user_ids": [] }, + "push_service_delivery": { + "enabled": true, + "config_version": 3, + "rollout_basis_points": 5000, + "rollout_salt": "push-service-delivery-v1", + "included_user_ids": ["1500000000000000002"], + "excluded_user_ids": [] + }, "experiment_delivery": {"poll_interval_seconds": 300, "poll_jitter_percent": 15}, "registration": { "mode": "open", diff --git a/fluxer_admin/tests/htmx_acceptance.rs b/fluxer_admin/tests/htmx_acceptance.rs index 5596c2eb0..c79265850 100644 --- a/fluxer_admin/tests/htmx_acceptance.rs +++ b/fluxer_admin/tests/htmx_acceptance.rs @@ -817,6 +817,9 @@ async fn spawn_mock_api() -> String { async fn mock_api(method: Method, uri: Uri) -> Response { let path = uri.path().to_owned(); + if method == Method::PATCH && path == "/admin/instance/config" { + return json_response(instance_config()); + } match (method, path.as_str()) { (Method::GET, "/admin/users/@me") => json_response(json!({ "user": admin_user() })), (Method::GET, "/admin/api-keys") => json_response(json!([])), diff --git a/fluxer_api/src/api/admin/controllers/InstanceConfigAdminController.ts b/fluxer_api/src/api/admin/controllers/InstanceConfigAdminController.ts index ce9411ff0..95176d26a 100644 --- a/fluxer_api/src/api/admin/controllers/InstanceConfigAdminController.ts +++ b/fluxer_api/src/api/admin/controllers/InstanceConfigAdminController.ts @@ -13,7 +13,11 @@ import {deriveSsoRedirectUri, normalizeAndValidateSsoConfig} from '@app/api/inst import {requireAdminACL} from '@app/api/middleware/AdminMiddleware'; import {RateLimitMiddleware} from '@app/api/middleware/RateLimitMiddleware'; import {OpenAPI} from '@app/api/middleware/ResponseTypeMiddleware'; -import {getGatewayRolloutConfigPublisher, getInstanceConfigRepository} from '@app/api/middleware/ServiceSingletons'; +import { + getGatewayRolloutConfigPublisher, + getInstanceConfigRepository, + getPushServiceDeliveryConfigPublisher, +} from '@app/api/middleware/ServiceSingletons'; import {RateLimitConfigs} from '@app/api/RateLimitConfig'; import type {HonoApp, HonoEnv} from '@app/api/types/HonoEnv'; import {Validator} from '@app/api/Validator'; @@ -31,6 +35,7 @@ import { RegistrationUrlIdParam, } from '@fluxer/schema/src/domains/admin/AdminSchemas'; import {GatewayRolloutConfigSchema} from '@fluxer/schema/src/domains/admin/GatewayRolloutSchemas'; +import {PushServiceDeliveryConfigSchema} from '@fluxer/schema/src/domains/admin/PushServiceDeliverySchemas'; import {ScreenShareDeliveryConfigSchema} from '@fluxer/schema/src/domains/admin/ScreenShareDeliverySchemas'; import {VoiceNoiseSuppressionConfigSchema} from '@fluxer/schema/src/domains/admin/VoiceNoiseSuppressionSchemas'; import {UserIdParam} from '@fluxer/schema/src/domains/common/CommonParamSchemas'; @@ -61,6 +66,7 @@ async function buildInstanceConfigResponse(): Promise { gatewayRollout, voiceNoiseSuppression, screenShareDelivery, + pushServiceDelivery, experimentDelivery, registrationConfig, registrationUrls, @@ -70,6 +76,7 @@ async function buildInstanceConfigResponse(): Promise { instanceConfigRepository.getGatewayRolloutConfig(), instanceConfigRepository.getVoiceNoiseSuppressionConfig(), instanceConfigRepository.getScreenShareDeliveryConfig(), + instanceConfigRepository.getPushServiceDeliveryConfig(), instanceConfigRepository.getExperimentDeliveryConfig(), instanceConfigRepository.getRegistrationConfig(), instanceConfigRepository.getRegistrationUrlsForAdmin(), @@ -102,6 +109,7 @@ async function buildInstanceConfigResponse(): Promise { gateway_rollout: gatewayRollout, voice_noise_suppression: voiceNoiseSuppression, screen_share_delivery: screenShareDelivery, + push_service_delivery: pushServiceDelivery, experiment_delivery: experimentDelivery, registration: { ...registrationConfig, @@ -247,43 +255,54 @@ export function InstanceConfigAdminController(app: HonoApp) { const shouldGrantSetupCompleterAdmin = appPublicBeforeUpdate !== null && completesInitialSetup(data, appPublicBeforeUpdate.setup.configured); if (data.gateway_rollout) { - const currentRollout = await instanceConfigRepository.getGatewayRolloutConfig(); - const merged = {...currentRollout, ...data.gateway_rollout}; - const validated = GatewayRolloutConfigSchema.parse(merged); - await instanceConfigRepository.setGatewayRolloutConfig(validated); - await getGatewayRolloutConfigPublisher().publish(validated); + const patch = data.gateway_rollout; + const landed = await instanceConfigRepository.updateGatewayRolloutConfig((current) => + GatewayRolloutConfigSchema.parse({...current, ...patch}), + ); + await getGatewayRolloutConfigPublisher().publish(landed); } if (data.voice_noise_suppression) { const patch = omitUndefinedFields(data.voice_noise_suppression); if (Object.keys(patch).length > 0) { - const currentNoiseSuppression = await instanceConfigRepository.getVoiceNoiseSuppressionConfig(); - const validated = VoiceNoiseSuppressionConfigSchema.parse({ - ...currentNoiseSuppression, - ...patch, - config_version: currentNoiseSuppression.config_version + 1, - }); - await instanceConfigRepository.setVoiceNoiseSuppressionConfig(validated); + await instanceConfigRepository.updateVoiceNoiseSuppressionConfig((current) => + VoiceNoiseSuppressionConfigSchema.parse({ + ...current, + ...patch, + config_version: current.config_version + 1, + }), + ); } } if (data.screen_share_delivery) { const patch = omitUndefinedFields(data.screen_share_delivery); if (Object.keys(patch).length > 0) { - const currentScreenShareDelivery = await instanceConfigRepository.getScreenShareDeliveryConfig(); - const validated = ScreenShareDeliveryConfigSchema.parse({ - ...currentScreenShareDelivery, - ...patch, - config_version: currentScreenShareDelivery.config_version + 1, - }); - await instanceConfigRepository.setScreenShareDeliveryConfig(validated); + await instanceConfigRepository.updateScreenShareDeliveryConfig((current) => + ScreenShareDeliveryConfigSchema.parse({ + ...current, + ...patch, + config_version: current.config_version + 1, + }), + ); + } + } + if (data.push_service_delivery) { + const patch = omitUndefinedFields(data.push_service_delivery); + if (Object.keys(patch).length > 0) { + const landed = await instanceConfigRepository.updatePushServiceDeliveryConfig((current) => + PushServiceDeliveryConfigSchema.parse({ + ...current, + ...patch, + config_version: current.config_version + 1, + }), + ); + await getPushServiceDeliveryConfigPublisher().publish(landed); } } if (data.experiment_delivery) { - const currentExperimentDelivery = await instanceConfigRepository.getExperimentDeliveryConfig(); - const validated = ExperimentDeliveryConfigSchema.parse({ - ...currentExperimentDelivery, - ...data.experiment_delivery, - }); - await instanceConfigRepository.setExperimentDeliveryConfig(validated); + const patch = data.experiment_delivery; + await instanceConfigRepository.updateExperimentDeliveryConfig((current) => + ExperimentDeliveryConfigSchema.parse({...current, ...patch}), + ); } if (data.sso) { const sso = data.sso; @@ -308,21 +327,22 @@ export function InstanceConfigAdminController(app: HonoApp) { const validated = await normalizeAndValidateSsoConfig(next, { testModeEnabled: Config.dev.testModeEnabled, }); + const supplied = (field: keyof typeof sso, value: T): T | undefined => + readOptionalField(sso, field) === undefined ? undefined : value; await instanceConfigRepository.setSsoConfig({ - enabled: validated.enabled, - enforced: validated.enforced, - displayName: next.displayName, - issuer: validated.issuer, - authorizationUrl: validated.authorizationUrl, - tokenUrl: validated.tokenUrl, - userInfoUrl: validated.userInfoUrl, - jwksUrl: validated.jwksUrl, - clientId: validated.clientId, + enabled: supplied('enabled', validated.enabled), + enforced: supplied('enforced', validated.enforced), + displayName: supplied('display_name', next.displayName), + issuer: supplied('issuer', validated.issuer), + authorizationUrl: supplied('authorization_url', validated.authorizationUrl), + tokenUrl: supplied('token_url', validated.tokenUrl), + userInfoUrl: supplied('userinfo_url', validated.userInfoUrl), + jwksUrl: supplied('jwks_url', validated.jwksUrl), + clientId: supplied('client_id', validated.clientId), clientSecret: readOptionalField(sso, 'client_secret'), - scope: next.scope, - allowedEmailDomains: validated.allowedEmailDomains, - autoProvision: next.autoProvision, - redirectUri: null, + scope: supplied('scope', next.scope), + allowedEmailDomains: supplied('allowed_domains', validated.allowedEmailDomains), + autoProvision: supplied('auto_provision', next.autoProvision), }); } if (data.registration) { @@ -625,7 +645,6 @@ export function InstanceConfigAdminController(app: HonoApp) { async (ctx) => { const userId = ctx.req.valid('param').user_id.toString(); const decision = ctx.req.valid('json').status === 'approved' ? 'approve' : 'reject'; - await instanceConfigRepository.getPendingRegistrations(); await updatePendingRegistrationUser(ctx, userId, decision); await instanceConfigRepository.removePendingRegistration(userId); return ctx.json(await buildInstanceConfigResponse()); @@ -638,27 +657,47 @@ async function applyInstancePolicyUpdate( policy: NonNullable, ): Promise { const instanceConfigRepository = getInstanceConfigRepository(); - const [current, appPublic] = await Promise.all([ - instanceConfigRepository.getInstancePolicyConfig(), - instanceConfigRepository.getAppPublicConfig(), - ]); + const appPublic = await instanceConfigRepository.getAppPublicConfig(); + const adminUser = + policy.single_community_enabled === true + ? await ctx.get('userRepository').findUnique(ctx.get('adminUserId')) + : null; + let enablesSingleCommunity = false; + await instanceConfigRepository.updateInstancePolicyConfig((current) => { + const planned = planInstancePolicyPatch(policy, current, { + setupConfigured: appPublic.setup.configured, + adminUserFound: adminUser !== null, + }); + enablesSingleCommunity = planned.enablesSingleCommunity; + return planned.patch; + }); + if (enablesSingleCommunity && adminUser) { + await ctx.get('singleCommunityService').ensureStockCommunity({ + owner: adminUser, + name: policy.single_community_name?.trim() || appPublic.branding.product_name, + }); + } + if (policy.premium_mode !== undefined) { + await ctx.get('limitConfigService').updatePolicyConfig({premium_mode: policy.premium_mode}); + } +} + +function planInstancePolicyPatch( + policy: NonNullable, + current: InstancePolicyConfig, + context: {setupConfigured: boolean; adminUserFound: boolean}, +): {patch: Partial; enablesSingleCommunity: boolean} { const patch: Partial = {}; + let enablesSingleCommunity = false; if ( policy.single_community_enabled !== undefined && policy.single_community_enabled !== current.single_community_enabled ) { if (policy.single_community_enabled) { - if (appPublic.setup.configured && current.single_community_guild_id == null) { + if ((context.setupConfigured && current.single_community_guild_id == null) || !context.adminUserFound) { throw new InstancePolicyTransitionNotAllowedError(); } - const adminUser = await ctx.get('userRepository').findUnique(ctx.get('adminUserId')); - if (!adminUser) { - throw new InstancePolicyTransitionNotAllowedError(); - } - await ctx.get('singleCommunityService').ensureStockCommunity({ - owner: adminUser, - name: policy.single_community_name?.trim() || appPublic.branding.product_name, - }); + enablesSingleCommunity = true; } else { patch.single_community_enabled = false; } @@ -679,9 +718,6 @@ async function applyInstancePolicyUpdate( patch.direct_messages_locked = true; } } - if (policy.premium_mode !== undefined) { - patch.premium_mode = policy.premium_mode; - } if (policy.services) { if (policy.services.gif_enabled !== undefined) { patch.gif_enabled = policy.services.gif_enabled ?? null; @@ -704,11 +740,7 @@ async function applyInstancePolicyUpdate( patch.deferred_phone_gate_member_threshold = policy.deferred_phone_gate.member_threshold; } } - if (patch.premium_mode !== undefined) { - await ctx.get('limitConfigService').updatePolicyConfig(patch); - } else if (Object.keys(patch).length > 0) { - await instanceConfigRepository.setInstancePolicyConfig(patch); - } + return {patch, enablesSingleCommunity}; } async function updatePendingRegistrationUser( diff --git a/fluxer_api/src/api/admin/tests/InstanceConfigAdminController.test.ts b/fluxer_api/src/api/admin/tests/InstanceConfigAdminController.test.ts new file mode 100644 index 000000000..c4642b6aa --- /dev/null +++ b/fluxer_api/src/api/admin/tests/InstanceConfigAdminController.test.ts @@ -0,0 +1,107 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +import type {AdminAuditLog} from '@app/api/admin/IAdminRepository'; +import type {TestAccount} from '@app/api/auth/tests/AuthTestUtils'; +import {createTestAccount, setUserACLs} from '@app/api/auth/tests/AuthTestUtils'; +import {setCassandraQueryExecutorForTesting} from '@app/api/database/CassandraQueryExecution'; +import {PushServiceDeliveryConfigPublisher} from '@app/api/instance/PushServiceDeliveryConfigPublisher'; +import {InstanceConfigWriteRaceExecutor} from '@app/api/instance/tests/InstanceConfigWriteRaceExecutor'; +import {getAdminRepository} from '@app/api/middleware/ServiceSingletons'; +import type {ApiTestHarness} from '@app/api/test/ApiTestHarness'; +import {createApiTestHarness} from '@app/api/test/ApiTestHarness'; +import {InMemoryCassandraQueryExecutor} from '@app/api/test/InMemoryCassandraQueryExecutor'; +import {HTTP_STATUS} from '@app/api/test/TestConstants'; +import {createBuilder} from '@app/api/test/TestRequestBuilder'; +import {AdminACLs} from '@fluxer/constants/src/AdminACLs'; +import {APIErrorCodes} from '@fluxer/constants/src/ApiErrorCodes'; +import type {InstanceConfigResponse} from '@fluxer/schema/src/domains/admin/AdminSchemas'; +import { + DEFAULT_PUSH_SERVICE_DELIVERY_CONFIG, + type PushServiceDeliveryConfig, +} from '@fluxer/schema/src/domains/admin/PushServiceDeliverySchemas'; +import {afterAll, afterEach, beforeAll, beforeEach, describe, expect, it, vi} from 'vitest'; + +const PUSH_SERVICE_DELIVERY_CONFIG_KEY = 'push_service_delivery_config'; + +describe('instance config admin PATCH under concurrent writes', () => { + let harness: ApiTestHarness; + let executor: InstanceConfigWriteRaceExecutor; + + beforeAll(async () => { + harness = await createApiTestHarness(); + executor = new InstanceConfigWriteRaceExecutor(new InMemoryCassandraQueryExecutor()); + setCassandraQueryExecutorForTesting(executor); + }); + + beforeEach(async () => { + await harness.reset(); + }); + + afterEach(() => { + vi.restoreAllMocks(); + }); + + afterAll(async () => { + await harness.shutdown(); + }); + + const createAdmin = async (): Promise => + await setUserACLs(harness, await createTestAccount(harness), [ + AdminACLs.AUTHENTICATE, + AdminACLs.INSTANCE_CONFIG_VIEW, + AdminACLs.INSTANCE_CONFIG_UPDATE, + ]); + + const patchConfig = (admin: TestAccount, body: Record) => + createBuilder(harness, admin.token).patch('/admin/instance/config').body(body); + + const spyOnPushDeliveryPublishes = () => + vi.spyOn(PushServiceDeliveryConfigPublisher.prototype, 'publish').mockResolvedValue(undefined); + + async function readStoredPushServiceDelivery(): Promise { + const raw = await executor.readDirectly(PUSH_SERVICE_DELIVERY_CONFIG_KEY); + if (raw === null) throw new Error('push service delivery config was never stored'); + return JSON.parse(raw) as PushServiceDeliveryConfig; + } + + async function listConfigUpdateAudits(): Promise> { + const logs = await getAdminRepository().listAllAuditLogsPaginated(100000); + return logs.filter((log) => log.action === 'update_instance_config'); + } + + it('answers with a conflict and neither writes, publishes nor audits once every attempt has lost the race', async () => { + const publish = spyOnPushDeliveryPublishes(); + const admin = await createAdmin(); + await patchConfig(admin, {push_service_delivery: {enabled: true, rollout_basis_points: 1000}}).execute(); + publish.mockClear(); + const auditsBefore = await listConfigUpdateAudits(); + executor.watch(PUSH_SERVICE_DELIVERY_CONFIG_KEY); + let competingWrites = 0; + executor.competeBeforeEachWrite(async () => { + competingWrites++; + await executor.writeDirectly( + PUSH_SERVICE_DELIVERY_CONFIG_KEY, + JSON.stringify({ + ...DEFAULT_PUSH_SERVICE_DELIVERY_CONFIG, + enabled: false, + rollout_basis_points: 1000, + config_version: 100 + competingWrites, + }), + ); + }); + + await patchConfig(admin, {push_service_delivery: {rollout_basis_points: 5000}}) + .expect(HTTP_STATUS.CONFLICT, APIErrorCodes.CONFLICT) + .execute(); + + expect(executor.events).not.toContain('write'); + expect(await readStoredPushServiceDelivery()).toEqual({ + ...DEFAULT_PUSH_SERVICE_DELIVERY_CONFIG, + enabled: false, + rollout_basis_points: 1000, + config_version: 100 + competingWrites, + }); + expect(publish).not.toHaveBeenCalled(); + expect(await listConfigUpdateAudits()).toHaveLength(auditsBefore.length); + }); +}); diff --git a/fluxer_api/src/api/admin/tests/InstanceConfigStaleSectionWrites.test.ts b/fluxer_api/src/api/admin/tests/InstanceConfigStaleSectionWrites.test.ts new file mode 100644 index 000000000..5a995640f --- /dev/null +++ b/fluxer_api/src/api/admin/tests/InstanceConfigStaleSectionWrites.test.ts @@ -0,0 +1,94 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +import type {TestAccount} from '@app/api/auth/tests/AuthTestUtils'; +import {createTestAccount, setUserACLs} from '@app/api/auth/tests/AuthTestUtils'; +import {setCassandraQueryExecutorForTesting} from '@app/api/database/CassandraQueryExecution'; +import {InstanceConfigWriteRaceExecutor} from '@app/api/instance/tests/InstanceConfigWriteRaceExecutor'; +import type {ApiTestHarness} from '@app/api/test/ApiTestHarness'; +import {createApiTestHarness} from '@app/api/test/ApiTestHarness'; +import {InMemoryCassandraQueryExecutor} from '@app/api/test/InMemoryCassandraQueryExecutor'; +import {HTTP_STATUS} from '@app/api/test/TestConstants'; +import {createBuilder} from '@app/api/test/TestRequestBuilder'; +import {AdminACLs} from '@fluxer/constants/src/AdminACLs'; +import {APIErrorCodes} from '@fluxer/constants/src/ApiErrorCodes'; +import type {InstanceConfigResponse} from '@fluxer/schema/src/domains/admin/AdminSchemas'; +import {afterAll, beforeAll, beforeEach, describe, expect, it} from 'vitest'; + +const INSTANCE_POLICY_CONFIG_KEY = 'instance_policy_config'; + +describe('instance config admin PATCH against state another node changed', () => { + let harness: ApiTestHarness; + let executor: InstanceConfigWriteRaceExecutor; + + beforeAll(async () => { + harness = await createApiTestHarness(); + executor = new InstanceConfigWriteRaceExecutor(new InMemoryCassandraQueryExecutor()); + setCassandraQueryExecutorForTesting(executor); + }); + + beforeEach(async () => { + await harness.reset(); + }); + + afterAll(async () => { + await harness.shutdown(); + }); + + const createAdmin = async (): Promise => + await setUserACLs(harness, await createTestAccount(harness), [ + AdminACLs.AUTHENTICATE, + AdminACLs.INSTANCE_CONFIG_VIEW, + AdminACLs.INSTANCE_CONFIG_UPDATE, + ]); + + const patchConfig = (admin: TestAccount, body: Record) => + createBuilder(harness, admin.token).patch('/admin/instance/config').body(body); + + it('keeps an SSO field another node changed when a patch changes a different one', async () => { + const admin = await createAdmin(); + await patchConfig(admin, {sso: {display_name: 'Before', client_id: 'client-before'}}).execute(); + await executor.writeDirectly('sso_display_name', 'Changed on another node'); + + await patchConfig(admin, {sso: {client_id: 'client-after'}}).execute(); + + expect(await executor.readDirectly('sso_display_name')).toBe('Changed on another node'); + expect(await executor.readDirectly('sso_client_id')).toBe('client-after'); + }); + + it('refuses to disable direct messages when their lock lands between the read and the write', async () => { + const admin = await createAdmin(); + await patchConfig(admin, {policy: {services: {gif_enabled: true}}}).execute(); + executor.watch(INSTANCE_POLICY_CONFIG_KEY); + let competed = false; + executor.competeBeforeEachWrite(async () => { + if (competed) return; + competed = true; + await executor.writeDirectly( + INSTANCE_POLICY_CONFIG_KEY, + JSON.stringify({direct_messages_disabled: false, direct_messages_locked: true, gif_enabled: true}), + ); + }); + + await patchConfig(admin, {policy: {direct_messages_disabled: true}}) + .expect(HTTP_STATUS.BAD_REQUEST, APIErrorCodes.INSTANCE_POLICY_TRANSITION_NOT_ALLOWED) + .execute(); + + const stored = JSON.parse((await executor.readDirectly(INSTANCE_POLICY_CONFIG_KEY)) ?? 'null'); + expect(stored).toMatchObject({direct_messages_disabled: false, direct_messages_locked: true, gif_enabled: true}); + }); + + it('applies the DM rule and a premium mode change from one request', async () => { + const admin = await createAdmin(); + await patchConfig(admin, {policy: {direct_messages_disabled: true}}).execute(); + + const updated = await patchConfig(admin, { + policy: {direct_messages_disabled: false, premium_mode: 'mirror'}, + }).execute(); + + expect(updated.policy).toMatchObject({ + direct_messages_disabled: false, + direct_messages_locked: true, + premium_mode: 'mirror', + }); + }); +}); diff --git a/fluxer_api/src/api/auth/AuthRegistration.ts b/fluxer_api/src/api/auth/AuthRegistration.ts index 24c56cd38..5525260f7 100644 --- a/fluxer_api/src/api/auth/AuthRegistration.ts +++ b/fluxer_api/src/api/auth/AuthRegistration.ts @@ -7,12 +7,14 @@ import * as AuthUtility from '@app/api/auth/AuthUtility'; import type {IRegistrationRiskEvaluator} from '@app/api/auth/services/IRegistrationRiskEvaluator'; import {createEmailVerificationToken, createInviteCode, createUserID, type UserID} from '@app/api/BrandedTypes'; import type {APIConfig} from '@app/api/config/APIConfig'; +import type {UserRow} from '@app/api/database/types/UserTypes'; import type {IDiscriminatorService} from '@app/api/infrastructure/DiscriminatorService'; import type {KVActivityTracker} from '@app/api/infrastructure/KVActivityTracker'; import { type InstanceConfigRepository, type InstanceRegistrationUrl, REGISTRATION_PENDING_APPROVAL_TRAIT, + type RegistrationUrlClaim, } from '@app/api/instance/InstanceConfigRepository'; import type {SingleCommunityService} from '@app/api/instance/SingleCommunityService'; import type {InviteService} from '@app/api/invite/InviteService'; @@ -135,9 +137,6 @@ export async function register( } const now = new Date(); const registrationAccess = await resolveRegistrationAccess(instanceConfigRepository, data.registration_url_code); - if (registrationAccess.pendingApproval) { - await instanceConfigRepository.getPendingRegistrations(); - } const clientIp = requireClientIp(request, { trustClientIpHeader: config.proxy.trust_client_ip_header, clientIpHeaderName: config.proxy.client_ip_header, @@ -228,7 +227,7 @@ export async function register( const userLocale = parseAcceptLanguage(acceptLanguage); const passwordHash = data.password ? await AuthPassword.hashPassword(ctx, data.password) : null; const flags = config.nodeEnv === 'development' ? UserFlags.STAFF : 0n; - let user = await users.create({ + const userRow: UserRow = { user_id: userId, username, discriminator, @@ -287,7 +286,39 @@ export async function register( mention_flags: null, last_voice_activity_sharing_change_at: null, version: 1, - }); + }; + const registrationUrlUse = await claimRegistrationUrlUse( + instanceConfigRepository, + registrationAccess.registrationUrl, + userId, + ); + let user: User; + let createAttempted = false; + try { + if (registrationAccess.pendingApproval) { + await instanceConfigRepository.addPendingRegistration({ + user_id: userId.toString(), + username: userRow.username, + discriminator: userRow.discriminator, + global_name: userRow.global_name, + email: rawEmail, + requested_at: now.toISOString(), + registration_url_id: registrationAccess.registrationUrl?.id ?? null, + client_ip: clientIp, + }); + } + createAttempted = true; + user = await users.create(userRow); + } catch (error) { + if (!createAttempted) { + await withdrawSignupOfUncreatedAccount(instanceConfigRepository, { + userId, + registrationUrlUse, + pendingApproval: registrationAccess.pendingApproval, + }); + } + throw error; + } await users.upsertSettings( UserSettings.getDefaultUserSettings({ userId, @@ -401,20 +432,7 @@ export async function register( } if (rawEmail && emailEnabled) await maybeSendVerificationEmail(ctx, {user, email: rawEmail}); await users.createAuthorizedIp(userId, clientIp); - if (registrationAccess.registrationUrl) { - await instanceConfigRepository.recordRegistrationUrlUse(registrationAccess.registrationUrl.id, user.id.toString()); - } if (registrationAccess.pendingApproval) { - await instanceConfigRepository.addPendingRegistration({ - user_id: user.id.toString(), - username: user.username, - discriminator: user.discriminator, - global_name: user.globalName, - email: rawEmail, - requested_at: now.toISOString(), - registration_url_id: registrationAccess.registrationUrl?.id ?? null, - client_ip: clientIp, - }); return { registration_pending_approval: true, user_id: user.id.toString(), @@ -469,6 +487,38 @@ function shouldAttemptBootstrapAdminGrant( ); } +async function claimRegistrationUrlUse( + instanceConfigRepository: InstanceConfigRepository, + registrationUrl: InstanceRegistrationUrl | null, + userId: UserID, +): Promise { + if (registrationUrl === null) return null; + const use = await instanceConfigRepository.claimRegistrationUrlUse(registrationUrl.id, userId.toString()); + if (use === null) { + throw new RegistrationUrlInvalidError(); + } + return use; +} + +async function withdrawSignupOfUncreatedAccount( + instanceConfigRepository: InstanceConfigRepository, + signup: {userId: UserID; registrationUrlUse: RegistrationUrlClaim | null; pendingApproval: boolean}, +): Promise { + try { + if (signup.registrationUrlUse !== null) { + await instanceConfigRepository.releaseRegistrationUrlUse(signup.registrationUrlUse); + } + if (signup.pendingApproval) { + await instanceConfigRepository.removePendingRegistration(signup.userId.toString()); + } + } catch (error) { + Logger.warn( + {userId: signup.userId.toString(), registrationUrlId: signup.registrationUrlUse?.registration_url_id, error}, + '[AuthRegistration] Failed to withdraw the registration URL use or pending approval of an account that was never created', + ); + } +} + async function resolveRegistrationAccess( instanceConfigRepository: InstanceConfigRepository, registrationUrlCode: string | null | undefined, diff --git a/fluxer_api/src/api/auth/services/SsoService.ts b/fluxer_api/src/api/auth/services/SsoService.ts index 62016bfa9..c6ef17a4f 100644 --- a/fluxer_api/src/api/auth/services/SsoService.ts +++ b/fluxer_api/src/api/auth/services/SsoService.ts @@ -382,21 +382,8 @@ export class SsoService { throw new RegistrationClosedError(); } const pendingApproval = registrationConfig.mode === 'approval'; - if (pendingApproval) { - await this.instanceConfigRepository.getPendingRegistrations(); - } const user = await this.provisionUserFromClaims(claims, config, {pendingApproval}); if (pendingApproval) { - await this.instanceConfigRepository.addPendingRegistration({ - user_id: user.id.toString(), - username: user.username, - discriminator: user.discriminator, - global_name: user.globalName, - email: user.email, - requested_at: new Date().toISOString(), - registration_url_id: null, - client_ip: null, - }); throw new RegistrationPendingApprovalError(); } return user; @@ -537,8 +524,22 @@ export class SsoService { version: 1, } as const; await this.claimSsoIdentity(userId, claims.sub, config); + let createAttempted = false; let userCreated = false; try { + if (options?.pendingApproval) { + await this.instanceConfigRepository.addPendingRegistration({ + user_id: userId.toString(), + username, + discriminator: discriminatorResult.discriminator, + global_name: globalName, + email: userRow.email, + requested_at: now.toISOString(), + registration_url_id: null, + client_ip: null, + }); + } + createAttempted = true; const user = await users.create(userRow); userCreated = true; await users.upsertSettings( @@ -557,6 +558,16 @@ export class SsoService { await this.ssoIdentityRepository.releaseIdentity(config.providerId, claims.sub).catch((releaseError) => { getLogger().error({releaseError}, 'Failed to release SSO identity after user provisioning failed'); }); + if (options?.pendingApproval && !createAttempted) { + await this.instanceConfigRepository + .removePendingRegistration(userId.toString()) + .catch((removeError: unknown) => { + getLogger().error( + {userId: userId.toString(), removeError}, + 'Failed to withdraw the pending approval of an SSO user that was never created', + ); + }); + } } throw error; } diff --git a/fluxer_api/src/api/auth/tests/RegistrationSignupRaces.test.ts b/fluxer_api/src/api/auth/tests/RegistrationSignupRaces.test.ts new file mode 100644 index 000000000..e87ad0f52 --- /dev/null +++ b/fluxer_api/src/api/auth/tests/RegistrationSignupRaces.test.ts @@ -0,0 +1,436 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +import {createHash} from 'node:crypto'; +import { + createAuthHarness, + createTestAccount, + createUniqueEmail, + createUniqueUsername, + enableSso, + setUserACLs, + type TestAccount, +} from '@app/api/auth/tests/AuthTestUtils'; +import {createUserID} from '@app/api/BrandedTypes'; +import type {UserRow} from '@app/api/database/types/UserTypes'; +import { + InstanceConfigRepository, + REGISTRATION_PENDING_APPROVAL_TRAIT, +} from '@app/api/instance/InstanceConfigRepository'; +import {getInstanceConfigRepository} from '@app/api/middleware/ServiceSingletons'; +import type {ApiTestHarness} from '@app/api/test/ApiTestHarness'; +import {HTTP_STATUS} from '@app/api/test/TestConstants'; +import {createBuilder, createBuilderWithoutAuth} from '@app/api/test/TestRequestBuilder'; +import {UserRepository} from '@app/api/user/repositories/UserRepository'; +import {AdminACLs} from '@fluxer/constants/src/AdminACLs'; +import {APIErrorCodes} from '@fluxer/constants/src/ApiErrorCodes'; +import type {InstanceConfigResponse} from '@fluxer/schema/src/domains/admin/AdminSchemas'; +import {afterAll, afterEach, beforeAll, beforeEach, describe, expect, it, vi} from 'vitest'; + +const REGISTRATION_URLS_KEY = 'registration_urls'; +const REGISTRATION_PENDING_APPROVALS_KEY = 'registration_pending_approvals'; + +interface RegistrationResponse { + user_id?: string; + token?: string; + registration_pending_approval?: true; + code?: string; +} + +function registrationBody(prefix: string, registrationUrlCode?: string): Record { + return { + email: createUniqueEmail(prefix), + username: createUniqueUsername(prefix), + global_name: 'Signup Race', + password: 'a-strong-password', + date_of_birth: '2000-01-01', + consent: true, + ...(registrationUrlCode === undefined ? {} : {registration_url_code: registrationUrlCode}), + }; +} + +describe('signups racing on registration URLs and pending approvals', () => { + let harness: ApiTestHarness; + let admin: TestAccount; + + beforeAll(async () => { + harness = await createAuthHarness(); + }); + + beforeEach(async () => { + await harness.reset(); + admin = await setUserACLs(harness, await createTestAccount(harness), [ + AdminACLs.AUTHENTICATE, + AdminACLs.INSTANCE_CONFIG_VIEW, + AdminACLs.INSTANCE_CONFIG_UPDATE, + ]); + }); + + afterEach(() => { + vi.restoreAllMocks(); + }); + + afterAll(async () => { + await harness?.shutdown(); + }); + + const register = (prefix: string, registrationUrlCode?: string) => + createBuilderWithoutAuth(harness) + .post('/auth/register') + .body(registrationBody(prefix, registrationUrlCode)) + .executeRaw(); + + const readAdminConfig = (): Promise => + createBuilder(harness, admin.token).get('/admin/instance/config').execute(); + + const completeSso = async (prefix: string) => { + const start = await createBuilderWithoutAuth<{state: string}>(harness) + .post('/auth/sso/start') + .body({redirect_to: '/me'}) + .execute(); + return createBuilderWithoutAuth(harness) + .post('/auth/sso/complete') + .body({code: createUniqueEmail(prefix), state: start.state}) + .executeRaw(); + }; + + const failCreateAfterTheUserRowIsWritten = () => { + const create = UserRepository.prototype.create; + vi.spyOn(UserRepository.prototype, 'create').mockImplementationOnce(async function ( + this: UserRepository, + row: UserRow, + ) { + await create.call(this, row); + throw new Error('the user indexes could not be written after the user row'); + }); + }; + + const failAfterThePendingApprovalIsStored = () => { + const addPendingRegistration = InstanceConfigRepository.prototype.addPendingRegistration; + vi.spyOn(InstanceConfigRepository.prototype, 'addPendingRegistration').mockImplementationOnce(async function ( + this: InstanceConfigRepository, + entry: Parameters[0], + ) { + await addPendingRegistration.call(this, entry); + throw new Error('the pending approval could not be published'); + }); + }; + + const expectOnePendingAccount = async () => { + const pending = (await readAdminConfig()).registration.pending_registrations; + expect(pending).toHaveLength(1); + const account = await new UserRepository().findUnique(createUserID(BigInt(pending[0]!.user_id))); + expect(account?.traits.has(REGISTRATION_PENDING_APPROVAL_TRAIT)).toBe(true); + }; + + it('never lets concurrent signups through a capped registration URL exceed max_uses', async () => { + const repository = getInstanceConfigRepository(); + await repository.setRegistrationConfig({mode: 'closed', admin_registration_urls_enabled: true}); + const {code, registrationUrl} = await repository.createRegistrationUrl({ + label: 'Capped', + createdByUserId: '1', + expiresAt: null, + maxUses: 2, + approvalRequired: false, + }); + + const attempts = await Promise.all(Array.from({length: 6}, (_, index) => register(`capped${index}`, code))); + + const admitted = attempts.filter((attempt) => attempt.response.status === HTTP_STATUS.OK); + const refused = attempts.filter((attempt) => attempt.response.status !== HTTP_STATUS.OK); + expect(admitted).toHaveLength(2); + for (const attempt of refused) { + expect(attempt.response.status).toBe(HTTP_STATUS.BAD_REQUEST); + expect(attempt.json.code).toBe(APIErrorCodes.REGISTRATION_URL_INVALID); + } + const stored = (await readAdminConfig()).registration.urls.find((url) => url.id === registrationUrl.id); + expect(stored?.use_count).toBe(2); + expect(admitted.map((attempt) => attempt.json.user_id)).toContain(stored?.last_used_by_user_id); + }); + + it('admits exactly max_uses when 120 signups race through a registration URL capped at 40', async () => { + const repository = getInstanceConfigRepository(); + await repository.setRegistrationConfig({mode: 'closed', admin_registration_urls_enabled: true}); + const {code, registrationUrl} = await repository.createRegistrationUrl({ + label: 'Capped at 40', + createdByUserId: '1', + expiresAt: null, + maxUses: 40, + approvalRequired: false, + }); + const registerUntilDecided = async (prefix: string) => { + for (let attempt = 0; attempt < 20; attempt += 1) { + const result = await register(`${prefix}r${attempt}`, code); + if (result.response.status !== HTTP_STATUS.SERVICE_UNAVAILABLE) return result; + } + throw new Error('a signup never reached a decision'); + }; + + const attempts = await Promise.all(Array.from({length: 120}, (_, index) => registerUntilDecided(`surge${index}`))); + + const admitted = attempts.filter((attempt) => attempt.response.status === HTTP_STATUS.OK); + expect(admitted).toHaveLength(40); + for (const attempt of attempts.filter((entry) => entry.response.status !== HTTP_STATUS.OK)) { + expect(attempt.response.status).toBe(HTTP_STATUS.BAD_REQUEST); + expect(attempt.json.code).toBe(APIErrorCodes.REGISTRATION_URL_INVALID); + } + const stored = (await readAdminConfig()).registration.urls.find((url) => url.id === registrationUrl.id); + expect(stored?.use_count).toBe(40); + }); + + it('counts every concurrent signup through an uncapped registration URL', async () => { + const repository = getInstanceConfigRepository(); + await repository.setRegistrationConfig({mode: 'closed', admin_registration_urls_enabled: true}); + const {code, registrationUrl} = await repository.createRegistrationUrl({ + label: 'Uncapped', + createdByUserId: '1', + expiresAt: null, + maxUses: null, + approvalRequired: false, + }); + + const attempts = await Promise.all(Array.from({length: 5}, (_, index) => register(`uncapped${index}`, code))); + + expect(attempts.map((attempt) => attempt.response.status)).toEqual(Array(5).fill(HTTP_STATUS.OK)); + const stored = (await readAdminConfig()).registration.urls.find((url) => url.id === registrationUrl.id); + expect(stored?.use_count).toBe(5); + }); + + it('gives the seat and the pending entry back when the signup failed before the account was created', async () => { + const repository = getInstanceConfigRepository(); + await repository.setRegistrationConfig({mode: 'closed', admin_registration_urls_enabled: true}); + const {code, registrationUrl} = await repository.createRegistrationUrl({ + label: 'Single use', + createdByUserId: '1', + expiresAt: null, + maxUses: 1, + approvalRequired: true, + }); + failAfterThePendingApprovalIsStored(); + + const failed = await register('seatreleased', code); + expect(failed.response.status).toBe(HTTP_STATUS.INTERNAL_SERVER_ERROR); + const withdrawn = await readAdminConfig(); + expect(withdrawn.registration.pending_registrations).toEqual([]); + expect(withdrawn.registration.urls.find((url) => url.id === registrationUrl.id)?.use_count).toBe(0); + + const retried = await register('seatreleasedretry', code); + + expect(retried.response.status).toBe(HTTP_STATUS.OK); + const stored = (await readAdminConfig()).registration.urls.find((url) => url.id === registrationUrl.id); + expect(stored).toMatchObject({use_count: 1, last_used_by_user_id: retried.json.user_id}); + }); + + it('keeps the seat when the account create itself failed, because the row may still have landed', async () => { + const repository = getInstanceConfigRepository(); + await repository.setRegistrationConfig({mode: 'closed', admin_registration_urls_enabled: true}); + const {code, registrationUrl} = await repository.createRegistrationUrl({ + label: 'Single use', + createdByUserId: '1', + expiresAt: null, + maxUses: 1, + approvalRequired: false, + }); + vi.spyOn(UserRepository.prototype, 'create').mockRejectedValueOnce(new Error('the user row write failed')); + + const failed = await register('seatkeptoncreate', code); + expect(failed.response.status).toBe(HTTP_STATUS.INTERNAL_SERVER_ERROR); + const second = await register('seatkeptcreate2', code); + + expect(second.response.status).toBe(HTTP_STATUS.BAD_REQUEST); + expect(second.json.code).toBe(APIErrorCodes.REGISTRATION_URL_INVALID); + const stored = (await readAdminConfig()).registration.urls.find((url) => url.id === registrationUrl.id); + expect(stored?.use_count).toBe(1); + }); + + it('keeps the seat of an account whose row was written before its creation failed', async () => { + const repository = getInstanceConfigRepository(); + await repository.setRegistrationConfig({mode: 'closed', admin_registration_urls_enabled: true}); + const {code, registrationUrl} = await repository.createRegistrationUrl({ + label: 'Single use', + createdByUserId: '1', + expiresAt: null, + maxUses: 1, + approvalRequired: false, + }); + failCreateAfterTheUserRowIsWritten(); + + const failed = await register('seatkept', code); + expect(failed.response.status).toBe(HTTP_STATUS.INTERNAL_SERVER_ERROR); + const second = await register('seatkeptsecond', code); + + expect(second.response.status).toBe(HTTP_STATUS.BAD_REQUEST); + expect(second.json.code).toBe(APIErrorCodes.REGISTRATION_URL_INVALID); + const stored = (await readAdminConfig()).registration.urls.find((url) => url.id === registrationUrl.id); + expect(stored?.use_count).toBe(1); + }); + + it('honours the use count and cap already stored on a registration URL', async () => { + const repository = getInstanceConfigRepository(); + await repository.setRegistrationConfig({mode: 'closed', admin_registration_urls_enabled: true}); + const id = 'b3c4f0b2-8a6e-4c41-9f55-3f0c2a7d1e90'; + await repository.setConfig( + REGISTRATION_URLS_KEY, + JSON.stringify([ + { + id, + label: 'Issued earlier', + code_hash: createHash('sha256').update(id).digest('hex'), + created_by_user_id: '1400000000000000001', + created_at: '2026-09-01T00:00:00.000Z', + expires_at: null, + max_uses: 2, + use_count: 1, + revoked_at: null, + approval_required: false, + last_used_at: '2026-09-02T00:00:00.000Z', + last_used_by_user_id: '1400000000000000002', + }, + ]), + ); + + const before = (await readAdminConfig()).registration.urls.find((url) => url.id === id); + expect(before).toMatchObject({use_count: 1, max_uses: 2, last_used_by_user_id: '1400000000000000002'}); + + const first = await register('storedinvite', id); + expect(first.response.status).toBe(HTTP_STATUS.OK); + const second = await register('storedinviteagain', id); + expect(second.response.status).toBe(HTTP_STATUS.BAD_REQUEST); + expect(second.json.code).toBe(APIErrorCodes.REGISTRATION_URL_INVALID); + + const after = (await readAdminConfig()).registration.urls.find((url) => url.id === id); + expect(after).toMatchObject({use_count: 2, max_uses: 2, last_used_by_user_id: first.json.user_id}); + }); + + it('keeps every pending approval when approval-mode signups race', async () => { + await getInstanceConfigRepository().setRegistrationConfig({mode: 'approval'}); + + const attempts = await Promise.all(Array.from({length: 5}, (_, index) => register(`pending${index}`))); + + expect(attempts.map((attempt) => attempt.json.registration_pending_approval)).toEqual(Array(5).fill(true)); + const pending = (await readAdminConfig()).registration.pending_registrations.map((entry) => entry.user_id); + expect(pending.toSorted()).toEqual(attempts.map((attempt) => attempt.json.user_id).toSorted()); + }); + + it('lists an approval-mode account whose signup failed after the account was created', async () => { + await getInstanceConfigRepository().setRegistrationConfig({mode: 'approval'}); + vi.spyOn(UserRepository.prototype, 'createAuthorizedIp').mockRejectedValueOnce( + new Error('the authorized IP write failed'), + ); + + const failed = await register('pendingstranded'); + + expect(failed.response.status).toBe(HTTP_STATUS.INTERNAL_SERVER_ERROR); + await expectOnePendingAccount(); + }); + + it('lists an approval-mode account whose row was written before its creation failed', async () => { + await getInstanceConfigRepository().setRegistrationConfig({mode: 'approval'}); + failCreateAfterTheUserRowIsWritten(); + + const failed = await register('pendingrowwritten'); + + expect(failed.response.status).toBe(HTTP_STATUS.INTERNAL_SERVER_ERROR); + await expectOnePendingAccount(); + }); + + it('keeps the pending approval of an approval-mode signup whose account create failed', async () => { + await getInstanceConfigRepository().setRegistrationConfig({mode: 'approval'}); + vi.spyOn(UserRepository.prototype, 'create').mockRejectedValueOnce(new Error('the user row write failed')); + + const failed = await register('pendingkept'); + + expect(failed.response.status).toBe(HTTP_STATUS.INTERNAL_SERVER_ERROR); + expect((await readAdminConfig()).registration.pending_registrations).toHaveLength(1); + }); + + it('lists no pending approval for an approval-mode signup that failed before the account was created', async () => { + await getInstanceConfigRepository().setRegistrationConfig({mode: 'approval'}); + failAfterThePendingApprovalIsStored(); + + const failed = await register('pendingnever'); + + expect(failed.response.status).toBe(HTTP_STATUS.INTERNAL_SERVER_ERROR); + expect((await readAdminConfig()).registration.pending_registrations).toEqual([]); + }); + + it('lists an SSO account provisioned in approval mode whose provisioning failed after the account was created', async () => { + await enableSso(harness, admin.token, {enforced: false}); + await getInstanceConfigRepository().setRegistrationConfig({mode: 'approval'}); + vi.spyOn(UserRepository.prototype, 'upsertSettings').mockRejectedValueOnce(new Error('the settings write failed')); + + const failed = await completeSso('ssopendingstranded'); + + expect(failed.response.status).toBe(HTTP_STATUS.INTERNAL_SERVER_ERROR); + await expectOnePendingAccount(); + }); + + it('lists an SSO account provisioned in approval mode whose row was written before its creation failed', async () => { + await enableSso(harness, admin.token, {enforced: false}); + await getInstanceConfigRepository().setRegistrationConfig({mode: 'approval'}); + failCreateAfterTheUserRowIsWritten(); + + const failed = await completeSso('ssopendingrowwritten'); + + expect(failed.response.status).toBe(HTTP_STATUS.INTERNAL_SERVER_ERROR); + await expectOnePendingAccount(); + }); + + it('keeps the pending approval of an SSO signup in approval mode whose account create failed', async () => { + await enableSso(harness, admin.token, {enforced: false}); + await getInstanceConfigRepository().setRegistrationConfig({mode: 'approval'}); + vi.spyOn(UserRepository.prototype, 'create').mockRejectedValueOnce(new Error('the user row write failed')); + + const failed = await completeSso('ssopendingkept'); + + expect(failed.response.status).toBe(HTTP_STATUS.INTERNAL_SERVER_ERROR); + expect((await readAdminConfig()).registration.pending_registrations).toHaveLength(1); + }); + + it('lists no pending approval for an SSO signup in approval mode that failed before the account was created', async () => { + await enableSso(harness, admin.token, {enforced: false}); + await getInstanceConfigRepository().setRegistrationConfig({mode: 'approval'}); + failAfterThePendingApprovalIsStored(); + + const failed = await completeSso('ssopendingnever'); + + expect(failed.response.status).toBe(HTTP_STATUS.INTERNAL_SERVER_ERROR); + expect((await readAdminConfig()).registration.pending_registrations).toEqual([]); + }); + + it('keeps a stored pending approval listed until an admin decides it', async () => { + const account = await createTestAccount(harness); + await getInstanceConfigRepository().setConfig( + REGISTRATION_PENDING_APPROVALS_KEY, + JSON.stringify([ + { + user_id: account.userId, + username: 'stored_pending', + discriminator: 1, + global_name: null, + email: account.email, + requested_at: '2026-09-01T00:00:00.000Z', + registration_url_id: null, + client_ip: '127.0.0.1', + }, + ]), + ); + await getInstanceConfigRepository().setRegistrationConfig({mode: 'approval'}); + const fresh = await register('pendingafter'); + + const listed = (await readAdminConfig()).registration.pending_registrations.map((entry) => entry.user_id); + expect(listed.toSorted()).toEqual([account.userId, fresh.json.user_id].toSorted()); + + const decided = await createBuilder(harness, admin.token) + .patch(`/admin/instance/pending-registrations/${account.userId}`) + .body({status: 'approved'}) + .expect(HTTP_STATUS.OK) + .execute(); + + expect(decided.registration.pending_registrations.map((entry) => entry.user_id)).toEqual([fresh.json.user_id]); + expect( + JSON.parse( + (await getInstanceConfigRepository().getConfig(REGISTRATION_PENDING_APPROVALS_KEY)) ?? 'null', + ) as Array<{user_id: string}>, + ).toEqual([expect.objectContaining({user_id: fresh.json.user_id})]); + }); +}); diff --git a/fluxer_api/src/api/instance/InstanceConfigRepository.test.ts b/fluxer_api/src/api/instance/InstanceConfigRepository.test.ts index 2260a5ec1..3f870c736 100644 --- a/fluxer_api/src/api/instance/InstanceConfigRepository.test.ts +++ b/fluxer_api/src/api/instance/InstanceConfigRepository.test.ts @@ -1,12 +1,21 @@ // SPDX-License-Identifier: AGPL-3.0-or-later +import {spawnSync} from 'node:child_process'; +import {createHash} from 'node:crypto'; +import {createServer} from 'node:net'; +import type {CassandraQueryExecutorForTesting} from '@app/api/database/CassandraQueryExecution'; import {setCassandraQueryExecutorForTesting} from '@app/api/database/CassandraQueryExecution'; import type {PreparedQuery} from '@app/api/database/CassandraTypes'; +import {ensurePostgresKvSchema, PostgresKvQueryExecutor} from '@app/api/database/PostgresKvQueryExecutor'; import { INSTANCE_CONFIG_REFRESH_CHANNEL, + INSTANCE_CONFIG_WRITE_ATTEMPTS, InstanceConfigRepository, + InstanceConfigWriteConflictError, type InstanceRegistrationConfig, } from '@app/api/instance/InstanceConfigRepository'; +import {InstanceConfigWriteRaceExecutor} from '@app/api/instance/tests/InstanceConfigWriteRaceExecutor'; +import {startDockerContainer} from '@app/api/test/DockerTestContainer'; import {InMemoryCassandraQueryExecutor} from '@app/api/test/InMemoryCassandraQueryExecutor'; import {MockKVProvider} from '@app/api/test/mocks/MockKVProvider'; import { @@ -21,7 +30,13 @@ import { DEFAULT_EXPERIMENT_DELIVERY_CONFIG, type ExperimentDeliveryConfig, } from '@fluxer/schema/src/domains/experiment/ExperimentSchemas'; -import {afterEach, describe, expect, it, vi} from 'vitest'; +import { + getDefaultPostgresClient, + type IPostgresClient, + initPostgres, + shutdownPostgres, +} from '@pkgs/postgres/src/Client'; +import {afterAll, afterEach, beforeAll, beforeEach, describe, expect, it, vi} from 'vitest'; const VOICE_NOISE_SUPPRESSION_CONFIG_KEY = 'voice_noise_suppression_config'; const SCREEN_SHARE_DELIVERY_CONFIG_KEY = 'screen_share_delivery_config'; @@ -29,6 +44,12 @@ const EXPERIMENT_DELIVERY_CONFIG_KEY = 'experiment_delivery_config'; const APP_PUBLIC_CONFIG_KEY = 'app_public_config'; const INSTANCE_POLICY_CONFIG_KEY = 'instance_policy_config'; const INSTANCE_INTEGRATIONS_CONFIG_KEY = 'instance_integrations_config'; +const REGISTRATION_CONFIG_KEY = 'registration_config'; +const REGISTRATION_URLS_KEY = 'registration_urls'; +const REGISTRATION_PENDING_APPROVALS_KEY = 'registration_pending_approvals'; +const POSTGRES_KV_TABLE = 'kv_instance_config_races'; +const POSTGRES_CONTAINER = `fluxer-instance-config-races-${process.pid.toString(36)}-${Date.now().toString(36)}`; +const dockerAvailable = spawnSync('docker', ['version'], {stdio: 'ignore'}).status === 0; class CountingInMemoryCassandraQueryExecutor extends InMemoryCassandraQueryExecutor { instanceConfigSelects = 0; @@ -512,3 +533,466 @@ describe('InstanceConfigRepository', () => { }); }); }); + +async function sleep(ms: number): Promise { + await new Promise((resolve) => setTimeout(resolve, ms)); +} + +async function freePort(): Promise { + return new Promise((resolve, reject) => { + const server = createServer(); + server.on('error', reject); + server.listen(0, '127.0.0.1', () => { + const address = server.address(); + if (typeof address === 'string' || address === null) { + reject(new Error('no port')); + return; + } + server.close(() => resolve(address.port)); + }); + }); +} + +function describeConcurrentInstanceConfigWrites(prepareBase: () => Promise): void { + const pods: Array = []; + let executor: InstanceConfigWriteRaceExecutor; + + beforeEach(async () => { + executor = new InstanceConfigWriteRaceExecutor(await prepareBase()); + setCassandraQueryExecutorForTesting(executor); + }); + + afterEach(async () => { + await Promise.all(pods.map((pod) => pod.shutdown())); + pods.length = 0; + }); + + function createPod(): InstanceConfigRepository { + const pod = new InstanceConfigRepository(new MockKVProvider()); + pods.push(pod); + return pod; + } + + async function readStoredRegistrationConfig(): Promise { + const raw = await executor.readDirectly(REGISTRATION_CONFIG_KEY); + return raw === null ? null : JSON.parse(raw); + } + + it('applies two concurrent patches on top of each other instead of dropping one', async () => { + const first = createPod(); + const second = createPod(); + await first.setRegistrationConfig({mode: 'open', admin_registration_urls_enabled: true}); + await second.getRegistrationConfig(); + executor.watch(REGISTRATION_CONFIG_KEY); + executor.pauseWritesUntil(2); + + await Promise.all([ + first.setRegistrationConfig({mode: 'closed'}), + second.setRegistrationConfig({admin_registration_urls_enabled: false}), + ]); + + expect(executor.events.filter((event) => event === 'write rejected')).toHaveLength(1); + expect(executor.events.filter((event) => event === 'write')).toHaveLength(2); + expect(await readStoredRegistrationConfig()).toEqual({mode: 'closed', admin_registration_urls_enabled: false}); + }); + + it('lets one of two concurrent first writes create the config and applies the other on top', async () => { + const first = createPod(); + const second = createPod(); + await first.getRegistrationConfig(); + await second.getRegistrationConfig(); + executor.watch(REGISTRATION_CONFIG_KEY); + executor.pauseWritesUntil(2); + + await Promise.all([ + first.setRegistrationConfig({mode: 'closed'}), + second.setRegistrationConfig({admin_registration_urls_enabled: false}), + ]); + + expect(executor.events.filter((event) => event === 'write rejected')).toHaveLength(1); + expect(executor.events.filter((event) => event === 'write')).toHaveLength(2); + expect(await readStoredRegistrationConfig()).toEqual({mode: 'closed', admin_registration_urls_enabled: false}); + }); + + it('re-reads the database, not its stale cache, when a concurrent write lands between its read and its write', async () => { + const stale = createPod(); + const other = createPod(); + await stale.setRegistrationConfig({mode: 'open', admin_registration_urls_enabled: true}); + await stale.getRegistrationConfig(); + await other.setRegistrationConfig({mode: 'approval'}); + expect(await stale.getRegistrationConfig()).toEqual({mode: 'open', admin_registration_urls_enabled: true}); + executor.watch(REGISTRATION_CONFIG_KEY); + let competed = false; + executor.competeBeforeEachWrite(async () => { + if (competed) return; + competed = true; + await executor.writeDirectly( + REGISTRATION_CONFIG_KEY, + JSON.stringify({mode: 'approval', admin_registration_urls_enabled: false}), + ); + }); + + await stale.setRegistrationConfig({mode: 'closed'}); + + expect(executor.events).toEqual(['read', 'write rejected', 'read', 'write']); + expect(await readStoredRegistrationConfig()).toEqual({mode: 'closed', admin_registration_urls_enabled: false}); + }); + + it('fails loudly and writes nothing once every attempt has lost the race', async () => { + const pod = createPod(); + await pod.setRegistrationConfig({mode: 'open', admin_registration_urls_enabled: true}); + executor.watch(REGISTRATION_CONFIG_KEY); + let competingWrites = 0; + executor.competeBeforeEachWrite(async () => { + competingWrites++; + await executor.writeDirectly( + REGISTRATION_CONFIG_KEY, + JSON.stringify({mode: 'approval', admin_registration_urls_enabled: competingWrites % 2 === 0}), + ); + }); + + const write = pod.setRegistrationConfig({mode: 'closed'}); + + await expect(write).rejects.toBeInstanceOf(InstanceConfigWriteConflictError); + await expect(write).rejects.toMatchObject({ + status: 409, + code: 'CONFLICT', + message: expect.stringContaining(REGISTRATION_CONFIG_KEY), + }); + expect(executor.events.filter((event) => event === 'write rejected')).toHaveLength(INSTANCE_CONFIG_WRITE_ATTEMPTS); + expect(executor.events).not.toContain('write'); + expect(await readStoredRegistrationConfig()).toEqual({ + mode: 'approval', + admin_registration_urls_enabled: INSTANCE_CONFIG_WRITE_ATTEMPTS % 2 === 0, + }); + }); + + it('keeps a pending registration another pod added while this pod held a stale list', async () => { + const first = createPod(); + const second = createPod(); + await first.getPendingRegistrations(); + await second.getPendingRegistrations(); + + await first.addPendingRegistration(pendingRegistration('1400000000000000011')); + await second.addPendingRegistration(pendingRegistration('1400000000000000012')); + + const listed = await createPod().getPendingRegistrations(); + expect(listed.map((entry) => entry.user_id)).toEqual(['1400000000000000011', '1400000000000000012']); + }); + + it('keeps a pending registration another pod stored and removes it once decided', async () => { + const pod = createPod(); + await executor.writeDirectly( + REGISTRATION_PENDING_APPROVALS_KEY, + JSON.stringify([pendingRegistration('1400000000000000021')]), + ); + await pod.addPendingRegistration(pendingRegistration('1400000000000000022')); + + expect((await createPod().getPendingRegistrations()).map((entry) => entry.user_id)).toEqual([ + '1400000000000000021', + '1400000000000000022', + ]); + + await pod.removePendingRegistration('1400000000000000021'); + await pod.removePendingRegistration('1400000000000000022'); + + expect(await createPod().getPendingRegistrations()).toEqual([]); + expect(await executor.readDirectly(REGISTRATION_PENDING_APPROVALS_KEY)).toBe('[]'); + }); + + it('keeps a registration URL another pod created while this pod held a stale list', async () => { + const first = createPod(); + const second = createPod(); + await first.getRegistrationUrls(); + await second.getRegistrationUrls(); + + const created = [ + await first.createRegistrationUrl(registrationUrlParams(null)), + await second.createRegistrationUrl(registrationUrlParams(null)), + ]; + + const listed = await createPod().getRegistrationUrlsForAdmin(); + expect(listed.map((url) => url.id).toSorted()).toEqual(created.map((entry) => entry.registrationUrl.id).toSorted()); + }); + + it('refuses a registration URL another pod revoked while this pod held a stale list', async () => { + const admin = createPod(); + const signup = createPod(); + const {code, registrationUrl} = await admin.createRegistrationUrl(registrationUrlParams(null)); + expect(await signup.resolveRegistrationUrlCode(code)).not.toBeNull(); + + await admin.revokeRegistrationUrl(registrationUrl.id); + + await expect(signup.resolveRegistrationUrlCode(code)).resolves.toBeNull(); + await expect(signup.claimRegistrationUrlUse(registrationUrl.id, '1400000000000000501')).resolves.toBeNull(); + const [listed] = await createPod().getRegistrationUrlsForAdmin(); + expect(listed?.use_count).toBe(0); + }); + + it('admits a registration URL another pod created while this pod held a stale list', async () => { + const admin = createPod(); + const signup = createPod(); + await signup.getRegistrationUrls(); + + const {code, registrationUrl} = await admin.createRegistrationUrl(registrationUrlParams(1)); + + await expect(signup.resolveRegistrationUrlCode(code)).resolves.toMatchObject({ + id: registrationUrl.id, + approval_required: false, + }); + await expect(signup.claimRegistrationUrlUse(registrationUrl.id, '1400000000000000601')).resolves.not.toBeNull(); + const [listed] = await createPod().getRegistrationUrlsForAdmin(); + expect(listed?.use_count).toBe(1); + }); + + it('never seats more signups than max_uses when pods claim the same registration URL at once', async () => { + const pods = [createPod(), createPod(), createPod()]; + const {code} = await pods[0]!.createRegistrationUrl(registrationUrlParams(2)); + const registrationUrl = await pods[0]!.resolveRegistrationUrlCode(code); + if (registrationUrl === null) throw new Error('registration URL did not resolve'); + + const claims = await Promise.all( + Array.from({length: 6}, (_, index) => + pods[index % pods.length]!.claimRegistrationUrlUse(registrationUrl.id, `14000000000000001${index}0`), + ), + ); + + expect(claims.filter((claim) => claim !== null)).toHaveLength(2); + const [listed] = await createPod().getRegistrationUrlsForAdmin(); + expect(listed?.use_count).toBe(2); + await expect(createPod().resolveRegistrationUrlCode(code)).resolves.toBeNull(); + }); + + it('refuses a claim whose retry finds the registration URL exhausted, rather than reporting the lost attempt', async () => { + const pod = createPod(); + const {code, registrationUrl} = await pod.createRegistrationUrl(registrationUrlParams(1)); + expect(await pod.resolveRegistrationUrlCode(code)).not.toBeNull(); + executor.watch(REGISTRATION_URLS_KEY); + let competed = false; + executor.competeBeforeEachWrite(async () => { + if (competed) return; + competed = true; + const stored = JSON.parse((await executor.readDirectly(REGISTRATION_URLS_KEY)) ?? 'null') as Array< + Record + >; + await executor.writeDirectly( + REGISTRATION_URLS_KEY, + JSON.stringify( + stored.map((entry) => ({ + ...entry, + use_count: 1, + last_used_at: '2026-09-20T00:00:00.000Z', + last_used_by_user_id: '1400000000000000901', + })), + ), + ); + }); + + await expect(pod.claimRegistrationUrlUse(registrationUrl.id, '1400000000000000902')).resolves.toBeNull(); + + expect(executor.events).toEqual(['read', 'write rejected', 'read']); + const [listed] = await createPod().getRegistrationUrlsForAdmin(); + expect(listed).toMatchObject({use_count: 1, last_used_by_user_id: '1400000000000000901'}); + }); + + it('counts concurrent uses of an uncapped registration URL without ever refusing one', async () => { + const pods = [createPod(), createPod()]; + const {code} = await pods[0]!.createRegistrationUrl(registrationUrlParams(null)); + const registrationUrl = await pods[0]!.resolveRegistrationUrlCode(code); + if (registrationUrl === null) throw new Error('registration URL did not resolve'); + + const claims = await Promise.all( + Array.from({length: 5}, (_, index) => + pods[index % pods.length]!.claimRegistrationUrlUse(registrationUrl.id, `14000000000000002${index}0`), + ), + ); + + expect(claims.every((claim) => claim !== null)).toBe(true); + const [listed] = await createPod().getRegistrationUrlsForAdmin(); + expect(listed?.use_count).toBe(5); + }); + + it('frees a released seat for the next signup', async () => { + const pod = createPod(); + const {code} = await pod.createRegistrationUrl(registrationUrlParams(1)); + const registrationUrl = await pod.resolveRegistrationUrlCode(code); + if (registrationUrl === null) throw new Error('registration URL did not resolve'); + + const failedSignup = await pod.claimRegistrationUrlUse(registrationUrl.id, '1400000000000000301'); + if (failedSignup === null) throw new Error('the first claim was refused'); + await expect(pod.claimRegistrationUrlUse(registrationUrl.id, '1400000000000000302')).resolves.toBeNull(); + await pod.releaseRegistrationUrlUse(failedSignup); + + await expect(pod.claimRegistrationUrlUse(registrationUrl.id, '1400000000000000303')).resolves.not.toBeNull(); + const [listed] = await createPod().getRegistrationUrlsForAdmin(); + expect(listed).toMatchObject({use_count: 1, last_used_by_user_id: '1400000000000000303'}); + }); + + it('enforces max_uses against the use count already stored in the blob', async () => { + const pod = createPod(); + const id = 'b3c4f0b2-8a6e-4c41-9f55-3f0c2a7d1e91'; + await executor.writeDirectly( + REGISTRATION_URLS_KEY, + JSON.stringify([ + { + id, + label: 'Issued earlier', + code_hash: createHash('sha256').update(id).digest('hex'), + created_by_user_id: '1400000000000000001', + created_at: '2026-09-01T00:00:00.000Z', + expires_at: null, + max_uses: 3, + use_count: 2, + revoked_at: null, + approval_required: true, + last_used_at: '2026-09-02T00:00:00.000Z', + last_used_by_user_id: '1400000000000000002', + }, + ]), + ); + + expect(await pod.getRegistrationUrlsForAdmin()).toEqual([ + expect.objectContaining({ + id, + use_count: 2, + max_uses: 3, + approval_required: true, + last_used_at: '2026-09-02T00:00:00.000Z', + last_used_by_user_id: '1400000000000000002', + }), + ]); + const registrationUrl = await pod.resolveRegistrationUrlCode(id); + if (registrationUrl === null) throw new Error('stored registration URL did not resolve'); + expect(registrationUrl).toMatchObject({id, approval_required: true}); + + await expect(pod.claimRegistrationUrlUse(registrationUrl.id, '1400000000000000401')).resolves.not.toBeNull(); + await expect(pod.claimRegistrationUrlUse(registrationUrl.id, '1400000000000000402')).resolves.toBeNull(); + + const [listed] = await createPod().getRegistrationUrlsForAdmin(); + expect(listed).toMatchObject({use_count: 3, max_uses: 3, last_used_by_user_id: '1400000000000000401'}); + await expect(pod.resolveRegistrationUrlCode(id)).resolves.toBeNull(); + expect(JSON.parse((await executor.readDirectly(REGISTRATION_URLS_KEY)) ?? 'null')[0]).toMatchObject({ + use_count: 3, + max_uses: 3, + }); + }); + + it('keeps an SSO field another pod changed while this pod held a stale snapshot', async () => { + const first = createPod(); + const second = createPod(); + await first.getSsoConfig(); + await second.getSsoConfig(); + + await first.setSsoConfig({displayName: 'Set by the first pod'}); + await second.setSsoConfig({clientId: 'set-by-the-second-pod'}); + + expect(await createPod().getSsoConfig()).toMatchObject({ + displayName: 'Set by the first pod', + clientId: 'set-by-the-second-pod', + }); + }); + + it('leaves an SSO row alone when another pod wrote it between this pod reading and writing it', async () => { + const pod = createPod(); + await pod.getSsoConfig(); + executor.watch('sso_enforced'); + let competed = false; + executor.competeBeforeEachWrite(async () => { + if (competed) return; + competed = true; + await executor.writeDirectly('sso_enforced', 'true'); + }); + + await pod.setSsoConfig({displayName: 'Only the display name'}); + + expect(await executor.readDirectly('sso_enforced')).toBe('true'); + expect(await executor.readDirectly('sso_display_name')).toBe('Only the display name'); + }); +} + +function pendingRegistration(userId: string) { + return { + user_id: userId, + username: `pending_${userId.slice(-3)}`, + discriminator: 1, + global_name: null, + email: `${userId}@example.com`, + requested_at: `2026-09-01T00:00:${userId.slice(-2)}.000Z`, + registration_url_id: null, + client_ip: '127.0.0.1', + }; +} + +function registrationUrlParams(maxUses: number | null) { + return { + label: maxUses === null ? 'Uncapped' : `Capped at ${maxUses}`, + createdByUserId: '1400000000000000001', + expiresAt: null, + maxUses, + approvalRequired: false, + }; +} + +describe('InstanceConfigRepository concurrent writes', () => { + describe('in memory', () => { + describeConcurrentInstanceConfigWrites(async () => new InMemoryCassandraQueryExecutor()); + }); + + describe.skipIf(!dockerAvailable)('on postgres', () => { + let client: IPostgresClient; + + beforeAll(async () => { + const port = await freePort(); + startDockerContainer([ + 'run', + '-d', + '--name', + POSTGRES_CONTAINER, + '-e', + 'POSTGRES_USER=fluxer', + '-e', + 'POSTGRES_PASSWORD=fluxer', + '-e', + 'POSTGRES_DB=fluxer', + '-p', + `127.0.0.1:${port}:5432`, + 'postgres:16-alpine', + '-c', + 'fsync=off', + ]); + let ready = false; + for (let attempt = 0; attempt < 180 && !ready; attempt += 1) { + await sleep(500); + const probe = spawnSync('docker', ['exec', POSTGRES_CONTAINER, 'pg_isready', '-U', 'fluxer', '-d', 'fluxer'], { + stdio: 'ignore', + }); + if (probe.status !== 0) continue; + try { + await initPostgres({ + url: `postgres://fluxer:fluxer@127.0.0.1:${port}/fluxer`, + maxConnections: 4, + kvTable: POSTGRES_KV_TABLE, + }); + await getDefaultPostgresClient().query('SELECT 1'); + ready = true; + } catch { + await shutdownPostgres().catch(() => {}); + } + } + if (!ready) throw new Error('postgres never came up'); + client = getDefaultPostgresClient(); + await ensurePostgresKvSchema(client); + }, 900_000); + + afterAll(async () => { + setCassandraQueryExecutorForTesting(new InMemoryCassandraQueryExecutor()); + await shutdownPostgres().catch(() => {}); + spawnSync('docker', ['rm', '-f', POSTGRES_CONTAINER], {stdio: 'ignore'}); + }); + + describeConcurrentInstanceConfigWrites(async () => { + await client.query(`DELETE FROM ${POSTGRES_KV_TABLE}`); + return new PostgresKvQueryExecutor(client); + }); + }); +}); diff --git a/fluxer_api/src/api/instance/InstanceConfigRepository.ts b/fluxer_api/src/api/instance/InstanceConfigRepository.ts index 448cea819..e9d6181de 100644 --- a/fluxer_api/src/api/instance/InstanceConfigRepository.ts +++ b/fluxer_api/src/api/instance/InstanceConfigRepository.ts @@ -3,7 +3,8 @@ import crypto from 'node:crypto'; import {Config} from '@app/api/Config'; import type {APIConfig, BlueskyOAuthConfig, BlueskyOAuthKeyConfig} from '@app/api/config/APIConfig'; -import {fetchMany, fetchOne, upsertOne} from '@app/api/database/CassandraQueryExecution'; +import {executeConditional, fetchMany, fetchOne, upsertOne} from '@app/api/database/CassandraQueryExecution'; +import {Db, type PreparedQuery} from '@app/api/database/CassandraTypes'; import type {InstanceConfigurationRow} from '@app/api/database/types/InstanceConfigTypes'; import { getDefaultDateOfBirthCollection, @@ -17,6 +18,9 @@ import {resolveDeferredPhoneGateEnabled, setCachedDeferredPhoneGateEnabled} from import {InstanceConfiguration} from '@app/api/Tables'; import {DEFAULT_DECAY_CONSTANTS, DEFAULT_RENEWAL_CONSTANTS} from '@app/api/utils/AttachmentDecay'; import {isJsonRecord} from '@app/api/utils/JsonBoundaryUtils'; +import {APIErrorCodes} from '@fluxer/constants/src/ApiErrorCodes'; +import {ConflictError} from '@fluxer/errors/src/domains/core/ConflictError'; +import {ServiceUnavailableError} from '@fluxer/errors/src/domains/core/ServiceUnavailableError'; import type {LimitConfigSnapshot} from '@fluxer/limits/src/LimitTypes'; import { InstanceConfigResponse, @@ -28,6 +32,10 @@ import { type GatewayRolloutConfig, GatewayRolloutConfigSchema, } from '@fluxer/schema/src/domains/admin/GatewayRolloutSchemas'; +import { + type PushServiceDeliveryConfig, + PushServiceDeliveryConfigSchema, +} from '@fluxer/schema/src/domains/admin/PushServiceDeliverySchemas'; import { type ScreenShareDeliveryConfig, ScreenShareDeliveryConfigSchema, @@ -59,6 +67,7 @@ import {z} from 'zod'; const GATEWAY_ROLLOUT_CONFIG_KEY = 'gateway_rollout_config'; const VOICE_NOISE_SUPPRESSION_CONFIG_KEY = 'voice_noise_suppression_config'; const SCREEN_SHARE_DELIVERY_CONFIG_KEY = 'screen_share_delivery_config'; +const PUSH_SERVICE_DELIVERY_CONFIG_KEY = 'push_service_delivery_config'; const EXPERIMENT_DELIVERY_CONFIG_KEY = 'experiment_delivery_config'; const REGISTRATION_CONFIG_KEY = 'registration_config'; const REGISTRATION_URLS_KEY = 'registration_urls'; @@ -72,6 +81,22 @@ const INSTANCE_MEDIA_CONFIG_KEY = 'instance_media_config'; export const INSTANCE_CONFIG_REFRESH_CHANNEL = 'instance-config-refresh'; export const REGISTRATION_PENDING_APPROVAL_TRAIT = 'registration_pending_approval'; export const REGISTRATION_REJECTED_TRAIT = 'registration_rejected'; +export const INSTANCE_CONFIG_WRITE_ATTEMPTS = 5; + +export class InstanceConfigWriteConflictError extends ConflictError { + constructor(key: string) { + super({ + code: APIErrorCodes.CONFLICT, + message: `Instance config "${key}" changed concurrently on all ${INSTANCE_CONFIG_WRITE_ATTEMPTS} write attempts. Nothing was written. Retry the change.`, + }); + this.name = 'InstanceConfigWriteConflictError'; + } +} + +interface StoredValueUpdate { + value: string | null; + result: T; +} export type InstanceRegistrationConfig = InstanceRegistration; @@ -277,6 +302,11 @@ export interface InstanceRegistrationUrl extends RegistrationUrlResponse { type InstanceRegistrationUrlPublic = RegistrationUrlResponse; type InstancePendingRegistration = PendingRegistrationResponse; +export interface RegistrationUrlClaim { + registration_url_id: string; + user_id: string; +} + const DEFAULT_REGISTRATION_CONFIG: InstanceRegistrationConfig = { mode: 'open', admin_registration_urls_enabled: true, @@ -345,6 +375,7 @@ type StoredConfigSection = | 'gateway rollout' | 'voice noise suppression' | 'screen share delivery' + | 'push service delivery' | 'experiment delivery' | 'instance policy' | 'integrations' @@ -487,6 +518,10 @@ function parseStoredScreenShareDeliveryConfig(raw: string | null): ScreenShareDe return parseStoredConfigOrDefault(ScreenShareDeliveryConfigSchema, raw, 'screen share delivery'); } +function parseStoredPushServiceDeliveryConfig(raw: string | null): PushServiceDeliveryConfig { + return parseStoredConfigOrDefault(PushServiceDeliveryConfigSchema, raw, 'push service delivery'); +} + function parseStoredExperimentDeliveryConfig(raw: string | null): ExperimentDeliveryConfig { return parseStoredConfigOrDefault(ExperimentDeliveryConfigSchema, raw, 'experiment delivery'); } @@ -910,6 +945,75 @@ function parseStoredSsoAllowedEmailDomains(raw: string | undefined, log = false) return Array.from(domains).slice(0, MAX_SSO_ALLOWED_DOMAINS); } +function readStoredSsoConfig( + configs: ReadonlyMap, + options?: {includeSecret?: boolean}, +): InstanceSsoConfig { + const flags = readStoredSsoFlags(configs); + const read = (key: string): string | null => { + const v = configs.get(key); + if (!v) return null; + const trimmed = v.trim(); + return trimmed.length === 0 ? null : trimmed; + }; + const allowedDomains = parseStoredSsoAllowedEmailDomains(configs.get('sso_allowed_domains')); + const clientSecret = read('sso_client_secret'); + return { + ...flags, + displayName: read('sso_display_name'), + issuer: read('sso_issuer'), + authorizationUrl: read('sso_authorization_url'), + tokenUrl: read('sso_token_url'), + userInfoUrl: read('sso_userinfo_url'), + jwksUrl: read('sso_jwks_url'), + clientId: read('sso_client_id'), + clientSecret: options?.includeSecret ? clientSecret : undefined, + clientSecretSet: Boolean(clientSecret), + scope: read('sso_scope'), + allowedEmailDomains: allowedDomains, + redirectUri: null, + }; +} + +interface SsoRowWrite { + key: string; + value: string | undefined; + unset: string; +} + +function ssoRow(key: string, value: T | undefined, current: T, format: (value: T) => string): SsoRowWrite { + return {key, value: value === undefined ? undefined : format(value), unset: format(current)}; +} + +function nextSsoRowValue(row: SsoRowWrite, raw: string | null): string | null { + const value = row.value ?? raw ?? row.unset; + return value === raw ? null : value; +} + +function formatSsoBoolean(value: boolean): string { + return value ? 'true' : 'false'; +} + +function formatSsoString(value: string | null): string { + return value ?? ''; +} + +function formatSsoDomains(value: Array): string { + return JSON.stringify(value); +} + +function normalizeSsoAllowedEmailDomainsForWrite(domains: Array, enabled: boolean): Array { + try { + return normalizeSsoAllowedEmailDomains(domains); + } catch (error) { + if (enabled) { + throw error; + } + Logger.warn({error}, 'Clearing invalid SSO allowed domain config while SSO is disabled'); + return []; + } +} + export class InstanceConfigRepository { private readonly kvClient: IKVProvider | null; private configCache: InstanceConfigCache; @@ -993,6 +1097,57 @@ export class InstanceConfigRepository { ); } + private async updateStoredConfig(key: string, next: (raw: string | null) => T): Promise { + const cache = this.configCache; + const {result} = await this.compareAndSetStoredValue(cache, key, (raw) => { + const config = next(raw); + return {value: JSON.stringify(config), result: config}; + }); + await this.publishRefresh(cache.sourceId); + return result; + } + + private async compareAndSetStoredValue( + cache: InstanceConfigCache, + key: string, + next: (raw: string | null) => StoredValueUpdate, + ): Promise<{result: T; written: boolean}> { + await cache.getSnapshot(); + for (let attempt = 0; attempt < INSTANCE_CONFIG_WRITE_ATTEMPTS; attempt++) { + cache.assertActive(); + const current = await this.fetchConfigForWrite(key); + cache.assertActive(); + const {value, result} = next(current); + if (value === null) return {result, written: false}; + if (await executeConditional(this.compareAndSetConfig(key, current, value))) { + cache.update(key, value); + return {result, written: true}; + } + } + Logger.error( + {key, attempts: INSTANCE_CONFIG_WRITE_ATTEMPTS}, + 'Instance config write lost to a concurrent write on every attempt', + ); + throw new InstanceConfigWriteConflictError(key); + } + + private compareAndSetConfig(key: string, current: string | null, value: string): PreparedQuery { + const updatedAt = new Date(); + if (current === null) { + return InstanceConfiguration.insertIfNotExists({key, value, updated_at: updatedAt}); + } + return InstanceConfiguration.conditionalPatchByPk( + {key}, + {value: Db.set(value), updated_at: Db.set(updatedAt)}, + {value: current}, + ); + } + + private async fetchConfigForWrite(key: string): Promise { + const [row] = await fetchMany(FETCH_CONFIG_QUERY, {key}, {consistency: 'serial'}); + return row?.value ?? null; + } + private async fetchConfigFromDatabase(key: string): Promise { const row = await fetchOne(FETCH_CONFIG_QUERY, {key}); return row?.value ?? null; @@ -1015,6 +1170,7 @@ export class InstanceConfigRepository { ); parseStoredVoiceNoiseSuppressionConfig(snapshot.get(VOICE_NOISE_SUPPRESSION_CONFIG_KEY) ?? null); parseStoredScreenShareDeliveryConfig(snapshot.get(SCREEN_SHARE_DELIVERY_CONFIG_KEY) ?? null); + parseStoredPushServiceDeliveryConfig(snapshot.get(PUSH_SERVICE_DELIVERY_CONFIG_KEY) ?? null); parseStoredExperimentDeliveryConfig(snapshot.get(EXPERIMENT_DELIVERY_CONFIG_KEY) ?? null); const policy = parseStoredInstancePolicyConfig(snapshot.get(INSTANCE_POLICY_CONFIG_KEY) ?? null); checkStoredConfig('registration', () => @@ -1082,8 +1238,12 @@ export class InstanceConfigRepository { return parseStoredGatewayRolloutConfig(raw); } - async setGatewayRolloutConfig(config: GatewayRolloutConfig): Promise { - await this.setConfig(GATEWAY_ROLLOUT_CONFIG_KEY, JSON.stringify(decodeGatewayRolloutConfig(config))); + updateGatewayRolloutConfig( + update: (current: GatewayRolloutConfig) => GatewayRolloutConfig, + ): Promise { + return this.updateStoredConfig(GATEWAY_ROLLOUT_CONFIG_KEY, (raw) => + decodeGatewayRolloutConfig(update(parseStoredGatewayRolloutConfig(raw))), + ); } async getVoiceNoiseSuppressionConfig(): Promise { @@ -1092,8 +1252,19 @@ export class InstanceConfigRepository { } async setVoiceNoiseSuppressionConfig(config: VoiceNoiseSuppressionConfig): Promise { - const validated = validateStoredConfig(VoiceNoiseSuppressionConfigSchema, config, 'voice noise suppression'); - await this.setConfig(VOICE_NOISE_SUPPRESSION_CONFIG_KEY, JSON.stringify(validated)); + await this.updateVoiceNoiseSuppressionConfig(() => config); + } + + updateVoiceNoiseSuppressionConfig( + update: (current: VoiceNoiseSuppressionConfig) => VoiceNoiseSuppressionConfig, + ): Promise { + return this.updateStoredConfig(VOICE_NOISE_SUPPRESSION_CONFIG_KEY, (raw) => + validateStoredConfig( + VoiceNoiseSuppressionConfigSchema, + update(parseStoredVoiceNoiseSuppressionConfig(raw)), + 'voice noise suppression', + ), + ); } async getScreenShareDeliveryConfig(): Promise { @@ -1102,8 +1273,36 @@ export class InstanceConfigRepository { } async setScreenShareDeliveryConfig(config: ScreenShareDeliveryConfig): Promise { - const validated = validateStoredConfig(ScreenShareDeliveryConfigSchema, config, 'screen share delivery'); - await this.setConfig(SCREEN_SHARE_DELIVERY_CONFIG_KEY, JSON.stringify(validated)); + await this.updateScreenShareDeliveryConfig(() => config); + } + + updateScreenShareDeliveryConfig( + update: (current: ScreenShareDeliveryConfig) => ScreenShareDeliveryConfig, + ): Promise { + return this.updateStoredConfig(SCREEN_SHARE_DELIVERY_CONFIG_KEY, (raw) => + validateStoredConfig( + ScreenShareDeliveryConfigSchema, + update(parseStoredScreenShareDeliveryConfig(raw)), + 'screen share delivery', + ), + ); + } + + async getPushServiceDeliveryConfig(): Promise { + const raw = await this.getConfig(PUSH_SERVICE_DELIVERY_CONFIG_KEY); + return parseStoredPushServiceDeliveryConfig(raw); + } + + updatePushServiceDeliveryConfig( + update: (current: PushServiceDeliveryConfig) => PushServiceDeliveryConfig, + ): Promise { + return this.updateStoredConfig(PUSH_SERVICE_DELIVERY_CONFIG_KEY, (raw) => + validateStoredConfig( + PushServiceDeliveryConfigSchema, + update(parseStoredPushServiceDeliveryConfig(raw)), + 'push service delivery', + ), + ); } async getExperimentDeliveryConfig(): Promise { @@ -1112,8 +1311,19 @@ export class InstanceConfigRepository { } async setExperimentDeliveryConfig(config: ExperimentDeliveryConfig): Promise { - const validated = validateStoredConfig(ExperimentDeliveryConfigSchema, config, 'experiment delivery'); - await this.setConfig(EXPERIMENT_DELIVERY_CONFIG_KEY, JSON.stringify(validated)); + await this.updateExperimentDeliveryConfig(() => config); + } + + updateExperimentDeliveryConfig( + update: (current: ExperimentDeliveryConfig) => ExperimentDeliveryConfig, + ): Promise { + return this.updateStoredConfig(EXPERIMENT_DELIVERY_CONFIG_KEY, (raw) => + validateStoredConfig( + ExperimentDeliveryConfigSchema, + update(parseStoredExperimentDeliveryConfig(raw)), + 'experiment delivery', + ), + ); } async readLimitConfigInputs(): Promise { @@ -1149,26 +1359,27 @@ export class InstanceConfigRepository { legal?: Partial; registration?: Partial; }): Promise { - const current = await this.getAppPublicConfig(); - const next = decodeAppPublicConfig({ - branding: { - ...current.branding, - ...(config.branding ?? {}), - }, - setup: { - ...current.setup, - ...(config.setup ?? {}), - }, - legal: { - ...current.legal, - ...(config.legal ?? {}), - }, - registration: { - ...current.registration, - ...(config.registration ?? {}), - }, + const next = await this.updateStoredConfig(APP_PUBLIC_CONFIG_KEY, (raw) => { + const current = parseStoredAppPublicConfig(raw); + return decodeAppPublicConfig({ + branding: { + ...current.branding, + ...(config.branding ?? {}), + }, + setup: { + ...current.setup, + ...(config.setup ?? {}), + }, + legal: { + ...current.legal, + ...(config.legal ?? {}), + }, + registration: { + ...current.registration, + ...(config.registration ?? {}), + }, + }); }); - await this.setConfig(APP_PUBLIC_CONFIG_KEY, JSON.stringify(next)); setCachedDateOfBirthCollection(next.registration.collect_date_of_birth); return next; } @@ -1188,10 +1399,21 @@ export class InstanceConfigRepository { return parseStoredInstancePolicyConfig(raw); } - async setInstancePolicyConfig(config: Partial): Promise { - const current = await this.readStoredInstancePolicyConfig(); - const next = decodeInstancePolicyConfig({...current, ...config}); - await this.setConfig(INSTANCE_POLICY_CONFIG_KEY, JSON.stringify(next)); + setInstancePolicyConfig(config: Partial): Promise { + return this.updateInstancePolicyConfig(() => config); + } + + async updateInstancePolicyConfig( + plan: (current: InstancePolicyConfig) => Partial, + ): Promise { + const cache = this.configCache; + const {result: next, written} = await this.compareAndSetStoredValue(cache, INSTANCE_POLICY_CONFIG_KEY, (raw) => { + const current = parseStoredInstancePolicyConfig(raw); + const patch = plan(current); + const config = decodeInstancePolicyConfig({...current, ...patch}); + return {value: Object.keys(patch).length === 0 ? null : JSON.stringify(config), result: config}; + }); + if (written) await this.publishRefresh(cache.sourceId); setCachedDeferredPhoneGateEnabled(resolveDeferredPhoneGateEnabled(next)); return next; } @@ -1201,37 +1423,37 @@ export class InstanceConfigRepository { return parseStoredInstanceIntegrationsConfig(raw); } - async setInstanceIntegrationsConfig(config: InstanceIntegrationsConfigPatch): Promise { - const current = await this.getInstanceIntegrationsConfig(); - const next = decodeInstanceIntegrationsConfig({ - gif: { - ...current.gif, - ...(config.gif ?? {}), - }, - youtube: { - ...current.youtube, - ...(config.youtube ?? {}), - }, - captcha: { - ...current.captcha, - ...(config.captcha ?? {}), - }, - email: { - ...current.email, - ...(config.email ?? {}), - smtp: { - ...current.email.smtp, - ...(config.email?.smtp ?? {}), + setInstanceIntegrationsConfig(config: InstanceIntegrationsConfigPatch): Promise { + return this.updateStoredConfig(INSTANCE_INTEGRATIONS_CONFIG_KEY, (raw) => { + const current = parseStoredInstanceIntegrationsConfig(raw); + return decodeInstanceIntegrationsConfig({ + gif: { + ...current.gif, + ...(config.gif ?? {}), }, - }, - bluesky: { - ...current.bluesky, - ...(config.bluesky ?? {}), - keys: config.bluesky?.keys ?? current.bluesky.keys, - }, + youtube: { + ...current.youtube, + ...(config.youtube ?? {}), + }, + captcha: { + ...current.captcha, + ...(config.captcha ?? {}), + }, + email: { + ...current.email, + ...(config.email ?? {}), + smtp: { + ...current.email.smtp, + ...(config.email?.smtp ?? {}), + }, + }, + bluesky: { + ...current.bluesky, + ...(config.bluesky ?? {}), + keys: config.bluesky?.keys ?? current.bluesky.keys, + }, + }); }); - await this.setConfig(INSTANCE_INTEGRATIONS_CONFIG_KEY, JSON.stringify(next)); - return next; } async getInstanceMediaConfig(): Promise { @@ -1239,16 +1461,16 @@ export class InstanceConfigRepository { return parseStoredInstanceMediaConfig(raw); } - async setInstanceMediaConfig(config: InstanceMediaConfigPatch): Promise { - const current = await this.getInstanceMediaConfig(); - const next = decodeInstanceMediaConfig({ - attachment_decay: { - ...current.attachment_decay, - ...(config.attachment_decay ?? {}), - }, + setInstanceMediaConfig(config: InstanceMediaConfigPatch): Promise { + return this.updateStoredConfig(INSTANCE_MEDIA_CONFIG_KEY, (raw) => { + const current = parseStoredInstanceMediaConfig(raw); + return decodeInstanceMediaConfig({ + attachment_decay: { + ...current.attachment_decay, + ...(config.attachment_decay ?? {}), + }, + }); }); - await this.setConfig(INSTANCE_MEDIA_CONFIG_KEY, JSON.stringify(next)); - return next; } async getEffectiveAttachmentDecayConfig(): Promise { @@ -1485,15 +1707,15 @@ export class InstanceConfigRepository { return parseStoredRegistrationConfig(raw); } - async setRegistrationConfig(config: Partial): Promise { - const current = await this.getRegistrationConfig(); - const next = decodeRegistrationConfig({ - mode: config.mode ?? current.mode, - admin_registration_urls_enabled: - config.admin_registration_urls_enabled ?? current.admin_registration_urls_enabled, + setRegistrationConfig(config: Partial): Promise { + return this.updateStoredConfig(REGISTRATION_CONFIG_KEY, (raw) => { + const current = parseStoredRegistrationConfig(raw); + return decodeRegistrationConfig({ + mode: config.mode ?? current.mode, + admin_registration_urls_enabled: + config.admin_registration_urls_enabled ?? current.admin_registration_urls_enabled, + }); }); - await this.setConfig(REGISTRATION_CONFIG_KEY, JSON.stringify(next)); - return next; } async getRegistrationPublicConfig(): Promise { @@ -1534,16 +1756,20 @@ export class InstanceConfigRepository { last_used_at: null, last_used_by_user_id: null, }; - const registrationUrls = await this.getRegistrationUrls(); - await this.setRegistrationUrls([registrationUrl, ...registrationUrls]); + await this.updateStoredConfig(REGISTRATION_URLS_KEY, (raw) => + validateStoredCollection( + StoredRegistrationUrlSchema, + [registrationUrl, ...parseStoredCollection(StoredRegistrationUrlSchema, raw, 'registration URLs')], + 'registration URLs', + ), + ); return {registrationUrl: redactRegistrationUrl(registrationUrl), code}; } async revokeRegistrationUrl(id: string): Promise { const now = new Date().toISOString(); - const registrationUrls = await this.getRegistrationUrls(); - await this.setRegistrationUrls( - registrationUrls.map((registrationUrl) => + await this.updateStoredConfig(REGISTRATION_URLS_KEY, (raw) => + parseStoredCollection(StoredRegistrationUrlSchema, raw, 'registration URLs').map((registrationUrl) => registrationUrl.id === id && !registrationUrl.revoked_at ? {...registrationUrl, revoked_at: now} : registrationUrl, @@ -1556,9 +1782,8 @@ export class InstanceConfigRepository { if (!normalizedCode) return null; const hash = this.hashRegistrationUrlCode(normalizedCode); const now = new Date(); - const registrationUrls = await this.getRegistrationUrls(); return ( - registrationUrls.find( + (await this.fetchRegistrationUrlDefinitions()).find( (registrationUrl) => (registrationUrl.id === normalizedCode || registrationUrl.code_hash === hash) && isRegistrationUrlUsable(registrationUrl, now), @@ -1566,21 +1791,68 @@ export class InstanceConfigRepository { ); } - async recordRegistrationUrlUse(id: string, userId: string): Promise { - const now = new Date().toISOString(); - const registrationUrls = await this.getRegistrationUrls(); - await this.setRegistrationUrls( - registrationUrls.map((registrationUrl) => - registrationUrl.id === id - ? { - ...registrationUrl, - use_count: registrationUrl.use_count + 1, - last_used_at: now, - last_used_by_user_id: userId, - } - : registrationUrl, - ), - ); + async claimRegistrationUrlUse(registrationUrlId: string, userId: string): Promise { + const cache = this.configCache; + let claimed: {result: RegistrationUrlClaim | null; written: boolean}; + try { + claimed = await this.compareAndSetStoredValue( + cache, + REGISTRATION_URLS_KEY, + (raw) => { + const registrationUrls = parseStoredCollection(StoredRegistrationUrlSchema, raw, 'registration URLs'); + const now = new Date(); + const claimable = registrationUrls.find( + (registrationUrl) => + registrationUrl.id === registrationUrlId && isRegistrationUrlUsable(registrationUrl, now), + ); + if (!claimable) return {value: null, result: null}; + const next = registrationUrls.map((registrationUrl) => + registrationUrl === claimable + ? { + ...registrationUrl, + use_count: registrationUrl.use_count + 1, + last_used_at: now.toISOString(), + last_used_by_user_id: userId, + } + : registrationUrl, + ); + return { + value: JSON.stringify(validateStoredCollection(StoredRegistrationUrlSchema, next, 'registration URLs')), + result: {registration_url_id: registrationUrlId, user_id: userId}, + }; + }, + ); + } catch (error) { + if (error instanceof InstanceConfigWriteConflictError) throw new ServiceUnavailableError(); + throw error; + } + if (claimed.written) await this.publishRefresh(cache.sourceId); + return claimed.result; + } + + async releaseRegistrationUrlUse(claim: RegistrationUrlClaim): Promise { + const cache = this.configCache; + try { + const {written} = await this.compareAndSetStoredValue(cache, REGISTRATION_URLS_KEY, (raw) => { + const registrationUrls = parseStoredCollection(StoredRegistrationUrlSchema, raw, 'registration URLs'); + const released = registrationUrls.find( + (registrationUrl) => registrationUrl.id === claim.registration_url_id && registrationUrl.use_count > 0, + ); + if (!released) return {value: null, result: null}; + const next = registrationUrls.map((registrationUrl) => + registrationUrl === released + ? {...registrationUrl, use_count: registrationUrl.use_count - 1} + : registrationUrl, + ); + return {value: JSON.stringify(next), result: null}; + }); + if (written) await this.publishRefresh(cache.sourceId); + } catch (error) { + Logger.warn( + {registrationUrlId: claim.registration_url_id, userId: claim.user_id, error}, + 'Releasing a registration URL use failed', + ); + } } async getPendingRegistrations(): Promise> { @@ -1591,104 +1863,82 @@ export class InstanceConfigRepository { } async addPendingRegistration(pendingRegistration: InstancePendingRegistration): Promise { - const pendingRegistrations = await this.getPendingRegistrations(); - const next = [ - pendingRegistration, - ...pendingRegistrations.filter((entry) => entry.user_id !== pendingRegistration.user_id), - ]; - await this.setPendingRegistrations(next); + await this.updateStoredConfig(REGISTRATION_PENDING_APPROVALS_KEY, (raw) => + validateStoredCollection( + StoredPendingRegistrationSchema, + [ + pendingRegistration, + ...parseStoredCollection(StoredPendingRegistrationSchema, raw, 'pending registrations').filter( + (entry) => entry.user_id !== pendingRegistration.user_id, + ), + ], + 'pending registrations', + ), + ); } async removePendingRegistration(userId: string): Promise { - const pendingRegistrations = await this.getPendingRegistrations(); - await this.setPendingRegistrations(pendingRegistrations.filter((entry) => entry.user_id !== userId)); + await this.updateStoredConfig(REGISTRATION_PENDING_APPROVALS_KEY, (raw) => + parseStoredCollection(StoredPendingRegistrationSchema, raw, 'pending registrations').filter( + (entry) => entry.user_id !== userId, + ), + ); } async getSsoConfig(options?: {includeSecret?: boolean}): Promise { - const configs = await this.getAllConfigs(); - const flags = readStoredSsoFlags(configs); - const read = (key: string): string | null => { - const v = configs.get(key); - if (!v) return null; - const trimmed = v.trim(); - return trimmed.length === 0 ? null : trimmed; - }; - const allowedDomains = parseStoredSsoAllowedEmailDomains(configs.get('sso_allowed_domains')); - const clientSecret = read('sso_client_secret'); - return { - ...flags, - displayName: read('sso_display_name'), - issuer: read('sso_issuer'), - authorizationUrl: read('sso_authorization_url'), - tokenUrl: read('sso_token_url'), - userInfoUrl: read('sso_userinfo_url'), - jwksUrl: read('sso_jwks_url'), - clientId: read('sso_client_id'), - clientSecret: options?.includeSecret ? clientSecret : undefined, - clientSecretSet: Boolean(clientSecret), - scope: read('sso_scope'), - allowedEmailDomains: allowedDomains, - redirectUri: null, - }; + return readStoredSsoConfig(await this.getAllConfigs(), options); } async setSsoConfig(config: Partial): Promise { - const current = await this.getSsoConfig({includeSecret: true}); - const definedConfig = Object.fromEntries( - Object.entries(config).filter(([, value]) => value !== undefined), - ) as Partial; - const next: InstanceSsoConfig = { - ...current, - ...definedConfig, - clientSecret: config.clientSecret !== undefined ? config.clientSecret : current.clientSecret, - }; - if (config.enabled === true && config.enforced === undefined && !current.enabled) { - next.enforced = true; - } - let allowedEmailDomains: Array; - try { - allowedEmailDomains = normalizeSsoAllowedEmailDomains(next.allowedEmailDomains); - } catch (error) { - if (next.enabled) { - throw error; - } - Logger.warn({error}, 'Clearing invalid SSO allowed domain config while SSO is disabled'); - allowedEmailDomains = []; - } - const entries: Array<[string, string]> = [ - ['sso_enabled', next.enabled ? 'true' : 'false'], - ['sso_enforced', next.enforced ? 'true' : 'false'], - ['sso_display_name', next.displayName ?? ''], - ['sso_issuer', next.issuer ?? ''], - ['sso_authorization_url', next.authorizationUrl ?? ''], - ['sso_token_url', next.tokenUrl ?? ''], - ['sso_userinfo_url', next.userInfoUrl ?? ''], - ['sso_jwks_url', next.jwksUrl ?? ''], - ['sso_client_id', next.clientId ?? ''], - ['sso_scope', next.scope ?? ''], - ['sso_allowed_domains', JSON.stringify(allowedEmailDomains)], - ['sso_auto_provision', next.autoProvision ? 'true' : 'false'], - ['sso_redirect_uri', ''], + const configs = await this.getAllConfigs(); + const current = readStoredSsoConfig(configs, {includeSecret: true}); + const enabled = config.enabled ?? current.enabled; + const allowedEmailDomains = + config.allowedEmailDomains === undefined + ? undefined + : normalizeSsoAllowedEmailDomainsForWrite(config.allowedEmailDomains, enabled); + const rows: Array = [ + ssoRow('sso_enabled', config.enabled, current.enabled, formatSsoBoolean), + ssoRow('sso_enforced', config.enforced, current.enforced, formatSsoBoolean), + ssoRow('sso_display_name', config.displayName, current.displayName, formatSsoString), + ssoRow('sso_issuer', config.issuer, current.issuer, formatSsoString), + ssoRow('sso_authorization_url', config.authorizationUrl, current.authorizationUrl, formatSsoString), + ssoRow('sso_token_url', config.tokenUrl, current.tokenUrl, formatSsoString), + ssoRow('sso_userinfo_url', config.userInfoUrl, current.userInfoUrl, formatSsoString), + ssoRow('sso_jwks_url', config.jwksUrl, current.jwksUrl, formatSsoString), + ssoRow('sso_client_id', config.clientId, current.clientId, formatSsoString), + ssoRow('sso_scope', config.scope, current.scope, formatSsoString), + ssoRow('sso_allowed_domains', allowedEmailDomains, current.allowedEmailDomains, formatSsoDomains), + ssoRow('sso_auto_provision', config.autoProvision, current.autoProvision, formatSsoBoolean), + ssoRow('sso_redirect_uri', undefined, null, formatSsoString), ]; if (config.clientSecret !== undefined) { - entries.push(['sso_client_secret', config.clientSecret ?? '']); + rows.push(ssoRow('sso_client_secret', config.clientSecret, current.clientSecret ?? null, formatSsoString)); } - await this.setConfigs(entries); + const cache = this.configCache; + const results = await Promise.allSettled( + rows + .filter((row) => row.value !== undefined || !configs.has(row.key)) + .map((row) => + this.compareAndSetStoredValue(cache, row.key, (raw) => ({value: nextSsoRowValue(row, raw), result: null})), + ), + ); + const errors: Array = results.flatMap((result) => (result.status === 'rejected' ? [result.reason] : [])); + if (results.some((result) => result.status === 'fulfilled' && result.value.written)) { + try { + await this.publishRefresh(cache.sourceId); + } catch (error) { + errors.push(error); + } + } + if (errors.length === 1) throw errors[0]; + if (errors.length > 1) throw new AggregateError(errors, 'Failed to write or publish the SSO config'); return this.getSsoConfig({includeSecret: true}); } - private async setRegistrationUrls(registrationUrls: Array): Promise { - const validated = validateStoredCollection(StoredRegistrationUrlSchema, registrationUrls, 'registration URLs'); - await this.setConfig(REGISTRATION_URLS_KEY, JSON.stringify(validated)); - } - - private async setPendingRegistrations(pendingRegistrations: Array): Promise { - const validated = validateStoredCollection( - StoredPendingRegistrationSchema, - pendingRegistrations, - 'pending registrations', - ); - await this.setConfig(REGISTRATION_PENDING_APPROVALS_KEY, JSON.stringify(validated)); + private async fetchRegistrationUrlDefinitions(): Promise> { + const raw = await this.fetchConfigFromDatabase(REGISTRATION_URLS_KEY); + return parseStoredCollection(StoredRegistrationUrlSchema, raw, 'registration URLs'); } private hashRegistrationUrlCode(code: string): string { diff --git a/fluxer_api/src/api/instance/PushServiceDeliveryConfigPublisher.ts b/fluxer_api/src/api/instance/PushServiceDeliveryConfigPublisher.ts new file mode 100644 index 000000000..3803123a8 --- /dev/null +++ b/fluxer_api/src/api/instance/PushServiceDeliveryConfigPublisher.ts @@ -0,0 +1,30 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +import type {PushServiceDeliveryConfig} from '@fluxer/schema/src/domains/admin/PushServiceDeliverySchemas'; +import type {INatsConnectionManager} from '@pkgs/nats/src/INatsConnectionManager'; + +const textEncoder = new TextEncoder(); + +export const PUSH_SERVICE_DELIVERY_CONFIG_NATS_SUBJECT = 'config.push.delivery'; + +interface PushServiceDeliveryConfigNatsMessage { + type: 'push_service_delivery_config'; + config: PushServiceDeliveryConfig; +} + +export class PushServiceDeliveryConfigPublisher { + constructor(private readonly connectionManager: INatsConnectionManager) {} + + async publish(config: PushServiceDeliveryConfig): Promise { + if (this.connectionManager.isClosed()) { + await this.connectionManager.connect(); + } + const connection = this.connectionManager.getConnection(); + const message: PushServiceDeliveryConfigNatsMessage = { + type: 'push_service_delivery_config', + config, + }; + connection.publish(PUSH_SERVICE_DELIVERY_CONFIG_NATS_SUBJECT, textEncoder.encode(JSON.stringify(message))); + await connection.flush(); + } +} diff --git a/fluxer_api/src/api/instance/tests/InstanceConfigWriteRaceExecutor.ts b/fluxer_api/src/api/instance/tests/InstanceConfigWriteRaceExecutor.ts new file mode 100644 index 000000000..8459541b6 --- /dev/null +++ b/fluxer_api/src/api/instance/tests/InstanceConfigWriteRaceExecutor.ts @@ -0,0 +1,100 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +import {getKvMeta} from '@app/api/database/CassandraMetaRegistry'; +import type {CassandraQueryExecutorForTesting} from '@app/api/database/CassandraQueryExecution'; +import type {CassandraParams, KvQueryMeta, PreparedQuery} from '@app/api/database/CassandraTypes'; +import type {InstanceConfigurationRow} from '@app/api/database/types/InstanceConfigTypes'; +import {InstanceConfiguration} from '@app/api/Tables'; + +type InstanceConfigWriteEvent = 'read' | 'write' | 'write rejected'; + +interface WriteGate { + size: number; + paused: Array<() => void>; +} + +const FETCH_ROW_QUERY = InstanceConfiguration.selectCql({ + where: InstanceConfiguration.where.eq('key'), + limit: 1, +}); + +export class InstanceConfigWriteRaceExecutor implements CassandraQueryExecutorForTesting { + readonly events: Array = []; + private watchedKey: string | null = null; + private gate: WriteGate | null = null; + private beforeEachWrite: (() => Promise) | null = null; + + constructor(private readonly base: CassandraQueryExecutorForTesting) {} + + watch(key: string): void { + this.watchedKey = key; + this.events.length = 0; + } + + pauseWritesUntil(size: number): void { + this.gate = {size, paused: []}; + } + + competeBeforeEachWrite(write: () => Promise): void { + this.beforeEachWrite = write; + } + + async writeDirectly(key: string, value: string): Promise { + await this.base.executeQuery(InstanceConfiguration.upsertAll({key, value, updated_at: new Date()})); + } + + async readDirectly(key: string): Promise { + const [row] = await this.base.executeQuery({cql: FETCH_ROW_QUERY, params: {key}}); + return row?.value ?? null; + } + + async executeQuery, P extends CassandraParams = CassandraParams>( + query: PreparedQuery

, + ): Promise> { + const meta = query.kvMeta ?? getKvMeta(query.cql); + if (this.watchedKey === null || !this.isWatched(meta, query.params)) { + return this.base.executeQuery(query); + } + if (meta?.action === 'select') { + this.events.push('read'); + return this.base.executeQuery(query); + } + await this.passGate(); + await this.beforeEachWrite?.(); + const rows = await this.base.executeQuery(query); + const applied = (rows[0] as {'[applied]'?: unknown} | undefined)?.['[applied]']; + this.events.push(applied === false ? 'write rejected' : 'write'); + return rows; + } + + executeBatch(queries: Array<{query: string; params: object; meta?: KvQueryMeta}>, atomic?: boolean): Promise { + return this.base.executeBatch(queries, atomic); + } + + reset(): void { + this.base.reset?.(); + this.watchedKey = null; + this.gate = null; + this.beforeEachWrite = null; + this.events.length = 0; + } + + async shutdown(): Promise { + await this.base.shutdown?.(); + } + + private isWatched(meta: KvQueryMeta | null | undefined, params: CassandraParams): boolean { + return meta?.table.name === InstanceConfiguration.name && params.key === this.watchedKey; + } + + private async passGate(): Promise { + const gate = this.gate; + if (gate === null) return; + await new Promise((release) => { + gate.paused.push(release); + if (gate.paused.length < gate.size) return; + this.gate = null; + for (const release of gate.paused) release(); + }); + } +} diff --git a/fluxer_api/src/api/middleware/ServiceSingletons.ts b/fluxer_api/src/api/middleware/ServiceSingletons.ts index 00528f1c4..17d368771 100644 --- a/fluxer_api/src/api/middleware/ServiceSingletons.ts +++ b/fluxer_api/src/api/middleware/ServiceSingletons.ts @@ -51,6 +51,7 @@ import {createUsersServiceClient} from '@app/api/infrastructure/UsersServiceClie import {VirusScanService} from '@app/api/infrastructure/VirusScanService'; import {GatewayRolloutConfigPublisher} from '@app/api/instance/GatewayRolloutConfigPublisher'; import {InstanceConfigRepository} from '@app/api/instance/InstanceConfigRepository'; +import {PushServiceDeliveryConfigPublisher} from '@app/api/instance/PushServiceDeliveryConfigPublisher'; import {InviteRepository} from '@app/api/invite/InviteRepository'; import {Logger} from '@app/api/Logger'; import {LimitConfigService} from '@app/api/limits/LimitConfigService'; @@ -155,6 +156,18 @@ export const getGatewayRolloutConfigPublisher = singleton( }), ), ); + +export const getPushServiceDeliveryConfigPublisher = singleton( + () => + new PushServiceDeliveryConfigPublisher( + new NatsConnectionManager({ + url: Config.nats.coreUrl, + token: Config.nats.authToken || undefined, + name: 'fluxer-api-push-service-delivery-config', + }), + ), +); + export const getVisionarySlotRepository = singleton(() => new VisionarySlotRepository()); export const getCacheService: () => ICacheService = singleton(() => new KVCacheProvider({client: getKVClient()})); export const getRateLimitService = singleton(() => new RateLimitService(getKVClient())); diff --git a/fluxer_api/src/api/openapi/openapi.json b/fluxer_api/src/api/openapi/openapi.json index b0b346bae..6869d16cc 100644 --- a/fluxer_api/src/api/openapi/openapi.json +++ b/fluxer_api/src/api/openapi/openapi.json @@ -17030,7 +17030,7 @@ "content": {"application/json": {"schema": {"$ref": "#/components/schemas/Error"}}} } }, - "description": "Registers a mobile push device token for APNs, Firebase Cloud Messaging, or UnifiedPush. UnifiedPush registrations include the endpoint URL plus Web Push encryption keys.", + "description": "Registers a mobile push device for APNs, Firebase Cloud Messaging, or UnifiedPush. A Web Push registration sends the endpoint URL with encryption_key and auth_secret. A raw registration sends the platform push token with no keys.", "security": [{"sessionToken": []}], "requestBody": { "required": true, @@ -22744,7 +22744,10 @@ "enum": ["android_fcm", "ios_apns", "android_unified_push"], "type": "string" }, - "token": {"description": "The platform-specific push notification token to unregister", "type": "string"}, + "token": { + "description": "The Web Push endpoint URL or raw platform push token used at registration", + "type": "string" + }, "app_id": { "description": "Client app channel or bundle mapping identifier, such as stable, beta, or canary", "type": "string" @@ -22808,7 +22811,10 @@ "enum": ["android_fcm", "ios_apns", "android_unified_push"], "type": "string" }, - "token": {"description": "The platform-specific push notification token or endpoint URL", "type": "string"}, + "token": { + "description": "The Web Push endpoint URL when encryption keys are supplied, otherwise the raw platform push token", + "type": "string" + }, "user_agent": {"description": "The user agent string identifying the device", "type": "string"}, "app_id": { "description": "Client app channel or bundle mapping identifier, such as stable, beta, or canary", @@ -22825,11 +22831,11 @@ "type": "string" }, "encryption_key": { - "description": "The P-256 ECDH public key for UnifiedPush encryption (base64url)", + "description": "The P-256 ECDH public key for Web Push encryption (base64url)", "type": "string" }, "auth_secret": { - "description": "The authentication secret for UnifiedPush encryption (base64url)", + "description": "The authentication secret for Web Push encryption (base64url)", "type": "string" } }, diff --git a/fluxer_api/src/api/rpc/RpcService.ts b/fluxer_api/src/api/rpc/RpcService.ts index a479556fc..385a05130 100644 --- a/fluxer_api/src/api/rpc/RpcService.ts +++ b/fluxer_api/src/api/rpc/RpcService.ts @@ -97,6 +97,7 @@ import {RateLimitError} from '@fluxer/errors/src/domains/core/RateLimitError'; import {UnauthorizedError} from '@fluxer/errors/src/domains/core/UnauthorizedError'; import {UnknownGuildError} from '@fluxer/errors/src/domains/guild/UnknownGuildError'; import {UnknownUserError} from '@fluxer/errors/src/domains/user/UnknownUserError'; +import {pushServiceDeliveryEnrols} from '@fluxer/schema/src/domains/admin/PushServiceDeliverySchemas'; import type {ChannelResponse} from '@fluxer/schema/src/domains/channel/ChannelSchemas'; import type {VoiceStateResponse} from '@fluxer/schema/src/domains/gateway/GatewaySchemas'; import type {GuildMemberResponse} from '@fluxer/schema/src/domains/guild/GuildMemberSchemas'; @@ -430,6 +431,13 @@ export class RpcService { }), }; case 'send_apns_push': { + const deliveryConfig = await this.instanceConfigRepository.getPushServiceDeliveryConfig(); + if (pushServiceDeliveryEnrols(deliveryConfig, request.user_id.toString())) { + Logger.warn( + {userId: request.user_id.toString(), configVersion: deliveryConfig.config_version}, + 'push service delivery path mismatch', + ); + } const result = await sendApnsPush({ userId: request.user_id.toString(), subscriptionId: request.subscription_id, @@ -635,6 +643,13 @@ export class RpcService { data: {config: rolloutConfig}, }; } + case 'get_push_service_delivery_config': { + const config = await this.instanceConfigRepository.getPushServiceDeliveryConfig(); + return { + type: 'get_push_service_delivery_config', + data: {config}, + }; + } default: { const exhaustiveCheck: never = request; throw new Error( diff --git a/fluxer_api/src/api/user/controllers/UserAccountController.ts b/fluxer_api/src/api/user/controllers/UserAccountController.ts index 94fbdb13b..c430b63c1 100644 --- a/fluxer_api/src/api/user/controllers/UserAccountController.ts +++ b/fluxer_api/src/api/user/controllers/UserAccountController.ts @@ -957,7 +957,7 @@ export function UserAccountController(app: HonoApp) { security: ['bearerToken', 'sessionToken'], tags: ['Users'], description: - 'Registers a mobile push device token for APNs, Firebase Cloud Messaging, or UnifiedPush. UnifiedPush registrations include the endpoint URL plus Web Push encryption keys.', + 'Registers a mobile push device for APNs, Firebase Cloud Messaging, or UnifiedPush. A Web Push registration sends the endpoint URL with encryption_key and auth_secret. A raw registration sends the platform push token with no keys.', }), async (ctx) => { const authSession = ctx.get('authSession'); diff --git a/fluxer_api/src/api/user/repositories/account/UserAccountRepository.ts b/fluxer_api/src/api/user/repositories/account/UserAccountRepository.ts index 387bebf7d..5d8ad13b3 100644 --- a/fluxer_api/src/api/user/repositories/account/UserAccountRepository.ts +++ b/fluxer_api/src/api/user/repositories/account/UserAccountRepository.ts @@ -3,6 +3,7 @@ import type {UserID} from '@app/api/BrandedTypes'; import {Db, type DbOp} from '@app/api/database/CassandraTypes'; import type {UserRow} from '@app/api/database/types/UserTypes'; +import {Logger} from '@app/api/Logger'; import {User} from '@app/api/models/User'; import { UserDataRepository, @@ -96,7 +97,12 @@ export class UserAccountRepository { return updatedUser; } catch (error) { if (!dataCommitted && emailClaim) { - await this.emailOwnershipRepo.abortEmailClaim(emailClaim); + await this.emailOwnershipRepo.abortEmailClaim(emailClaim).catch((abortError: unknown) => { + Logger.warn( + {userId: userId.toString(), abortError}, + 'Failed to abort the email claim of a user write that did not commit', + ); + }); } throw error; } diff --git a/fluxer_api/src/api/user/services/UserContentService.ts b/fluxer_api/src/api/user/services/UserContentService.ts index 386e9ebc0..8dbe93ab9 100644 --- a/fluxer_api/src/api/user/services/UserContentService.ts +++ b/fluxer_api/src/api/user/services/UserContentService.ts @@ -104,6 +104,27 @@ function assertPublicPushEndpoint(endpoint: string, fieldName: string): void { } } +function isPushEndpointUrl(token: string): boolean { + const normalized = token.trim().toLowerCase(); + return normalized.startsWith('https://') || normalized.startsWith('http://'); +} + +function resolveMobileWebPushKeys(device: RegisterMobileDeviceRequest): {p256dh: string; auth: string} | null { + const p256dh = device.encryption_key; + const auth = device.auth_secret; + if (p256dh && auth) return {p256dh, auth}; + if (p256dh || auth) { + throw InputValidationError.create( + p256dh ? 'auth_secret' : 'encryption_key', + 'Web Push registrations require encryption_key and auth_secret', + ); + } + if (isPushEndpointUrl(device.token)) { + throw InputValidationError.create('token', 'Endpoint URL registrations require encryption_key and auth_secret'); + } + return null; +} + function normalizeMobileAppId(appId: string | undefined): string { const normalized = appId?.trim(); return normalized && normalized.length > 0 ? normalized : DEFAULT_MOBILE_APP_ID; @@ -400,7 +421,8 @@ export class UserContentService { async registerMobileDevice(params: RegisterMobileDeviceParams): Promise { const {userId, authSessionIdHash, device} = params; - if (device.platform === 'android_unified_push') { + const webPushKeys = resolveMobileWebPushKeys(device); + if (webPushKeys) { assertPublicPushEndpoint(device.token, 'token'); } const appId = normalizeMobileAppId(device.app_id); @@ -411,8 +433,8 @@ export class UserContentService { subscription_id: subscriptionId, auth_session_id_hash: authSessionIdHash ?? null, endpoint: device.token, - p256dh_key: device.platform === 'android_unified_push' ? (device.encryption_key ?? null) : null, - auth_key: device.platform === 'android_unified_push' ? (device.auth_secret ?? null) : null, + p256dh_key: webPushKeys?.p256dh ?? null, + auth_key: webPushKeys?.auth ?? null, user_agent: device.user_agent ?? null, platform: device.platform, app_id: appId, diff --git a/fluxer_api/src/api/user/tests/PushSubscriptionLifecycle.test.ts b/fluxer_api/src/api/user/tests/PushSubscriptionLifecycle.test.ts index 91a784d7f..682cc8754 100644 --- a/fluxer_api/src/api/user/tests/PushSubscriptionLifecycle.test.ts +++ b/fluxer_api/src/api/user/tests/PushSubscriptionLifecycle.test.ts @@ -6,9 +6,12 @@ import { loginAccount, logoutSpecificSessions, } from '@app/api/auth/tests/AuthTestUtils'; +import {createUserID} from '@app/api/BrandedTypes'; +import type {PushSubscription} from '@app/api/models/PushSubscription'; import {type ApiTestHarness, createApiTestHarness} from '@app/api/test/ApiTestHarness'; import {HTTP_STATUS} from '@app/api/test/TestConstants'; import {createBuilder} from '@app/api/test/TestRequestBuilder'; +import {PushSubscriptionRepository} from '@app/api/user/repositories/PushSubscriptionRepository'; import { deleteMobileDevice, deletePushSubscription, @@ -21,6 +24,15 @@ import { import type {AuthSessionResponse} from '@fluxer/schema/src/domains/auth/AuthSchemas'; import {beforeEach, describe, expect, test} from 'vitest'; +async function findStoredSubscription(userId: string, subscriptionId: string): Promise { + const subscriptions = await new PushSubscriptionRepository().listPushSubscriptions(createUserID(BigInt(userId))); + const subscription = subscriptions.find((entry) => entry.subscriptionId === subscriptionId); + if (!subscription) { + throw new Error(`Stored push subscription ${subscriptionId} not found`); + } + return subscription; +} + describe('Push Subscription Lifecycle', () => { let harness: ApiTestHarness; beforeEach(async () => { @@ -144,6 +156,155 @@ describe('Push Subscription Lifecycle', () => { expect(mobileDevices.devices[0].device_id).toBe(registered.device_id); expect(mobileDevices.devices[0].platform).toBe('android_unified_push'); }); + test('APNs Web Push registration stores the endpoint and encryption keys', async () => { + const account = await createTestAccount(harness); + const endpoint = 'https://relay.example.com/apns/device-1'; + const registered = await registerMobileDevice(harness, account.token, { + platform: 'ios_apns', + token: endpoint, + encryption_key: 'relay-p256dh-key', + auth_secret: 'relay-auth-secret', + app_id: 'stable', + }); + const subscription = await findStoredSubscription(account.userId, registered.device_id); + expect(subscription.platform).toBe('ios_apns'); + expect(subscription.endpoint).toBe(endpoint); + expect(subscription.p256dhKey).toBe('relay-p256dh-key'); + expect(subscription.authKey).toBe('relay-auth-secret'); + }); + test('FCM Web Push registration stores the endpoint and encryption keys', async () => { + const account = await createTestAccount(harness); + const endpoint = 'https://relay.example.com/fcm/device-1'; + const registered = await registerMobileDevice(harness, account.token, { + platform: 'android_fcm', + token: endpoint, + encryption_key: 'fcm-relay-p256dh-key', + auth_secret: 'fcm-relay-auth-secret', + }); + const subscription = await findStoredSubscription(account.userId, registered.device_id); + expect(subscription.platform).toBe('android_fcm'); + expect(subscription.endpoint).toBe(endpoint); + expect(subscription.p256dhKey).toBe('fcm-relay-p256dh-key'); + expect(subscription.authKey).toBe('fcm-relay-auth-secret'); + }); + test('raw token registration stores the token without encryption keys', async () => { + const account = await createTestAccount(harness); + const registered = await registerMobileDevice(harness, account.token, { + platform: 'ios_apns', + token: '0123456789abcdef', + provider_environment: 'production', + }); + const subscription = await findStoredSubscription(account.userId, registered.device_id); + expect(subscription.platform).toBe('ios_apns'); + expect(subscription.endpoint).toBe('0123456789abcdef'); + expect(subscription.p256dhKey).toBeNull(); + expect(subscription.authKey).toBeNull(); + }); + test('Web Push and raw token registrations coexist for one platform', async () => { + const account = await createTestAccount(harness); + const rawDevice = await registerMobileDevice(harness, account.token, { + platform: 'android_fcm', + token: 'fcm-legacy-token', + }); + const webPushDevice = await registerMobileDevice(harness, account.token, { + platform: 'android_fcm', + token: 'https://relay.example.com/fcm/device-2', + encryption_key: 'coexist-p256dh-key', + auth_secret: 'coexist-auth-secret', + }); + expect(rawDevice.device_id).not.toBe(webPushDevice.device_id); + const rawSubscription = await findStoredSubscription(account.userId, rawDevice.device_id); + const webPushSubscription = await findStoredSubscription(account.userId, webPushDevice.device_id); + expect(rawSubscription.p256dhKey).toBeNull(); + expect(rawSubscription.authKey).toBeNull(); + expect(webPushSubscription.p256dhKey).toBe('coexist-p256dh-key'); + expect(webPushSubscription.authKey).toBe('coexist-auth-secret'); + const mobileDevices = await listMobileDevices(harness, account.token); + expect(mobileDevices.devices).toHaveLength(2); + }); + test('endpoint registration without encryption keys is rejected', async () => { + const account = await createTestAccount(harness); + await createBuilder(harness, account.token) + .post('/users/@me/mobile-devices') + .body({ + platform: 'ios_apns', + token: 'https://relay.example.com/apns/no-keys', + }) + .expect(HTTP_STATUS.BAD_REQUEST) + .execute(); + }); + test('endpoint registration with only one encryption key is rejected', async () => { + const account = await createTestAccount(harness); + await createBuilder(harness, account.token) + .post('/users/@me/mobile-devices') + .body({ + platform: 'android_fcm', + token: 'https://relay.example.com/fcm/half-keys', + encryption_key: 'half-p256dh-key', + }) + .expect(HTTP_STATUS.BAD_REQUEST) + .execute(); + }); + test('raw token registration with encryption keys is rejected', async () => { + const account = await createTestAccount(harness); + await createBuilder(harness, account.token) + .post('/users/@me/mobile-devices') + .body({ + platform: 'ios_apns', + token: '0123456789abcdef', + encryption_key: 'raw-p256dh-key', + auth_secret: 'raw-auth-secret', + }) + .expect(HTTP_STATUS.BAD_REQUEST) + .execute(); + }); + test('Web Push registration rejects an endpoint that is not publicly routable', async () => { + const account = await createTestAccount(harness); + const response = await createBuilder(harness, account.token) + .post('/users/@me/mobile-devices') + .body({ + platform: 'ios_apns', + token: 'https://127.0.0.1/apns/device', + encryption_key: 'local-p256dh-key', + auth_secret: 'local-auth-secret', + }) + .expect(HTTP_STATUS.BAD_REQUEST) + .execute(); + expect(JSON.stringify(response)).toContain('URL_NOT_PUBLICLY_ROUTABLE'); + }); + test('unregister removes a Web Push mobile registration', async () => { + const account = await createTestAccount(harness); + const endpoint = 'https://relay.example.com/apns/unregister'; + await registerMobileDevice(harness, account.token, { + platform: 'ios_apns', + token: endpoint, + encryption_key: 'unregister-p256dh-key', + auth_secret: 'unregister-auth-secret', + app_id: 'stable', + provider_environment: 'production', + }); + await unregisterMobileDevice(harness, account.token, { + platform: 'ios_apns', + token: endpoint, + app_id: 'stable', + provider_environment: 'production', + }); + const mobileDevices = await listMobileDevices(harness, account.token); + expect(mobileDevices.devices).toHaveLength(0); + }); + test('mobile Web Push registrations stay out of the web push subscription list', async () => { + const account = await createTestAccount(harness); + await registerMobileDevice(harness, account.token, { + platform: 'android_fcm', + token: 'https://relay.example.com/fcm/separate', + encryption_key: 'separate-p256dh-key', + auth_secret: 'separate-auth-secret', + }); + const webSubscriptions = await listPushSubscriptions(harness, account.token); + const mobileDevices = await listMobileDevices(harness, account.token); + expect(webSubscriptions.subscriptions).toHaveLength(0); + expect(mobileDevices.devices).toHaveLength(1); + }); test('list subscriptions returns multiple subscriptions', async () => { const account = await createTestAccount(harness); const first = await subscribePush(harness, account.token, 'https://push.example.com/multi-1'); diff --git a/fluxer_app/src/features/platform/service_worker/Worker.ts b/fluxer_app/src/features/platform/service_worker/Worker.ts index 2313f9834..fabde8765 100644 --- a/fluxer_app/src/features/platform/service_worker/Worker.ts +++ b/fluxer_app/src/features/platform/service_worker/Worker.ts @@ -22,7 +22,9 @@ import { matchesPushChannelNotification, normalizePushPayload, resolvePushChannelId, + resolvePushMessageId, resolvePushNotificationTag, + shouldRenotifyPushNotification, shouldSilenceNonMobilePushNotification, } from '@app/features/platform/service_worker/WorkerPushPayload'; @@ -214,6 +216,17 @@ const closePushNotifications = async (tag: string | undefined): Promise return 0; } }; +const getShownPushNotifications = async (tag: string): Promise> => { + if (typeof self.registration.getNotifications !== 'function') { + return []; + } + try { + return await self.registration.getNotifications({tag}); + } catch (error) { + await log('error', 'push: failed to read shown notifications', {tag, error: describeError(error)}); + return []; + } +}; const closePushNotificationsForChannel = async (channelId: string): Promise => { if (typeof self.registration.getNotifications !== 'function') { return 0; @@ -347,6 +360,9 @@ self.addEventListener('push', (event: PushEvent) => { })) as ReadonlyArray; clientState = getPushNotificationClientState(clientList); } catch {} + const renotify = + tag !== undefined && + shouldRenotifyPushNotification(resolvePushMessageId(payload), await getShownPushNotifications(tag)); const options: NotificationOptions & { renotify?: boolean; } = { @@ -355,7 +371,7 @@ self.addEventListener('push', (event: PushEvent) => { badge: payload.badge ?? undefined, data: payload.data ?? undefined, tag, - renotify: tag !== undefined, + renotify, ...getNotificationAlertOptions({ mobileOrTablet: isMobileOrTabletUserAgent(workerNavigator.userAgent, workerNavigator.maxTouchPoints ?? 0), silentOnNonMobile: shouldSilenceNonMobilePushNotification(clientState), @@ -365,6 +381,7 @@ self.addEventListener('push', (event: PushEvent) => { title, hasBody: Boolean(payload.body), tag, + renotify, hasData: payload.data !== undefined, badgeCount, hasWindowClient: clientState.hasWindowClient, diff --git a/fluxer_app/src/features/platform/service_worker/WorkerPushPayload.ts b/fluxer_app/src/features/platform/service_worker/WorkerPushPayload.ts index 6e176676b..06dd5070b 100644 --- a/fluxer_app/src/features/platform/service_worker/WorkerPushPayload.ts +++ b/fluxer_app/src/features/platform/service_worker/WorkerPushPayload.ts @@ -53,6 +53,19 @@ export const resolvePushNotificationTag = (payload: PushPayload): string | undef } return undefined; }; +export const resolvePushMessageId = (payload: PushPayload): string | undefined => { + const messageId = payload.data?.message_id; + if (typeof messageId === 'string' && messageId.length > 0) { + return messageId; + } + return undefined; +}; +export const shouldRenotifyPushNotification = ( + messageId: string | undefined, + shownWithSameTag: ReadonlyArray<{readonly data?: unknown}>, +): boolean => + messageId === undefined || + !shownWithSameTag.some((notification) => isRecord(notification.data) && notification.data.message_id === messageId); export const resolvePushChannelId = (payload: PushPayload): string | undefined => { const channelId = payload.data?.channel_id; if (typeof channelId === 'string' && channelId.length > 0) { diff --git a/fluxer_docs/astro.config.ts b/fluxer_docs/astro.config.ts index 32aa71380..66fe665b7 100644 --- a/fluxer_docs/astro.config.ts +++ b/fluxer_docs/astro.config.ts @@ -200,6 +200,7 @@ export default defineConfig({ 'http-api/users/relationships', 'http-api/users/notes', 'http-api/users/private-channels', + 'http-api/users/push-notifications', 'http-api/users/content', 'http-api/users/gifts', 'http-api/users/data-harvest', diff --git a/fluxer_docs/scripts/VerifyDocsCoverage.ts b/fluxer_docs/scripts/VerifyDocsCoverage.ts index 79386833c..eefda12cd 100644 --- a/fluxer_docs/scripts/VerifyDocsCoverage.ts +++ b/fluxer_docs/scripts/VerifyDocsCoverage.ts @@ -102,8 +102,6 @@ const MAIN_SPEC_EXEMPT = new Map2 | object | `branding`, `setup`, `legal`, and `registration` sub-objects, each merged field by field | @@ -566,6 +585,8 @@ The body has one optional object for each section. Fluxer leaves an absent secti `screen_share_delivery` works the same way, over the [screen share delivery configuration](#screen-share-delivery-configuration-object) fields and its own `config_version`. +`push_service_delivery` works the same way, over the [push service delivery configuration](#push-service-delivery-configuration-object) fields and its own `config_version`. + `experiment_delivery` takes both [experiment delivery configuration](#experiment-delivery-configuration-object) fields, each bound as documented there. It is a section of its own, so a write to it changes no `config_version` and changes no assignment, only the cadence on which clients ask for one. 3 A secret such as `klipy_api_key`, `api_key`, `hcaptcha_secret_key`, `turnstile_secret_key`, or the SMTP `password` is written when supplied and left alone when absent. `integrations.bluesky.keys` is the only way to write the Bluesky signing keys counted as `bluesky.key_count`. It takes up to 8 entries of `kid` (1-255 characters) and nullable `private_key` (up to 10000 characters), and replaces the stored key set outright @@ -602,7 +623,7 @@ Fluxer skips URL validation while the merged configuration leaves single sign-on | 400 | [error response](/admin-api/#error-response) | A policy transition is refused, returned as `INSTANCE_POLICY_TRANSITION_NOT_ALLOWED` | :::caution[Sections are applied one after another] -The order is `gateway_rollout`, `voice_noise_suppression`, `screen_share_delivery`, `experiment_delivery`, `sso`, `registration`, `app_public` branding, legal, and registration fields, `integrations`, `media`, `policy`, and finally `app_public.setup`. A failure part way through leaves the earlier sections written. +The order is `gateway_rollout`, `voice_noise_suppression`, `screen_share_delivery`, `push_service_delivery`, `experiment_delivery`, `sso`, `registration`, `app_public` branding, legal, and registration fields, `integrations`, `media`, `policy`, and finally `app_public.setup`. A failure part way through leaves the earlier sections written. ::: ### Side effects diff --git a/fluxer_docs/src/content/docs/http-api/users/push-notifications.mdx b/fluxer_docs/src/content/docs/http-api/users/push-notifications.mdx new file mode 100644 index 000000000..42609c180 --- /dev/null +++ b/fluxer_docs/src/content/docs/http-api/users/push-notifications.mdx @@ -0,0 +1,152 @@ +--- +# SPDX-License-Identifier: AGPL-3.0-or-later +title: Push notifications +description: Registering a device for push delivery and decrypting what arrives. +--- + +import RouteHeader from '@/components/RouteHeader.astro'; + +Fluxer delivers a notification to a registered device as [RFC 8291](https://datatracker.ietf.org/doc/html/rfc8291) Web Push. The client generates a P-256 key pair and an auth secret, registers the public half of the pair, and decrypts each delivery with the private half. + +Every route here requires a user session. A bot or OAuth2 bearer credential is refused with 403 `ACCESS_DENIED`. An account with an outstanding required action is refused with 403 `ACCOUNT_SUSPICIOUS_ACTIVITY`. + +## Registration shapes + +A registration takes one of two shapes on every platform. + +| Shape | What `token` holds | Keys | +| --- | --- | --- | +| Web Push | A publicly routable endpoint URL | Both `encryption_key` and `auth_secret` | +| Legacy | A raw vendor device token | Neither key is sent | + +Fluxer reads the shape from the body rather than from `platform`. A `token` that parses as a URL without both keys is refused, and so is a pair of keys sent with a raw vendor token. + +`platform` names the transport the device was reached on. + +| Value | Meaning | +| --- | --- | +| `android_fcm` | Firebase Cloud Messaging | +| `ios_apns` | Apple Push Notification service | +| `android_unified_push` | UnifiedPush, on an Android build without Google services | + +`android_unified_push` is always a Web Push registration. Sending it with no keys is refused. + +## Device registration object + +The identifier Fluxer assigns to one registration. + +### Structure + +| Field | Type | Description | +| --- | --- | --- | +| device_id | string | The registration identifier, 32 lowercase hexadecimal characters | + +Fluxer derives the identifier from `platform`, `app_id`, `provider_environment`, and `token`. The same four values always produce the same identifier, and registering them twice replaces the stored entry. + +## Register mobile push device + + + +Stores a push registration for the current account and returns its [device registration](#device-registration-object) object. + +### JSON body + +| Field | Type | Description | +| --- | --- | --- | +| platform | string | The [platform value](#registration-shapes) the device was reached on | +| token1 | string | The endpoint URL, or the raw vendor token on a legacy registration | +| encryption_key?2 | string | The base64url P-256 public key (1-1024 characters) | +| auth_secret?2 | string | The base64url auth secret (1-1024 characters) | +| app_id? | string | The client build, such as `stable`, `beta`, or `canary` (default `stable`) | +| provider_environment? | string | `production` or `development` | +| user_agent? | string | A user agent string describing the device (1-1024 characters) | + +1 1 to 4096 characters. A value that parses as a URL is a Web Push registration and needs both keys, and a value that does not must be sent with neither + +2 Sent together or not at all. Sending one alone is refused at the missing field + +An `ios_apns` registration with no `provider_environment` is stored as `production`. Every other platform stores no environment. + +Register an `https` endpoint. A Web Push registration whose `token` is not a valid URL is refused at `token`, and one whose host is a private or reserved address is refused with `URL_NOT_PUBLICLY_ROUTABLE`. + +### Response + +| Status | Body | Condition | +| --- | --- | --- | +| 200 | [device registration](#device-registration-object) object | The registration was stored | +| 400 | [error response](/http-api/#error-response) | The body matches neither shape and the request returns `INVALID_FORM_BODY` | + +### Rate limit + +20 requests per minute for each authenticated user, on the `user:push:subscribe` bucket. + +## Unregister mobile push device + + + +Removes the registration named by the values the client already holds. Returns 200 whether or not a registration was there. + +The four values below identify the registration the same way [Register mobile push device](#register-mobile-push-device) does. A value that differs from the one sent at registration names a different registration and removes nothing. + +### JSON body + +| Field | Type | Description | +| --- | --- | --- | +| platform | string | The [platform value](#registration-shapes) sent at registration | +| token | string | The endpoint URL or raw vendor token sent at registration (1-4096 characters) | +| app_id? | string | The client build sent at registration (default `stable`) | +| provider_environment? | string | The environment sent at registration | + +### Response body + +| Field | Type | Description | +| --- | --- | --- | +| success | boolean | Whether the removal ran, always true | + +### Response + +| Status | Body | Condition | +| --- | --- | --- | +| 200 | response body | The registration was removed, or there was none | +| 400 | [error response](/http-api/#error-response) | `platform` or `token` is missing or malformed and the request returns `INVALID_FORM_BODY` | + +### Rate limit + +40 requests per minute for each authenticated user, on the `user:push:unsubscribe` bucket. + +## Handling an incoming push + +Fluxer posts one encrypted record to the registered endpoint for each notification. The push service that owns the endpoint hands that record to the client. + +| Header | Value | +| --- | --- | +| Content-Encoding | Always `aes128gcm` | +| Content-Type | Always `application/octet-stream` | +| TTL | `86400` on a notification and `3600` on a clear | +| Urgency | `high` on a notification and `low` on a clear | +| Authorization | A VAPID token and the instance public key | + +The body is one `aes128gcm` record encrypted to the `encryption_key` and `auth_secret` the client registered. The client decrypts it locally with the private half of its key pair and its auth secret. Fluxer holds no key that opens the record after it is sealed. + +A record is 2816 bytes and its plaintext is at most 2713 bytes of JSON. A notification too large for that is shrunk before it is encrypted, one step at a time, until it fits. + +| Step | Effect | +| --- | --- | +| First | Media fields are dropped | +| Second | Icon fields are dropped | +| Third | The body text is shortened | +| Last | Only a minimal payload is left | + +A client has to tolerate a missing field. + +Two kinds of payload arrive. A notification payload describes something to show. A clear payload sets `type` to `notification_clear` and `action` to `clear_channel`, and asks the client to dismiss what it already showed for one channel. + +An endpoint that answers 404 or 410 removes the registration. Fluxer retries a transient failure and keeps the registration. + +### When decryption fails + +A record that does not decrypt cannot be recovered. Discard it and show nothing. + +Decryption fails when the registered keys no longer match the pair the client holds. Regenerating the key pair without registering again does that. Fluxer sees none of it. The push service already answered 2xx and the registration stays live. + +The client is the only party that can repair it. Unregister the stale entry, then register again with the current public key and auth secret. diff --git a/fluxer_docs/src/content/docs/operator/configuration.mdx b/fluxer_docs/src/content/docs/operator/configuration.mdx index e1d50714b..5aeef8e61 100644 --- a/fluxer_docs/src/content/docs/operator/configuration.mdx +++ b/fluxer_docs/src/content/docs/operator/configuration.mdx @@ -813,11 +813,11 @@ The Gateway reads the same names. A malformed pair, or a private key that does n ## Mobile push -Both `api` and `gateway` read the APNs and FCM names. None appear in `.env.example` or in `docker-compose.yml`, so configuring mobile push means editing the Compose file. All are optional. +`api`, `gateway`, and `push` read the APNs and FCM names. None appear in `.env.example` or in `docker-compose.yml`, so configuring mobile push means editing the Compose file. All are optional. #### `FLUXER_PUSH_APNS_ENABLED` -Default `false`. The APNs switch. Must be set on both `api` and `gateway`. +Default `false`. The APNs switch. Must be set on `api`, `gateway`, and `push`. #### `FLUXER_PUSH_APNS_TEAM_ID` @@ -845,7 +845,7 @@ Default `[]`. Per-app APNs configuration. JSON array. An entry with no `app_id` #### `FLUXER_PUSH_FCM_ENABLED` -Default `false`. The FCM switch. Must be set on both `api` and `gateway`. +Default `false`. The FCM switch. Must be set on `api`, `gateway`, and `push`. #### `FLUXER_PUSH_FCM_PROJECT_ID` @@ -875,6 +875,34 @@ Default `https://oauth2.googleapis.com/token`. The OAuth token endpoint. Change Default `[]`. Per-app FCM configuration. JSON array, under the same `app_id` rule as APNs. +## Push service settings + +`push` is the push notification service in the bundled stack. It reads the VAPID pair from [Web push](#web-push), the APNs and FCM names from [Mobile push](#mobile-push), and `FLUXER_SVC_NATS_URL` with `FLUXER_NATS_AUTH_TOKEN` from [Message bus and internal services](#message-bus-and-internal-services). It refuses to start without `FLUXER_INTERNAL_API_ENDPOINT` and `FLUXER_GATEWAY_RPC_AUTH_TOKEN`. Compose supplies both. The names below belong to `push`. All are optional. + +#### `FLUXER_PUSH_SERVICE_HOST` + +Default `0.0.0.0`. The bind address. It must parse as an IP address. Also settable with `--bind-host`. Compose sets `0.0.0.0`. + +#### `FLUXER_PUSH_SERVICE_PORT` + +Default `8126`. The listen port for the health and metrics endpoints. Also settable with `--port`. Compose sets `8126`. + +#### `FLUXER_PUSH_SERVICE_QUEUE_CAPACITY` + +Default `10000`. Notification jobs `push` holds at once. Accepts 1 to 1000000. Compose passes it through from `.env`. An empty value keeps the default. + +#### `FLUXER_PUSH_SERVICE_SEND_CONCURRENCY` + +Default `256`. Provider requests in flight at once, across every job. Accepts 1 to 65536. Compose passes it through from `.env`. An empty value keeps the default. + +#### `FLUXER_PUSH_SERVICE_APNS_BASE_URL` + +No default. Replaces the APNs host in both environments. Set it only for a test double. + +#### `FLUXER_PUSH_SERVICE_FCM_BASE_URL` + +Default `https://fcm.googleapis.com`. The FCM host. Set it only for a proxy or a test double. + ## Payments Stripe billing, which the shipped stack keeps off. All are optional. @@ -1077,7 +1105,7 @@ Defaults to `debug` in development, `info` otherwise. The Node log level. Read b #### `RUST_LOG` -Default `info`. The Rust log filter. Read by `media-proxy`, `app-proxy`, `admin`, and the internal services. Not in `.env.example` or the Compose file. +Default `info`. The Rust log filter. Read by `media-proxy`, `app-proxy`, `admin`, `push`, and the internal services. Not in `.env.example` or the Compose file. #### `FLUXER_GATEWAY_LOGGER_LEVEL` @@ -1111,7 +1139,7 @@ Default `development`. The runtime mode. `development`, `production`, or `test`. Default `false`. The self-host switch. Compose sets `true`. It relaxes the production Postgres SSL requirement, seeds the limit tier, gates registration, billing and discovery controllers, and turns blocklist feeds off. -`/_metrics` on `api`, `media-proxy`, and `gateway`, plus the Gateway's `/_health/ready`, `/_health/drain`, and `/_health/undrain`, are gated to loopback peers, so no proxy reaches them. The probes that work from outside are `/api/_health`, `/gateway/_health`, `/media/_health`, and the edge's own `/_health`. +`/_metrics` on `api`, `media-proxy`, `gateway`, and `push`, plus the Gateway's `/_health/ready`, `/_health/drain`, and `/_health/undrain`, are gated to loopback peers, so no proxy reaches them. The probes that work from outside are `/api/_health`, `/gateway/_health`, `/media/_health`, and the edge's own `/_health`. ## Feature flags and development switches @@ -1663,7 +1691,7 @@ See [Bluesky connections](#bluesky-connections) for configuration. #### `FLUXER_PUSH_APNS_` and `FLUXER_PUSH_FCM_` -Mobile push cannot be configured at all from the example. +Mobile push cannot be configured at all from the example. `api`, `gateway`, and `push` all read these names. An override has to reach all three. #### `RUST_LOG`, `LOG_LEVEL` and `LOGGER_LEVEL` @@ -1742,6 +1770,7 @@ The stack runs its containers on one Docker bridge network, which is private to | worker | fluxer-api | Background lanes and the cron scheduler | | gateway | fluxer-gateway | The Gateway WebSocket | | media-proxy | fluxer-media-proxy | Uploads, transforms, and media delivery | +| push | fluxer-push | Push notification delivery | | admin | fluxer-admin | The admin dashboard | | snowflakes, snowflakes-shard | fluxer-snowflakes | Identifier allocation | | users, users-shard | fluxer-users | User reads and writes | @@ -1758,7 +1787,7 @@ The stack runs its containers on one Docker bridge network, which is private to The edge and LiveKit are the only services that publish ports. The edge publishes 80/tcp, 443/tcp, 443/udp, or one plain-HTTP port under the overlay. LiveKit publishes 7881/tcp and 7882/udp. Everything else is reachable only over the bridge network. -`api` is the one service an operator configures directly, through the shared environment block. `worker`, `gateway`, `app-proxy`, and `media-proxy` are touched rarely, `worker` for lane concurrency, `app-proxy` for CSP extras, and `media-proxy` for transform limits. The edge takes only the `FLUXER_EDGE_` variables, `postgres` only the password, `meilisearch` only the master key, `valkey` only the `FLUXER_VALKEY_` tuning values, and `livekit` only the key pair and the port variables. `static-proxy` reads no environment variables, and the remaining services need none. +`api` is the one service an operator configures directly, through the shared environment block. `worker`, `gateway`, `app-proxy`, `media-proxy`, and `push` are touched rarely, `worker` for lane concurrency, `app-proxy` for CSP extras, `media-proxy` for transform limits, and `push` for queue capacity and send concurrency. The edge takes only the `FLUXER_EDGE_` variables, `postgres` only the password, `meilisearch` only the master key, `valkey` only the `FLUXER_VALKEY_` tuning values, and `livekit` only the key pair and the port variables. `static-proxy` reads no environment variables, and the remaining services need none. The internal services each run a router, which takes requests and holds no state, and one shard, which holds the caches and the database connections. The shipped stack fixes `FLUXER_SVC_SHARD_COUNT` at `1`. @@ -1766,7 +1795,7 @@ The internal services each run a router, which takes requests and holds no state Every service has a memory limit and four also have a memory reservation, all under `deploy.resources`. Compose reads `gb` as 1024 MiB and `mb` as 1 MiB, so `5gb` is 5368709120 bytes. Plain `docker compose up` applies both keys on a single host, with no Swarm and no `--compatibility` flag. The engine rejects any limit below `6mb`, and rejects a limit lower than the same service's reservation with `Minimum memory limit can not be less than memory reservation limit`. -A limit is a ceiling. The limits below sum to 16.75 GiB and the stack does not need a host that large, because each container uses only the memory it allocates, up to its limit. +A limit is a ceiling. The limits below sum to 18.5 GiB and the stack does not need a host that large, because each container uses only the memory it allocates, up to its limit. `deploy.resources.reservations.memory` becomes the container's cgroup v2 `memory.low`, which biases kernel reclaim towards other containers under host pressure. It reserves nothing on its own. @@ -1828,6 +1857,10 @@ Default `1gb`. The ceiling for `gateway`. The BEAM has no heap ceiling of its ow Default `512mb`. The ceiling for `media-proxy`. Image and video transforms decode into this ceiling, so a large upload is what reaches this limit. +#### `FLUXER_PUSH_MEMORY_LIMIT` + +Default `256mb`. The ceiling for `push`. Raise it alongside `FLUXER_PUSH_SERVICE_QUEUE_CAPACITY` or `FLUXER_PUSH_SERVICE_SEND_CONCURRENCY`. + #### `FLUXER_STATIC_PROXY_MEMORY_LIMIT` Default `256mb`. The ceiling for `static-proxy`. The service reads no environment variables and serves files only. diff --git a/fluxer_docs/src/content/docs/operator/reverse-proxy.mdx b/fluxer_docs/src/content/docs/operator/reverse-proxy.mdx index 4eec2f65d..e4cbbdc59 100644 --- a/fluxer_docs/src/content/docs/operator/reverse-proxy.mdx +++ b/fluxer_docs/src/content/docs/operator/reverse-proxy.mdx @@ -402,7 +402,7 @@ Pass the query string on `/gateway` through untouched. Clients always send `?v=` Apple and Google require the association files at those fixed paths for saved-password autofill and app links. A proxy that forwards all paths needs no extra rules. Include them explicitly if you use a path allowlist. -`/_metrics` on the API, Media Proxy, and Gateway, plus `/_health/ready`, `/_health/drain`, and `/_health/undrain` on the Gateway, are gated to loopback and are unreachable through any proxy. The probes that work through a proxy are `/_health`, `/api/_health`, `/gateway/_health`, and `/media/_health`. +`/_metrics` on the API, Media Proxy, Gateway, and push service, plus `/_health/ready`, `/_health/drain`, and `/_health/undrain` on the Gateway, are gated to loopback and are unreachable through any proxy. The push service has no public path. The probes that work through a proxy are `/_health`, `/api/_health`, `/gateway/_health`, and `/media/_health`. ## Trusted proxies diff --git a/fluxer_gateway/src/gateway/fluxer_gateway_config.erl b/fluxer_gateway/src/gateway/fluxer_gateway_config.erl index b38c08a79..05095de5a 100644 --- a/fluxer_gateway/src/gateway/fluxer_gateway_config.erl +++ b/fluxer_gateway/src/gateway/fluxer_gateway_config.erl @@ -78,6 +78,12 @@ env_gateway_base_config() -> <<"gateway_role">> => env_optional_binary("FLUXER_GATEWAY_ROLE"), <<"rpc_auth_token">> => env_binary("FLUXER_GATEWAY_RPC_AUTH_TOKEN", <<>>), <<"push_enabled">> => env_bool("FLUXER_GATEWAY_PUSH_ENABLED", true), + <<"push_clear_notifications_enabled">> => env_bool( + "FLUXER_GATEWAY_PUSH_CLEAR_NOTIFICATIONS_ENABLED", true + ), + <<"push_outbox_request_timeout_ms">> => env_int( + "FLUXER_GATEWAY_PUSH_OUTBOX_REQUEST_TIMEOUT_MS", 100000 + ), <<"logger_level">> => env_binary("FLUXER_GATEWAY_LOGGER_LEVEL", <<"info">>), <<"api_rpc_endpoint">> => env_optional_binary("FLUXER_GATEWAY_API_RPC_ENDPOINT"), <<"cluster_enabled">> => env_bool("FLUXER_GATEWAY_CLUSTER_ENABLED", false), @@ -246,7 +252,16 @@ build_push_config(Service, Public) -> push_dispatcher_max_inflight => get_int( Service, <<"push_dispatcher_max_inflight">>, 16 ), - push_dispatcher_max_queue => get_int(Service, <<"push_dispatcher_max_queue">>, 2048) + push_dispatcher_max_queue => get_int(Service, <<"push_dispatcher_max_queue">>, 2048), + push_clear_notifications_enabled => get_bool( + Service, <<"push_clear_notifications_enabled">>, true + ), + push_outbox_max_queue => get_int(Service, <<"push_outbox_max_queue">>, 10000), + push_outbox_max_inflight => get_int(Service, <<"push_outbox_max_inflight">>, 64), + push_outbox_request_timeout_ms => get_int( + Service, <<"push_outbox_request_timeout_ms">>, 100000 + ), + push_outbox_max_age_ms => get_int(Service, <<"push_outbox_max_age_ms">>, 300000) }. -spec build_sharding_config(map()) -> config(). diff --git a/fluxer_gateway/src/gateway/fluxer_gateway_sup.erl b/fluxer_gateway/src/gateway/fluxer_gateway_sup.erl index e2b953955..115f043a3 100644 --- a/fluxer_gateway/src/gateway/fluxer_gateway_sup.erl +++ b/fluxer_gateway/src/gateway/fluxer_gateway_sup.erl @@ -38,7 +38,8 @@ common_children() -> child_spec(gateway_nats_pool, gateway_nats_pool), child_spec(gateway_event_pause, gateway_event_pause), child_spec(gateway_concurrency, gateway_concurrency), - child_spec(gateway_rollout_config, gateway_rollout_config) + child_spec(gateway_rollout_config, gateway_rollout_config), + child_spec(push_delivery_config, push_delivery_config) ] ++ cluster_children() ++ [ child_spec(gateway_dispatch_relay, gateway_dispatch_relay), @@ -99,6 +100,7 @@ role_specs(calls, Role) -> role_specs(push, _Role) -> [ child_spec(push_dispatcher, push_dispatcher), + child_spec(push_outbox, push_outbox), child_spec(push, push) ]. diff --git a/fluxer_gateway/src/gateway/gateway_handler.erl b/fluxer_gateway/src/gateway/gateway_handler.erl index 6696f3ebf..41f7f7c61 100644 --- a/fluxer_gateway/src/gateway/gateway_handler.erl +++ b/fluxer_gateway/src/gateway/gateway_handler.erl @@ -149,11 +149,18 @@ websocket_info(_, State) -> {ok, State}. -spec terminate(term(), cowboy_req:req(), state() | term()) -> ok. -terminate(_Reason, _Req, State) when is_map(State) -> - terminate_with_state(eqwalizer:dynamic_cast(State)); +terminate(Reason, _Req, State) when is_map(State) -> + terminate_with_state(eqwalizer:dynamic_cast(State)), + exit_on_client_close(Reason); terminate(_Reason, _Req, _State) -> ok. +-spec exit_on_client_close(term()) -> ok. +exit_on_client_close({remote, Code, _Payload}) when Code =:= 1000; Code =:= 1001 -> + exit({shutdown, client_closed}); +exit_on_client_close(_Reason) -> + ok. + -spec terminate_with_state(state()) -> ok. terminate_with_state(#{compress_ctx := CompressCtx, session_pid := SessionPid} = State) -> gateway_rollout_config:unsubscribe_changes(self()), diff --git a/fluxer_gateway/src/gateway/gateway_nats_rpc.erl b/fluxer_gateway/src/gateway/gateway_nats_rpc.erl index fcba4fb24..827bf73f5 100644 --- a/fluxer_gateway/src/gateway/gateway_nats_rpc.erl +++ b/fluxer_gateway/src/gateway/gateway_nats_rpc.erl @@ -346,6 +346,7 @@ subscribe_recorded_extra_subject( ) -> case nats:sub(Conn, Subject, subscription_opts(QueueGroup)) of {ok, Sid} -> + notify_resubscribed(Subject), Acc#{Key => Sub#{sid => Sid}}; {error, Reason} -> logger:error("Gateway NATS RPC failed to resubscribe extra subject", #{ @@ -381,11 +382,50 @@ handle_msg( true -> dispatch_rpc(Subject, Payload, MsgOpts, HC, MaxH, HRefs, State); false -> - ReplyTo = maps:get(reply_to, MsgOpts, undefined), - gateway_rollout_config ! {nats_msg, Subject, Payload, ReplyTo}, + deliver_subject(Subject, Payload, maps:get(reply_to, MsgOpts, undefined)), State end. +-spec subject_owner(binary()) -> atom() | undefined. +subject_owner(<<"config.gateway.rollout">>) -> gateway_rollout_config; +subject_owner(<<"config.push.delivery">>) -> push_delivery_config; +subject_owner(_Subject) -> undefined. + +-spec notify_resubscribed(binary()) -> ok. +notify_resubscribed(Subject) -> + case subject_owner(Subject) of + undefined -> ok; + Owner -> notify_owner(whereis(Owner), Subject) + end. + +-spec notify_owner(pid() | undefined, binary()) -> ok. +notify_owner(Pid, Subject) when is_pid(Pid) -> + Pid ! {nats_resubscribed, Subject}, + ok; +notify_owner(undefined, _Subject) -> + ok. + +-spec deliver_subject(binary(), binary(), binary() | undefined) -> ok. +deliver_subject(Subject, Payload, ReplyTo) -> + case subject_owner(Subject) of + undefined -> + logger:warning("Gateway NATS RPC has no owner for subject", #{subject => Subject}); + Owner -> + deliver_to_owner(Owner, Subject, Payload, ReplyTo) + end. + +-spec deliver_to_owner(atom(), binary(), binary(), binary() | undefined) -> ok. +deliver_to_owner(Owner, Subject, Payload, ReplyTo) -> + case whereis(Owner) of + undefined -> + logger:warning("Gateway NATS RPC subject owner is not running", #{ + subject => Subject, owner => Owner + }); + Pid -> + Pid ! {nats_msg, Subject, Payload, ReplyTo}, + ok + end. + -spec dispatch_rpc( binary(), binary(), diff --git a/fluxer_gateway/src/gateway/gateway_rpc_presence.erl b/fluxer_gateway/src/gateway/gateway_rpc_presence.erl index c6f46686f..37383bb0f 100644 --- a/fluxer_gateway/src/gateway/gateway_rpc_presence.erl +++ b/fluxer_gateway/src/gateway/gateway_rpc_presence.erl @@ -30,9 +30,23 @@ handle_dispatch(#{<<"user_id">> := UserIdBin, <<"event">> := Event, <<"data">> : case dispatch_to_owner(UserId, EventAtom, Data) of ok -> true; {error, not_found} -> handle_offline_dispatch(EventAtom, UserId, Data); + {error, unavailable} -> handle_unreachable_dispatch(EventAtom, UserId, Data); _ -> gateway_rpc_error:raise(<<"presence_dispatch_error">>) end. +-spec handle_unreachable_dispatch(atom(), integer(), map()) -> no_return(). +handle_unreachable_dispatch(message_create, UserId, Data) -> + push_unreachable_dispatch(push_delivery_config:is_enrolled(UserId), UserId, Data); +handle_unreachable_dispatch(_EventAtom, _UserId, _Data) -> + gateway_rpc_error:raise(<<"presence_dispatch_error">>). + +-spec push_unreachable_dispatch(boolean(), integer(), map()) -> no_return(). +push_unreachable_dispatch(true, UserId, Data) -> + _ = handle_offline_dispatch(message_create, UserId, Data), + gateway_rpc_error:raise(<<"presence_dispatch_error">>); +push_unreachable_dispatch(false, _UserId, _Data) -> + gateway_rpc_error:raise(<<"presence_dispatch_error">>). + -spec dispatch_event_atom_or_error(term()) -> atom(). dispatch_event_atom_or_error(Event) when is_binary(Event) -> normalize_dispatch_event_atom(Event); diff --git a/fluxer_gateway/src/gateway/metrics_handler.erl b/fluxer_gateway/src/gateway/metrics_handler.erl index 7c5925719..146201f39 100644 --- a/fluxer_gateway/src/gateway/metrics_handler.erl +++ b/fluxer_gateway/src/gateway/metrics_handler.erl @@ -45,6 +45,8 @@ render_metrics() -> render_cluster_counters(), render_process_counts(), render_push_dispatcher_stats(), + render_push_delivery_gate_stats(), + render_push_outbox_stats(safe_apply_map(fun push_outbox:stats/0)), render_vm_metrics() ]. @@ -277,6 +279,153 @@ render_push_dispatcher_stats() -> ] end. +-spec render_push_delivery_gate_stats() -> iolist(). +render_push_delivery_gate_stats() -> + ConfigVersion = safe_apply_int(fun push_delivery_config:config_version/0), + [ + format_metric( + <<"fluxer_gateway_push_delivery_config_version">>, + <<"gauge">>, + <<"Push delivery config version in effect">>, + integer_to_binary(ConfigVersion) + ), + render_push_delivery_config_updates( + safe_apply_map(fun push_delivery_config:update_counts/0) + ), + render_push_delivery_gate_counters(safe_apply_map(fun push:delivery_gate_counters/0)) + ]. + +-spec render_push_delivery_config_updates(map()) -> iolist(). +render_push_delivery_config_updates(Counts) -> + format_labeled_series( + <<"fluxer_gateway_push_delivery_config_updates_total">>, + <<"counter">>, + <<"Push delivery config reads by outcome">>, + [ + {<<"result=\"updated\"">>, gate_counter(updated, Counts)}, + {<<"result=\"unchanged\"">>, gate_counter(unchanged, Counts)}, + {<<"result=\"stale\"">>, gate_counter(stale, Counts)}, + {<<"result=\"rejected\"">>, gate_counter(rejected, Counts)} + ] + ). + +-spec render_push_delivery_gate_counters(map()) -> iolist(). +render_push_delivery_gate_counters(Counters) when map_size(Counters) =:= 0 -> + []; +render_push_delivery_gate_counters(Counters) -> + [ + format_metric( + <<"fluxer_gateway_push_delivery_service_users_total">>, + <<"counter">>, + <<"Recipients routed to the push service">>, + gate_counter(delivery_gate_service_users, Counters) + ), + format_metric( + <<"fluxer_gateway_push_delivery_gateway_users_total">>, + <<"counter">>, + <<"Recipients kept on the gateway push path">>, + gate_counter(delivery_gate_gateway_users, Counters) + ), + format_metric( + <<"fluxer_gateway_push_delivery_jobs_published_total">>, + <<"counter">>, + <<"Push jobs published to the push service">>, + gate_counter(delivery_gate_jobs_published, Counters) + ), + format_metric( + <<"fluxer_gateway_push_delivery_publish_failed_total">>, + <<"counter">>, + <<"Push job publishes that fell back to the gateway path">>, + gate_counter(delivery_gate_publish_failed, Counters) + ) + ]. + +-spec render_push_outbox_stats(map()) -> iolist(). +render_push_outbox_stats(Stats) when map_size(Stats) =:= 0 -> + []; +render_push_outbox_stats(Stats) -> + [render_push_outbox_queue_stats(Stats), render_push_outbox_hand_back_stats(Stats)]. + +-spec render_push_outbox_queue_stats(map()) -> iolist(). +render_push_outbox_queue_stats(Stats) -> + [ + format_metric( + <<"fluxer_gateway_push_outbox_depth">>, + <<"gauge">>, + <<"Push jobs queued in the outbox">>, + gate_counter(depth, Stats) + ), + format_metric( + <<"fluxer_gateway_push_outbox_inflight">>, + <<"gauge">>, + <<"Push job requests awaiting a reply">>, + gate_counter(inflight, Stats) + ), + format_metric( + <<"fluxer_gateway_push_outbox_delivered_total">>, + <<"counter">>, + <<"Push jobs acknowledged by the push service">>, + gate_counter(delivered, Stats) + ), + format_metric( + <<"fluxer_gateway_push_outbox_retries_total">>, + <<"counter">>, + <<"Push job requests scheduled for retry">>, + gate_counter(retries, Stats) + ), + format_metric( + <<"fluxer_gateway_push_outbox_sheds_total">>, + <<"counter">>, + <<"Earliest queued push jobs shed at outbox capacity">>, + gate_counter(sheds, Stats) + ), + format_metric( + <<"fluxer_gateway_push_outbox_truncations_total">>, + <<"counter">>, + <<"Queued recipients dropped because they read the channel">>, + gate_counter(truncations, Stats) + ), + format_metric( + <<"fluxer_gateway_push_outbox_skipped_active_total">>, + <<"counter">>, + <<"Queued recipients skipped because they became active">>, + gate_counter(skipped_active, Stats) + ) + ]. + +-spec render_push_outbox_hand_back_stats(map()) -> iolist(). +render_push_outbox_hand_back_stats(Stats) -> + [ + format_metric( + <<"fluxer_gateway_push_outbox_fallbacks_total">>, + <<"counter">>, + <<"Push jobs handed back to the gateway push path">>, + gate_counter(fallbacks, Stats) + ), + format_metric( + <<"fluxer_gateway_push_outbox_lost_total">>, + <<"counter">>, + <<"Push job hand-backs whose gateway push path send crashed">>, + gate_counter(lost, Stats) + ), + format_metric( + <<"fluxer_gateway_push_outbox_fallback_backlog">>, + <<"gauge">>, + <<"Push job hand-backs waiting for a runner">>, + gate_counter(fallback_backlog, Stats) + ), + format_metric( + <<"fluxer_gateway_push_outbox_fallback_runners">>, + <<"gauge">>, + <<"Push job hand-backs running on the gateway push path">>, + gate_counter(fallback_runners, Stats) + ) + ]. + +-spec gate_counter(atom(), map()) -> binary(). +gate_counter(Key, Counters) -> + integer_to_binary(maps:get(Key, Counters, 0)). + -spec render_vm_metrics() -> iolist(). render_vm_metrics() -> Memory = safe_apply_list(fun erlang:memory/0), diff --git a/fluxer_gateway/src/guild/guild_dispatch_push.erl b/fluxer_gateway/src/guild/guild_dispatch_push.erl index befd95c83..ffc507a28 100644 --- a/fluxer_gateway/src/guild/guild_dispatch_push.erl +++ b/fluxer_gateway/src/guild/guild_dispatch_push.erl @@ -5,19 +5,26 @@ -export([ maybe_send_push_notifications/4, - collect_and_send_push_notifications/3 + collect_and_send_push_notifications/3, + push_counters/0 ]). -define(MAX_FORMAT_MEMBERS, 50). -define(CONCURRENCY_LIMIT_KEY, guild_push_concurrency_limit). -define(DEFAULT_PUSH_CONCURRENCY, 2). -define(MAX_PUSH_CONCURRENCY, 8). +-define(QUEUE_LIMIT_KEY, guild_push_queue_limit). +-define(DEFAULT_PUSH_QUEUE_LIMIT, 64). +-define(MAX_PUSH_QUEUE_LIMIT, 1024). -define(PUSH_WORKER_MAX_AGE_MS, 60000). +-define(GRACE_RECHECK_KEY, guild_push_offline_grace_recheck_ms). +-define(DEFAULT_GRACE_RECHECK_MS, 7000). -define(PUSH_COUNTERS, guild_push_counters). -define(PUSH_COUNTER_KEYS, [ worker_started, worker_completed, worker_failed, + queued, dropped_at_limit, spawn_failed, slot_reclaimed, @@ -32,6 +39,10 @@ -type guild_id() :: integer(). -type user_id() :: integer(). -type push_worker() :: {integer(), pid(), integer()}. +-type push_gate() :: push_worker() | undefined. +-type grace_sessions() :: #{user_id() => [pid()]}. +-type grace_hold() :: {grace_sessions(), integer()} | none. +-type held_push() :: {[user_id()], grace_hold(), term()}. -export_type([event/0, event_data/0, guild_state/0, guild_id/0]). -spec maybe_send_push_notifications(event(), event_data(), guild_id(), guild_state()) -> ok. @@ -46,11 +57,50 @@ maybe_send_push_notifications(_Event, _FinalData, _GuildId, _UpdatedState) -> -spec maybe_spawn_push(event_data(), guild_id(), guild_state()) -> ok. maybe_spawn_push(FinalData, GuildId, UpdatedState) -> Limit = push_concurrency_limit(), + QueueLimit = push_queue_limit(), Workers = live_push_workers(), put(push_inflight_workers, Workers), - case length(Workers) < Limit of - true -> spawn_push(FinalData, GuildId, UpdatedState); - false -> count_push_event(dropped_at_limit) + case push_admission(length(Workers), Limit, QueueLimit) of + run -> + spawn_push(FinalData, GuildId, UpdatedState, undefined); + queue -> + count_push_event(queued), + spawn_push(FinalData, GuildId, UpdatedState, lists:nth(Limit, Workers)); + overflow -> + push_queue_overflow(GuildId, Limit, QueueLimit) + end. + +-spec push_admission(non_neg_integer(), pos_integer(), non_neg_integer()) -> + run | queue | overflow. +push_admission(Tracked, Limit, _QueueLimit) when Tracked < Limit -> run; +push_admission(Tracked, Limit, QueueLimit) when Tracked - Limit < QueueLimit -> queue; +push_admission(_Tracked, _Limit, _QueueLimit) -> overflow. + +-spec push_queue_overflow(guild_id(), pos_integer(), non_neg_integer()) -> ok. +push_queue_overflow(GuildId, Limit, QueueLimit) -> + count_push_event(dropped_at_limit), + logger:warning( + "guild_push_dropped_at_limit: guild_id=~p concurrency=~p queue_limit=~p", + [GuildId, Limit, QueueLimit] + ). + +-spec push_queue_limit() -> non_neg_integer(). +push_queue_limit() -> + case application:get_env(fluxer_gateway, ?QUEUE_LIMIT_KEY, undefined) of + Value when is_integer(Value), Value >= 0 -> min(Value, ?MAX_PUSH_QUEUE_LIMIT); + _ -> ?DEFAULT_PUSH_QUEUE_LIMIT + end. + +-spec push_counters() -> #{atom() => non_neg_integer()}. +push_counters() -> + try ets:tab2list(?PUSH_COUNTERS) of + Rows -> + maps:from_list([ + {Key, Count} + || {Key, Count} <- Rows, is_atom(Key), is_integer(Count) + ]) + catch + error:badarg -> #{} end. -spec local_process_alive(pid()) -> boolean(). @@ -185,29 +235,31 @@ ensure_push_counter_key(Key) -> error:badarg -> ok end. --spec spawn_push(event_data(), guild_id(), guild_state()) -> ok. -spawn_push(FinalData, GuildId, UpdatedState) -> +-spec spawn_push(event_data(), guild_id(), guild_state(), push_gate()) -> ok. +spawn_push(FinalData, GuildId, UpdatedState, Gate) -> Data = maps:get(data, UpdatedState, #{}), case maps:get(members_ets, Data, undefined) of MembersTab when is_reference(MembersTab) -> CompactState = compact_push_state( eqwalizer:dynamic_cast(MembersTab), Data, GuildId, UpdatedState ), - spawn_compact_push(FinalData, GuildId, CompactState); + spawn_compact_push(FinalData, GuildId, CompactState, Gate); _ -> - missing_members_table(FinalData, GuildId, UpdatedState) + missing_members_table(FinalData, GuildId, UpdatedState, Gate) end. -spec compact_push_state(ets:tid(), map(), guild_id(), guild_state()) -> guild_state(). compact_push_state(MembersTab, Data, GuildId, UpdatedState) -> Sessions = maps:get(sessions, UpdatedState, #{}), + SessionEligibility = build_push_session_eligibility(Sessions, UpdatedState), #{ id => maps:get(id, UpdatedState, GuildId), data => compact_push_data(Data), virtual_channel_access => maps:get(virtual_channel_access, UpdatedState, #{}), members_ets => MembersTab, member_presence => maps:get(member_presence, UpdatedState, undefined), - session_eligibility => build_push_session_eligibility(Sessions, UpdatedState), + session_eligibility => SessionEligibility, + grace_hold => grace_hold(Sessions, SessionEligibility), member_count => maps:get(member_count, UpdatedState, undefined) }. @@ -227,32 +279,34 @@ compact_push_data(Data) -> Data ). --spec spawn_compact_push(event_data(), guild_id(), guild_state()) -> ok. -spawn_compact_push(FinalData, GuildId, CompactState) -> +-spec spawn_compact_push(event_data(), guild_id(), guild_state(), push_gate()) -> ok. +spawn_compact_push(FinalData, GuildId, CompactState, Gate) -> spawn_push_worker( fun() -> collect_and_send_compact_push_notifications(FinalData, GuildId, CompactState) end, - GuildId + GuildId, + Gate ). --spec missing_members_table(event_data(), guild_id(), guild_state()) -> ok. -missing_members_table(FinalData, GuildId, UpdatedState) -> +-spec missing_members_table(event_data(), guild_id(), guild_state(), push_gate()) -> ok. +missing_members_table(FinalData, GuildId, UpdatedState, Gate) -> count_push_event(members_table_missing), logger:warning( "guild_push_members_table_unavailable: guild_id=~p phase=spawn", [GuildId] ), - spawn_legacy_push(FinalData, GuildId, UpdatedState). + spawn_legacy_push(FinalData, GuildId, UpdatedState, Gate). --spec spawn_legacy_push(event_data(), guild_id(), guild_state()) -> ok. -spawn_legacy_push(FinalData, GuildId, UpdatedState) -> +-spec spawn_legacy_push(event_data(), guild_id(), guild_state(), push_gate()) -> ok. +spawn_legacy_push(FinalData, GuildId, UpdatedState, Gate) -> LegacyState = legacy_push_state(GuildId, UpdatedState), spawn_push_worker( fun() -> collect_and_send_push_notifications(FinalData, GuildId, LegacyState) end, - GuildId + GuildId, + Gate ). -spec legacy_push_state(guild_id(), guild_state()) -> guild_state(). @@ -266,9 +320,12 @@ legacy_push_state(GuildId, UpdatedState) -> member_count => maps:get(member_count, UpdatedState, undefined) }. --spec spawn_push_worker(fun(() -> ok), guild_id()) -> ok. -spawn_push_worker(Worker, GuildId) -> - Counted = fun() -> run_counted_push_worker(Worker) end, +-spec spawn_push_worker(fun(() -> ok), guild_id(), push_gate()) -> ok. +spawn_push_worker(Worker, GuildId, Gate) -> + Counted = fun() -> + ok = wait_for_push_slot(Gate), + run_counted_push_worker(Worker) + end, case try_spawn_push_worker(Counted, GuildId) of {ok, Pid} -> put(push_inflight, Pid), @@ -278,6 +335,18 @@ spawn_push_worker(Worker, GuildId) -> count_push_event(spawn_failed) end. +-spec wait_for_push_slot(push_gate()) -> ok. +wait_for_push_slot(undefined) -> + ok; +wait_for_push_slot({_Gen, Pid, _StartedMs}) -> + Ref = erlang:monitor(process, Pid), + receive + {'DOWN', Ref, process, Pid, _Reason} -> ok + after ?PUSH_WORKER_MAX_AGE_MS -> + erlang:demonitor(Ref, [flush]), + ok + end. + -spec run_counted_push_worker(fun(() -> ok)) -> ok. run_counted_push_worker(Worker) -> try Worker() of @@ -348,7 +417,7 @@ scan_and_send_compact_push(MessageData, GuildId, ChannelId, State) -> Context = compact_scan_context(MessageData, ChannelId, State), MembersTab = eqwalizer:dynamic_cast(maps:get(members_ets, State)), case scan_push_members(MembersTab, Context, State) of - {ok, #{eligible_user_ids := []}} -> + {ok, #{eligible_user_ids := [], held_user_ids := []}} -> ok; {ok, Scan} -> send_compact_scanned_push(MessageData, GuildId, Scan, State); @@ -374,7 +443,8 @@ compact_scan_context(MessageData, ChannelId, State) -> mention_roles => mention_role_id_set(MentionRoles), format_members => format_member_id_set(MessageData), channel_id => ChannelId, - session_eligibility => maps:get(session_eligibility, State) + session_eligibility => maps:get(session_eligibility, State), + grace_sessions => held_sessions(maps:get(grace_hold, State, none)) }. -spec scan_push_members(ets:tid(), map(), guild_state()) -> @@ -388,7 +458,12 @@ scan_push_members(MembersTab, Context, State) -> -spec initial_scan_acc() -> map(). initial_scan_acc() -> - #{eligible_user_ids => [], user_roles => #{}, format_members => #{}}. + #{ + eligible_user_ids => [], + held_user_ids => [], + user_roles => #{}, + format_members => #{} + }. -spec member_id_snapshot(ets:tid()) -> [user_id()]. member_id_snapshot(MembersTab) -> @@ -414,26 +489,51 @@ scan_push_member({UserId, Member}, Context, State, Acc) when is_integer(UserId), UserId > 0, is_map(Member) -> Acc1 = maybe_collect_format_member(UserId, Member, Context, Acc), - case scanned_member_is_eligible(UserId, Member, Context, State) of - true -> add_scanned_member(UserId, Member, Acc1); - false -> Acc1 + case scanned_member_route(UserId, Member, Context, State) of + skip -> Acc1; + Key -> add_scanned_member(Key, UserId, Member, Acc1) end; scan_push_member(_Row, _Context, _State, Acc) -> Acc. --spec scanned_member_is_eligible(user_id(), map(), map(), guild_state()) -> boolean(). -scanned_member_is_eligible(UserId, Member, Context, State) -> - is_push_candidate(UserId, Member, Context) andalso - guild_permissions:can_view_channel( - UserId, maps:get(channel_id, Context), Member, State - ). +-spec scanned_member_route(user_id(), map(), map(), guild_state()) -> + eligible_user_ids | held_user_ids | skip. +scanned_member_route(UserId, Member, Context, State) -> + case member_push_route(UserId, Member, Context) of + skip -> + skip; + Key -> + visible_member_route( + Key, + guild_permissions:can_view_channel( + UserId, maps:get(channel_id, Context), Member, State + ) + ) + end. --spec add_scanned_member(user_id(), map(), map()) -> map(). -add_scanned_member(UserId, Member, Acc) -> - Eligible = maps:get(eligible_user_ids, Acc), +-spec member_push_route(user_id(), map(), map()) -> + eligible_user_ids | held_user_ids | skip. +member_push_route(UserId, Member, Context) -> + case is_push_candidate(UserId, Member, Context) of + true -> eligible_user_ids; + false -> grace_member_route(maps:is_key(UserId, maps:get(grace_sessions, Context, #{}))) + end. + +-spec grace_member_route(boolean()) -> held_user_ids | skip. +grace_member_route(true) -> held_user_ids; +grace_member_route(false) -> skip. + +-spec visible_member_route(eligible_user_ids | held_user_ids, boolean()) -> + eligible_user_ids | held_user_ids | skip. +visible_member_route(Key, true) -> Key; +visible_member_route(_Key, false) -> skip. + +-spec add_scanned_member(eligible_user_ids | held_user_ids, user_id(), map(), map()) -> map(). +add_scanned_member(Key, UserId, Member, Acc) -> + UserIds = maps:get(Key, Acc), UserRoles = maps:get(user_roles, Acc), Acc#{ - eligible_user_ids := [UserId | Eligible], + Key := [UserId | UserIds], user_roles := UserRoles#{UserId => extract_role_ids(Member)} }. @@ -473,7 +573,12 @@ send_compact_scanned_push(MessageData, GuildId, Scan, State) -> maps:get(user_roles, Scan), maps:get(session_eligibility, State), FormatData, - large_guild_meta(State) + large_guild_meta(State), + { + lists:reverse(maps:get(held_user_ids, Scan, [])), + maps:get(grace_hold, State, none), + maps:get(member_presence, State, undefined) + } ). -spec compact_format_data(map(), map()) -> map(). @@ -563,6 +668,7 @@ send_push_notifications(MessageData, GuildId, State) -> CandidateUserIds, ChannelId, SessionEligibility, + grace_hold(Sessions, SessionEligibility), Data, State ) @@ -575,6 +681,7 @@ send_push_notifications(MessageData, GuildId, State) -> [user_id()], integer(), map(), + grace_hold(), map(), guild_state() ) -> ok. @@ -585,14 +692,21 @@ send_to_eligible( CandidateUserIds, ChannelId, SessionEligibility, + Hold, Data, State ) -> - case find_eligible_users_for_push(Members, CandidateUserIds, ChannelId, State) of - [] -> + EligibleUserIds = find_eligible_users_for_push( + Members, CandidateUserIds, ChannelId, State + ), + HeldUserIds = find_eligible_users_for_push( + Members, held_candidate_user_ids(Hold, CandidateUserIds), ChannelId, State + ), + case {EligibleUserIds, HeldUserIds} of + {[], []} -> ok; - EligibleUserIds -> - UserRolesMap = build_user_roles_map(Members, EligibleUserIds), + _ -> + UserRolesMap = build_user_roles_map(Members, EligibleUserIds ++ HeldUserIds), send_push_to_eligible_users( MessageData, GuildId, @@ -600,10 +714,18 @@ send_to_eligible( UserRolesMap, SessionEligibility, Data, - large_guild_meta(State) + large_guild_meta(State), + {HeldUserIds, Hold, maps:get(member_presence, State, undefined)} ) end. +-spec held_candidate_user_ids(grace_hold(), [user_id()]) -> [user_id()]. +held_candidate_user_ids(none, _CandidateUserIds) -> + []; +held_candidate_user_ids({Held, _RecheckAt}, CandidateUserIds) -> + Candidates = maps:from_keys(CandidateUserIds, true), + [UserId || UserId <- maps:keys(Held), not maps:is_key(UserId, Candidates)]. + -spec push_candidate_user_ids(map(), map(), event_data()) -> [user_id()]. push_candidate_user_ids(Members, SessionEligibility, MessageData) -> push_candidate_user_ids( @@ -852,7 +974,7 @@ accumulate_session_eligibility(Session, Acc) -> end. -spec send_push_to_eligible_users( - event_data(), guild_id(), [user_id()], map(), map(), map(), map() | undefined + event_data(), guild_id(), [user_id()], map(), map(), map(), map() | undefined, held_push() ) -> ok. send_push_to_eligible_users( MessageData, @@ -861,7 +983,8 @@ send_push_to_eligible_users( UserRolesMap, ConnectedUsers, Data, - LargeGuildMeta + LargeGuildMeta, + Held ) -> AuthorIdBin = maps:get(<<"id">>, maps:get(<<"author">>, MessageData, #{}), undefined), case guild_dispatch_decorate:parse_snowflake(<<"author.id">>, AuthorIdBin) of @@ -871,10 +994,9 @@ send_push_to_eligible_users( ChannelIdBin = maps:get(<<"channel_id">>, MessageData), ChannelName = find_channel_name(ChannelIdBin, Data), RoleNames = build_role_names_map(Data), - do_send_push( + Params = push_params( MessageData, GuildId, - EligibleUserIds, UserRolesMap, ConnectedUsers, ChannelName, @@ -882,13 +1004,14 @@ send_push_to_eligible_users( Data, AuthorId, LargeGuildMeta - ) + ), + ok = send_push_now(EligibleUserIds, Params), + hold_push_through_grace(Held, Params) end. --spec do_send_push( +-spec push_params( event_data(), guild_id(), - [user_id()], map(), map(), binary(), @@ -896,11 +1019,10 @@ send_push_to_eligible_users( map(), integer(), map() | undefined -) -> ok. -do_send_push( +) -> map(). +push_params( MessageData, GuildId, - EligibleUserIds, UserRolesMap, ConnectedUsers, ChannelName, @@ -912,9 +1034,8 @@ do_send_push( Guild = maps:get(<<"guild">>, Data), DefaultMessageNotifications = maps:get(<<"default_message_notifications">>, Guild, 0), GuildName = maps:get(<<"name">>, Guild, <<"Unknown">>), - push:handle_message_create(#{ + #{ message_data => MessageData, - user_ids => EligibleUserIds, guild_id => GuildId, author_id => AuthorId, guild_default_notifications => DefaultMessageNotifications, @@ -929,7 +1050,128 @@ do_send_push( connected_users => ConnectedUsers, guild_member_count => meta_member_count(LargeGuildMeta), guild_features => meta_features(LargeGuildMeta) - }). + }. + +-spec send_push_now([user_id()], map()) -> ok. +send_push_now([], _Params) -> + ok; +send_push_now(UserIds, Params) -> + push:handle_message_create(Params#{user_ids => UserIds}). + +-spec hold_push_through_grace(held_push(), map()) -> ok. +hold_push_through_grace({[_ | _] = UserIds, {Sessions, RecheckAt}, Presences}, Params) -> + HeldParams = Params#{ + user_roles => maps:with(UserIds, maps:get(user_roles, Params)), + connected_users => maps:with(UserIds, maps:get(connected_users, Params)) + }, + HeldSessions = maps:with(UserIds, Sessions), + _ = spawn(fun() -> + deliver_after_grace(HeldParams, HeldSessions, RecheckAt, Presences) + end), + ok; +hold_push_through_grace(_Held, _Params) -> + ok. + +-spec deliver_after_grace(map(), grace_sessions(), integer(), term()) -> ok. +deliver_after_grace(Params, Sessions, RecheckAt, Presences) -> + ok = apply_push_worker_priority(), + LiveSessions = await_grace_recheck(monitor_grace_sessions(Sessions), RecheckAt), + send_push_now( + users_left_offline(lists:usort(maps:values(LiveSessions)), Presences), Params + ). + +-spec monitor_grace_sessions(grace_sessions()) -> #{reference() => user_id()}. +monitor_grace_sessions(Sessions) -> + maps:fold(fun monitor_user_sessions/3, #{}, Sessions). + +-spec monitor_user_sessions(user_id(), [pid()], #{reference() => user_id()}) -> + #{reference() => user_id()}. +monitor_user_sessions(UserId, Pids, Monitors) -> + lists:foldl( + fun(Pid, Acc) -> Acc#{erlang:monitor(process, Pid) => UserId} end, Monitors, Pids + ). + +-spec await_grace_recheck(#{reference() => user_id()}, integer()) -> + #{reference() => user_id()}. +await_grace_recheck(Monitors, RecheckAt) -> + Remaining = max(0, RecheckAt - erlang:monotonic_time(millisecond)), + receive + {'DOWN', Ref, process, _Pid, _Reason} when is_map_key(Ref, Monitors) -> + await_grace_recheck(maps:remove(Ref, Monitors), RecheckAt) + after Remaining -> + Monitors + end. + +-spec users_left_offline([user_id()], term()) -> [user_id()]. +users_left_offline(UserIds, Presences) -> + [UserId || UserId <- UserIds, presence_is_offline(UserId, Presences)]. + +-spec presence_is_offline(user_id(), term()) -> boolean(). +presence_is_offline(UserId, Presences) -> + case lookup_presence_safe(UserId, Presences) of + undefined -> false; + Presence -> maps:get(<<"status">>, Presence, <<"offline">>) =:= <<"offline">> + end. + +-spec grace_hold(map(), #{user_id() => boolean()}) -> grace_hold(). +grace_hold(Sessions, SessionEligibility) -> + case suppressed_enrolled_sessions(Sessions, SessionEligibility) of + Held when map_size(Held) =:= 0 -> + none; + Held -> + {Held, erlang:monotonic_time(millisecond) + grace_recheck_ms()} + end. + +-spec held_sessions(grace_hold()) -> grace_sessions(). +held_sessions(none) -> #{}; +held_sessions({Held, _RecheckAt}) -> Held. + +-spec suppressed_enrolled_sessions(map(), #{user_id() => boolean()}) -> grace_sessions(). +suppressed_enrolled_sessions(Sessions, SessionEligibility) -> + maps:fold( + fun(_Sid, Session, Acc) -> maybe_hold_session(Session, SessionEligibility, Acc) end, + #{}, + Sessions + ). + +-spec maybe_hold_session(term(), #{user_id() => boolean()}, grace_sessions()) -> + grace_sessions(). +maybe_hold_session(Session, SessionEligibility, Acc) when is_map(Session) -> + hold_suppressed_session( + maps:get(user_id, Session, undefined), + maps:get(pid, Session, undefined), + SessionEligibility, + Acc + ); +maybe_hold_session(_Session, _SessionEligibility, Acc) -> + Acc. + +-spec hold_suppressed_session(term(), term(), #{user_id() => boolean()}, grace_sessions()) -> + grace_sessions(). +hold_suppressed_session(UserId, Pid, SessionEligibility, Acc) when + is_integer(UserId), is_pid(Pid) +-> + case maps:get(UserId, SessionEligibility, true) of + false -> + hold_enrolled_session(push_delivery_config:is_enrolled(UserId), UserId, Pid, Acc); + true -> + Acc + end; +hold_suppressed_session(_UserId, _Pid, _SessionEligibility, Acc) -> + Acc. + +-spec hold_enrolled_session(boolean(), user_id(), pid(), grace_sessions()) -> grace_sessions(). +hold_enrolled_session(true, UserId, Pid, Acc) -> + Acc#{UserId => [Pid | maps:get(UserId, Acc, [])]}; +hold_enrolled_session(false, _UserId, _Pid, Acc) -> + Acc. + +-spec grace_recheck_ms() -> pos_integer(). +grace_recheck_ms() -> + case application:get_env(fluxer_gateway, ?GRACE_RECHECK_KEY, undefined) of + Value when is_integer(Value), Value > 0 -> Value; + _ -> ?DEFAULT_GRACE_RECHECK_MS + end. -spec meta_member_count(map() | undefined) -> non_neg_integer() | undefined. meta_member_count(#{member_count := MemberCount}) -> MemberCount; @@ -1243,7 +1485,7 @@ send_push_to_eligible_users_uses_full_data_for_channel_name_test() -> ?assertEqual( ok, send_push_to_eligible_users( - MessageData, 10, [1], #{1 => []}, #{}, Data, undefined + MessageData, 10, [1], #{1 => []}, #{}, Data, undefined, {[], none, undefined} ) ), receive @@ -1388,7 +1630,9 @@ format_member_id_set_includes_author_test() -> spawn_push_without_members_table_falls_back_to_legacy_push_test() -> reset_push_worker_state(), try - ?assertEqual(ok, spawn_push(#{}, 7, #{id => 7, data => #{}, sessions => #{}})), + ?assertEqual( + ok, spawn_push(#{}, 7, #{id => 7, data => #{}, sessions => #{}}, undefined) + ), ?assert(is_pid(get(push_inflight))) after reset_push_worker_state() @@ -1443,6 +1687,7 @@ inflight_pid_without_a_slot_still_counts_against_the_limit_test() -> Blocker = blocking_push_worker(), put(push_inflight, Blocker), ok = application:set_env(fluxer_gateway, ?CONCURRENCY_LIMIT_KEY, 1), + ok = application:set_env(fluxer_gateway, ?QUEUE_LIMIT_KEY, 0), try ?assertEqual([Blocker], worker_pids(live_push_workers())), ?assertEqual(ok, maybe_spawn_push(#{}, 7, legacy_test_state())), @@ -1510,6 +1755,7 @@ stale_push_worker_slot_is_reclaimed_and_counted_test() -> push_drop_is_readable_from_named_ets_table_test() -> reset_push_worker_state(), ok = application:set_env(fluxer_gateway, ?CONCURRENCY_LIMIT_KEY, 1), + ok = application:set_env(fluxer_gateway, ?QUEUE_LIMIT_KEY, 0), Blocker = blocking_push_worker(), put(push_inflight_workers, [{1, Blocker, erlang:monotonic_time(millisecond)}]), Before = read_push_counter(dropped_at_limit), @@ -1575,7 +1821,8 @@ spawned_worker_reports_started_and_completed_test() -> Self ! {ran, self()}, ok end, - 7 + 7, + undefined ) ), Pid = get(push_inflight), @@ -1653,6 +1900,7 @@ bounded_push_concurrency_spawns_second_worker_under_limit_test() -> bounded_push_concurrency_counts_drops_at_limit_test() -> reset_push_worker_state(), ok = application:set_env(fluxer_gateway, ?CONCURRENCY_LIMIT_KEY, 2), + ok = application:set_env(fluxer_gateway, ?QUEUE_LIMIT_KEY, 0), First = blocking_push_worker(), Second = blocking_push_worker(), Now = erlang:monotonic_time(millisecond), @@ -1690,6 +1938,7 @@ bounded_push_concurrency_reaps_finished_workers_test() -> bounded_worker_list_stays_bounded_by_the_limit_test() -> reset_push_worker_state(), ok = application:set_env(fluxer_gateway, ?CONCURRENCY_LIMIT_KEY, 2), + ok = application:set_env(fluxer_gateway, ?QUEUE_LIMIT_KEY, 0), Blockers = [blocking_push_worker() || _ <- lists:seq(1, 6)], Now = erlang:monotonic_time(millisecond), put(push_inflight_workers, [ diff --git a/fluxer_gateway/src/presence/presence_update.erl b/fluxer_gateway/src/presence/presence_update.erl index 231878eae..7936d0183 100644 --- a/fluxer_gateway/src/presence/presence_update.erl +++ b/fluxer_gateway/src/presence/presence_update.erl @@ -10,7 +10,8 @@ handle_message_create_event/2, handle_message_ack_event/2, flush_push_buffer/1, - maybe_update_push_eligibility/1 + maybe_update_push_eligibility/1, + push_buffer_counters/0 ]). -export_type([state/0]). @@ -25,6 +26,8 @@ -define(DEFAULT_PUSH_BUFFER_MAX_BYTES, 1048576). -define(PUSH_BUFFER_MAX_ENTRIES_CONFIG_KEY, presence_push_buffer_max_entries). -define(PUSH_BUFFER_MAX_BYTES_CONFIG_KEY, presence_push_buffer_max_bytes). +-define(PUSH_BUFFER_COUNTERS, presence_push_buffer_counters). +-define(PUSH_BUFFER_OVERFLOW, push_buffer_overflow). -spec maybe_handle_custom_status(map(), state()) -> {map(), state()}. maybe_handle_custom_status(Request, State) -> @@ -152,12 +155,45 @@ flush_push_buffer(#{push_buffer := Buffer} = State) -> -spec maybe_update_push_eligibility(state()) -> state(). maybe_update_push_eligibility(State) -> - Sessions = maps:get(sessions, State, #{}), - case {is_push_eligible(Sessions), maps:get(push_buffer, State, [])} of + update_push_eligibility(enrolled_in_push_delivery(State), State). + +-spec update_push_eligibility(boolean(), state()) -> state(). +update_push_eligibility(true, State) -> + Eligible = no_session_holds_push(maps:get(sessions, State, #{})), + flush_when_eligible(Eligible, record_push_eligibility(Eligible, State)); +update_push_eligibility(false, State) -> + flush_when_eligible(is_push_eligible(maps:get(sessions, State, #{})), State). + +-spec flush_when_eligible(boolean(), state()) -> state(). +flush_when_eligible(Eligible, State) -> + case {Eligible, maps:get(push_buffer, State, [])} of {true, [_ | _]} -> flush_push_buffer(State); _ -> State end. +-spec enrolled_in_push_delivery(state()) -> boolean(). +enrolled_in_push_delivery(State) -> + case maps:get(user_id, State, undefined) of + UserId when is_integer(UserId) -> push_delivery_config:is_enrolled(UserId); + _ -> false + end. + +-spec record_push_eligibility(boolean(), state()) -> state(). +record_push_eligibility(true, State) -> + State#{push_eligible => true}; +record_push_eligibility(false, State) -> + case maps:get(push_eligible, State, true) of + true -> announce_session_active(maps:get(user_id, State, undefined)); + false -> ok + end, + State#{push_eligible => false}. + +-spec announce_session_active(user_id() | undefined) -> ok. +announce_session_active(UserId) when is_integer(UserId) -> + push_outbox:note_session_active(UserId); +announce_session_active(_UserId) -> + ok. + -spec compare_and_validate(map(), map(), state()) -> {map(), state()}. compare_and_validate(CustomStatus, Request, State) -> PreviousCustomStatus = maps:get(custom_status, State, null), @@ -206,8 +242,7 @@ field_or_null(Map, Key) -> -spec route_push_notification(map(), state()) -> state(). route_push_notification(Params, State) -> - Sessions = maps:get(sessions, State, #{}), - case is_push_eligible(Sessions) of + case push_eligible(State) of true -> FlushedState = flush_push_buffer(State), push:handle_message_create(Params), @@ -216,6 +251,16 @@ route_push_notification(Params, State) -> buffer_push_notification(Params, State) end. +-spec push_eligible(state()) -> boolean(). +push_eligible(State) -> + push_eligible(enrolled_in_push_delivery(State), maps:get(sessions, State, #{})). + +-spec push_eligible(boolean(), map()) -> boolean(). +push_eligible(true, Sessions) -> + no_session_holds_push(Sessions); +push_eligible(false, Sessions) -> + is_push_eligible(Sessions). + -spec build_push_create_params(user_id(), map()) -> map() | undefined. build_push_create_params(UserId, Data) -> AuthorIdBin = maps:get(<<"id">>, maps:get(<<"author">>, Data, #{}), undefined), @@ -237,8 +282,45 @@ buffer_push_notification(Params, State) -> undefined -> State; Entry -> - Buffer = maps:get(push_buffer, State, []), - State#{push_buffer := cap_push_buffer([Entry | Buffer])} + Buffer = [Entry | maps:get(push_buffer, State, [])], + Capped = cap_push_buffer(Buffer), + ok = count_push_buffer_overflow( + maps:get(user_id, State, undefined), length(Buffer) - length(Capped) + ), + State#{push_buffer := Capped} + end. + +-spec count_push_buffer_overflow(term(), non_neg_integer()) -> ok. +count_push_buffer_overflow(_UserId, 0) -> + ok; +count_push_buffer_overflow(UserId, Dropped) -> + ok = guild_ets_utils:ensure_table(?PUSH_BUFFER_COUNTERS, [ + named_table, public, set, {write_concurrency, true} + ]), + try + _ = ets:update_counter( + ?PUSH_BUFFER_COUNTERS, + ?PUSH_BUFFER_OVERFLOW, + {2, Dropped}, + {?PUSH_BUFFER_OVERFLOW, 0} + ), + ok + catch + error:badarg -> ok + end, + logger:warning( + "presence_push_buffer_overflow: user_id=~p dropped=~p", [UserId, Dropped] + ). + +-spec push_buffer_counters() -> #{atom() => non_neg_integer()}. +push_buffer_counters() -> + try ets:lookup(?PUSH_BUFFER_COUNTERS, ?PUSH_BUFFER_OVERFLOW) of + [{?PUSH_BUFFER_OVERFLOW, Count}] when is_integer(Count), Count >= 0 -> + #{?PUSH_BUFFER_OVERFLOW => Count}; + _ -> + #{?PUSH_BUFFER_OVERFLOW => 0} + catch + error:badarg -> #{?PUSH_BUFFER_OVERFLOW => 0} end. -spec cap_push_buffer([push_buffer_entry()]) -> [push_buffer_entry()]. @@ -335,6 +417,14 @@ is_push_eligible(Sessions) -> all_sessions_afk(Sessions) -> lists:all(fun(S) -> maps:get(afk, S, false) end, maps:values(Sessions)). +-spec no_session_holds_push(map()) -> boolean(). +no_session_holds_push(Sessions) -> + not lists:any(fun session_holds_push/1, maps:values(Sessions)). + +-spec session_holds_push(map()) -> boolean(). +session_holds_push(Session) -> + not maps:get(afk, Session, false) andalso maps:get(status, Session, online) =/= offline. + -spec extract_snowflake(binary(), map()) -> integer() | undefined. extract_snowflake(FieldName, Data) -> parse_snowflake(FieldName, maps:get(FieldName, Data, undefined)). diff --git a/fluxer_gateway/src/push/push.erl b/fluxer_gateway/src/push/push.erl index eefc28a1b..f9967d44a 100644 --- a/fluxer_gateway/src/push/push.erl +++ b/fluxer_gateway/src/push/push.erl @@ -19,7 +19,7 @@ invalidate_user_badge_counts_local/1, clear_channel_notifications/3 ]). --export([get_cache_stats/0]). +-export([get_cache_stats/0, delivery_gate_counters/0]). -export([push_owner_key/1]). -define(EVICT_INTERVAL_MS, 60000). @@ -27,12 +27,10 @@ -define(PUSH_COUNTER_TABLE, push_worker_counter). -define(CNT_WORKER_POOL, push_loss_worker_pool). --define(CNT_WARM_INFLIGHT, push_blocked_ids_warm_inflight). -define(CNT_FETCH_ATTEMPTS, push_blocked_ids_fetch_attempts). -define(CNT_FETCH_FAILURES, push_blocked_ids_fetch_failures). -define(CNT_SUPPRESSED, push_blocked_ids_suppressed). -define(CNT_BUDGET_EXHAUSTED, push_blocked_ids_budget_exhausted). --define(CNT_WARM_DROPPED, push_blocked_ids_warm_dropped). -define(CNT_DISPATCH_DROPPED, push_loss_dispatch_dropped). -define(CNT_DISPATCH_DROPPED_USERS, push_loss_dispatch_dropped_users). -define(CNT_CLEAR_DROPPED, push_loss_clear_dropped). @@ -46,8 +44,11 @@ -define(CNT_RESTART_DISCARDED, push_dispatcher_restart_discarded). -define(CNT_QUEUE_ENQUEUED, push_dispatcher_queue_enqueued). -define(CNT_QUEUE_DEQUEUED, push_dispatcher_queue_dequeued). +-define(CNT_GATE_SERVICE_USERS, push_delivery_gate_service_users). +-define(CNT_GATE_GATEWAY_USERS, push_delivery_gate_gateway_users). +-define(CNT_GATE_JOBS_PUBLISHED, push_delivery_gate_jobs_published). +-define(CNT_GATE_PUBLISH_FAILED, push_delivery_gate_publish_failed). --define(MAX_WARM_INFLIGHT, 4). -define(MAX_FETCH_RPCS, 8). -define(DEFAULT_FETCH_USERS, 2000). -define(MAX_FETCH_USERS, 5000). @@ -125,7 +126,6 @@ handle_cast(_Msg, State) -> -spec handle_info(term(), state()) -> {noreply, state()}. handle_info(evict_caches, State) -> - unstick_warm_gate(), MaxEntries = maps:get(max_entries, State), push_ets_cache:evict_tables(#{ user_guild_settings => MaxEntries, @@ -229,20 +229,33 @@ is_push_noop() -> -spec clear_channel_notifications(integer(), integer(), integer()) -> ok. clear_channel_notifications(UserId, ChannelId, MessageId) -> - case is_push_active() andalso clear_notifications_enabled() of + case is_push_active() of true -> - cast_to_push_owner( - UserId, {clear_channel_notifications, UserId, ChannelId, MessageId} + clear_for_enrolment( + push_delivery_config:is_enrolled(UserId), UserId, ChannelId, MessageId ); false -> ok end. +-spec clear_for_enrolment(boolean(), integer(), integer(), integer()) -> ok. +clear_for_enrolment(true, UserId, ChannelId, MessageId) -> + ok = push_outbox:truncate_read(UserId, ChannelId, MessageId), + cast_clear_if_enabled(clear_notifications_enabled(), UserId, ChannelId, MessageId); +clear_for_enrolment(false, UserId, ChannelId, MessageId) -> + cast_clear_if_enabled(clear_notifications_enabled(), UserId, ChannelId, MessageId). + +-spec cast_clear_if_enabled(boolean(), integer(), integer(), integer()) -> ok. +cast_clear_if_enabled(true, UserId, ChannelId, MessageId) -> + cast_to_push_owner(UserId, {clear_channel_notifications, UserId, ChannelId, MessageId}); +cast_clear_if_enabled(false, _UserId, _ChannelId, _MessageId) -> + ok. + -spec clear_notifications_enabled() -> boolean(). clear_notifications_enabled() -> case persistent_term:get(push_clear_notifications_enabled, undefined) of - Value when is_boolean(Value) -> Value; - _ -> env_boolean(push_clear_notifications_enabled, false) + OperatorChoice when is_boolean(OperatorChoice) -> OperatorChoice; + _ -> env_boolean(push_clear_notifications_enabled, true) end. -spec get_cache_stats() -> {ok, map()}. @@ -342,8 +355,8 @@ do_handle_message_create_context(Context, State) -> channel_id => ChannelId, eligible_count => length(EligibleUsers) }), - dispatch_if_eligible( - EligibleUsers, + route_partitioned_users( + split_for_delivery(EligibleUsers), MessageData, MarkdownContext, GuildId, @@ -354,6 +367,147 @@ do_handle_message_create_context(Context, State) -> State ). +-spec split_for_delivery([integer()]) -> {non_neg_integer(), [integer()], [integer()]}. +split_for_delivery(UserIds) -> + Config = push_delivery_config:config(), + {ServiceUsers, GatewayUsers} = push_delivery_config:partition_users(Config, UserIds), + {maps:get(config_version, Config), ServiceUsers, GatewayUsers}. + +-spec route_partitioned_users( + {non_neg_integer(), [integer()], [integer()]}, + map(), + map(), + integer(), + integer(), + integer(), + binary() | undefined, + binary() | undefined, + worker_state() +) -> ok. +route_partitioned_users( + {ConfigVersion, ServiceUsers, GatewayUsers}, + MessageData, + MarkdownContext, + GuildId, + ChannelId, + MessageId, + GuildName, + ChannelName, + State +) -> + ok = route_service_users( + ServiceUsers, + ConfigVersion, + MessageData, + MarkdownContext, + GuildId, + ChannelId, + MessageId, + GuildName, + ChannelName, + State + ), + count_gateway_users(length(GatewayUsers)), + dispatch_if_eligible( + GatewayUsers, + MessageData, + MarkdownContext, + GuildId, + ChannelId, + MessageId, + GuildName, + ChannelName, + State + ). + +-spec route_service_users( + [integer()], + non_neg_integer(), + map(), + map(), + integer(), + integer(), + integer(), + binary() | undefined, + binary() | undefined, + worker_state() +) -> ok. +route_service_users( + [], + _ConfigVersion, + _MessageData, + _MarkdownContext, + _GuildId, + _ChannelId, + _MessageId, + _GuildName, + _ChannelName, + _State +) -> + ok; +route_service_users( + ServiceUsers, + ConfigVersion, + MessageData, + MarkdownContext, + GuildId, + ChannelId, + MessageId, + GuildName, + ChannelName, + State +) -> + Fallback = fun(RemainingUsers) -> + dispatch_if_eligible( + RemainingUsers, + MessageData, + MarkdownContext, + GuildId, + ChannelId, + MessageId, + GuildName, + ChannelName, + State + ) + end, + case + push_job_publisher:publish_message( + ServiceUsers, + MessageData, + MarkdownContext, + GuildId, + ChannelId, + MessageId, + GuildName, + ChannelName, + ConfigVersion, + Fallback + ) + of + ok -> + bump_counter(?CNT_GATE_SERVICE_USERS, length(ServiceUsers)), + bump_counter(?CNT_GATE_JOBS_PUBLISHED); + {error, _Reason} -> + bump_counter(?CNT_GATE_PUBLISH_FAILED), + dispatch_if_eligible( + ServiceUsers, + MessageData, + MarkdownContext, + GuildId, + ChannelId, + MessageId, + GuildName, + ChannelName, + State + ) + end. + +-spec count_gateway_users(non_neg_integer()) -> ok. +count_gateway_users(0) -> + ok; +count_gateway_users(Count) -> + bump_counter(?CNT_GATE_GATEWAY_USERS, Count). + -spec filter_eligible_users( [integer()], integer(), @@ -377,7 +531,7 @@ filter_eligible_users( SuppliedMetadata ) -> LargeGuildMetadata = resolve_large_guild_metadata(GuildId, SuppliedMetadata), - Candidates = drop_blocked_recipients(UserIds, AuthorId), + Candidates = drop_blocked_recipients(UserIds, AuthorId, #{}), push_eligibility:prefetch_user_guild_settings(Candidates, AuthorId, GuildId), EligibleUsers = lists:filter( fun(UserId) -> @@ -395,8 +549,7 @@ filter_eligible_users( end, Candidates ), - warm_blocked_ids(EligibleUsers), - EligibleUsers. + drop_blocked_recipients(EligibleUsers, AuthorId, fetch_missing_blocked_ids(EligibleUsers)). -spec resolve_large_guild_metadata(integer(), map() | undefined) -> map() | undefined. resolve_large_guild_metadata(0, _SuppliedMetadata) -> @@ -416,72 +569,61 @@ supplied_or_local_metadata(GuildId, _SuppliedMetadata) -> large_guild_metadata_local(GuildId) -> push_eligibility_checks:get_guild_large_metadata(GuildId). --spec drop_blocked_recipients([integer()], integer()) -> [integer()]. -drop_blocked_recipients(UserIds, AuthorId) -> +-spec drop_blocked_recipients([integer()], integer(), #{integer() => [integer()]}) -> + [integer()]. +drop_blocked_recipients(UserIds, AuthorId, Fetched) -> Kept = lists:filter( - fun(UserId) -> not push_eligibility:is_user_blocked(UserId, AuthorId) end, + fun(UserId) -> not is_blocked_recipient(UserId, AuthorId, Fetched) end, UserIds ), count_suppressed(length(UserIds) - length(Kept)), Kept. +-spec is_blocked_recipient(integer(), integer(), #{integer() => [integer()]}) -> boolean(). +is_blocked_recipient(UserId, AuthorId, Fetched) -> + lists:member(AuthorId, maps:get(UserId, Fetched, [])) orelse + push_eligibility:is_user_blocked(UserId, AuthorId). + -spec count_suppressed(non_neg_integer()) -> ok. count_suppressed(0) -> ok; count_suppressed(Suppressed) -> bump_counter(?CNT_SUPPRESSED, Suppressed). --spec warm_blocked_ids([integer()]) -> ok. -warm_blocked_ids([]) -> - ok; -warm_blocked_ids(EligibleUsers) -> - case missing_blocked_ids(EligibleUsers) of - [] -> ok; - Missing -> spawn_blocked_ids_warm(Missing) +-spec fetch_missing_blocked_ids([integer()]) -> #{integer() => [integer()]}. +fetch_missing_blocked_ids(UserIds) -> + case missing_blocked_ids(UserIds) of + [] -> #{}; + Missing -> fetch_blocked_ids_within_budget(Missing) end. -spec missing_blocked_ids([integer()]) -> [integer()]. -missing_blocked_ids(EligibleUsers) -> - lists:usort(lists:filter(fun is_blocked_ids_cache_miss/1, EligibleUsers)). +missing_blocked_ids(UserIds) -> + lists:usort(lists:filter(fun is_blocked_ids_cache_miss/1, UserIds)). -spec is_blocked_ids_cache_miss(integer()) -> boolean(). is_blocked_ids_cache_miss(UserId) -> push_ets_cache:get_blocked_ids(UserId) =:= undefined. --spec spawn_blocked_ids_warm([integer()]) -> ok. -spawn_blocked_ids_warm(Missing) -> - case claim_warm_slot() of - ok -> - _ = spawn(fun() -> run_blocked_ids_warm(Missing) end), - ok; - full -> - bump_counter(?CNT_WARM_DROPPED) - end. - --spec run_blocked_ids_warm([integer()]) -> ok. -run_blocked_ids_warm(Missing) -> - try - warm_within_budget(Missing) - catch - throw:Reason -> warm_crashed(throw, Reason); - error:Reason -> warm_crashed(error, Reason); - exit:Reason -> warm_crashed(exit, Reason) - after - release_warm_slot() - end. - --spec warm_crashed(throw | error | exit, term()) -> ok. -warm_crashed(Class, Reason) -> - bump_counter(?CNT_FETCH_FAILURES), - logger:debug("Push: blocked id warm crashed", #{class => Class, reason => Reason}), - ok. - --spec warm_within_budget([integer()]) -> ok. -warm_within_budget(Missing) -> +-spec fetch_blocked_ids_within_budget([integer()]) -> #{integer() => [integer()]}. +fetch_blocked_ids_within_budget(Missing) -> {Budget, ChunkSize} = blocked_ids_fetch_budget(), Budgeted = lists:sublist(Missing, Budget), count_budget_exhausted(length(Missing) - length(Budgeted)), - fetch_blocked_ids_chunks(chunk_user_ids(Budgeted, ChunkSize, [])). + Fill = push_ets_cache:reserve_blocked_ids(Budgeted), + try + fetch_blocked_ids_chunks(chunk_user_ids(Budgeted, ChunkSize, []), Fill, #{}) + catch + Class:Reason -> blocked_ids_fetch_crashed(Class, Reason) + after + push_ets_cache:release(Fill) + end. + +-spec blocked_ids_fetch_crashed(throw | error | exit, term()) -> #{}. +blocked_ids_fetch_crashed(Class, Reason) -> + bump_counter(?CNT_FETCH_FAILURES), + logger:debug("Push: blocked id fetch crashed", #{class => Class, reason => Reason}), + #{}. -spec blocked_ids_fetch_budget() -> {pos_integer(), pos_integer()}. blocked_ids_fetch_budget() -> @@ -495,13 +637,17 @@ count_budget_exhausted(0) -> count_budget_exhausted(Skipped) -> bump_counter(?CNT_BUDGET_EXHAUSTED, Skipped). --spec fetch_blocked_ids_chunks([[integer()]]) -> ok. -fetch_blocked_ids_chunks([]) -> - ok; -fetch_blocked_ids_chunks([Chunk | Rest]) -> - case fetch_blocked_ids_chunk(Chunk) of - ok -> fetch_blocked_ids_chunks(Rest); - error -> ok +-spec fetch_blocked_ids_chunks( + [[integer()]], push_ets_cache:fill(), #{integer() => [integer()]} +) -> #{integer() => [integer()]}. +fetch_blocked_ids_chunks([], _Fill, Fetched) -> + Fetched; +fetch_blocked_ids_chunks([Chunk | Rest], Fill, Fetched) -> + case fetch_blocked_ids_chunk(Chunk, Fill) of + {ok, ChunkFetched} -> + fetch_blocked_ids_chunks(Rest, Fill, maps:merge(Fetched, ChunkFetched)); + error -> + Fetched end. -spec chunk_user_ids([integer()], pos_integer(), [[integer()]]) -> [[integer()]]. @@ -519,10 +665,11 @@ drop_prefix([], _N) -> drop_prefix([_UserId | Rest], N) -> drop_prefix(Rest, N - 1). --spec fetch_blocked_ids_chunk([integer()]) -> ok | error. -fetch_blocked_ids_chunk([]) -> - ok; -fetch_blocked_ids_chunk(UserIds) -> +-spec fetch_blocked_ids_chunk([integer()], push_ets_cache:fill()) -> + {ok, #{integer() => [integer()]}} | error. +fetch_blocked_ids_chunk([], _Fill) -> + {ok, #{}}; +fetch_blocked_ids_chunk(UserIds, Fill) -> bump_counter(?CNT_FETCH_ATTEMPTS), Request = #{ <<"type">> => <<"get_user_blocked_ids">>, @@ -530,7 +677,7 @@ fetch_blocked_ids_chunk(UserIds) -> }, case rpc_client:call(Request) of {ok, Data} -> - cache_blocked_ids_response(UserIds, Data); + {ok, cache_blocked_ids_response(UserIds, Data, Fill)}; {error, Reason} -> fetch_blocked_ids_failed(Reason, length(UserIds)) end. @@ -543,18 +690,16 @@ fetch_blocked_ids_failed(Reason, UserCount) -> }), error. --spec cache_blocked_ids_response([integer()], map()) -> ok. -cache_blocked_ids_response(UserIds, Data) -> - lists:foreach( - fun(UserId) -> cache_blocked_ids_entry(UserId, Data) end, - UserIds - ), - ok. +-spec cache_blocked_ids_response([integer()], map(), push_ets_cache:fill()) -> + #{integer() => [integer()]}. +cache_blocked_ids_response(UserIds, Data, Fill) -> + maps:from_list([{UserId, cache_blocked_ids_entry(UserId, Data, Fill)} || UserId <- UserIds]). --spec cache_blocked_ids_entry(integer(), map()) -> ok. -cache_blocked_ids_entry(UserId, Data) -> - Raw = maps:get(integer_to_binary(UserId), Data, []), - push_ets_cache:put_blocked_ids_fetched(UserId, blocked_ids_from_response(Raw)). +-spec cache_blocked_ids_entry(integer(), map(), push_ets_cache:fill()) -> [integer()]. +cache_blocked_ids_entry(UserId, Data, Fill) -> + BlockedIds = blocked_ids_from_response(maps:get(integer_to_binary(UserId), Data, [])), + ok = push_ets_cache:put_blocked_ids_fetched(UserId, BlockedIds, Fill), + BlockedIds. -spec blocked_ids_from_response(term()) -> [integer()]. blocked_ids_from_response(Values) when is_list(Values) -> @@ -585,55 +730,9 @@ blocked_ids_counters() -> blocked_ids_fetch_attempts => read_counter(?CNT_FETCH_ATTEMPTS), blocked_ids_fetch_failures => read_counter(?CNT_FETCH_FAILURES), blocked_ids_suppressed => read_counter(?CNT_SUPPRESSED), - blocked_ids_budget_exhausted => read_counter(?CNT_BUDGET_EXHAUSTED), - blocked_ids_warm_dropped => read_counter(?CNT_WARM_DROPPED), - blocked_ids_warm_inflight => read_counter(?CNT_WARM_INFLIGHT) + blocked_ids_budget_exhausted => read_counter(?CNT_BUDGET_EXHAUSTED) }. --spec claim_warm_slot() -> ok | full. -claim_warm_slot() -> - try ets:update_counter(?PUSH_COUNTER_TABLE, ?CNT_WARM_INFLIGHT, {2, 1}) of - Value when is_integer(Value), Value =< ?MAX_WARM_INFLIGHT -> - ok; - _Value -> - release_warm_slot(), - full - catch - error:badarg -> claim_first_warm_slot() - end. - --spec claim_first_warm_slot() -> ok | full. -claim_first_warm_slot() -> - try ets:insert_new(?PUSH_COUNTER_TABLE, {?CNT_WARM_INFLIGHT, 1}) of - true -> ok; - false -> full - catch - error:badarg -> full - end. - --spec release_warm_slot() -> ok. -release_warm_slot() -> - try ets:update_counter(?PUSH_COUNTER_TABLE, ?CNT_WARM_INFLIGHT, {2, -1, 0, 0}) of - _Value -> ok - catch - error:badarg -> ok - end. - --spec unstick_warm_gate() -> ok. -unstick_warm_gate() -> - case read_counter(?CNT_WARM_INFLIGHT) of - Value when is_integer(Value), Value >= ?MAX_WARM_INFLIGHT -> reset_warm_gate(); - _Value -> ok - end. - --spec reset_warm_gate() -> ok. -reset_warm_gate() -> - try ets:insert(?PUSH_COUNTER_TABLE, {?CNT_WARM_INFLIGHT, 0}) of - _Value -> ok - catch - error:badarg -> ok - end. - -spec dispatch_if_eligible( [integer()], map(), @@ -732,6 +831,30 @@ handle_message_create_cast(Params, State) -> -spec handle_clear_channel_notifications(integer(), integer(), integer(), state()) -> {noreply, state()}. handle_clear_channel_notifications(UserId, ChannelId, MessageId, State) -> + case split_for_delivery([UserId]) of + {ConfigVersion, [UserId], []} -> + clear_via_service(UserId, ChannelId, MessageId, ConfigVersion, State); + {_ConfigVersion, _ServiceUsers, _GatewayUsers} -> + clear_via_dispatcher(UserId, ChannelId, MessageId, State) + end. + +-spec clear_via_service(integer(), integer(), integer(), non_neg_integer(), state()) -> + {noreply, state()}. +clear_via_service(UserId, ChannelId, MessageId, ConfigVersion, State) -> + Fallback = fun(_UserIds) -> clear_via_dispatcher(UserId, ChannelId, MessageId, State) end, + case + push_job_publisher:publish_clear(UserId, ChannelId, MessageId, ConfigVersion, Fallback) + of + ok -> + bump_counter(?CNT_GATE_JOBS_PUBLISHED), + {noreply, State}; + {error, _Reason} -> + bump_counter(?CNT_GATE_PUBLISH_FAILED), + clear_via_dispatcher(UserId, ChannelId, MessageId, State) + end. + +-spec clear_via_dispatcher(integer(), integer(), integer(), state()) -> {noreply, state()}. +clear_via_dispatcher(UserId, ChannelId, MessageId, State) -> BadgeCountsTtl = maps:get(badge_counts_ttl_seconds, State), case push_dispatcher:enqueue_clear_notifications( @@ -781,7 +904,31 @@ log_worker_pool_drop(false, MessageId, ChannelId) -> -spec cache_stats_with_counters() -> map(). cache_stats_with_counters() -> Base = maps:merge(push_ets_cache:cache_stats(), blocked_ids_counters()), - maps:merge(Base, push_loss_counters()). + WithLossCounters = maps:merge(Base, push_loss_counters()), + maps:merge(WithLossCounters, delivery_gate_counters()). + +-spec delivery_gate_counters() -> map(). +delivery_gate_counters() -> + case counter_table_status() of + live -> live_delivery_gate_counters(); + unavailable -> #{} + end. + +-spec live_delivery_gate_counters() -> map(). +live_delivery_gate_counters() -> + #{ + delivery_gate_service_users => counter_or_zero(?CNT_GATE_SERVICE_USERS), + delivery_gate_gateway_users => counter_or_zero(?CNT_GATE_GATEWAY_USERS), + delivery_gate_jobs_published => counter_or_zero(?CNT_GATE_JOBS_PUBLISHED), + delivery_gate_publish_failed => counter_or_zero(?CNT_GATE_PUBLISH_FAILED) + }. + +-spec counter_or_zero(atom()) -> non_neg_integer(). +counter_or_zero(Key) -> + case read_counter(Key) of + Value when is_integer(Value) -> Value; + unavailable -> 0 + end. -spec push_loss_counters() -> map(). push_loss_counters() -> @@ -981,8 +1128,10 @@ sync_user_blocked_ids_local_updates_local_cache_test() -> invalidate_user_badge_counts_local_deletes_every_cached_entry_test() -> push_ets_cache:init(), - push_ets_cache:put_badge_count(10, 5, 1000), - push_ets_cache:put_badge_count(11, 7, 1000), + ok = seed_badge_count(10, 5, 1000), + ok = seed_badge_count(11, 7, 1000), + ?assertEqual({5, 1000}, push_ets_cache:get_badge_count(10)), + ?assertEqual({7, 1000}, push_ets_cache:get_badge_count(11)), with_registered_push(fun() -> ok = invalidate_user_badge_counts_local([10, 11]) end), @@ -991,7 +1140,7 @@ invalidate_user_badge_counts_local_deletes_every_cached_entry_test() -> invalidate_user_badge_counts_local_ignores_untyped_ids_test() -> push_ets_cache:init(), - push_ets_cache:put_badge_count(12, 5, 1000), + ok = seed_badge_count(12, 5, 1000), with_registered_push(fun() -> ok = invalidate_user_badge_counts_local([<<"12">>]) end), @@ -1000,7 +1149,11 @@ invalidate_user_badge_counts_local_ignores_untyped_ids_test() -> invalidate_user_subscriptions_local_deletes_local_cache_test() -> push_ets_cache:init(), - push_ets_cache:put_subscriptions(10, [#{<<"endpoint">> => <<"test">>}]), + Subscriptions = [#{<<"endpoint">> => <<"test">>}], + ok = push_ets_cache:put_subscriptions( + 10, Subscriptions, push_ets_cache:reserve_subscriptions([10]) + ), + ?assertEqual(Subscriptions, push_ets_cache:get_subscriptions(10)), with_registered_push(fun() -> ok = invalidate_user_subscriptions_local(10), ?assertEqual(undefined, push_ets_cache:get_subscriptions(10)) @@ -1072,7 +1225,7 @@ synced_blocked_recipients_are_dropped_test() -> fetched_blocked_recipients_are_dropped_test() -> push_ets_cache:init(), - ok = push_ets_cache:put_blocked_ids_fetched(5031, [999]), + ok = seed_fetched_blocked_ids(5031, [999]), ?assertEqual([], filter_dm_recipients([5031], 999)). block_suppression_drops_blocked_recipient_test() -> @@ -1081,50 +1234,21 @@ block_suppression_drops_blocked_recipient_test() -> ok = push_ets_cache:put_blocked_ids(5003, []), ?assertEqual([5003], filter_dm_recipients([5002, 5003], 999)). -blocked_ids_fetch_stops_after_a_failing_chunk_test() -> - ?assertEqual(ok, fetch_blocked_ids_chunks([])), - ?assertEqual(ok, fetch_blocked_ids_chunks([[], []])). - -cached_recipients_need_no_blocked_ids_warm_test() -> +cached_recipients_need_no_blocked_ids_fetch_test() -> push_ets_cache:init(), - ok = push_ets_cache:put_blocked_ids_fetched(5004, []), - ?assertEqual([], missing_blocked_ids([5004])), - ?assertEqual(ok, warm_blocked_ids([5004])). + ok = seed_fetched_blocked_ids(5004, []), + ?assertEqual([], missing_blocked_ids([5004])). + +blocked_ids_fetch_stops_after_a_failing_chunk_test() -> + push_ets_cache:init(), + Fill = push_ets_cache:reserve_blocked_ids([]), + ?assertEqual(#{}, fetch_blocked_ids_chunks([], Fill, #{})), + ?assertEqual(#{}, fetch_blocked_ids_chunks([[], []], Fill, #{})), + ok = push_ets_cache:release(Fill). blocked_ids_counters_are_unavailable_without_the_shared_table_test() -> delete_counter_table(), - ?assertEqual(unavailable, read_counter(?CNT_SUPPRESSED)), - ?assertEqual(full, claim_warm_slot()), - ?assertEqual(ok, release_warm_slot()), - ?assertEqual(ok, unstick_warm_gate()). - -blocked_ids_warm_gate_bounds_concurrency_test() -> - with_counter_table(fun() -> - Claims = [claim_warm_slot() || _ <- lists:seq(1, ?MAX_WARM_INFLIGHT + 2)], - ?assertEqual(?MAX_WARM_INFLIGHT, length([ok || ok <- Claims])), - ?assertEqual(?MAX_WARM_INFLIGHT, read_counter(?CNT_WARM_INFLIGHT)), - lists:foreach( - fun(_) -> release_warm_slot() end, lists:seq(1, ?MAX_WARM_INFLIGHT + 2) - ), - ?assertEqual(0, read_counter(?CNT_WARM_INFLIGHT)) - end). - -blocked_ids_warm_gate_unsticks_in_the_eviction_sweep_test() -> - with_counter_table(fun() -> - lists:foreach(fun(_) -> claim_warm_slot() end, lists:seq(1, ?MAX_WARM_INFLIGHT)), - ?assertEqual(full, claim_warm_slot()), - unstick_warm_gate(), - ?assertEqual(0, read_counter(?CNT_WARM_INFLIGHT)), - ?assertEqual(ok, claim_warm_slot()) - end). - -blocked_ids_warm_is_dropped_when_the_gate_is_full_test() -> - push_ets_cache:init(), - with_counter_table(fun() -> - lists:foreach(fun(_) -> claim_warm_slot() end, lists:seq(1, ?MAX_WARM_INFLIGHT)), - ?assertEqual(ok, warm_blocked_ids([5099])), - ?assertEqual(1, read_counter(?CNT_WARM_DROPPED)) - end). + ?assertEqual(unavailable, read_counter(?CNT_SUPPRESSED)). blocked_ids_counters_are_exposed_in_cache_stats_test() -> push_ets_cache:init(), @@ -1144,9 +1268,7 @@ assert_cache_stats_expose_blocked_ids() -> blocked_ids_fetch_attempts, blocked_ids_fetch_failures, blocked_ids_suppressed, - blocked_ids_budget_exhausted, - blocked_ids_warm_dropped, - blocked_ids_warm_inflight + blocked_ids_budget_exhausted ] ). @@ -1231,6 +1353,16 @@ delete_counter_table() -> error:badarg -> ok end. +seed_badge_count(UserId, Count, CachedAt) -> + push_ets_cache:put_badge_count( + UserId, Count, CachedAt, push_ets_cache:reserve_badge_counts([UserId]) + ). + +seed_fetched_blocked_ids(UserId, BlockedIds) -> + push_ets_cache:put_blocked_ids_fetched( + UserId, BlockedIds, push_ets_cache:reserve_blocked_ids([UserId]) + ). + filter_dm_recipients(UserIds, AuthorId) -> MessageData = #{<<"channel_type">> => 1}, filter_eligible_users(UserIds, AuthorId, 0, 10, MessageData, 0, #{}, #{}, undefined). diff --git a/fluxer_gateway/src/push/push_delivery_config.erl b/fluxer_gateway/src/push/push_delivery_config.erl new file mode 100644 index 000000000..fdc03b2fc --- /dev/null +++ b/fluxer_gateway/src/push/push_delivery_config.erl @@ -0,0 +1,437 @@ +%% SPDX-License-Identifier: AGPL-3.0-or-later + +-module(push_delivery_config). +-typing([eqwalizer]). +-behaviour(gen_server). + +-export([ + start_link/0, + config/0, + config_version/0, + update_counts/0, + partition_users/1, + partition_users/2, + is_enrolled/1, + bucket/2 +]). +-export([init/1, handle_call/3, handle_cast/2, handle_info/2, terminate/2, code_change/3]). + +-define(PERSISTENT_TERM_KEY, push_delivery_config). +-define(UPDATE_COUNTS_KEY, {push_delivery_config, update_counts}). +-define(NATS_SUBJECT, <<"config.push.delivery">>). +-define(FETCH_DELAY_MS, 2000). +-define(RECONCILE_INTERVAL_MS, 30000). +-define(NATS_SUBSCRIBE_RETRY_MS, 2000). +-define(RESOLUTION, 10000). +-define(MAX_TARGETED_USERS, 1000). +-define(MAX_SALT_BYTES, 64). +-define(MAX_USER_ID_DIGITS, 20). +-define(DEFAULT_SALT, <<"push-service-delivery-v1">>). +-define(FNV_OFFSET_BASIS_32, 16#811c9dc5). +-define(FNV_PRIME_32, 16#01000193). + +-type user_id_set() :: #{binary() => true}. +-type config() :: #{ + enabled := boolean(), + config_version := non_neg_integer(), + rollout_basis_points := non_neg_integer(), + rollout_salt := binary(), + included := user_id_set(), + excluded := user_id_set() +}. +-type state() :: #{ + nats_subscription := term(), + nats_monitor := reference() | undefined, + fetch_timer := reference() | undefined +}. +-type store_result() :: updated | unchanged | stale | rejected. + +-export_type([config/0]). + +-spec start_link() -> gen_server:start_ret(). +start_link() -> + gen_server:start_link({local, ?MODULE}, ?MODULE, [], []). + +-spec config() -> config(). +config() -> + case persistent_term:get(?PERSISTENT_TERM_KEY, undefined) of + Config when is_map(Config) -> Config; + _ -> default_config() + end. + +-spec config_version() -> non_neg_integer(). +config_version() -> + maps:get(config_version, config(), 0). + +-spec update_counts() -> #{store_result() => non_neg_integer()}. +update_counts() -> + Counters = persistent_term:get(?UPDATE_COUNTS_KEY, undefined), + maps:from_list([{Result, read_count(Counters, Result)} || Result <- store_results()]). + +-spec partition_users([integer()]) -> {[integer()], [integer()]}. +partition_users(UserIds) -> + partition_users(config(), UserIds). + +-spec partition_users(config(), [integer()]) -> {[integer()], [integer()]}. +partition_users(Config, UserIds) -> + partition_enrolled(maps:get(enabled, Config, false), Config, UserIds). + +-spec partition_enrolled(boolean(), config(), [integer()]) -> {[integer()], [integer()]}. +partition_enrolled(true, Config, UserIds) -> + lists:partition(fun(UserId) -> enrolled(Config, integer_to_binary(UserId)) end, UserIds); +partition_enrolled(_Enabled, _Config, UserIds) -> + {[], UserIds}. + +-spec is_enrolled(integer()) -> boolean(). +is_enrolled(UserId) -> + Config = config(), + is_enrolled(maps:get(enabled, Config, false), Config, UserId). + +-spec is_enrolled(boolean(), config(), integer()) -> boolean(). +is_enrolled(true, Config, UserId) -> + enrolled(Config, integer_to_binary(UserId)); +is_enrolled(_Enabled, _Config, _UserId) -> + false. + +-spec bucket(binary(), binary()) -> non_neg_integer(). +bucket(UserId, Salt) -> + hash(<>, ?FNV_OFFSET_BASIS_32) rem ?RESOLUTION. + +-spec enrolled(config(), binary()) -> boolean(). +enrolled(Config, UserId) -> + case maps:is_key(UserId, maps:get(excluded, Config, #{})) of + true -> + false; + false -> + maps:is_key(UserId, maps:get(included, Config, #{})) orelse + bucket(UserId, maps:get(rollout_salt, Config, ?DEFAULT_SALT)) < + maps:get(rollout_basis_points, Config, 0) + end. + +-spec hash(binary(), non_neg_integer()) -> non_neg_integer(). +hash(<<>>, Hash) -> + Hash; +hash(<>, Hash) -> + hash(Rest, ((Hash bxor Byte) * ?FNV_PRIME_32) band 16#ffffffff). + +-spec init([]) -> {ok, state()}. +init([]) -> + erlang:process_flag(fullsweep_after, 50), + persistent_term:put(?PERSISTENT_TERM_KEY, config()), + self() ! subscribe_nats, + {ok, #{ + nats_subscription => undefined, + nats_monitor => undefined, + fetch_timer => erlang:send_after(?FETCH_DELAY_MS, self(), fetch_config) + }}. + +-spec handle_call(term(), gen_server:from(), state()) -> {reply, term(), state()}. +handle_call(_Request, _From, State) -> + {reply, ok, State}. + +-spec handle_cast(term(), state()) -> {noreply, state()}. +handle_cast(_Msg, State) -> + {noreply, State}. + +-spec handle_info(term(), state()) -> {noreply, state()}. +handle_info(subscribe_nats, State) -> + {noreply, subscribe_to_nats(State)}; +handle_info(fetch_config, State) -> + count(fetch_config_from_api()), + {noreply, State#{ + fetch_timer => erlang:send_after(?RECONCILE_INTERVAL_MS, self(), fetch_config) + }}; +handle_info({nats_resubscribed, ?NATS_SUBJECT}, State) -> + count(fetch_config_from_api()), + {noreply, State}; +handle_info({nats_msg, ?NATS_SUBJECT, Payload, _ReplyTo}, State) when is_binary(Payload) -> + count(apply_nats_payload(Payload)), + {noreply, State}; +handle_info({'DOWN', MonRef, process, _Pid, _Reason}, #{nats_monitor := MonRef} = State) -> + erlang:send_after(?NATS_SUBSCRIBE_RETRY_MS, self(), subscribe_nats), + {noreply, State#{nats_subscription => undefined, nats_monitor => undefined}}; +handle_info(_Info, State) -> + {noreply, State}. + +-spec terminate(term(), state()) -> ok. +terminate(_Reason, _State) -> + ok. + +-spec code_change(term(), state(), term()) -> {ok, state()}. +code_change(_OldVsn, State, _Extra) -> + erlang:garbage_collect(), + {ok, State}. + +-spec default_config() -> config(). +default_config() -> + #{ + enabled => false, + config_version => 0, + rollout_basis_points => 0, + rollout_salt => ?DEFAULT_SALT, + included => #{}, + excluded => #{} + }. + +-spec fetch_config_from_api() -> store_result(). +fetch_config_from_api() -> + RpcRequest = #{<<"type">> => <<"get_push_service_delivery_config">>}, + case api_rpc_client:call(RpcRequest) of + {ok, #{<<"config">> := Config}} when is_map(Config) -> + store_valid_config(Config, api); + {ok, _Other} -> + logger:warning("Push delivery config: unexpected API response format"), + rejected; + {error, Reason} -> + logger:warning("Push delivery config failed to fetch from API", #{reason => Reason}), + rejected + end. + +-spec apply_nats_payload(binary()) -> store_result(). +apply_nats_payload(Payload) -> + try json:decode(Payload) of + #{<<"type">> := <<"push_service_delivery_config">>, <<"config">> := Config} when + is_map(Config) + -> + store_valid_config(Config, nats); + #{<<"config">> := Config} when is_map(Config) -> + store_valid_config(Config, nats); + Config when is_map(Config) -> + store_valid_config(Config, nats); + _Other -> + logger:warning("Push delivery config: unexpected NATS payload format"), + rejected + catch + Class:Reason -> + logger:warning("Push delivery config failed to decode NATS payload", #{ + class => Class, reason => Reason + }), + rejected + end. + +-spec store_valid_config(map(), api | nats) -> store_result(). +store_valid_config(Config, Source) -> + case validate_config(Config) of + {ok, Validated} -> + store_validated_config(config(), Validated, Source); + {error, Reason} -> + logger:warning("Push delivery config rejected invalid config: ~p", [Reason]), + rejected + end. + +-spec store_validated_config(config(), config(), api | nats) -> store_result(). +store_validated_config(Previous, Previous, _Source) -> + unchanged; +store_validated_config( + #{config_version := Held}, #{config_version := Offered, enabled := true}, Source +) when + Offered < Held +-> + logger:warning("Push delivery config ignored a lower config_version", #{ + source => Source, held => Held, offered => Offered + }), + stale; +store_validated_config(Previous, Current, Source) -> + log_config_transitions(Previous, Current), + persistent_term:put(?PERSISTENT_TERM_KEY, Current), + logger:info("Push delivery config updated", #{source => Source}), + ok = push_outbox:delivery_config_changed(), + updated. + +-spec validate_config(map()) -> {ok, config()} | {error, term()}. +validate_config(Wire) -> + lists:foldl( + fun(Field, Acc) -> validate_field(Field, Wire, Acc) end, + {ok, default_config()}, + config_fields() + ). + +-spec config_fields() -> [{atom(), binary(), fun((term()) -> {ok, term()} | error)}]. +config_fields() -> + [ + {enabled, <<"enabled">>, fun validate_enabled/1}, + {config_version, <<"config_version">>, fun validate_config_version/1}, + {rollout_basis_points, <<"rollout_basis_points">>, fun validate_basis_points/1}, + {rollout_salt, <<"rollout_salt">>, fun validate_salt/1}, + {included, <<"included_user_ids">>, fun validate_user_ids/1}, + {excluded, <<"excluded_user_ids">>, fun validate_user_ids/1} + ]. + +-spec validate_field( + {atom(), binary(), fun((term()) -> {ok, term()} | error)}, + map(), + {ok, config()} | {error, term()} +) -> {ok, config()} | {error, term()}. +validate_field(_Field, _Wire, {error, _} = Error) -> + Error; +validate_field({Key, WireKey, Validate}, Wire, {ok, Acc}) -> + case maps:find(WireKey, Wire) of + error -> {ok, Acc}; + {ok, Value} -> store_validated_field(Key, WireKey, Value, Validate(Value), Acc) + end. + +-spec store_validated_field(atom(), binary(), term(), {ok, term()} | error, config()) -> + {ok, config()} | {error, term()}. +store_validated_field(Key, _WireKey, _Value, {ok, Normalised}, Acc) -> + {ok, Acc#{Key => Normalised}}; +store_validated_field(_Key, WireKey, Value, error, _Acc) -> + {error, {invalid_field, WireKey, Value}}. + +-spec validate_enabled(term()) -> {ok, boolean()} | error. +validate_enabled(Value) when is_boolean(Value) -> + {ok, Value}; +validate_enabled(_Value) -> + error. + +-spec validate_config_version(term()) -> {ok, non_neg_integer()} | error. +validate_config_version(Value) when is_integer(Value), Value >= 0 -> + {ok, Value}; +validate_config_version(_Value) -> + error. + +-spec validate_basis_points(term()) -> {ok, non_neg_integer()} | error. +validate_basis_points(Value) when is_integer(Value), Value >= 0, Value =< ?RESOLUTION -> + {ok, Value}; +validate_basis_points(_Value) -> + error. + +-spec validate_salt(term()) -> {ok, binary()} | error. +validate_salt(Value) when is_binary(Value) -> + validate_salt_bytes(Value, byte_size(Value)); +validate_salt(_Value) -> + error. + +-spec validate_salt_bytes(binary(), non_neg_integer()) -> {ok, binary()} | error. +validate_salt_bytes(Value, Size) when Size >= 1, Size =< ?MAX_SALT_BYTES -> + validate_printable_ascii(Value, Value); +validate_salt_bytes(_Value, _Size) -> + error. + +-spec validate_printable_ascii(binary(), binary()) -> {ok, binary()} | error. +validate_printable_ascii(<<>>, Value) -> + {ok, Value}; +validate_printable_ascii(<>, Value) when Byte >= 16#20, Byte =< 16#7e -> + validate_printable_ascii(Rest, Value); +validate_printable_ascii(_Remaining, _Value) -> + error. + +-spec validate_user_ids(term()) -> {ok, user_id_set()} | error. +validate_user_ids(Value) when is_list(Value), length(Value) =< ?MAX_TARGETED_USERS -> + collect_user_ids(Value, #{}); +validate_user_ids(_Value) -> + error. + +-spec collect_user_ids([term()], user_id_set()) -> {ok, user_id_set()} | error. +collect_user_ids([], Acc) -> + {ok, Acc}; +collect_user_ids([UserId | Rest], Acc) when is_binary(UserId) -> + collect_valid_user_id(UserId, is_user_id(UserId, byte_size(UserId)), Rest, Acc); +collect_user_ids(_UserIds, _Acc) -> + error. + +-spec collect_valid_user_id(binary(), boolean(), [term()], user_id_set()) -> + {ok, user_id_set()} | error. +collect_valid_user_id(UserId, true, Rest, Acc) -> + collect_user_ids(Rest, Acc#{UserId => true}); +collect_valid_user_id(_UserId, false, _Rest, _Acc) -> + error. + +-spec is_user_id(binary(), non_neg_integer()) -> boolean(). +is_user_id(UserId, Size) when Size >= 1, Size =< ?MAX_USER_ID_DIGITS -> + is_all_digits(UserId); +is_user_id(_UserId, _Size) -> + false. + +-spec is_all_digits(binary()) -> boolean(). +is_all_digits(<<>>) -> + true; +is_all_digits(<>) when Byte >= $0, Byte =< $9 -> + is_all_digits(Rest); +is_all_digits(_Remaining) -> + false. + +-spec subscribe_to_nats(state()) -> state(). +subscribe_to_nats(#{nats_subscription := Sid} = State) when Sid =/= undefined -> + State; +subscribe_to_nats(State) -> + case subscribe_to_delivery_subject() of + {ok, Sid} -> + MonRef = monitor_nats_rpc(), + logger:info("Push delivery config subscribed to NATS", #{subject => ?NATS_SUBJECT}), + count(fetch_config_from_api()), + State#{nats_subscription => Sid, nats_monitor => MonRef}; + {error, Reason} -> + logger:debug("Push delivery config waiting for NATS subscription", #{ + subject => ?NATS_SUBJECT, reason => Reason + }), + erlang:send_after(?NATS_SUBSCRIBE_RETRY_MS, self(), subscribe_nats), + State + end. + +-spec subscribe_to_delivery_subject() -> {ok, term()} | {error, term()}. +subscribe_to_delivery_subject() -> + try gateway_nats_rpc:subscribe(?NATS_SUBJECT, <<>>) of + {ok, Sid} -> {ok, Sid}; + {error, Reason} -> {error, Reason} + catch + Class:Reason -> {error, {Class, Reason}} + end. + +-spec monitor_nats_rpc() -> reference() | undefined. +monitor_nats_rpc() -> + case whereis(gateway_nats_rpc) of + Pid when is_pid(Pid) -> erlang:monitor(process, Pid); + _ -> undefined + end. + +-spec count(store_result()) -> ok. +count(Result) -> + counters:add(update_counters(), result_index(Result), 1). + +-spec update_counters() -> counters:counters_ref(). +update_counters() -> + case persistent_term:get(?UPDATE_COUNTS_KEY, undefined) of + undefined -> + Counters = counters:new(length(store_results()), [atomics]), + persistent_term:put(?UPDATE_COUNTS_KEY, Counters), + Counters; + Counters -> + Counters + end. + +-spec read_count(counters:counters_ref() | undefined, store_result()) -> non_neg_integer(). +read_count(undefined, _Result) -> + 0; +read_count(Counters, Result) -> + counters:get(Counters, result_index(Result)). + +-spec store_results() -> [store_result()]. +store_results() -> + [updated, unchanged, stale, rejected]. + +-spec result_index(store_result()) -> pos_integer(). +result_index(updated) -> 1; +result_index(unchanged) -> 2; +result_index(stale) -> 3; +result_index(rejected) -> 4. + +-spec log_config_transitions(config(), config()) -> ok. +log_config_transitions(Previous, Current) -> + lists:foreach( + fun(Key) -> log_key_transition(Key, Previous, Current) end, + [enabled, rollout_basis_points, config_version] + ). + +-spec log_key_transition(atom(), config(), config()) -> ok. +log_key_transition(Key, Previous, Current) -> + PreviousValue = maps:get(Key, Previous, undefined), + CurrentValue = maps:get(Key, Current, undefined), + case PreviousValue =:= CurrentValue of + true -> + ok; + false -> + logger:notice( + "Push delivery config transition: key=~s previous=~p current=~p", + [Key, PreviousValue, CurrentValue] + ) + end. diff --git a/fluxer_gateway/src/push/push_eligibility.erl b/fluxer_gateway/src/push/push_eligibility.erl index c309fcaf1..824d83ca5 100644 --- a/fluxer_gateway/src/push/push_eligibility.erl +++ b/fluxer_gateway/src/push/push_eligibility.erl @@ -4,12 +4,10 @@ -typing([eqwalizer]). -export([is_user_blocked/2]). --export([check_user_guild_settings/7]). -export([check_user_guild_settings/8]). -export([prefetch_user_guild_settings/3]). -export([should_allow_notification/6]). -export([is_user_mentioned/5]). --export([is_eligible_for_push/8]). -export([is_eligible_for_push/9]). -export([get_setting/3]). @@ -17,42 +15,6 @@ -define(MESSAGE_NOTIFICATIONS_ONLY_MENTIONS, 1). -define(SETTINGS_PREFETCH_CHUNK_SIZE, 200). --spec is_eligible_for_push( - integer(), integer(), integer(), integer(), map(), integer(), map(), map() -) -> boolean(). -is_eligible_for_push( - UserId, - UserId, - _GuildId, - _ChannelId, - _MessageData, - _GuildDefaultNotifications, - _UserRoles, - _ConnectedUsers -) -> - false; -is_eligible_for_push( - UserId, - AuthorId, - GuildId, - ChannelId, - MessageData, - GuildDefaultNotifications, - UserRolesMap, - ConnectedUsers -) -> - is_eligible_for_push( - UserId, - AuthorId, - GuildId, - ChannelId, - MessageData, - GuildDefaultNotifications, - UserRolesMap, - ConnectedUsers, - push_eligibility_checks:get_guild_large_metadata(GuildId) - ). - -spec is_eligible_for_push( integer(), integer(), @@ -127,46 +89,6 @@ is_user_blocked(UserId, AuthorId) -> BlockedIds -> lists:member(AuthorId, BlockedIds) end. --spec check_user_guild_settings( - integer(), integer(), integer(), map(), integer(), map(), map() -) -> boolean(). -check_user_guild_settings( - _UserId, - 0, - _ChannelId, - _MessageData, - _GuildDefaultNotifications, - _UserRolesMap, - _ConnectedUsers -) -> - true; -check_user_guild_settings( - UserId, - GuildId, - ChannelId, - MessageData, - GuildDefaultNotifications, - UserRolesMap, - ConnectedUsers -) -> - Settings = fetch_settings(UserId, GuildId), - MobilePush = get_boolean_setting(mobile_push, Settings, true), - case MobilePush of - false -> - false; - true -> - push_eligibility_checks:check_muted_and_notifications( - UserId, - ChannelId, - MessageData, - GuildDefaultNotifications, - UserRolesMap, - Settings, - GuildId, - ConnectedUsers - ) - end. - -spec check_user_guild_settings( integer(), integer(), integer(), map(), integer(), map(), map(), map() | undefined ) -> boolean(). @@ -287,15 +209,18 @@ prefetch_settings_chunk_rpc(UserIds, GuildId) -> "Push: prefetching user guild settings via RPC", #{user_count => length(UserIds), guild_id => GuildId} ), - case rpc_client:call(Req) of + Fill = push_ets_cache:reserve_user_guild_settings(UserIds, GuildId), + try rpc_client:call(Req) of {ok, Data} -> - cache_prefetched_settings(UserIds, GuildId, settings_list(Data)); + cache_prefetched_settings(UserIds, GuildId, settings_list(Data), Fill); {error, Reason} -> logger:debug( "Push: RPC failed to prefetch user guild settings", #{user_count => length(UserIds), guild_id => GuildId, reason => Reason} ), ok + after + push_ets_cache:release(Fill) end. -spec settings_list(map()) -> [term()]. @@ -305,19 +230,19 @@ settings_list(Data) -> _ -> [] end. --spec cache_prefetched_settings([integer()], integer(), [term()]) -> ok. -cache_prefetched_settings(UserIds, GuildId, Settings) when +-spec cache_prefetched_settings([integer()], integer(), [term()], push_ets_cache:fill()) -> ok. +cache_prefetched_settings(UserIds, GuildId, Settings, Fill) when length(UserIds) =:= length(Settings) -> lists:foreach( fun({UserId, UserSettings}) -> push_ets_cache:put_user_guild_settings( - UserId, GuildId, settings_map(UserSettings) + UserId, GuildId, settings_map(UserSettings), Fill ) end, lists:zip(UserIds, Settings) ); -cache_prefetched_settings(UserIds, GuildId, Settings) -> +cache_prefetched_settings(UserIds, GuildId, Settings, _Fill) -> logger:debug( "Push: prefetched user guild settings did not match requested users", #{ @@ -391,17 +316,16 @@ extract_mention_flags(UserId, MessageData, Settings, ConnectedUsers) -> -spec evaluate_mentions( boolean(), boolean(), boolean(), integer(), map(), map() ) -> boolean(). -evaluate_mentions(true, false, _SuppressRoles, _UserId, _MessageData, _UserRolesMap) -> - true; -evaluate_mentions(true, true, _SuppressRoles, _UserId, _MessageData, _UserRolesMap) -> - false; -evaluate_mentions(_EvMention, _SuppEv, SuppressRoles, UserId, MessageData, UserRolesMap) -> +evaluate_mentions( + EveryoneMention, SuppressEveryone, SuppressRoles, UserId, MessageData, UserRolesMap +) -> Mentions = maps:get(<<"mentions">>, MessageData, []), MentionRoles = maps:get(<<"mention_roles">>, MessageData, []), UserRoles = maps:get(UserId, UserRolesMap, []), - InMentions = push_eligibility_checks:is_user_in_mentions(UserId, Mentions), - HasRole = push_eligibility_checks:has_mentioned_role(UserRoles, MentionRoles), - InMentions orelse (not SuppressRoles andalso HasRole). + push_eligibility_checks:is_user_in_mentions(UserId, Mentions) orelse + (EveryoneMention andalso not SuppressEveryone) orelse + (not SuppressRoles andalso + push_eligibility_checks:has_mentioned_role(UserRoles, MentionRoles)). -spec get_setting(atom(), term(), term()) -> term(). get_setting(Key, Settings, Default) when is_atom(Key), is_map(Settings) -> @@ -444,18 +368,17 @@ get_setting_default_test() -> ?assertEqual(default, get_setting(mobile_push, not_a_map, default)). is_eligible_same_user_test() -> - ?assertEqual(false, is_eligible_for_push(123, 123, 0, 0, #{}, 0, #{}, #{})). + ?assertEqual(false, is_eligible_for_push(123, 123, 0, 0, #{}, 0, #{}, #{}, undefined)). is_user_blocked_test() -> push_ets_cache:init(), push_ets_cache:put_blocked_ids(123, [456, 789]), - ?assertEqual(true, is_user_blocked(123, 456)), - ?assertEqual(false, is_user_blocked(123, 999)), - ?assertEqual(false, is_user_blocked(999, 456)), try - ets:delete(push_blocked_ids) - catch - error:badarg -> ok + ?assertEqual(true, is_user_blocked(123, 456)), + ?assertEqual(false, is_user_blocked(123, 999)), + ?assertEqual(false, is_user_blocked(999, 456)) + after + ets:delete(push_blocked_ids, 123) end. get_setting_json_null_treated_as_missing_test() -> diff --git a/fluxer_gateway/src/push/push_eligibility_checks.erl b/fluxer_gateway/src/push/push_eligibility_checks.erl index c13d6c20b..30128b8b4 100644 --- a/fluxer_gateway/src/push/push_eligibility_checks.erl +++ b/fluxer_gateway/src/push/push_eligibility_checks.erl @@ -3,7 +3,6 @@ -module(push_eligibility_checks). -typing([eqwalizer]). --export([check_muted_and_notifications/8]). -export([check_muted_and_notifications/9]). -export([is_private_channel/1]). -export([is_user_in_mentions/2]). @@ -13,13 +12,11 @@ -export([resolve_message_notifications/3]). -export([resolve_guild_notification/2]). -export([normalize_notification_level/1]). --export([override_for_large_guild/2]). -export([enforce_only_mentions/1]). -export([is_large_guild/2]). -export([large_guild_threshold/0]). -export([has_large_guild_override/1]). -export([get_guild_large_metadata/1]). --export([check_temp_muted/1]). -define(LARGE_GUILD_THRESHOLD, 2500). -define(LARGE_GUILD_OVERRIDE_FEATURE, <<"LARGE_GUILD_OVERRIDE">>). @@ -33,31 +30,6 @@ -define(LARGE_METADATA_MAILBOX_SHED_THRESHOLD, 100). -define(LARGE_METADATA_CALL_TIMEOUT_MS, 200). --spec check_muted_and_notifications( - integer(), integer(), map(), integer(), map(), map(), integer(), map() -) -> boolean(). -check_muted_and_notifications( - UserId, - ChannelId, - MessageData, - GuildDefaultNotifications, - UserRolesMap, - Settings, - GuildId, - ConnectedUsers -) -> - check_muted_and_notifications( - UserId, - ChannelId, - MessageData, - GuildDefaultNotifications, - UserRolesMap, - Settings, - GuildId, - ConnectedUsers, - get_guild_large_metadata(GuildId) - ). - -spec check_muted_and_notifications( integer(), integer(), map(), integer(), map(), map(), integer(), map(), map() | undefined ) -> boolean(). @@ -72,14 +44,9 @@ check_muted_and_notifications( ConnectedUsers, LargeGuildMetadata ) -> - Muted = boolean_setting(muted, Settings, false), ChannelOverrides = map_setting(channel_overrides, Settings), ChannelOverride = channel_override(ChannelId, ChannelOverrides, #{}), - ChannelMuted = optional_boolean_setting(muted, ChannelOverride), - ActualMuted = resolve_actual_muted(ChannelMuted, Muted), - MuteConfig = push_eligibility:get_setting(mute_config, Settings, undefined), - IsTempMuted = check_temp_muted(MuteConfig), - case ActualMuted orelse IsTempMuted of + case is_mute_active(Settings) orelse is_mute_active(ChannelOverride) of true -> false; false -> @@ -92,23 +59,29 @@ check_muted_and_notifications( ) end. --spec resolve_actual_muted(boolean() | undefined, boolean()) -> boolean(). -resolve_actual_muted(undefined, Muted) -> Muted; -resolve_actual_muted(ChannelMuted, _Muted) -> ChannelMuted. +-spec is_mute_active(term()) -> boolean(). +is_mute_active(Config) -> + boolean_setting(muted, Config, false) andalso + mute_unexpired(push_eligibility:get_setting(mute_config, Config, undefined)). --spec check_temp_muted(term()) -> boolean(). -check_temp_muted(undefined) -> - false; -check_temp_muted(#{<<"end_time">> := EndTimeStr}) -> - case push_utils:parse_timestamp(EndTimeStr) of - undefined -> - false; - EndTime -> - Now = erlang:system_time(millisecond), - Now < EndTime +-spec mute_unexpired(term()) -> boolean(). +mute_unexpired(MuteConfig) when is_map(MuteConfig) -> + case mute_end_ms(push_eligibility:get_setting(end_time, MuteConfig, undefined)) of + undefined -> true; + EndMs -> erlang:system_time(millisecond) < EndMs end; -check_temp_muted(_) -> - false. +mute_unexpired(_MuteConfig) -> + true. + +-spec mute_end_ms(term()) -> integer() | undefined. +mute_end_ms(EndTime) when is_binary(EndTime) -> + try calendar:rfc3339_to_system_time(binary_to_list(EndTime), [{unit, millisecond}]) of + EndMs -> EndMs + catch + _:_ -> undefined + end; +mute_end_ms(_EndTime) -> + undefined. -spec is_private_channel(map()) -> boolean(). is_private_channel(MessageData) -> @@ -189,10 +162,6 @@ normalize_notification_level(?MESSAGE_NOTIFICATIONS_NO_MESSAGES) -> normalize_notification_level(_) -> ?MESSAGE_NOTIFICATIONS_ALL. --spec override_for_large_guild(integer(), integer()) -> integer(). -override_for_large_guild(GuildId, CurrentLevel) -> - override_for_large_guild_metadata(get_guild_large_metadata(GuildId), CurrentLevel). - -spec override_for_large_guild_metadata(map() | undefined, integer()) -> integer(). override_for_large_guild_metadata(undefined, CurrentLevel) -> CurrentLevel; @@ -281,14 +250,6 @@ boolean_setting(Key, Settings, Default) -> _ -> Default end. --spec optional_boolean_setting(atom(), term()) -> boolean() | undefined. -optional_boolean_setting(Key, Settings) -> - case push_eligibility:get_setting(Key, Settings, undefined) of - true -> true; - false -> false; - _ -> undefined - end. - -spec notification_level_setting(term(), integer()) -> integer(). notification_level_setting(Settings, Default) -> case push_eligibility:get_setting(message_notifications, Settings, Default) of @@ -328,7 +289,8 @@ muted_channel_suppresses_push_test() -> UserRolesMap, Settings, GuildId, - ConnectedUsers + ConnectedUsers, + undefined ) ). @@ -336,12 +298,16 @@ guild_muted_suppresses_push_test() -> ?assertEqual(false, check_with_settings(#{muted => true})). temp_muted_suppresses_push_test() -> - FutureMs = integer_to_binary(erlang:system_time(millisecond) + 60000), - ?assertEqual(false, check_with_settings(#{mute_config => #{<<"end_time">> => FutureMs}})). + MuteConfig = #{<<"end_time">> => rfc3339_in_ms(60000)}, + ?assertEqual(false, check_with_settings(#{muted => true, mute_config => MuteConfig})). expired_temp_mute_allows_push_test() -> - PastMs = integer_to_binary(erlang:system_time(millisecond) - 60000), - ?assertEqual(true, check_with_settings(#{mute_config => #{<<"end_time">> => PastMs}})). + MuteConfig = #{<<"end_time">> => rfc3339_in_ms(-60000)}, + ?assertEqual(true, check_with_settings(#{muted => true, mute_config => MuteConfig})). + +rfc3339_in_ms(OffsetMs) -> + Ms = erlang:system_time(millisecond) + OffsetMs, + list_to_binary(calendar:system_time_to_rfc3339(Ms, [{unit, millisecond}, {offset, "Z"}])). check_with_settings(Settings) -> check_muted_and_notifications( @@ -352,7 +318,8 @@ check_with_settings(Settings) -> #{}, Settings, 1, - #{} + #{}, + undefined ). is_user_in_mentions_test() -> diff --git a/fluxer_gateway/src/push/push_endpoint_guard.erl b/fluxer_gateway/src/push/push_endpoint_guard.erl new file mode 100644 index 000000000..2fdf55837 --- /dev/null +++ b/fluxer_gateway/src/push/push_endpoint_guard.erl @@ -0,0 +1,259 @@ +%% SPDX-License-Identifier: AGPL-3.0-or-later + +-module(push_endpoint_guard). +-typing([eqwalizer]). + +-export([check/1, check/2]). + +-export_type([resolver/0]). + +-define(RESOLVE_TIMEOUT_MS, 3000). +-define(MAX_HOST_LENGTH, 253). +-define(MAX_LABEL_LENGTH, 63). + +-type resolver() :: fun((string()) -> {ok, [inet:ip_address()]} | {error, term()}). + +-spec check(binary()) -> ok | {error, term()}. +check(Endpoint) -> + check(Endpoint, fun resolve/1). + +-spec check(binary(), resolver()) -> ok | {error, term()}. +check(Endpoint, Resolver) -> + case parse_endpoint(Endpoint) of + {ok, Host} -> check_host(Host, Resolver); + {error, Reason} -> {error, Reason} + end. + +-spec parse_endpoint(binary()) -> {ok, string()} | {error, term()}. +parse_endpoint(Endpoint) when is_binary(Endpoint) -> + case safe_parse(Endpoint) of + {ok, Parsed} -> validate_parsed(Parsed); + error -> {error, endpoint_rejected} + end; +parse_endpoint(_Endpoint) -> + {error, endpoint_rejected}. + +-spec safe_parse(binary()) -> {ok, map()} | error. +safe_parse(Endpoint) -> + try uri_string:parse(binary_to_list(Endpoint)) of + Parsed when is_map(Parsed) -> {ok, Parsed}; + _ -> error + catch + _:_ -> error + end. + +-spec validate_parsed(map()) -> {ok, string()} | {error, term()}. +validate_parsed(Parsed) -> + Scheme = to_lower(to_string(maps:get(scheme, Parsed, ""))), + Host = to_lower(to_string(maps:get(host, Parsed, ""))), + Port = maps:get(port, Parsed, undefined), + HasUserinfo = maps:is_key(userinfo, Parsed), + Allowed = + Scheme =:= "https" andalso + not HasUserinfo andalso + allowed_port(Port) andalso + Host =/= "", + case Allowed of + true -> {ok, Host}; + false -> {error, endpoint_rejected} + end. + +-spec allowed_port(term()) -> boolean(). +allowed_port(undefined) -> true; +allowed_port(80) -> true; +allowed_port(443) -> true; +allowed_port(_Port) -> false. + +-spec check_host(string(), resolver()) -> ok | {error, term()}. +check_host(Host, Resolver) -> + case inet:parse_address(Host) of + {ok, Address} -> check_addresses([Address]); + {error, _Reason} -> check_hostname(Host, Resolver) + end. + +-spec check_hostname(string(), resolver()) -> ok | {error, term()}. +check_hostname(Host, Resolver) -> + case is_fqdn(Host) of + true -> resolve_and_check(Host, Resolver); + false -> {error, endpoint_rejected} + end. + +-spec resolve_and_check(string(), resolver()) -> ok | {error, term()}. +resolve_and_check(Host, Resolver) -> + case Resolver(Host) of + {ok, Addresses} -> check_addresses(Addresses); + {error, Reason} -> {error, Reason} + end. + +-spec resolve(string()) -> {ok, [inet:ip_address()]} | {error, term()}. +resolve(Host) -> + merge_addresses( + inet:getaddrs(Host, inet, ?RESOLVE_TIMEOUT_MS), + inet:getaddrs(Host, inet6, ?RESOLVE_TIMEOUT_MS) + ). + +-spec merge_addresses( + {ok, [inet:ip_address()]} | {error, term()}, + {ok, [inet:ip_address()]} | {error, term()} +) -> {ok, [inet:ip_address()]} | {error, term()}. +merge_addresses({ok, V4}, {ok, V6}) -> {ok, V4 ++ V6}; +merge_addresses({ok, V4}, {error, _Reason}) -> {ok, V4}; +merge_addresses({error, _Reason}, {ok, V6}) -> {ok, V6}; +merge_addresses({error, Reason}, {error, _Other}) -> {error, Reason}. + +-spec check_addresses([inet:ip_address()]) -> ok | {error, term()}. +check_addresses([]) -> + {error, nxdomain}; +check_addresses(Addresses) -> + case lists:all(fun address_allowed/1, Addresses) of + true -> ok; + false -> {error, endpoint_blocked} + end. + +-spec address_allowed(inet:ip_address()) -> boolean(). +address_allowed({_, _, _, _} = Address) -> + not blocked_v4(Address); +address_allowed({_, _, _, _, _, _, _, _} = Address) -> + case embedded_v4(Address) of + {ok, Embedded} -> not blocked_v4(Embedded); + none -> not blocked_v6(Address) + end. + +-spec blocked_v4(inet:ip4_address()) -> boolean(). +blocked_v4(Address) -> + Value = v4_value(Address), + lists:any( + fun({Network, Prefix}) -> + masked(Value, Prefix, 32) =:= masked(v4_value(Network), Prefix, 32) + end, + blocked_v4_ranges() + ). + +-spec blocked_v6(inet:ip6_address()) -> boolean(). +blocked_v6(Address) -> + Value = v6_value(Address), + lists:any( + fun({Network, Prefix}) -> + masked(Value, Prefix, 128) =:= masked(v6_value(Network), Prefix, 128) + end, + blocked_v6_ranges() + ). + +-spec blocked_v4_ranges() -> [{inet:ip4_address(), non_neg_integer()}]. +blocked_v4_ranges() -> + [ + {{0, 0, 0, 0}, 8}, + {{10, 0, 0, 0}, 8}, + {{100, 64, 0, 0}, 10}, + {{127, 0, 0, 0}, 8}, + {{169, 254, 0, 0}, 16}, + {{172, 16, 0, 0}, 12}, + {{192, 0, 0, 0}, 24}, + {{192, 0, 2, 0}, 24}, + {{192, 88, 99, 0}, 24}, + {{192, 168, 0, 0}, 16}, + {{198, 18, 0, 0}, 15}, + {{198, 51, 100, 0}, 24}, + {{203, 0, 113, 0}, 24}, + {{224, 0, 0, 0}, 4}, + {{240, 0, 0, 0}, 4} + ]. + +-spec blocked_v6_ranges() -> [{inet:ip6_address(), non_neg_integer()}]. +blocked_v6_ranges() -> + [ + {{0, 0, 0, 0, 0, 0, 0, 0}, 128}, + {{0, 0, 0, 0, 0, 0, 0, 1}, 128}, + {{16#2001, 16#0db8, 0, 0, 0, 0, 0, 0}, 32}, + {{16#fc00, 0, 0, 0, 0, 0, 0, 0}, 7}, + {{16#fe80, 0, 0, 0, 0, 0, 0, 0}, 10}, + {{16#ff00, 0, 0, 0, 0, 0, 0, 0}, 8} + ]. + +-spec embedded_v4(inet:ip6_address()) -> {ok, inet:ip4_address()} | none. +embedded_v4({0, 0, 0, 0, 0, 16#ffff, High, Low}) -> {ok, quad(High, Low)}; +embedded_v4({16#0064, 16#ff9b, 0, 0, 0, 0, High, Low}) -> {ok, quad(High, Low)}; +embedded_v4({0, 0, 0, 0, 0, 0, High, Low}) -> {ok, quad(High, Low)}; +embedded_v4({16#2002, High, Low, _, _, _, _, _}) -> {ok, quad(High, Low)}; +embedded_v4(_Address) -> none. + +-spec quad(non_neg_integer(), non_neg_integer()) -> inet:ip4_address(). +quad(High, Low) -> + {High bsr 8, High band 16#ff, Low bsr 8, Low band 16#ff}. + +-spec v4_value(inet:ip4_address()) -> non_neg_integer(). +v4_value({A, B, C, D}) -> + (A bsl 24) bor (B bsl 16) bor (C bsl 8) bor D. + +-spec v6_value(inet:ip6_address()) -> non_neg_integer(). +v6_value({A, B, C, D, E, F, G, H}) -> + lists:foldl(fun(Word, Acc) -> (Acc bsl 16) bor Word end, 0, [A, B, C, D, E, F, G, H]). + +-spec masked(non_neg_integer(), non_neg_integer(), pos_integer()) -> non_neg_integer(). +masked(Value, Prefix, Bits) -> + Value band (((1 bsl Prefix) - 1) bsl (Bits - Prefix)). + +-spec is_fqdn(string()) -> boolean(). +is_fqdn(Host0) -> + Host = strip_trailing_dot(Host0), + Labels = split_labels(Host), + Host =/= "" andalso + length(Host) =< ?MAX_HOST_LENGTH andalso + length(Labels) > 1 andalso + lists:all(fun is_label/1, Labels) andalso + not is_all_digits(lists:last(Labels)). + +-spec strip_trailing_dot(string()) -> string(). +strip_trailing_dot(Host) -> + case lists:reverse(Host) of + [$. | Rest] -> lists:reverse(Rest); + _ -> Host + end. + +-spec split_labels(string()) -> [string()]. +split_labels(Host) -> + split_labels(Host, [], []). + +-spec split_labels(string(), string(), [string()]) -> [string()]. +split_labels([], Current, Acc) -> + lists:reverse([lists:reverse(Current) | Acc]); +split_labels([$. | Rest], Current, Acc) -> + split_labels(Rest, [], [lists:reverse(Current) | Acc]); +split_labels([Char | Rest], Current, Acc) -> + split_labels(Rest, [Char | Current], Acc). + +-spec is_label(string()) -> boolean(). +is_label([]) -> + false; +is_label(Label) when length(Label) > ?MAX_LABEL_LENGTH -> + false; +is_label(Label) -> + lists:all(fun is_label_char/1, Label) andalso + is_alphanumeric(hd(Label)) andalso + is_alphanumeric(lists:last(Label)). + +-spec is_label_char(char()) -> boolean(). +is_label_char($-) -> true; +is_label_char(Char) -> is_alphanumeric(Char). + +-spec is_alphanumeric(char()) -> boolean(). +is_alphanumeric(Char) when Char >= $a, Char =< $z -> true; +is_alphanumeric(Char) when Char >= $0, Char =< $9 -> true; +is_alphanumeric(_Char) -> false. + +-spec is_all_digits(string()) -> boolean(). +is_all_digits("") -> false; +is_all_digits(Label) -> lists:all(fun(Char) -> Char >= $0 andalso Char =< $9 end, Label). + +-spec to_lower(string()) -> string(). +to_lower(Value) -> + [lower_char(Char) || Char <- Value]. + +-spec lower_char(char()) -> char(). +lower_char(Char) when Char >= $A, Char =< $Z -> Char + 32; +lower_char(Char) -> Char. + +-spec to_string(term()) -> string(). +to_string(Value) when is_list(Value) -> Value; +to_string(Value) when is_binary(Value) -> binary_to_list(Value); +to_string(_Value) -> "". diff --git a/fluxer_gateway/src/push/push_ets_cache.erl b/fluxer_gateway/src/push/push_ets_cache.erl index 3e216ead1..e018644e1 100644 --- a/fluxer_gateway/src/push/push_ets_cache.erl +++ b/fluxer_gateway/src/push/push_ets_cache.erl @@ -7,17 +7,25 @@ init/0, get_user_guild_settings/2, put_user_guild_settings/3, + put_user_guild_settings/4, delete_user_guild_settings/2, + reserve_user_guild_settings/2, get_subscriptions/1, get_subscriptions_many/1, - put_subscriptions/2, + put_subscriptions/3, delete_subscriptions/1, + reserve_subscriptions/1, get_blocked_ids/1, put_blocked_ids/2, - put_blocked_ids_fetched/2, + put_blocked_ids_fetched/3, + reserve_blocked_ids/1, get_badge_count/1, - put_badge_count/3, + put_badge_count/4, delete_badge_count/1, + reserve_badge_counts/1, + get_bearer_token/1, + put_bearer_token/3, + release/1, rebalance/0, rebalance_async/0, cache_stats/0, @@ -25,14 +33,19 @@ table_size/1 ]). +-export_type([fill/0]). + -define(USER_GUILD_SETTINGS, push_user_guild_settings). -define(SUBSCRIPTIONS, push_subscriptions). -define(BLOCKED_IDS, push_blocked_ids). -define(BADGE_COUNTS, push_badge_counts). +-define(BEARER_TOKENS, push_bearer_tokens). -define(MAX_TABLE_ENTRIES, 500000). +-define(MAX_BEARER_TOKENS, 10000). -define(EVICT_BATCH, 4096). -define(MAX_EVICT_RESEEKS, 8). +-define(RESERVATION_TTL_MS, 120000). -define(DEFAULT_BLOCKED_IDS_TTL, 300). -define(MIN_BLOCKED_IDS_TTL, 30). -define(MAX_BLOCKED_IDS_TTL, 3600). @@ -41,22 +54,21 @@ named_table, public, set, {read_concurrency, true}, {write_concurrency, true} ]). +-type fill() :: {atom(), pos_integer(), [term()]}. + -spec init() -> ok. init() -> ensure_table(?USER_GUILD_SETTINGS), ensure_table(?SUBSCRIPTIONS), ensure_table(?BLOCKED_IDS), ensure_table(?BADGE_COUNTS), + ensure_table(?BEARER_TOKENS), ok. -spec get_user_guild_settings(integer(), integer()) -> map() | undefined. get_user_guild_settings(UserId, GuildId) -> - get_user_guild_settings_ets(UserId, GuildId). - --spec get_user_guild_settings_ets(integer(), integer()) -> map() | undefined. -get_user_guild_settings_ets(UserId, GuildId) -> try ets:lookup(?USER_GUILD_SETTINGS, {UserId, GuildId}) of - [{{UserId, GuildId}, Settings}] -> Settings; + [{{UserId, GuildId}, Settings}] when is_map(Settings) -> Settings; _ -> undefined catch error:badarg -> undefined @@ -64,28 +76,24 @@ get_user_guild_settings_ets(UserId, GuildId) -> -spec put_user_guild_settings(integer(), integer(), map()) -> ok. put_user_guild_settings(UserId, GuildId, Settings) -> - guard_table_size(?USER_GUILD_SETTINGS), - ets:insert(?USER_GUILD_SETTINGS, {{UserId, GuildId}, Settings}), - ok. + write(?USER_GUILD_SETTINGS, {{UserId, GuildId}, Settings}). + +-spec put_user_guild_settings(integer(), integer(), map(), fill()) -> ok. +put_user_guild_settings(UserId, GuildId, Settings, {?USER_GUILD_SETTINGS, _, _} = Fill) -> + fill(Fill, {{UserId, GuildId}, Settings}). -spec delete_user_guild_settings(integer(), integer()) -> ok. delete_user_guild_settings(UserId, GuildId) -> - try ets:delete(?USER_GUILD_SETTINGS, {UserId, GuildId}) of - _ -> ok - catch - throw:_ -> ok; - error:_ -> ok; - exit:_ -> ok - end. + safe_delete(?USER_GUILD_SETTINGS, {UserId, GuildId}). + +-spec reserve_user_guild_settings([integer()], integer()) -> fill(). +reserve_user_guild_settings(UserIds, GuildId) -> + reserve(?USER_GUILD_SETTINGS, [{UserId, GuildId} || UserId <- UserIds]). -spec get_subscriptions(integer()) -> list() | undefined. get_subscriptions(UserId) -> - get_subscriptions_ets(UserId). - --spec get_subscriptions_ets(integer()) -> list() | undefined. -get_subscriptions_ets(UserId) -> try ets:lookup(?SUBSCRIPTIONS, UserId) of - [{UserId, Subs}] -> Subs; + [{UserId, Subs}] when is_list(Subs) -> Subs; _ -> undefined catch error:badarg -> undefined @@ -109,32 +117,26 @@ add_cached_subscriptions(UserId, {CachedAcc, MissingAcc}) -> {CachedAcc, [UserId | MissingAcc]} end. --spec put_subscriptions(integer(), list()) -> ok. -put_subscriptions(UserId, Subscriptions) -> - guard_table_size(?SUBSCRIPTIONS), - ets:insert(?SUBSCRIPTIONS, {UserId, Subscriptions}), - ok. +-spec put_subscriptions(integer(), list(), fill()) -> ok. +put_subscriptions(UserId, Subscriptions, {?SUBSCRIPTIONS, _, _} = Fill) -> + fill(Fill, {UserId, Subscriptions}). -spec delete_subscriptions(integer()) -> ok. delete_subscriptions(UserId) -> - try ets:delete(?SUBSCRIPTIONS, UserId) of - _ -> ok - catch - error:badarg -> ok - end. + safe_delete(?SUBSCRIPTIONS, UserId). + +-spec reserve_subscriptions([integer()]) -> fill(). +reserve_subscriptions(UserIds) -> + reserve(?SUBSCRIPTIONS, UserIds). -spec get_blocked_ids(integer()) -> [integer()] | undefined. get_blocked_ids(UserId) -> - get_blocked_ids_ets(UserId). - --spec get_blocked_ids_ets(integer()) -> [integer()] | undefined. -get_blocked_ids_ets(UserId) -> try ets:lookup(?BLOCKED_IDS, UserId) of - [{UserId, BlockedIds}] -> + [{UserId, BlockedIds}] when is_list(BlockedIds) -> BlockedIds; - [{UserId, BlockedIds, infinity}] -> + [{UserId, BlockedIds, infinity}] when is_list(BlockedIds) -> BlockedIds; - [{UserId, BlockedIds, ExpiresAt}] when is_integer(ExpiresAt) -> + [{UserId, BlockedIds, ExpiresAt}] when is_list(BlockedIds), is_integer(ExpiresAt) -> live_fetched_blocked_ids(BlockedIds, ExpiresAt); _ -> undefined @@ -151,21 +153,16 @@ live_fetched_blocked_ids(BlockedIds, ExpiresAt) -> -spec put_blocked_ids(integer(), [integer()]) -> ok. put_blocked_ids(UserId, BlockedIds) -> - insert_blocked_ids(UserId, BlockedIds, infinity). + write(?BLOCKED_IDS, {UserId, BlockedIds, infinity}). --spec put_blocked_ids_fetched(integer(), [integer()]) -> ok. -put_blocked_ids_fetched(UserId, BlockedIds) -> +-spec put_blocked_ids_fetched(integer(), [integer()], fill()) -> ok. +put_blocked_ids_fetched(UserId, BlockedIds, {?BLOCKED_IDS, _, _} = Fill) -> ExpiresAt = erlang:system_time(second) + blocked_ids_ttl_seconds(), - insert_blocked_ids(UserId, BlockedIds, ExpiresAt). + fill(Fill, {UserId, BlockedIds, ExpiresAt}). --spec insert_blocked_ids(integer(), [integer()], infinity | integer()) -> ok. -insert_blocked_ids(UserId, BlockedIds, ExpiresAt) -> - guard_table_size(?BLOCKED_IDS), - try ets:insert(?BLOCKED_IDS, {UserId, BlockedIds, ExpiresAt}) of - _ -> ok - catch - error:badarg -> ok - end. +-spec reserve_blocked_ids([integer()]) -> fill(). +reserve_blocked_ids(UserIds) -> + reserve(?BLOCKED_IDS, UserIds). -spec blocked_ids_ttl_seconds() -> pos_integer(). blocked_ids_ttl_seconds() -> @@ -181,75 +178,99 @@ app_pos_integer(Key, Default) -> -spec get_badge_count(integer()) -> {non_neg_integer(), integer()} | undefined. get_badge_count(UserId) -> - get_badge_count_ets(UserId). - --spec get_badge_count_ets(integer()) -> {non_neg_integer(), integer()} | undefined. -get_badge_count_ets(UserId) -> try ets:lookup(?BADGE_COUNTS, UserId) of - [{UserId, Count, CachedAt}] -> {Count, CachedAt}; - _ -> undefined + [{UserId, Count, CachedAt}] when is_integer(Count), Count >= 0, is_integer(CachedAt) -> + {Count, CachedAt}; + _ -> + undefined catch error:badarg -> undefined end. --spec put_badge_count(integer(), non_neg_integer(), integer()) -> ok. -put_badge_count(UserId, Count, CachedAt) -> - guard_table_size(?BADGE_COUNTS), - try - case ets:insert_new(?BADGE_COUNTS, {UserId, Count, CachedAt}) of - true -> - ok; - false -> - conditional_replace_badge_count(UserId, Count, CachedAt) - end - catch - error:badarg -> ok - end, - ok. - --spec conditional_replace_badge_count(integer(), non_neg_integer(), integer()) -> ok. -conditional_replace_badge_count(UserId, Count, CachedAt) -> - MatchSpec = [ - { - {UserId, '$1', '$2'}, - [{'=<', '$2', CachedAt}], - [{{{const, UserId}, {const, Count}, {const, CachedAt}}}] - } - ], - try ets:select_replace(?BADGE_COUNTS, MatchSpec) of - Replaced when is_integer(Replaced), Replaced > 0 -> - ok; - _ -> - retry_conditional_replace_badge_count(UserId, Count, CachedAt) - catch - error:badarg -> ok - end. - --spec retry_conditional_replace_badge_count(integer(), non_neg_integer(), integer()) -> ok. -retry_conditional_replace_badge_count(UserId, Count, CachedAt) -> - case get_badge_count_ets(UserId) of - {_ExistingCount, ExistingCachedAt} when ExistingCachedAt > CachedAt -> - ok; - _ -> - insert_badge_count_safely(UserId, Count, CachedAt) - end. - --spec insert_badge_count_safely(integer(), non_neg_integer(), integer()) -> ok. -insert_badge_count_safely(UserId, Count, CachedAt) -> - try ets:insert(?BADGE_COUNTS, {UserId, Count, CachedAt}) of - _ -> ok - catch - error:badarg -> ok - end. +-spec put_badge_count(integer(), non_neg_integer(), integer(), fill()) -> ok. +put_badge_count(UserId, Count, CachedAt, {?BADGE_COUNTS, _, _} = Fill) -> + fill(Fill, {UserId, Count, CachedAt}). -spec delete_badge_count(integer()) -> ok. delete_badge_count(UserId) -> - try ets:delete(?BADGE_COUNTS, UserId) of + safe_delete(?BADGE_COUNTS, UserId). + +-spec reserve_badge_counts([integer()]) -> fill(). +reserve_badge_counts(UserIds) -> + reserve(?BADGE_COUNTS, UserIds). + +-spec get_bearer_token(term()) -> {ok, binary(), integer()} | undefined. +get_bearer_token(Key) -> + try ets:lookup(?BEARER_TOKENS, Key) of + [{_, Token, ExpiresAt}] when is_binary(Token), is_integer(ExpiresAt) -> + {ok, Token, ExpiresAt}; + _ -> + undefined + catch + error:badarg -> undefined + end. + +-spec put_bearer_token(term(), binary(), integer()) -> ok. +put_bearer_token(Key, Token, ExpiresAt) when is_binary(Token), is_integer(ExpiresAt) -> + guard_table_size(?BEARER_TOKENS, ?MAX_BEARER_TOKENS), + try ets:insert(?BEARER_TOKENS, {Key, Token, ExpiresAt}) of _ -> ok catch error:badarg -> ok end. +-spec write(atom(), tuple()) -> ok. +write(Table, Row) -> + guard_table_size(Table, ?MAX_TABLE_ENTRIES), + try ets:insert(Table, Row) of + _ -> ok + catch + error:badarg -> ok + end. + +-spec reserve(atom(), [term()]) -> fill(). +reserve(Table, Keys) -> + guard_table_size(Table, ?MAX_TABLE_ENTRIES), + Token = erlang:unique_integer([positive]), + ReservedAt = erlang:monotonic_time(millisecond), + lists:foreach(fun(Key) -> reserve_key(Table, Key, Token, ReservedAt) end, Keys), + {Table, Token, Keys}. + +-spec reserve_key(atom(), term(), pos_integer(), integer()) -> ok. +reserve_key(Table, Key, Token, ReservedAt) -> + try + _ = ets:select_delete(Table, stale_rows(Table, Key)), + _ = ets:insert_new(Table, {Key, pending, Token, ReservedAt}), + ok + catch + error:badarg -> ok + end. + +-spec stale_rows(atom(), term()) -> ets:match_spec(). +stale_rows(?BLOCKED_IDS, Key) -> + Now = erlang:system_time(second), + [{{Key, '_', '$1'}, [{is_integer, '$1'}, {'=<', '$1', Now}], [true]}]; +stale_rows(?BADGE_COUNTS, Key) -> + [{{Key, '_', '_'}, [], [true]}]; +stale_rows(_Table, _Key) -> + []. + +-spec fill(fill(), tuple()) -> ok. +fill({Table, Token, _Keys}, Row) -> + Key = element(1, Row), + try ets:select_replace(Table, [{{Key, pending, Token, '_'}, [], [{const, Row}]}]) of + _ -> ok + catch + error:badarg -> ok + end. + +-spec release(fill()) -> ok. +release({Table, Token, Keys}) -> + lists:foreach( + fun(Key) -> select_delete(Table, [{{Key, pending, Token, '_'}, [], [true]}]) end, + Keys + ). + -spec rebalance_async() -> ok. rebalance_async() -> _ = spawn(fun rebalance/0), @@ -270,32 +291,47 @@ cache_stats() -> user_guild_settings_size => table_size(?USER_GUILD_SETTINGS), push_subscriptions_size => table_size(?SUBSCRIPTIONS), blocked_ids_size => table_size(?BLOCKED_IDS), - badge_counts_size => table_size(?BADGE_COUNTS) + badge_counts_size => table_size(?BADGE_COUNTS), + bearer_tokens_size => table_size(?BEARER_TOKENS) }. -spec evict_tables(map()) -> ok. evict_tables(MaxEntries) -> - _ = expire_blocked_ids(), + Now = erlang:system_time(second), + select_delete(?BLOCKED_IDS, expired_rows(Now)), + select_delete(?BEARER_TOKENS, expired_rows(Now)), + lists:foreach( + fun expire_reservations/1, + [?USER_GUILD_SETTINGS, ?SUBSCRIPTIONS, ?BLOCKED_IDS, ?BADGE_COUNTS] + ), evict_table(?USER_GUILD_SETTINGS, maps:get(user_guild_settings, MaxEntries, undefined)), evict_table(?SUBSCRIPTIONS, maps:get(subscriptions, MaxEntries, undefined)), evict_table(?BLOCKED_IDS, maps:get(blocked_ids, MaxEntries, undefined)), evict_table(?BADGE_COUNTS, maps:get(badge_counts, MaxEntries, undefined)), + evict_table(?BEARER_TOKENS, ?MAX_BEARER_TOKENS), ok. --spec expire_blocked_ids() -> non_neg_integer(). -expire_blocked_ids() -> - Now = erlang:system_time(second), - MatchSpec = [{{'_', '_', '$1'}, [{is_integer, '$1'}, {'=<', '$1', Now}], [true]}], - try ets:select_delete(?BLOCKED_IDS, MatchSpec) of - Deleted -> Deleted +-spec expired_rows(integer()) -> ets:match_spec(). +expired_rows(Now) -> + [{{'_', '_', '$1'}, [{is_integer, '$1'}, {'=<', '$1', Now}], [true]}]. + +-spec expire_reservations(atom()) -> ok. +expire_reservations(Table) -> + Cutoff = erlang:monotonic_time(millisecond) - ?RESERVATION_TTL_MS, + select_delete(Table, [{{'_', pending, '_', '$1'}, [{'<', '$1', Cutoff}], [true]}]). + +-spec select_delete(atom(), ets:match_spec()) -> ok. +select_delete(Table, MatchSpec) -> + try ets:select_delete(Table, MatchSpec) of + _ -> ok catch - error:badarg -> 0 + error:badarg -> ok end. --spec guard_table_size(atom()) -> ok. -guard_table_size(Table) -> - case table_size(Table) >= ?MAX_TABLE_ENTRIES of - true -> evict_table(Table, ?MAX_TABLE_ENTRIES - ?EVICT_BATCH); +-spec guard_table_size(atom(), non_neg_integer()) -> ok. +guard_table_size(Table, MaxEntries) -> + case table_size(Table) >= MaxEntries of + true -> evict_table(Table, max(0, MaxEntries - ?EVICT_BATCH)); false -> ok end. @@ -462,7 +498,7 @@ synced_blocked_ids_never_expire_test() -> fetched_blocked_ids_expire_and_are_reclaimed_test() -> init(), try - ok = put_blocked_ids_fetched(4004, [4005]), + ok = put_blocked_ids_fetched(4004, [4005], reserve_blocked_ids([4004])), ?assertEqual([4005], get_blocked_ids(4004)), Stale = erlang:system_time(second) - 1, true = ets:insert(?BLOCKED_IDS, {4004, [4005], Stale}), @@ -473,7 +509,7 @@ fetched_blocked_ids_expire_and_are_reclaimed_test() -> safe_delete(?BLOCKED_IDS, 4004) end. -legacy_untimed_blocked_ids_are_still_honoured_test() -> +untimed_blocked_ids_are_still_honoured_test() -> init(), try true = ets:insert(?BLOCKED_IDS, {4008, [4009]}), diff --git a/fluxer_gateway/src/push/push_fcm.erl b/fluxer_gateway/src/push/push_fcm.erl index 380c87aff..328fa235f 100644 --- a/fluxer_gateway/src/push/push_fcm.erl +++ b/fluxer_gateway/src/push/push_fcm.erl @@ -119,7 +119,7 @@ resolve_access_token() -> -spec get_or_fetch_token(map(), term(), integer(), binary()) -> {ok, binary()} | {error, term()}. get_or_fetch_token(ServiceAccount, CacheKey, Now, TokenUri) -> - case push_token_cache:get(CacheKey) of + case push_ets_cache:get_bearer_token(CacheKey) of {ok, Token, ExpiresAt} when ExpiresAt - ?ACCESS_TOKEN_SKEW_SECONDS > Now -> {ok, Token}; _ -> @@ -175,7 +175,7 @@ parse_token_response(CacheKey, Now, ResponseBody) -> case decode_json_map(ResponseBody) of #{<<"access_token">> := AccessToken} = Response when is_binary(AccessToken) -> ExpiresIn = normalize_expires_in(maps:get(<<"expires_in">>, Response, 3600)), - push_token_cache:put(CacheKey, AccessToken, Now + ExpiresIn), + push_ets_cache:put_bearer_token(CacheKey, AccessToken, Now + ExpiresIn), {ok, AccessToken}; _ -> {error, invalid_token_response} diff --git a/fluxer_gateway/src/push/push_fcm_payload.erl b/fluxer_gateway/src/push/push_fcm_payload.erl index 9796c03b0..c4bd6ce24 100644 --- a/fluxer_gateway/src/push/push_fcm_payload.erl +++ b/fluxer_gateway/src/push/push_fcm_payload.erl @@ -214,24 +214,24 @@ stringify_list_value(Value) -> -spec stringify_json_value(term()) -> binary(). stringify_json_value(Value) -> - iolist_to_binary(json:encode(json_compatible_value(Value))). + iolist_to_binary(json:encode(json_encodable_value(Value))). --spec json_compatible_value(term()) -> json:encode_value(). -json_compatible_value(Value) when is_binary(Value) -> Value; -json_compatible_value(Value) when is_integer(Value) -> Value; -json_compatible_value(Value) when is_float(Value) -> Value; -json_compatible_value(Value) when is_atom(Value) -> Value; -json_compatible_value(Value) when is_list(Value) -> - [json_compatible_value(Item) || Item <- Value]; -json_compatible_value(Value) when is_map(Value) -> +-spec json_encodable_value(term()) -> json:encode_value(). +json_encodable_value(Value) when is_binary(Value) -> Value; +json_encodable_value(Value) when is_integer(Value) -> Value; +json_encodable_value(Value) when is_float(Value) -> Value; +json_encodable_value(Value) when is_atom(Value) -> Value; +json_encodable_value(Value) when is_list(Value) -> + [json_encodable_value(Item) || Item <- Value]; +json_encodable_value(Value) when is_map(Value) -> maps:fold( fun(Key, Item, Acc) -> - Acc#{push_utils:normalize_binary(Key, <<>>) => json_compatible_value(Item)} + Acc#{push_utils:normalize_binary(Key, <<>>) => json_encodable_value(Item)} end, #{}, Value ); -json_compatible_value(Value) -> +json_encodable_value(Value) -> iolist_to_binary(io_lib:format("~p", [Value])). -spec first_binary(list()) -> binary() | undefined. diff --git a/fluxer_gateway/src/push/push_job_publisher.erl b/fluxer_gateway/src/push/push_job_publisher.erl new file mode 100644 index 000000000..f50184f09 --- /dev/null +++ b/fluxer_gateway/src/push/push_job_publisher.erl @@ -0,0 +1,249 @@ +%% SPDX-License-Identifier: AGPL-3.0-or-later + +-module(push_job_publisher). +-typing([eqwalizer]). + +-export([publish_message/8, publish_message/10, publish_clear/3, publish_clear/5, request/3]). + +-define(SUBJECT_MESSAGE, <<"push.job.message">>). +-define(SUBJECT_CLEAR, <<"push.job.clear">>). +-define(JOB_VERSION, 1). +-define(NATS_MAX_PAYLOAD_BYTES, 1048576). + +-type meta() :: #{ + kind := message | clear, + user_ids := [integer()], + channel_id := integer(), + message_id := integer(), + fallback := push_outbox:fallback() +}. + +-spec publish_message( + [integer()], + map(), + map(), + integer(), + integer(), + integer(), + binary() | undefined, + binary() | undefined +) -> ok | {error, term()}. +publish_message( + UserIds, + MessageData, + MarkdownContext, + GuildId, + ChannelId, + MessageId, + GuildName, + ChannelName +) -> + publish_message( + UserIds, + MessageData, + MarkdownContext, + GuildId, + ChannelId, + MessageId, + GuildName, + ChannelName, + push_delivery_config:config_version(), + fun ignore_fallback/1 + ). + +-spec publish_message( + [integer()], + map(), + map(), + integer(), + integer(), + integer(), + binary() | undefined, + binary() | undefined, + non_neg_integer(), + push_outbox:fallback() +) -> ok | {error, term()}. +publish_message( + UserIds, + MessageData, + MarkdownContext, + GuildId, + ChannelId, + MessageId, + GuildName, + ChannelName, + ConfigVersion, + Fallback +) -> + ChannelIdBin = integer_to_binary(ChannelId), + MessageIdBin = integer_to_binary(MessageId), + Job = #{ + <<"v">> => ?JOB_VERSION, + <<"config_version">> => ConfigVersion, + <<"guild_id">> => integer_to_binary(GuildId), + <<"channel_id">> => ChannelIdBin, + <<"message_id">> => MessageIdBin, + <<"notification">> => notification_fields( + MessageData, + MarkdownContext, + GuildId, + ChannelId, + MessageId, + GuildName, + ChannelName + ), + <<"user_ids">> => [integer_to_binary(UserId) || UserId <- UserIds] + }, + publish(?SUBJECT_MESSAGE, Job, #{ + kind => message, + user_ids => UserIds, + channel_id => ChannelId, + message_id => MessageId, + fallback => Fallback + }). + +-spec publish_clear(integer(), integer(), integer()) -> ok | {error, term()}. +publish_clear(UserId, ChannelId, MessageId) -> + publish_clear( + UserId, + ChannelId, + MessageId, + push_delivery_config:config_version(), + fun ignore_fallback/1 + ). + +-spec publish_clear( + integer(), integer(), integer(), non_neg_integer(), push_outbox:fallback() +) -> + ok | {error, term()}. +publish_clear(UserId, ChannelId, MessageId, ConfigVersion, Fallback) -> + Job = #{ + <<"v">> => ?JOB_VERSION, + <<"config_version">> => ConfigVersion, + <<"user_id">> => integer_to_binary(UserId), + <<"channel_id">> => integer_to_binary(ChannelId), + <<"message_id">> => integer_to_binary(MessageId) + }, + publish(?SUBJECT_CLEAR, Job, #{ + kind => clear, + user_ids => [UserId], + channel_id => ChannelId, + message_id => MessageId, + fallback => Fallback + }). + +-spec request(binary(), binary(), pos_integer()) -> ok | {error, term()}. +request(Subject, Body, Timeout) -> + case gateway_nats_pool_conn:get_pool_conn() of + {ok, Conn} -> reply_result(nats:request(Conn, Subject, Body, #{timeout => Timeout})); + {error, Reason} -> {error, Reason} + end. + +-spec reply_result({ok, {iodata(), map()}} | {error, term()}) -> ok | {error, term()}. +reply_result({ok, {Payload, _MsgOpts}}) -> + decode_reply(Payload); +reply_result({error, Reason}) -> + {error, Reason}. + +-spec decode_reply(iodata()) -> ok | {error, term()}. +decode_reply(Payload) -> + try json:decode(iolist_to_binary(Payload)) of + #{<<"ok">> := true} -> ok; + #{<<"ok">> := false} = Reply -> {error, {rejected, maps:get(<<"error">>, Reply, null)}}; + _ -> {error, invalid_reply} + catch + _:_ -> {error, invalid_reply} + end. + +-spec ignore_fallback([integer()]) -> ok. +ignore_fallback(_UserIds) -> + ok. + +-spec notification_fields( + map(), map(), integer(), integer(), integer(), binary() | undefined, binary() | undefined +) -> map(). +notification_fields( + MessageData, MarkdownContext, GuildId, ChannelId, MessageId, GuildName, ChannelName +) -> + AuthorData = maps:get(<<"author">>, MessageData, #{}), + AuthorUsername = maps:get(<<"username">>, AuthorData, <<"Unknown">>), + AuthorName = push_notification_format:resolve_author_name( + MessageData, MarkdownContext, AuthorUsername + ), + ChannelIdBin = integer_to_binary(ChannelId), + MessageIdBin = integer_to_binary(MessageId), + #{ + <<"title">> => push_notification:build_notification_title( + AuthorName, MessageData, GuildId, GuildName, ChannelName + ), + <<"body">> => push_notification_format:build_content_preview( + MessageData, MarkdownContext + ), + <<"icon">> => push_notification_format:resolve_author_avatar_url(AuthorData), + <<"badge">> => push_utils:construct_static_asset_url( + <<"marketing/branding/symbol-white.svg">> + ), + <<"tag">> => <<"channel:", ChannelIdBin/binary, ":", MessageIdBin/binary>>, + <<"notification_tag">> => <<"channel:", ChannelIdBin/binary>>, + <<"url">> => push_notification_format:build_url(GuildId, ChannelId, MessageId), + <<"image_url">> => nullable(push_notification_format:extract_image_url(MessageData)) + }. + +-spec nullable(binary() | undefined) -> binary() | null. +nullable(undefined) -> + null; +nullable(Value) -> + Value. + +-spec publish(binary(), map(), meta()) -> ok | {error, term()}. +publish(Subject, Job, Meta) -> + case encode(Job) of + {ok, Body} -> + publish_bounded(Subject, Job, Body, Meta); + {error, Reason} -> + logger:warning("Push job encode failed", #{subject => Subject, reason => Reason}), + {error, Reason} + end. + +-spec encode(map()) -> {ok, binary()} | {error, term()}. +encode(Job) -> + try + {ok, iolist_to_binary(json:encode(Job))} + catch + Class:Reason -> {error, {encode_failed, Class, Reason}} + end. + +-spec publish_bounded(binary(), map(), binary(), meta()) -> ok | {error, term()}. +publish_bounded(Subject, _Job, Body, _Meta) when byte_size(Body) > ?NATS_MAX_PAYLOAD_BYTES -> + logger:warning("Push job exceeds the NATS payload limit", #{ + subject => Subject, bytes => byte_size(Body), limit => ?NATS_MAX_PAYLOAD_BYTES + }), + {error, {payload_too_large, byte_size(Body)}}; +publish_bounded(Subject, Job, Body, Meta) -> + case push_outbox:enqueue(outbox_job(Subject, Job, Body, Meta)) of + ok -> + ok; + {error, Reason} -> + logger:warning("Push job publish failed", #{subject => Subject, reason => Reason}), + {error, Reason} + end. + +-spec outbox_job(binary(), map(), binary(), meta()) -> push_outbox:job(). +outbox_job(Subject, Job, Body, Meta) -> + #{ + kind := Kind, + user_ids := UserIds, + channel_id := ChannelId, + message_id := MessageId, + fallback := Fallback + } = Meta, + #{ + kind => Kind, + subject => Subject, + job => Job, + body => Body, + user_ids => UserIds, + channel_id => ChannelId, + message_id => MessageId, + fallback => Fallback + }. diff --git a/fluxer_gateway/src/push/push_normalize.erl b/fluxer_gateway/src/push/push_normalize.erl index 6d7edb26c..bfcec1ef2 100644 --- a/fluxer_gateway/src/push/push_normalize.erl +++ b/fluxer_gateway/src/push/push_normalize.erl @@ -17,9 +17,11 @@ notification_level(undefined) -> 0; notification_level(null) -> 0; +notification_level(Value) when is_integer(Value), Value >= -1, Value =< 3 -> + Value; notification_level(Value) -> case guild_data_normalize_schema:int(Value) of - Level when Level >= -1, Level =< 3 -> Level; + Level when is_integer(Level), Level =< 3 -> Level; _ -> undefined end. diff --git a/fluxer_gateway/src/push/push_notification.erl b/fluxer_gateway/src/push/push_notification.erl index 71fa239ec..b43e27eb3 100644 --- a/fluxer_gateway/src/push/push_notification.erl +++ b/fluxer_gateway/src/push/push_notification.erl @@ -6,11 +6,36 @@ -export([ build_notification_title/5, build_notification_payload/1, - build_clear_notification_payload/4 + build_clear_notification_payload/4, + fit_payload_json/2, + is_clear/1 ]). -export_type([notification_input/0]). +-define(WEB_PUSH_MARKER, 8030). +-define(CLEAR_TYPE, <<"notification_clear">>). +-define(CLEAR_ACTION, <<"clear_channel">>). +-define(FALLBACK_TITLE, <<"Fluxer">>). +-define(FALLBACK_TAG, <<"fluxer-message">>). +-define(SHRUNK_BODY_MAX_BYTES, 40). +-define(MINIMAL_TITLE_MAX_BYTES, 120). +-define(SHRINK_STEPS, [media, icons, body, minimal]). +-define(BLOCK_KEYS, [<<"notification">>, <<"data">>]). +-define(MEDIA_KEYS, [<<"image_url">>, <<"image">>]). +-define(ICON_KEYS, [<<"icon">>, <<"badge">>, <<"author_avatar_url">>]). +-define(MINIMAL_DATA_KEYS, [ + <<"channel_id">>, + <<"message_id">>, + <<"guild_id">>, + <<"target_user_id">>, + <<"notification_tag">>, + <<"url">>, + <<"badge_count">> +]). + +-type shrink_step() :: media | icons | body | minimal. + -type push_ctx() :: #{ channel_id := integer(), message_id := integer(), @@ -137,7 +162,7 @@ assemble_payload( Notification = build_notification_body(Ctx, Title, ContentPreview, AuthorAvatarUrl, Data), maps:merge( #{ - <<"web_push">> => 8030, + <<"web_push">> => ?WEB_PUSH_MARKER, <<"notification">> => Notification, <<"title">> => Title, <<"body">> => ContentPreview, @@ -215,8 +240,8 @@ build_clear_notification_payload(TargetUserId, ChannelId, MessageId, BadgeCount) BadgeValue = max(0, BadgeCount), Tag = build_channel_tag(ChannelId), Data = #{ - <<"type">> => <<"notification_clear">>, - <<"action">> => <<"clear_channel">>, + <<"type">> => ?CLEAR_TYPE, + <<"action">> => ?CLEAR_ACTION, <<"channel_id">> => integer_to_binary(ChannelId), <<"message_id">> => integer_to_binary(MessageId), <<"target_user_id">> => integer_to_binary(TargetUserId), @@ -225,14 +250,14 @@ build_clear_notification_payload(TargetUserId, ChannelId, MessageId, BadgeCount) <<"badge_count">> => BadgeValue }, #{ - <<"type">> => <<"notification_clear">>, - <<"action">> => <<"clear_channel">>, + <<"type">> => ?CLEAR_TYPE, + <<"action">> => ?CLEAR_ACTION, <<"silent">> => true, <<"tag">> => Tag, <<"notification_tag">> => Tag, <<"data">> => Data, <<"badge_count">> => BadgeValue, - <<"web_push">> => 8030, + <<"web_push">> => ?WEB_PUSH_MARKER, <<"notification">> => #{ <<"tag">> => Tag, <<"data">> => Data, @@ -241,6 +266,175 @@ build_clear_notification_payload(TargetUserId, ChannelId, MessageId, BadgeCount) } }. +-spec is_clear(map()) -> boolean(). +is_clear(Payload) -> + maps:get(<<"type">>, Payload, undefined) =:= ?CLEAR_TYPE orelse + maps:get(<<"action">>, Payload, undefined) =:= ?CLEAR_ACTION. + +-spec fit_payload_json(map(), non_neg_integer()) -> binary(). +fit_payload_json(Payload, Budget) -> + Encoded = encode_payload(Payload), + case byte_size(Encoded) =< Budget of + true -> Encoded; + false -> shrink_to_budget(Payload, Budget, ?SHRINK_STEPS) + end. + +-spec shrink_to_budget(map(), non_neg_integer(), [shrink_step()]) -> binary(). +shrink_to_budget(Payload, _Budget, []) -> + encode_payload(Payload); +shrink_to_budget(Payload, Budget, [Step | RemainingSteps]) -> + Shrunk = shrink_payload(Payload, Step, Budget), + Encoded = encode_payload(Shrunk), + case byte_size(Encoded) =< Budget of + true -> Encoded; + false -> shrink_to_budget(Shrunk, Budget, RemainingSteps) + end. + +-spec shrink_payload(map(), shrink_step(), non_neg_integer()) -> map(). +shrink_payload(Payload, media, _Budget) -> + drop_block_keys(Payload, ?MEDIA_KEYS); +shrink_payload(Payload, icons, _Budget) -> + drop_block_keys(Payload, ?ICON_KEYS); +shrink_payload(Payload, body, _Budget) -> + truncate_block_bodies(Payload); +shrink_payload(Payload, minimal, Budget) -> + minimal_payload(Payload, Budget). + +-spec drop_block_keys(map(), [binary()]) -> map(). +drop_block_keys(Payload, Keys) -> + map_blocks(Payload, fun(Block) -> maps:without(Keys, Block) end). + +-spec truncate_block_bodies(map()) -> map(). +truncate_block_bodies(Payload) -> + map_blocks(Payload, fun truncate_block_body/1). + +-spec truncate_block_body(map()) -> map(). +truncate_block_body(Block) -> + case maps:get(<<"body">>, Block, undefined) of + Body when is_binary(Body) -> + Block#{ + <<"body">> => push_notification_format:truncate_bytes( + Body, ?SHRUNK_BODY_MAX_BYTES + ) + }; + _ -> + Block + end. + +-spec map_blocks(map(), fun((map()) -> map())) -> map(). +map_blocks(Block, Apply) -> + Apply( + lists:foldl(fun(Key, Acc) -> map_nested_block(Key, Acc, Apply) end, Block, ?BLOCK_KEYS) + ). + +-spec map_nested_block(binary(), map(), fun((map()) -> map())) -> map(). +map_nested_block(Key, Block, Apply) -> + case maps:get(Key, Block, undefined) of + Nested when is_map(Nested) -> Block#{Key => map_blocks(Nested, Apply)}; + _ -> Block + end. + +-spec minimal_payload(map(), non_neg_integer()) -> map(). +minimal_payload(Payload, Budget) -> + Title = push_notification_format:truncate_bytes( + first_text(Payload, <<"title">>, ?FALLBACK_TITLE), ?MINIMAL_TITLE_MAX_BYTES + ), + Tag = first_text(Payload, <<"tag">>, ?FALLBACK_TAG), + Url = minimal_url(Payload), + Data = minimal_data(Payload), + first_payload_within_budget( + [ + minimal_envelope(Title, Tag, Url, Data), + minimal_envelope(Title, Tag, <<>>, Data), + minimal_envelope(Title, Tag, <<>>, #{}), + minimal_envelope(Title, <<>>, <<>>, #{}) + ], + Budget, + Title + ). + +-spec first_payload_within_budget([map()], non_neg_integer(), binary()) -> map(). +first_payload_within_budget([], Budget, Title) -> + title_only_payload(Title, Budget); +first_payload_within_budget([Candidate | Rest], Budget, Title) -> + case byte_size(encode_payload(Candidate)) =< Budget of + true -> Candidate; + false -> first_payload_within_budget(Rest, Budget, Title) + end. + +-spec minimal_envelope(binary(), binary(), binary(), map()) -> map(). +minimal_envelope(Title, Tag, Url, Data) -> + #{ + <<"web_push">> => ?WEB_PUSH_MARKER, + <<"title">> => Title, + <<"tag">> => Tag, + <<"data">> => Data, + <<"notification">> => #{ + <<"title">> => Title, + <<"tag">> => Tag, + <<"navigate">> => Url, + <<"data">> => Data + } + }. + +-spec title_only_payload(binary(), non_neg_integer()) -> map(). +title_only_payload(Title, Budget) -> + Candidate = #{<<"web_push">> => ?WEB_PUSH_MARKER, <<"title">> => Title}, + case byte_size(Title) =:= 0 orelse byte_size(encode_payload(Candidate)) =< Budget of + true -> + Candidate; + false -> + title_only_payload( + push_notification_format:truncate_bytes(Title, byte_size(Title) - 1), Budget + ) + end. + +-spec minimal_url(map()) -> binary(). +minimal_url(Payload) -> + case first_text(Payload, <<"navigate">>, <<>>) of + <<>> -> data_text(Payload, <<"url">>); + Url -> Url + end. + +-spec minimal_data(map()) -> map(). +minimal_data(Payload) -> + case maps:get(<<"data">>, Payload, undefined) of + Data when is_map(Data) -> maps:with(?MINIMAL_DATA_KEYS, Data); + _ -> #{} + end. + +-spec data_text(map(), binary()) -> binary(). +data_text(Payload, Key) -> + case maps:get(<<"data">>, Payload, undefined) of + Data when is_map(Data) -> text_or(maps:get(Key, Data, undefined), <<>>); + _ -> <<>> + end. + +-spec first_text(map(), binary(), binary()) -> binary(). +first_text(Payload, Key, Fallback) -> + case text_or(maps:get(Key, Payload, undefined), <<>>) of + <<>> -> + nested_text(maps:get(<<"notification">>, Payload, undefined), Key, Fallback); + Value -> + Value + end. + +-spec nested_text(term(), binary(), binary()) -> binary(). +nested_text(Notification, Key, Fallback) when is_map(Notification) -> + text_or(maps:get(Key, Notification, undefined), Fallback); +nested_text(_Notification, _Key, Fallback) -> + Fallback. + +-spec text_or(term(), binary()) -> binary(). +text_or(Value, _Fallback) when is_binary(Value), byte_size(Value) > 0 -> + Value; +text_or(_Value, Fallback) -> + Fallback. + +-spec encode_payload(map()) -> binary(). +encode_payload(Payload) -> + iolist_to_binary(json:encode(Payload)). + -ifdef(TEST). -include_lib("eunit/include/eunit.hrl"). @@ -411,6 +605,75 @@ build_clear_notification_payload_test() -> ?assertEqual(<<"channel:456">>, maps:get(<<"notification_tag">>, Data)), ?assertEqual(2, maps:get(<<"badge_count">>, Data)). +is_clear_test() -> + ?assertEqual(true, is_clear(build_clear_notification_payload(999, 456, 789, 2))), + ?assertEqual( + false, + is_clear( + test_notification_payload(#{<<"content">> => <<"hi">>}, 0, undefined, undefined) + ) + ). + +fit_payload_json_leaves_a_payload_within_budget_untouched_test() -> + Payload = test_notification_payload(#{<<"content">> => <<"hi">>}, 0, undefined, undefined), + Encoded = encode_payload(Payload), + ?assertEqual(Encoded, fit_payload_json(Payload, byte_size(Encoded))). + +fit_payload_json_drops_media_first_test() -> + ImageUrl = <<"https://media.example/external/", (binary:copy(<<"i">>, 400))/binary>>, + MessageData = #{ + <<"content">> => <<"Photo">>, + <<"mentions">> => [], + <<"attachments">> => [ + #{<<"content_type">> => <<"image/png">>, <<"proxy_url">> => ImageUrl} + ] + }, + Payload = test_notification_payload(MessageData, 123, <<"Server">>, <<"general">>), + ?assert(byte_size(encode_payload(Payload)) > 2713), + Fitted = fit_payload_json(Payload, 2713), + ?assert(byte_size(Fitted) =< 2713), + ?assertEqual(nomatch, binary:match(Fitted, ImageUrl)), + Decoded = json:decode(Fitted), + ?assertEqual(<<"Photo">>, maps:get(<<"body">>, Decoded)), + ?assertEqual(<<"http://avatar">>, maps:get(<<"icon">>, Decoded)). + +fit_payload_json_falls_back_to_a_minimal_payload_test() -> + Payload = test_notification_payload( + #{<<"content">> => <<"hi">>, <<"mentions">> => []}, + 123, + binary:copy(<<"g">>, 4000), + binary:copy(<<"c">>, 4000) + ), + Fitted = fit_payload_json(Payload, 2713), + ?assert(byte_size(Fitted) =< 2713), + Decoded = json:decode(Fitted), + ?assertEqual(<<"channel:456:789">>, maps:get(<<"tag">>, Decoded)), + ?assertEqual(?MINIMAL_TITLE_MAX_BYTES, byte_size(maps:get(<<"title">>, Decoded))). + +fit_payload_json_reaches_a_title_only_payload_test() -> + Payload = test_notification_payload( + #{<<"content">> => <<"hi">>, <<"mentions">> => []}, + 123, + binary:copy(<<"g">>, 4000), + binary:copy(<<"c">>, 4000) + ), + Fitted = fit_payload_json(Payload, 60), + ?assert(byte_size(Fitted) =< 60), + ?assertEqual([<<"title">>, <<"web_push">>], lists:sort(maps:keys(json:decode(Fitted)))). + +fit_payload_json_never_splits_a_utf8_character_test() -> + Emoji = binary:copy(<<"\xF0\x9F\x98\x80">>, 2000), + Payload = test_notification_payload( + #{<<"content">> => <<"hi">>, <<"mentions">> => []}, + 123, + Emoji, + <<"a", Emoji/binary>> + ), + Fitted = fit_payload_json(Payload, 2713), + ?assert(byte_size(Fitted) =< 2713), + Title = maps:get(<<"title">>, json:decode(Fitted)), + ?assertMatch(Bin when is_binary(Bin), unicode:characters_to_binary(Title, utf8, utf8)). + test_notification_payload(MessageData, GuildId, GuildName, ChannelName) -> test_notification_payload(MessageData, GuildId, GuildName, ChannelName, #{}). diff --git a/fluxer_gateway/src/push/push_notification_format.erl b/fluxer_gateway/src/push/push_notification_format.erl index 027098d6e..a584075d4 100644 --- a/fluxer_gateway/src/push/push_notification_format.erl +++ b/fluxer_gateway/src/push/push_notification_format.erl @@ -8,12 +8,15 @@ build_content_preview/2, build_markdown_context/4, resolve_author_name/3, + resolve_author_avatar_url/1, extract_image_url/1, maybe_image_fields/1, - build_url/3 + build_url/3, + truncate_bytes/2 ]). -define(MAX_MENTIONS_FOR_PUSH, 50). +-define(MAX_PREVIEW_BYTES, 100). -define(CHANNEL_TYPE_GUILD_TEXT, 0). -define(CHANNEL_TYPE_GUILD_VOICE, 2). -define(CHANNEL_TYPE_GUILD_CATEGORY, 4). @@ -112,6 +115,32 @@ user_nicknames_from_context_or_message(MessageData, MarkdownContext) when user_nicknames_from_context_or_message(MessageData, _MarkdownContext) -> group_dm_user_nicknames(MessageData). +-spec resolve_author_avatar_url(map()) -> binary(). +resolve_author_avatar_url(AuthorData) -> + resolve_avatar_url(AuthorData, maps:get(<<"avatar">>, AuthorData, null)). + +-spec resolve_avatar_url(map(), binary() | null) -> binary(). +resolve_avatar_url(AuthorData, null) -> + default_avatar_url(author_id_binary(AuthorData)); +resolve_avatar_url(AuthorData, Hash) -> + case author_id_binary(AuthorData) of + undefined -> default_avatar_url(undefined); + UserId -> push_utils:construct_avatar_url(UserId, Hash) + end. + +-spec author_id_binary(map()) -> binary() | undefined. +author_id_binary(AuthorData) -> + case snowflake_id:parse_optional(maps:get(<<"id">>, AuthorData, undefined)) of + undefined -> undefined; + UserId -> integer_to_binary(UserId) + end. + +-spec default_avatar_url(binary() | undefined) -> binary(). +default_avatar_url(undefined) -> + push_utils:get_default_avatar_url(<<>>); +default_avatar_url(UserId) -> + push_utils:get_default_avatar_url(UserId). + -spec user_nicknames(map(), non_neg_integer(), map()) -> map(). user_nicknames(MessageData, 0, _GuildData) -> group_dm_user_nicknames(MessageData); @@ -325,9 +354,13 @@ first_nonempty_binary([Value | Rest]) -> end. -spec truncate_preview(binary()) -> binary(). -truncate_preview(Content) when byte_size(Content) > 100 -> - valid_utf8_prefix(binary:part(Content, 0, 100)); truncate_preview(Content) -> + truncate_bytes(Content, ?MAX_PREVIEW_BYTES). + +-spec truncate_bytes(binary(), non_neg_integer()) -> binary(). +truncate_bytes(Content, MaxBytes) when byte_size(Content) > MaxBytes -> + valid_utf8_prefix(binary:part(Content, 0, MaxBytes)); +truncate_bytes(Content, _MaxBytes) -> valid_utf8_prefix(Content). -spec valid_utf8_prefix(binary()) -> binary(). diff --git a/fluxer_gateway/src/push/push_outbox.erl b/fluxer_gateway/src/push/push_outbox.erl new file mode 100644 index 000000000..e546154f3 --- /dev/null +++ b/fluxer_gateway/src/push/push_outbox.erl @@ -0,0 +1,661 @@ +%% SPDX-License-Identifier: AGPL-3.0-or-later + +-module(push_outbox). +-typing([eqwalizer]). +-behaviour(gen_server). + +-export([ + start_link/0, + enqueue/1, + truncate_read/3, + note_session_active/1, + delivery_config_changed/0, + stats/0, + request_timeout_ms/0 +]). +-export([init/1, handle_call/3, handle_cast/2, handle_info/2, terminate/2, code_change/3]). +-export_type([job/0, fallback/0]). + +-define(DEFAULT_MAX_QUEUE, 10000). +-define(DEFAULT_MAX_INFLIGHT, 64). +-define(DEFAULT_REQUEST_TIMEOUT_MS, 100000). +-define(DEFAULT_MAX_AGE_MS, 300000). +-define(DEFAULT_RETRY_BASE_MS, 1000). +-define(DEFAULT_MAX_FALLBACK_RUNNERS, 16). +-define(RETRY_MAX_MS, 30000). +-define(ENQUEUE_TIMEOUT_MS, 1000). +-define(STATS_TIMEOUT_MS, 1000). +-define(PRUNE_INTERVAL_MS, 30000). + +-type kind() :: message | clear. +-type fallback() :: fun(([integer()]) -> term()). +-type job() :: #{ + kind := kind(), + subject := binary(), + job := map(), + body := binary(), + user_ids := [integer()], + channel_id := integer(), + message_id := integer(), + fallback := fallback() +}. +-type entry() :: #{ + kind := kind(), + subject := binary(), + job := map(), + body := binary(), + user_ids := [integer()], + channel_id := integer(), + message_id := integer(), + fallback := fallback(), + seq := non_neg_integer(), + enqueued_at := integer(), + attempts := non_neg_integer(), + config_version := non_neg_integer() | undefined +}. +-type settled() :: {keep, entry(), state()} | {none, state()}. +-type counter() :: + delivered + | retries + | sheds + | truncations + | skipped_active + | fallbacks + | lost + | enqueued. +-type state() :: #{ + jobs := gb_trees:tree(non_neg_integer(), entry()), + ready := queue:queue(non_neg_integer()), + inflight := #{pid() => {reference(), reference(), entry()}}, + fallback_backlog := queue:queue({fallback(), [integer()]}), + fallback_runners := #{pid() => reference()}, + next_seq := non_neg_integer(), + reads := #{{integer(), integer()} => {integer(), integer()}}, + active := #{integer() => {non_neg_integer(), integer()}}, + counters := #{counter() => non_neg_integer()}, + max_queue := pos_integer(), + max_inflight := pos_integer(), + max_fallback_runners := pos_integer(), + request_timeout_ms := pos_integer(), + max_age_ms := pos_integer(), + retry_base_ms := pos_integer() +}. + +-spec start_link() -> {ok, pid()} | {error, term()} | ignore. +start_link() -> + gen_server:start_link({local, ?MODULE}, ?MODULE, [], []). + +-spec enqueue(job()) -> ok | {error, term()}. +enqueue(Job) -> + try gen_server:call(?MODULE, {enqueue, Job}, ?ENQUEUE_TIMEOUT_MS) of + ok -> ok + catch + exit:{noproc, _} -> {error, outbox_unavailable}; + exit:{Reason, _} -> {error, {outbox_unavailable, Reason}} + end. + +-spec truncate_read(integer(), integer(), integer()) -> ok. +truncate_read(UserId, ChannelId, MessageId) -> + broadcast({truncate_read, UserId, ChannelId, MessageId}). + +-spec note_session_active(integer()) -> ok. +note_session_active(UserId) -> + broadcast({session_active, UserId}). + +-spec delivery_config_changed() -> ok. +delivery_config_changed() -> + gen_server:cast(?MODULE, delivery_config_changed). + +-spec stats() -> map(). +stats() -> + try gen_server:call(?MODULE, stats, ?STATS_TIMEOUT_MS) of + Stats when is_map(Stats) -> Stats + catch + exit:_ -> #{} + end. + +-spec request_timeout_ms() -> pos_integer(). +request_timeout_ms() -> + env_pos_integer(push_outbox_request_timeout_ms, ?DEFAULT_REQUEST_TIMEOUT_MS). + +-spec init([]) -> {ok, state()}. +init([]) -> + erlang:process_flag(fullsweep_after, 10), + schedule_prune(), + {ok, #{ + jobs => gb_trees:empty(), + ready => queue:new(), + inflight => #{}, + fallback_backlog => queue:new(), + fallback_runners => #{}, + next_seq => 0, + reads => #{}, + active => #{}, + counters => #{}, + max_queue => env_pos_integer(push_outbox_max_queue, ?DEFAULT_MAX_QUEUE), + max_inflight => env_pos_integer(push_outbox_max_inflight, ?DEFAULT_MAX_INFLIGHT), + max_fallback_runners => app_pos_integer( + push_outbox_max_fallback_runners, ?DEFAULT_MAX_FALLBACK_RUNNERS + ), + request_timeout_ms => request_timeout_ms(), + max_age_ms => env_pos_integer(push_outbox_max_age_ms, ?DEFAULT_MAX_AGE_MS), + retry_base_ms => app_pos_integer(push_outbox_retry_base_ms, ?DEFAULT_RETRY_BASE_MS) + }}. + +-spec handle_call(term(), gen_server:from(), state()) -> {reply, term(), state()}. +handle_call({enqueue, Job}, _From, State) when + is_map_key(kind, Job), + is_map_key(subject, Job), + is_map_key(body, Job), + is_map_key(user_ids, Job), + is_map_key(channel_id, Job), + is_map_key(message_id, Job), + is_map_key(fallback, Job) +-> + {reply, ok, pump(admit(Job, State))}; +handle_call(stats, _From, State) -> + {reply, build_stats(State), State}; +handle_call(_Request, _From, State) -> + {reply, {error, unknown_request}, State}. + +-spec handle_cast(term(), state()) -> {noreply, state()}. +handle_cast({truncate_read, UserId, ChannelId, MessageId}, State) when + is_integer(UserId), is_integer(ChannelId), is_integer(MessageId) +-> + {noreply, apply_read(UserId, ChannelId, MessageId, State)}; +handle_cast({session_active, UserId}, State) when is_integer(UserId) -> + {noreply, record_active(UserId, State)}; +handle_cast(delivery_config_changed, State) -> + {noreply, drain(State)}; +handle_cast(_Msg, State) -> + {noreply, State}. + +-spec handle_info(term(), state()) -> {noreply, state()}. +handle_info({push_outbox_reply, Pid, Result}, State) when is_pid(Pid) -> + {noreply, pump(finish_worker(Pid, reply_result(Result), State))}; +handle_info({'DOWN', _MRef, process, Pid, Reason}, #{fallback_runners := Runners} = State) when + is_map_key(Pid, Runners) +-> + {noreply, finish_fallback(Pid, Reason, State)}; +handle_info({'DOWN', _MRef, process, Pid, Reason}, State) when is_pid(Pid) -> + {noreply, pump(finish_worker(Pid, down_result(Reason), State))}; +handle_info({request_deadline, Pid}, State) when is_pid(Pid) -> + kill_expired_worker(Pid, State), + {noreply, State}; +handle_info({retry, Seq}, State) when is_integer(Seq), Seq >= 0 -> + {noreply, pump(make_ready(Seq, State))}; +handle_info(prune, State) -> + schedule_prune(), + {noreply, prune(State)}; +handle_info(_Info, State) -> + {noreply, State}. + +-spec terminate(term(), state()) -> ok. +terminate(_Reason, _State) -> + ok. + +-spec code_change(term(), state(), term()) -> {ok, state()}. +code_change(_OldVsn, State, _Extra) -> + {ok, State}. + +-spec admit(job(), state()) -> state(). +admit(Job, #{next_seq := Seq} = State) -> + Entry = maps:merge(Job, #{ + seq => Seq, + enqueued_at => now_ms(), + attempts => 0, + config_version => job_config_version(Job) + }), + State1 = bump(enqueued, 1, State#{next_seq := Seq + 1}), + queue_settled(settle_if_stale(Entry, State1)). + +-spec job_config_version(job()) -> non_neg_integer() | undefined. +job_config_version(#{job := #{<<"config_version">> := Version}}) when + is_integer(Version), Version >= 0 +-> + Version; +job_config_version(_Job) -> + undefined. + +-spec queue_settled(settled()) -> state(). +queue_settled({keep, Entry, State}) -> + shed_to_capacity(insert(Entry, State)); +queue_settled({none, State}) -> + State. + +-spec insert(entry(), state()) -> state(). +insert(#{seq := Seq} = Entry, #{jobs := Jobs, ready := Ready} = State) -> + State#{jobs := gb_trees:enter(Seq, Entry, Jobs), ready := queue:in(Seq, Ready)}. + +-spec shed_to_capacity(state()) -> state(). +shed_to_capacity(#{jobs := Jobs, max_queue := MaxQueue} = State) -> + case gb_trees:size(Jobs) > MaxQueue of + true -> + {_Seq, Shed, Rest} = gb_trees:take_smallest(Jobs), + shed_to_capacity(shed(Shed, State#{jobs := Rest})); + false -> + State + end. + +-spec shed(entry(), state()) -> state(). +shed(Entry, State) -> + log_shed(Entry), + bump(sheds, 1, State). + +-spec log_shed(entry()) -> ok. +log_shed(#{kind := Kind, channel_id := ChannelId, message_id := MessageId}) -> + logger:warning( + "Push outbox at capacity, dropping the earliest queued job", + #{kind => Kind, channel_id => ChannelId, message_id => MessageId} + ). + +-spec pump(state()) -> state(). +pump(#{inflight := Inflight, max_inflight := MaxInflight} = State) when + map_size(Inflight) >= MaxInflight +-> + State; +pump(#{ready := Ready, jobs := Jobs} = State) -> + case queue:out(Ready) of + {empty, _} -> + State; + {{value, Seq}, Rest} -> + pump(take_ready(gb_trees:lookup(Seq, Jobs), Seq, State#{ready := Rest})) + end. + +-spec take_ready(none | {value, entry()}, non_neg_integer(), state()) -> state(). +take_ready(none, _Seq, State) -> + State; +take_ready({value, Entry}, Seq, #{jobs := Jobs} = State) -> + dispatch(Entry, State#{jobs := gb_trees:delete(Seq, Jobs)}). + +-spec dispatch(entry(), state()) -> state(). +dispatch(Entry, State) -> + case prepare(Entry, State) of + {skip, State1} -> State1; + {send, Prepared, State1} -> send_or_fall_back(Prepared, State1) + end. + +-spec send_or_fall_back(entry(), state()) -> state(). +send_or_fall_back(Entry, State) -> + case is_expired(Entry, State) of + true -> fall_back(Entry, State); + false -> start_worker(Entry, State) + end. + +-spec prepare(entry(), state()) -> {skip, state()} | {send, entry(), state()}. +prepare(#{kind := clear} = Entry, State) -> + {send, Entry, State}; +prepare(#{user_ids := UserIds} = Entry, #{reads := Reads, active := Active} = State) -> + #{channel_id := ChannelId, message_id := MessageId, seq := Seq} = Entry, + {Read, Unread} = lists:partition( + fun(UserId) -> is_read(UserId, ChannelId, MessageId, Reads) end, UserIds + ), + {Activated, Kept} = lists:partition( + fun(UserId) -> became_active(UserId, Seq, Active) end, Unread + ), + State1 = bump(skipped_active, length(Activated), bump(truncations, length(Read), State)), + case Kept of + [] -> {skip, State1}; + UserIds -> {send, Entry, State1}; + _ -> {send, with_user_ids(Kept, Entry), State1} + end. + +-spec is_read(integer(), integer(), integer(), #{ + {integer(), integer()} => {integer(), integer()} +}) -> + boolean(). +is_read(UserId, ChannelId, MessageId, Reads) -> + case maps:find({UserId, ChannelId}, Reads) of + {ok, {ReadMessageId, _At}} -> MessageId =< ReadMessageId; + error -> false + end. + +-spec became_active(integer(), non_neg_integer(), #{integer() => {non_neg_integer(), integer()}}) -> + boolean(). +became_active(UserId, Seq, Active) -> + case maps:find(UserId, Active) of + {ok, {ActiveSeq, _At}} -> ActiveSeq > Seq; + error -> false + end. + +-spec with_user_ids([integer()], entry()) -> entry(). +with_user_ids(UserIds, #{job := Job} = Entry) -> + Rewritten = Job#{<<"user_ids">> => [integer_to_binary(UserId) || UserId <- UserIds]}, + Entry#{ + user_ids := UserIds, + job := Rewritten, + body := iolist_to_binary(json:encode(Rewritten)) + }. + +-spec apply_read(integer(), integer(), integer(), state()) -> state(). +apply_read(UserId, ChannelId, MessageId, #{reads := Reads, jobs := Jobs} = State) -> + Key = {UserId, ChannelId}, + Watermark = + case maps:find(Key, Reads) of + {ok, {Previous, _At}} -> max(Previous, MessageId); + error -> MessageId + end, + State1 = State#{reads := Reads#{Key => {Watermark, now_ms()}}}, + lists:foldl( + fun({Seq, Entry}, Acc) -> + truncate_entry(Seq, Entry, UserId, ChannelId, MessageId, Acc) + end, + State1, + gb_trees:to_list(Jobs) + ). + +-spec truncate_entry(non_neg_integer(), entry(), integer(), integer(), integer(), state()) -> + state(). +truncate_entry( + Seq, + #{kind := message, channel_id := ChannelId, message_id := JobMessageId} = Entry, + UserId, + ChannelId, + MessageId, + State +) when JobMessageId =< MessageId -> + remove_reader(Seq, Entry, UserId, State); +truncate_entry(_Seq, _Entry, _UserId, _ChannelId, _MessageId, State) -> + State. + +-spec remove_reader(non_neg_integer(), entry(), integer(), state()) -> state(). +remove_reader(Seq, #{user_ids := UserIds} = Entry, UserId, #{jobs := Jobs} = State) -> + case lists:member(UserId, UserIds) of + false -> + State; + true -> + State1 = bump(truncations, 1, State), + case lists:delete(UserId, UserIds) of + [] -> + State1#{jobs := gb_trees:delete(Seq, Jobs)}; + Remaining -> + State1#{jobs := gb_trees:update(Seq, with_user_ids(Remaining, Entry), Jobs)} + end + end. + +-spec record_active(integer(), state()) -> state(). +record_active(UserId, #{active := Active, next_seq := Seq} = State) -> + State#{active := Active#{UserId => {Seq, now_ms()}}, next_seq := Seq + 1}. + +-spec start_worker(entry(), state()) -> state(). +start_worker(Entry, #{inflight := Inflight, request_timeout_ms := Timeout} = State) -> + #{subject := Subject, body := Body} = Entry, + Outbox = self(), + {Pid, MRef} = spawn_monitor(fun() -> + Outbox ! {push_outbox_reply, self(), send(Subject, Body, Timeout)} + end), + TRef = erlang:send_after(Timeout, Outbox, {request_deadline, Pid}), + State#{inflight := Inflight#{Pid => {MRef, TRef, Entry}}}. + +-spec send(binary(), binary(), pos_integer()) -> ok | {error, term()}. +send(Subject, Body, Timeout) -> + try push_job_publisher:request(Subject, Body, Timeout) of + ok -> ok; + {error, Reason} -> {error, Reason} + catch + Class:Reason -> {error, {Class, Reason}} + end. + +-spec kill_expired_worker(pid(), state()) -> ok. +kill_expired_worker(Pid, #{inflight := Inflight}) -> + case maps:is_key(Pid, Inflight) of + true -> + exit(Pid, kill), + ok; + false -> + ok + end. + +-spec finish_worker(pid(), ok | {error, term()}, state()) -> state(). +finish_worker(Pid, Result, #{inflight := Inflight} = State) -> + case maps:take(Pid, Inflight) of + {{MRef, TRef, Entry}, Rest} -> + erlang:demonitor(MRef, [flush]), + _ = erlang:cancel_timer(TRef, [{async, true}, {info, false}]), + handle_result(Result, Entry, State#{inflight := Rest}); + error -> + State + end. + +-spec reply_result(term()) -> ok | {error, term()}. +reply_result(ok) -> ok; +reply_result({error, Reason}) -> {error, Reason}; +reply_result(Other) -> {error, {invalid_worker_reply, Other}}. + +-spec down_result(term()) -> {error, term()}. +down_result(killed) -> {error, timeout}; +down_result(Reason) -> {error, {worker_down, Reason}}. + +-spec handle_result(ok | {error, term()}, entry(), state()) -> state(). +handle_result(ok, _Entry, State) -> + bump(delivered, 1, State); +handle_result({error, Reason}, Entry, State) -> + logger:debug("Push outbox request not delivered", #{ + reason => Reason, + kind => maps:get(kind, Entry), + message_id => maps:get(message_id, Entry), + attempts => maps:get(attempts, Entry) + }), + case is_expired(Entry, State) of + true -> fall_back(Entry, State); + false -> retry_settled(settle_if_stale(Entry, State)) + end. + +-spec retry_settled(settled()) -> state(). +retry_settled({keep, Entry, State}) -> + schedule_retry(Entry, State); +retry_settled({none, State}) -> + State. + +-spec schedule_retry(entry(), state()) -> state(). +schedule_retry(#{seq := Seq, attempts := Attempts} = Entry, #{jobs := Jobs} = State) -> + NextAttempts = Attempts + 1, + _ = erlang:send_after(retry_delay(NextAttempts, State), self(), {retry, Seq}), + State1 = State#{jobs := gb_trees:enter(Seq, Entry#{attempts := NextAttempts}, Jobs)}, + shed_to_capacity(bump(retries, 1, State1)). + +-spec retry_delay(pos_integer(), state()) -> pos_integer(). +retry_delay(Attempts, #{retry_base_ms := Base}) -> + Delay = min(?RETRY_MAX_MS, Base bsl min(Attempts - 1, 16)), + Delay + rand:uniform(max(1, Delay div 4)). + +-spec make_ready(non_neg_integer(), state()) -> state(). +make_ready(Seq, #{jobs := Jobs, ready := Ready} = State) -> + case gb_trees:is_defined(Seq, Jobs) of + true -> State#{ready := queue:in(Seq, Ready)}; + false -> State + end. + +-spec is_expired(entry(), state()) -> boolean(). +is_expired(#{enqueued_at := EnqueuedAt}, #{max_age_ms := MaxAge}) -> + now_ms() - EnqueuedAt >= MaxAge. + +-spec fall_back(entry(), state()) -> state(). +fall_back(#{user_ids := UserIds} = Entry, State) -> + logger:warning("Push outbox job expired undelivered, falling back to the gateway path", #{ + kind => maps:get(kind, Entry), + channel_id => maps:get(channel_id, Entry), + message_id => maps:get(message_id, Entry), + attempts => maps:get(attempts, Entry), + user_count => length(UserIds) + }), + hand_back(UserIds, Entry, State). + +-spec hand_back([integer()], entry(), state()) -> state(). +hand_back(UserIds, #{fallback := Fallback}, #{fallback_backlog := Backlog} = State) -> + Queued = State#{fallback_backlog := queue:in({Fallback, UserIds}, Backlog)}, + run_fallbacks(bump(fallbacks, 1, Queued)). + +-spec run_fallbacks(state()) -> state(). +run_fallbacks(#{fallback_runners := Runners, max_fallback_runners := MaxRunners} = State) when + map_size(Runners) >= MaxRunners +-> + State; +run_fallbacks(#{fallback_backlog := Backlog, fallback_runners := Runners} = State) -> + case queue:out(Backlog) of + {empty, _} -> + State; + {{value, {Fallback, UserIds}}, Rest} -> + {Pid, MRef} = spawn_monitor(fun() -> run_fallback(Fallback, UserIds) end), + run_fallbacks(State#{ + fallback_backlog := Rest, fallback_runners := Runners#{Pid => MRef} + }) + end. + +-spec finish_fallback(pid(), term(), state()) -> state(). +finish_fallback(Pid, Reason, #{fallback_runners := Runners} = State) -> + run_fallbacks( + count_fallback_exit(Reason, State#{fallback_runners := maps:remove(Pid, Runners)}) + ). + +-spec count_fallback_exit(term(), state()) -> state(). +count_fallback_exit(normal, State) -> + State; +count_fallback_exit(_Reason, State) -> + bump(lost, 1, State). + +-spec drain(state()) -> state(). +drain(#{jobs := Jobs} = State) -> + Drained = lists:foldl(fun resettle/2, State, gb_trees:to_list(Jobs)), + log_drain(count(fallbacks, Drained) - count(fallbacks, State), Drained), + Drained. + +-spec resettle({non_neg_integer(), entry()}, state()) -> state(). +resettle({Seq, Entry}, State) -> + case settle(Entry, State) of + {keep, Kept, #{jobs := Jobs} = State1} -> + State1#{jobs := gb_trees:update(Seq, Kept, Jobs)}; + {none, #{jobs := Jobs} = State1} -> + State1#{jobs := gb_trees:delete(Seq, Jobs)} + end. + +-spec settle_if_stale(entry(), state()) -> settled(). +settle_if_stale(#{config_version := Version} = Entry, State) -> + case push_delivery_config:config_version() of + Version -> {keep, Entry, State}; + _Changed -> settle(Entry, State) + end. + +-spec settle(entry(), state()) -> settled(). +settle(Entry, State) -> + case prepare(Entry, State) of + {skip, State1} -> + {none, State1}; + {send, Prepared, State1} -> + hand_back_unenrolled(push_delivery_config:config(), Prepared, State1) + end. + +-spec hand_back_unenrolled(push_delivery_config:config(), entry(), state()) -> settled(). +hand_back_unenrolled(Config, #{user_ids := UserIds} = Entry, State) -> + Settled = Entry#{config_version := maps:get(config_version, Config)}, + case push_delivery_config:partition_users(Config, UserIds) of + {UserIds, []} -> + {keep, Settled, State}; + {[], Unenrolled} -> + {none, hand_back(Unenrolled, Entry, State)}; + {Enrolled, Unenrolled} -> + {keep, with_user_ids(Enrolled, Settled), hand_back(Unenrolled, Entry, State)} + end. + +-spec log_drain(integer(), state()) -> ok. +log_drain(HandedBack, _State) when HandedBack =< 0 -> + ok; +log_drain(HandedBack, #{jobs := Jobs}) -> + logger:notice( + "Push outbox handed queued jobs back to the gateway path after a delivery config change", + #{ + handed_back => HandedBack, + config_version => push_delivery_config:config_version(), + depth => gb_trees:size(Jobs) + } + ). + +-spec run_fallback(fallback(), [integer()]) -> ok. +run_fallback(Fallback, UserIds) -> + try Fallback(UserIds) of + _ -> ok + catch + Class:Reason -> + logger:error("Push outbox fallback crashed", #{class => Class, reason => Reason}), + exit(fallback_crashed) + end. + +-spec prune(state()) -> state(). +prune(#{reads := Reads, active := Active, max_age_ms := MaxAge} = State) -> + Cutoff = now_ms() - MaxAge, + State#{ + reads := maps:filter(fun(_Key, {_MessageId, At}) -> At >= Cutoff end, Reads), + active := maps:filter(fun(_UserId, {_Seq, At}) -> At >= Cutoff end, Active) + }. + +-spec build_stats(state()) -> map(). +build_stats(#{ + jobs := Jobs, + inflight := Inflight, + fallback_backlog := Backlog, + fallback_runners := Runners, + counters := Counters +}) -> + maps:merge( + #{ + delivered => 0, + retries => 0, + sheds => 0, + truncations => 0, + skipped_active => 0, + fallbacks => 0, + lost => 0, + enqueued => 0 + }, + Counters#{ + depth => gb_trees:size(Jobs), + inflight => map_size(Inflight), + fallback_backlog => queue:len(Backlog), + fallback_runners => map_size(Runners) + } + ). + +-spec count(counter(), state()) -> non_neg_integer(). +count(Counter, #{counters := Counters}) -> + maps:get(Counter, Counters, 0). + +-spec bump(counter(), non_neg_integer(), state()) -> state(). +bump(_Counter, 0, State) -> + State; +bump(Counter, Increment, #{counters := Counters} = State) -> + State#{counters := Counters#{Counter => maps:get(Counter, Counters, 0) + Increment}}. + +-spec broadcast(term()) -> ok. +broadcast(Msg) -> + abcast = gen_server:abcast(push_nodes(), ?MODULE, Msg), + ok. + +-spec push_nodes() -> [node()]. +push_nodes() -> + try gateway_node_router:active_nodes(push) of + Nodes -> Nodes + catch + _:_ -> [node()] + end. + +-spec schedule_prune() -> reference(). +schedule_prune() -> + erlang:send_after(?PRUNE_INTERVAL_MS, self(), prune). + +-spec now_ms() -> integer(). +now_ms() -> + erlang:monotonic_time(millisecond). + +-spec env_pos_integer(atom(), pos_integer()) -> pos_integer(). +env_pos_integer(Key, Default) -> + case fluxer_gateway_env:get_optional(Key) of + Value when is_integer(Value), Value > 0 -> Value; + _ -> Default + end. + +-spec app_pos_integer(atom(), pos_integer()) -> pos_integer(). +app_pos_integer(Key, Default) -> + case application:get_env(fluxer_gateway, Key, undefined) of + Value when is_integer(Value), Value > 0 -> Value; + _ -> Default + end. diff --git a/fluxer_gateway/src/push/push_sender.erl b/fluxer_gateway/src/push/push_sender.erl index 7d7c05bcd..ec3eea68d 100644 --- a/fluxer_gateway/src/push/push_sender.erl +++ b/fluxer_gateway/src/push/push_sender.erl @@ -57,8 +57,7 @@ notification_payload(UserId, SendContext) -> } = SendContext, AuthorData = maps:get(<<"author">>, MessageData, #{}), AuthorUsername = maps:get(<<"username">>, AuthorData, <<"Unknown">>), - AuthorAvatar = maps:get(<<"avatar">>, AuthorData, null), - AuthorAvatarUrl = resolve_avatar_url(AuthorData, AuthorAvatar), + AuthorAvatarUrl = push_notification_format:resolve_author_avatar_url(AuthorData), push_notification:build_notification_payload(#{ message_data => MessageData, guild_id => GuildId, @@ -162,28 +161,6 @@ send_clear_channel_notifications(UserId, ChannelId, MessageId, BadgeCountsTtlSec ), ok. --spec resolve_avatar_url(map(), binary() | null) -> binary(). -resolve_avatar_url(AuthorData, null) -> - default_avatar_url(author_id_binary(AuthorData)); -resolve_avatar_url(AuthorData, Hash) -> - case author_id_binary(AuthorData) of - undefined -> default_avatar_url(undefined); - UserId -> push_utils:construct_avatar_url(UserId, Hash) - end. - --spec author_id_binary(map()) -> binary() | undefined. -author_id_binary(AuthorData) -> - case snowflake_id:parse_optional(maps:get(<<"id">>, AuthorData, undefined)) of - undefined -> undefined; - UserId -> integer_to_binary(UserId) - end. - --spec default_avatar_url(binary() | undefined) -> binary(). -default_avatar_url(undefined) -> - push_utils:get_default_avatar_url(<<>>); -default_avatar_url(UserId) -> - push_utils:get_default_avatar_url(UserId). - -spec handle_failed_subscriptions(integer(), list()) -> ok. handle_failed_subscriptions(_UserId, []) -> ok; @@ -208,28 +185,46 @@ send_subscriptions(UserId, Payload, [Subscription | Rest], FailedAcc) -> -spec send_notification_to_subscription(integer(), map(), map()) -> false | {true, map()}. send_notification_to_subscription(UserId, Subscription, Payload) -> + Platform = subscription_platform(Subscription), + WebPushShape = web_push_subscription(Subscription), logger:debug("Push: sending to subscription", #{ user_id => UserId, endpoint => maps:get(<<"endpoint">>, Subscription, undefined), - platform => subscription_platform(Subscription) + platform => Platform, + web_push_shape => WebPushShape }), - case subscription_platform(Subscription) of - <<"web_push">> -> + case WebPushShape of + true -> push_sender_delivery:send_webpush_notification(UserId, Subscription, Payload); - <<"android_unified_push">> -> - push_sender_delivery:send_webpush_notification(UserId, Subscription, Payload); - <<"android_fcm">> -> - push_fcm:send(UserId, Subscription, Payload); - <<"ios_apns">> -> - push_apns:send(UserId, Subscription, Payload); - Platform -> - logger:warning( - "Push: unsupported subscription platform", - #{user_id => UserId, platform => Platform} - ), - false + false -> + send_platform_notification(UserId, Platform, Subscription, Payload) end. +-spec send_platform_notification(integer(), binary(), map(), map()) -> false | {true, map()}. +send_platform_notification(UserId, <<"web_push">>, Subscription, Payload) -> + push_sender_delivery:send_webpush_notification(UserId, Subscription, Payload); +send_platform_notification(UserId, <<"android_unified_push">>, Subscription, Payload) -> + push_sender_delivery:send_webpush_notification(UserId, Subscription, Payload); +send_platform_notification(UserId, <<"android_fcm">>, Subscription, Payload) -> + push_fcm:send(UserId, Subscription, Payload); +send_platform_notification(UserId, <<"ios_apns">>, Subscription, Payload) -> + push_apns:send(UserId, Subscription, Payload); +send_platform_notification(UserId, Platform, _Subscription, _Payload) -> + logger:warning( + "Push: unsupported subscription platform", + #{user_id => UserId, platform => Platform} + ), + false. + +-spec web_push_subscription(map()) -> boolean(). +web_push_subscription(Subscription) -> + subscription_key_present(maps:get(<<"p256dh_key">>, Subscription, null)) andalso + subscription_key_present(maps:get(<<"auth_key">>, Subscription, null)). + +-spec subscription_key_present(term()) -> boolean(). +subscription_key_present(Value) -> + is_binary(Value) andalso byte_size(Value) > 0. + -spec subscription_platform(map()) -> binary(). subscription_platform(Subscription) -> Platform = maps:get(<<"platform">>, Subscription, <<"web_push">>), @@ -304,27 +299,27 @@ badge_batched_user_count([Batch | Rest], Acc) -> -spec fetch_next_badge_batch([[integer()]], badge_batch_acc(), integer(), integer()) -> badge_batch_acc(). -fetch_next_badge_batch( - [Batch | Rest], {Counts, FailedBatches, DefaultedUsers, Consecutive}, CachedAt, Deadline -) -> +fetch_next_badge_batch([Batch | Rest], Acc, CachedAt, Deadline) -> + fetch_badge_count_batches( + Rest, fetch_badge_batch(Batch, Acc, CachedAt), CachedAt, Deadline + ). + +-spec fetch_badge_batch([integer()], badge_batch_acc(), integer()) -> badge_batch_acc(). +fetch_badge_batch(Batch, {Counts, FailedBatches, DefaultedUsers, Consecutive}, CachedAt) -> Request = #{ <<"type">> => <<"get_badge_counts">>, <<"user_ids">> => [integer_to_binary(UserId) || UserId <- Batch] }, - case rpc_client:call(Request) of + Fill = push_ets_cache:reserve_badge_counts(Batch), + try rpc_client:call(Request) of {ok, Data} -> BadgeData = maps:get(<<"badge_counts">>, Data, #{}), - Merged = merge_badge_data(Batch, BadgeData, Counts, CachedAt), - fetch_badge_count_batches( - Rest, {Merged, FailedBatches, DefaultedUsers, 0}, CachedAt, Deadline - ); + Merged = merge_badge_data(Batch, BadgeData, Counts, CachedAt, Fill), + {Merged, FailedBatches, DefaultedUsers, 0}; {error, _Reason} -> - fetch_badge_count_batches( - Rest, - {Counts, FailedBatches + 1, DefaultedUsers + length(Batch), Consecutive + 1}, - CachedAt, - Deadline - ) + {Counts, FailedBatches + 1, DefaultedUsers + length(Batch), Consecutive + 1} + after + push_ets_cache:release(Fill) end. -spec report_badge_fetch_failures(non_neg_integer(), non_neg_integer()) -> ok. @@ -386,13 +381,13 @@ badge_fetch_budget_ms() -> _ -> ?DEFAULT_BADGE_FETCH_BUDGET_MS end. --spec merge_badge_data([integer()], map(), map(), integer()) -> map(). -merge_badge_data(UserIds, BadgeData, Counts, CachedAt) -> +-spec merge_badge_data([integer()], map(), map(), integer(), push_ets_cache:fill()) -> map(). +merge_badge_data(UserIds, BadgeData, Counts, CachedAt, Fill) -> lists:foldl( fun(UserId, Acc) -> UserIdBin = integer_to_binary(UserId), Count = normalize_badge_count(maps:get(UserIdBin, BadgeData, 0)), - push_ets_cache:put_badge_count(UserId, Count, CachedAt), + push_ets_cache:put_badge_count(UserId, Count, CachedAt, Fill), Acc#{UserId => Count} end, Counts, @@ -406,6 +401,99 @@ normalize_badge_count(_) -> 0. -ifdef(TEST). -include_lib("eunit/include/eunit.hrl"). +web_push_row(Platform) -> + #{ + <<"subscription_id">> => <<"sub-web">>, + <<"endpoint">> => <<"https://relay.fluxer.app/push/abc">>, + <<"p256dh_key">> => <<"p256dh">>, + <<"auth_key">> => <<"auth">>, + <<"platform">> => Platform + }. + +legacy_row(Platform) -> + #{ + <<"subscription_id">> => <<"sub-legacy">>, + <<"endpoint">> => <<"raw-vendor-device-token">>, + <<"p256dh_key">> => null, + <<"auth_key">> => null, + <<"platform">> => Platform + }. + +routed_target(Subscription) -> + ok = meck:new(push_sender_delivery, [passthrough, no_link]), + ok = meck:new(push_fcm, [passthrough, no_link]), + ok = meck:new(push_apns, [passthrough, no_link]), + try + ok = meck:expect(push_sender_delivery, send_webpush_notification, fun(_U, _S, _P) -> + false + end), + ok = meck:expect(push_fcm, send, fun(_U, _S, _P) -> false end), + ok = meck:expect(push_apns, send, fun(_U, _S, _P) -> false end), + ?assertEqual(false, send_notification_to_subscription(7, Subscription, #{})), + routed_target_counts() + after + meck:unload(push_apns), + meck:unload(push_fcm), + meck:unload(push_sender_delivery) + end. + +routed_target_counts() -> + Counts = { + meck:num_calls(push_sender_delivery, send_webpush_notification, '_'), + meck:num_calls(push_fcm, send, '_'), + meck:num_calls(push_apns, send, '_') + }, + case Counts of + {1, 0, 0} -> web_push; + {0, 1, 0} -> fcm; + {0, 0, 1} -> apns; + Other -> Other + end. + +web_push_row_takes_web_push_path_on_ios_apns_test() -> + ?assertEqual(web_push, routed_target(web_push_row(<<"ios_apns">>))). + +web_push_row_takes_web_push_path_on_android_fcm_test() -> + ?assertEqual(web_push, routed_target(web_push_row(<<"android_fcm">>))). + +web_push_row_takes_web_push_path_on_android_unified_push_test() -> + ?assertEqual(web_push, routed_target(web_push_row(<<"android_unified_push">>))). + +legacy_row_takes_apns_path_on_ios_apns_test() -> + ?assertEqual(apns, routed_target(legacy_row(<<"ios_apns">>))). + +legacy_row_takes_fcm_path_on_android_fcm_test() -> + ?assertEqual(fcm, routed_target(legacy_row(<<"android_fcm">>))). + +legacy_row_takes_web_push_path_on_android_unified_push_test() -> + ?assertEqual(web_push, routed_target(legacy_row(<<"android_unified_push">>))). + +only_p256dh_key_does_not_take_web_push_path_test() -> + Subscription = maps:put(<<"auth_key">>, null, web_push_row(<<"ios_apns">>)), + ?assertEqual(apns, routed_target(Subscription)). + +only_auth_key_does_not_take_web_push_path_test() -> + Subscription = maps:put(<<"p256dh_key">>, null, web_push_row(<<"android_fcm">>)), + ?assertEqual(fcm, routed_target(Subscription)). + +missing_key_fields_do_not_take_web_push_path_test() -> + Subscription = maps:without( + [<<"p256dh_key">>, <<"auth_key">>], web_push_row(<<"ios_apns">>) + ), + ?assertEqual(apns, routed_target(Subscription)). + +empty_key_does_not_take_web_push_path_test() -> + Subscription = maps:put(<<"auth_key">>, <<>>, web_push_row(<<"android_fcm">>)), + ?assertEqual(fcm, routed_target(Subscription)). + +web_push_row_takes_web_push_path_on_unknown_platform_test() -> + ?assertEqual(web_push, routed_target(web_push_row(<<"desktop_widget">>))). + +web_push_subscription_predicate_test() -> + ?assertEqual(true, web_push_subscription(web_push_row(<<"ios_apns">>))), + ?assertEqual(false, web_push_subscription(legacy_row(<<"ios_apns">>))), + ?assertEqual(false, web_push_subscription(#{})). + chunk_badge_user_ids_uses_bounded_batches_test() -> ?assertEqual([[1, 2], [3, 4], [5]], chunk_badge_user_ids([1, 2, 3, 4, 5], 2, [])), ?assertEqual([], chunk_badge_user_ids([], 2, [])). @@ -415,7 +503,9 @@ fetch_badge_counts_in_batches_keeps_successful_batches_test() -> ok = meck:new(push_ets_cache, [passthrough, no_link]), application:set_env(fluxer_gateway, push_badge_fetch_batch_size, 2), try - ok = meck:expect(push_ets_cache, put_badge_count, fun(_UserId, _Count, _At) -> ok end), + ok = meck:expect(push_ets_cache, put_badge_count, fun(_UserId, _Count, _At, _Fill) -> + ok + end), ok = meck:expect(rpc_client, call, fun(#{<<"user_ids">> := Ids}) -> case Ids of [<<"1">>, <<"2">>] -> @@ -468,7 +558,9 @@ badge_fetch_defaults_missing_users_to_zero_test() -> ok = meck:new(rpc_client, [passthrough, no_link]), ok = meck:new(push_ets_cache, [passthrough, no_link]), try - ok = meck:expect(push_ets_cache, put_badge_count, fun(_UserId, _Count, _At) -> ok end), + ok = meck:expect(push_ets_cache, put_badge_count, fun(_UserId, _Count, _At, _Fill) -> + ok + end), ok = meck:expect(rpc_client, call, fun(#{<<"user_ids">> := Ids}) -> ?assertEqual([<<"1">>, <<"2">>, <<"3">>], Ids), {ok, #{<<"badge_counts">> => #{<<"1">> => 1}}} diff --git a/fluxer_gateway/src/push/push_sender_delivery.erl b/fluxer_gateway/src/push/push_sender_delivery.erl index e0f7d52b3..9ba1dd78b 100644 --- a/fluxer_gateway/src/push/push_sender_delivery.erl +++ b/fluxer_gateway/src/push/push_sender_delivery.erl @@ -11,6 +11,8 @@ -export_type([push_response/0]). -define(PUSH_TTL, <<"86400">>). +-define(ALERT_URGENCY, <<"high">>). +-define(CLEAR_URGENCY, <<"low">>). -define(MAX_TRANSIENT_RETRIES, 2). -define(MAX_OVERLOAD_RETRIES, 3). -define(BASE_RETRY_DELAY_MS, 200). @@ -26,11 +28,32 @@ send_webpush_notification(UserId, Subscription, Payload) -> case extract_subscription_fields(Subscription) of {ok, Endpoint, P256dhKey, AuthKey, SubscriptionId} -> - send_with_vapid(UserId, Endpoint, P256dhKey, AuthKey, SubscriptionId, Payload); + send_to_allowed_endpoint( + UserId, Endpoint, P256dhKey, AuthKey, SubscriptionId, Payload + ); {error, _Reason} -> false end. +-spec send_to_allowed_endpoint(integer(), binary(), binary(), binary(), binary(), map()) -> + false | {true, map()}. +send_to_allowed_endpoint(UserId, Endpoint, P256dhKey, AuthKey, SubscriptionId, Payload) -> + case push_endpoint_guard:check(Endpoint) of + ok -> + send_with_vapid(UserId, Endpoint, P256dhKey, AuthKey, SubscriptionId, Payload); + {error, Reason} -> + log_endpoint_rejected(UserId, SubscriptionId, Reason), + false + end. + +-spec log_endpoint_rejected(integer(), binary(), term()) -> ok. +log_endpoint_rejected(UserId, SubscriptionId, Reason) -> + logger:debug( + "Push: endpoint rejected", + #{user_id => UserId, subscription_id => SubscriptionId, reason => Reason} + ), + ok. + -spec send_with_vapid(integer(), binary(), binary(), binary(), binary(), map()) -> false | {true, map()}. send_with_vapid(UserId, Endpoint, P256dhKey, AuthKey, SubscriptionId, Payload) -> @@ -39,18 +62,16 @@ send_with_vapid(UserId, Endpoint, P256dhKey, AuthKey, SubscriptionId, Payload) - Aud = push_utils:extract_origin(Endpoint), {ok, VapidToken} ?= cached_vapid_token(Aud, VapidEmail, VapidPublicKey, VapidPrivateKey), - Headers = build_push_headers(VapidToken, VapidPublicKey), - PayloadJson = iolist_to_binary(json:encode(Payload)), - InitialRecordSize = push_sender_retry:initial_record_size_for_endpoint(Endpoint), + Headers = build_push_headers(VapidToken, VapidPublicKey, push_urgency(Payload)), send_encrypted_push( UserId, SubscriptionId, Endpoint, Headers, - PayloadJson, + Payload, P256dhKey, AuthKey, - InitialRecordSize, + push_sender_retry:initial_record_size(), 0 ) else @@ -76,7 +97,7 @@ cached_vapid_token(Aud, VapidEmail, VapidPublicKey, VapidPrivateKey) -> {ok, binary()} | {error, term()}. cached_vapid_token(Aud, VapidEmail, VapidPublicKey, VapidPrivateKey, Now) -> CacheKey = vapid_cache_key(Aud, VapidEmail, VapidPublicKey), - case push_token_cache:get(CacheKey) of + case push_ets_cache:get_bearer_token(CacheKey) of {ok, Token, ExpiresAt} when is_binary(Token), ExpiresAt - ?VAPID_TOKEN_SKEW_SECONDS > Now -> @@ -99,7 +120,7 @@ generate_cached_vapid_token(CacheKey, Aud, VapidEmail, VapidPublicKey, VapidPriv }, case safe_generate_vapid_token(VapidClaims, VapidPublicKey, VapidPrivateKey) of {ok, Token} -> - push_token_cache:put(CacheKey, Token, ExpiresAt), + push_ets_cache:put_bearer_token(CacheKey, Token, ExpiresAt), {ok, Token}; {error, Reason} -> {error, Reason} @@ -124,7 +145,7 @@ safe_generate_vapid_token(VapidClaims, VapidPublicKey, VapidPrivateKey) -> binary(), binary(), [{binary(), binary()}], - binary(), + map(), binary(), binary(), pos_integer(), @@ -135,12 +156,13 @@ send_encrypted_push( SubscriptionId, Endpoint, Headers, - PayloadJson, + Payload, P256dhKey, AuthKey, RecordSize, Attempt ) -> + PayloadJson = fit_payload_json(Payload, RecordSize), case push_utils:encrypt_payload(PayloadJson, P256dhKey, AuthKey, RecordSize) of {ok, EncryptedBody} -> Response = request_push_endpoint(Endpoint, Headers, EncryptedBody), @@ -149,7 +171,7 @@ send_encrypted_push( SubscriptionId, Endpoint, Headers, - PayloadJson, + Payload, P256dhKey, AuthKey, RecordSize, @@ -160,12 +182,23 @@ send_encrypted_push( false end. +-spec fit_payload_json(map(), pos_integer()) -> binary(). +fit_payload_json(Payload, RecordSize) -> + push_notification:fit_payload_json(Payload, push_utils:plaintext_budget(RecordSize)). + +-spec push_urgency(map()) -> binary(). +push_urgency(Payload) -> + case push_notification:is_clear(Payload) of + true -> ?CLEAR_URGENCY; + false -> ?ALERT_URGENCY + end. + -spec handle_encrypted_response( integer(), binary(), binary(), [{binary(), binary()}], - binary(), + map(), binary(), binary(), pos_integer(), @@ -177,7 +210,7 @@ handle_encrypted_response( SubscriptionId, Endpoint, Headers, - PayloadJson, + Payload, P256dhKey, AuthKey, RecordSize, @@ -189,12 +222,12 @@ handle_encrypted_response( subscription_id => SubscriptionId, endpoint => Endpoint, headers => Headers, - payload_json => PayloadJson, + payload => Payload, p256dh_key => P256dhKey, auth_key => AuthKey }, RetryResult = push_sender_retry:maybe_retry_with_smaller_record_size( - Endpoint, Response, RecordSize, Attempt + Response, RecordSize, Attempt ), case RetryResult of {retry, NextRecordSize} -> @@ -256,13 +289,14 @@ retry_overload( #{ endpoint := Endpoint, headers := Headers, - payload_json := PayloadJson, + payload := Payload, p256dh_key := P256dhKey, auth_key := AuthKey } = Ctx, RecordSize, OverloadAttempt ) -> + PayloadJson = fit_payload_json(Payload, RecordSize), case push_utils:encrypt_payload(PayloadJson, P256dhKey, AuthKey, RecordSize) of {ok, EncryptedBody} -> Response = request_push_endpoint(Endpoint, Headers, EncryptedBody), @@ -312,7 +346,7 @@ retry_encrypted_push( subscription_id := SubscriptionId, endpoint := Endpoint, headers := Headers, - payload_json := PayloadJson, + payload := Payload, p256dh_key := P256dhKey, auth_key := AuthKey }, @@ -324,7 +358,7 @@ retry_encrypted_push( SubscriptionId, Endpoint, Headers, - PayloadJson, + Payload, P256dhKey, AuthKey, RecordSize, @@ -460,10 +494,11 @@ extract_subscription_fields(Subscription) -> {error, "missing keys"} end. --spec build_push_headers(binary(), binary()) -> [{binary(), binary()}]. -build_push_headers(VapidToken, VapidPublicKey) -> +-spec build_push_headers(binary(), binary(), binary()) -> [{binary(), binary()}]. +build_push_headers(VapidToken, VapidPublicKey, Urgency) -> [ {<<"TTL">>, ?PUSH_TTL}, + {<<"Urgency">>, Urgency}, {<<"Content-Type">>, <<"application/octet-stream">>}, {<<"Content-Encoding">>, <<"aes128gcm">>}, {<<"Authorization">>, <<"vapid t=", VapidToken/binary, ", k=", VapidPublicKey/binary>>} @@ -532,13 +567,13 @@ retry_delay_capped_test() -> maybe_retry_transient_waits_before_retry_test() -> Self = self(), Endpoint = <<"https://push.example/sub-1">>, - RecordSize = push_sender_retry:initial_record_size_for_endpoint(Endpoint), + RecordSize = push_sender_retry:initial_record_size(), Ctx = #{ user_id => 42, subscription_id => <<"sub-1">>, endpoint => Endpoint, headers => [], - payload_json => <<"{}">>, + payload => #{}, p256dh_key => <<"p256dh">>, auth_key => <<"auth">> }, @@ -647,11 +682,103 @@ receive_retried_push_request(ExpectedEndpoint) -> end. erase_token_cache(Key) -> - ok = push_token_cache:init(), - try ets:delete(push_bearer_tokens, Key) of - _ -> ok - catch - error:badarg -> ok + ok = push_ets_cache:init(), + true = ets:delete(push_bearer_tokens, Key), + ok. + +relay_endpoint_gets_the_shared_record_size_test() -> + {Body, _Headers} = capture_web_push(alert_payload(<<"Hello">>)), + ?assertEqual(2816, byte_size(Body)). + +alert_notification_sends_high_urgency_test() -> + {_Body, Headers} = capture_web_push(alert_payload(<<"Hello">>)), + ?assertEqual(?ALERT_URGENCY, header_value(<<"Urgency">>, Headers)). + +clear_notification_sends_low_urgency_test() -> + Payload = push_notification:build_clear_notification_payload(999, 456, 789, 2), + {_Body, Headers} = capture_web_push(Payload), + ?assertEqual(?CLEAR_URGENCY, header_value(<<"Urgency">>, Headers)). + +payload_over_the_budget_is_shrunk_instead_of_dropped_test() -> + Budget = push_utils:plaintext_budget(push_sender_retry:initial_record_size()), + Payload = alert_payload(binary:copy(<<"x">>, Budget * 2)), + ?assert(byte_size(iolist_to_binary(json:encode(Payload))) > Budget), + {Body, _Headers} = capture_web_push(Payload), + ?assertEqual(2816, byte_size(Body)). + +payload_at_the_budget_boundary_is_sent_unshrunk_test() -> + Budget = push_utils:plaintext_budget(push_sender_retry:initial_record_size()), + Payload = boundary_payload(Budget), + ?assertEqual(2713, Budget), + ?assertEqual(Budget, byte_size(iolist_to_binary(json:encode(Payload)))), + {Body, _Headers} = capture_web_push(Payload), + ?assertEqual(2816, byte_size(Body)). + +boundary_payload(Budget) -> + Base = #{<<"web_push">> => 8030, <<"title">> => <<>>}, + Overhead = byte_size(iolist_to_binary(json:encode(Base))), + Base#{<<"title">> => binary:copy(<<"x">>, Budget - Overhead)}. + +alert_payload(Body) -> + #{ + <<"web_push">> => 8030, + <<"title">> => <<"Alice">>, + <<"body">> => Body, + <<"tag">> => <<"channel:456:789">>, + <<"data">> => #{ + <<"channel_id">> => <<"456">>, + <<"message_id">> => <<"789">>, + <<"url">> => <<"/channels/123/456/789">> + } + }. + +header_value(Name, Headers) -> + proplists:get_value(Name, Headers). + +capture_web_push(Payload) -> + Endpoint = << + "https://push.fluxer.app/relay/v1/apns/stable/production/", + (binary:copy(<<"a">>, 64))/binary + >>, + {PeerPub, _PeerPriv} = crypto:generate_key(ecdh, prime256v1), + Subscription = #{ + <<"endpoint">> => Endpoint, + <<"p256dh_key">> => push_utils:base64url_encode(PeerPub), + <<"auth_key">> => push_utils:base64url_encode(crypto:strong_rand_bytes(16)), + <<"subscription_id">> => <<"sub-1">> + }, + ok = push_ets_cache:init(), + ok = meck:new(fluxer_gateway_env, [passthrough, no_link]), + ok = meck:new(push_utils, [passthrough, no_link]), + ok = meck:new(gateway_http_client, [passthrough, no_link]), + ok = meck:new(push_endpoint_guard, [passthrough, no_link]), + try + ok = meck:expect(push_endpoint_guard, check, fun(_Endpoint) -> ok end), + ok = meck:expect(fluxer_gateway_env, get, fun vapid_env_meck/1), + ok = meck:expect(push_utils, generate_vapid_token, fun(_Claims, _Public, _Private) -> + <<"vapid-token">> + end), + ok = meck:expect(gateway_http_client, request, fun capture_request_meck/6), + ?assertEqual(false, send_webpush_notification(42, Subscription, Payload)), + receive + {captured_push, CapturedHeaders, CapturedBody} -> {CapturedBody, CapturedHeaders} + after 1000 -> + erlang:error(no_push_request) + end + after + meck:unload(push_endpoint_guard), + meck:unload(gateway_http_client), + meck:unload(push_utils), + meck:unload(fluxer_gateway_env) end. +vapid_env_meck(vapid_email) -> <<"ops@example.com">>; +vapid_env_meck(vapid_public_key) -> <<"public-key">>; +vapid_env_meck(vapid_private_key) -> <<"private-key">>; +vapid_env_meck(Key) -> meck:passthrough([Key]). + +capture_request_meck(push, post, _Endpoint, Headers, Body, _Opts) -> + self() ! {captured_push, Headers, Body}, + {ok, 201, [], <<>>}. + -endif. diff --git a/fluxer_gateway/src/push/push_sender_retry.erl b/fluxer_gateway/src/push/push_sender_retry.erl index b6fa5056f..a556b6c3a 100644 --- a/fluxer_gateway/src/push/push_sender_retry.erl +++ b/fluxer_gateway/src/push/push_sender_retry.erl @@ -4,60 +4,53 @@ -typing([eqwalizer]). -export([ - maybe_retry_with_smaller_record_size/4, - initial_record_size_for_endpoint/1 + maybe_retry_with_smaller_record_size/3, + initial_record_size/0 ]). -export_type([push_response/0]). --define(STANDARD_PUSH_RECORD_SIZE, 4096). --define(MOZILLA_COMPAT_PUSH_RECORD_SIZE, 2820). --define(MOZILLA_CONSTRAINED_PUSH_RECORD_SIZE, 2048). +-define(PUSH_RECORD_SIZE, 2816). +-define(CONSTRAINED_PUSH_RECORD_SIZE, 2048). -define(MIN_PUSH_RECORD_SIZE, 1024). -define(MAX_PAYLOAD_RETRY_ATTEMPTS, 2). -type push_response() :: {ok, integer(), term(), binary()} | {error, term()}. --spec initial_record_size_for_endpoint(binary()) -> pos_integer(). -initial_record_size_for_endpoint(Endpoint) -> - case is_mozilla_push_endpoint(Endpoint) of - true -> ?MOZILLA_COMPAT_PUSH_RECORD_SIZE; - false -> ?STANDARD_PUSH_RECORD_SIZE - end. +-spec initial_record_size() -> pos_integer(). +initial_record_size() -> + ?PUSH_RECORD_SIZE. -spec maybe_retry_with_smaller_record_size( - binary(), push_response(), pos_integer(), non_neg_integer() + push_response(), pos_integer(), non_neg_integer() ) -> no_retry | {retry, pos_integer()}. -maybe_retry_with_smaller_record_size(_Endpoint, _Response, _CurrentRecordSize, Attempt) when +maybe_retry_with_smaller_record_size(_Response, _CurrentRecordSize, Attempt) when Attempt >= ?MAX_PAYLOAD_RETRY_ATTEMPTS -> no_retry; maybe_retry_with_smaller_record_size( - Endpoint, {ok, 413, _ResponseHeaders, ResponseBody}, CurrentRecordSize, _Attempt + {ok, 413, _ResponseHeaders, ResponseBody}, CurrentRecordSize, _Attempt ) -> - case next_record_size_for_payload_too_large(CurrentRecordSize, Endpoint, ResponseBody) of + case next_record_size_for_payload_too_large(CurrentRecordSize, ResponseBody) of undefined -> no_retry; NextRecordSize -> {retry, NextRecordSize} end; -maybe_retry_with_smaller_record_size(_Endpoint, _Response, _CurrentRecordSize, _Attempt) -> +maybe_retry_with_smaller_record_size(_Response, _CurrentRecordSize, _Attempt) -> no_retry. --spec next_record_size_for_payload_too_large(pos_integer(), binary(), binary()) -> +-spec next_record_size_for_payload_too_large(pos_integer(), binary()) -> pos_integer() | undefined. -next_record_size_for_payload_too_large(CurrentRecordSize, Endpoint, ResponseBody) -> +next_record_size_for_payload_too_large(CurrentRecordSize, ResponseBody) -> case parse_constrained_overage_bytes(ResponseBody) of OverageBytes when is_integer(OverageBytes), OverageBytes > 0 -> sanitize_next_record_size(CurrentRecordSize - OverageBytes, CurrentRecordSize); _ -> - FallbackRecordSize = fallback_record_size_for_endpoint(CurrentRecordSize, Endpoint), - sanitize_next_record_size(FallbackRecordSize, CurrentRecordSize) + sanitize_next_record_size(?CONSTRAINED_PUSH_RECORD_SIZE, CurrentRecordSize) end. --spec sanitize_next_record_size(integer() | undefined, pos_integer()) -> +-spec sanitize_next_record_size(integer(), pos_integer()) -> pos_integer() | undefined. -sanitize_next_record_size(undefined, _CurrentRecordSize) -> - undefined; sanitize_next_record_size(CandidateRecordSize, CurrentRecordSize) when is_integer(CandidateRecordSize) -> @@ -67,17 +60,6 @@ sanitize_next_record_size(CandidateRecordSize, CurrentRecordSize) when false -> undefined end. --spec fallback_record_size_for_endpoint(pos_integer(), binary()) -> pos_integer() | undefined. -fallback_record_size_for_endpoint(CurrentRecordSize, Endpoint) -> - case is_mozilla_push_endpoint(Endpoint) of - true when CurrentRecordSize > ?MOZILLA_COMPAT_PUSH_RECORD_SIZE -> - ?MOZILLA_COMPAT_PUSH_RECORD_SIZE; - true when CurrentRecordSize > ?MOZILLA_CONSTRAINED_PUSH_RECORD_SIZE -> - ?MOZILLA_CONSTRAINED_PUSH_RECORD_SIZE; - _ -> - undefined - end. - -spec parse_constrained_overage_bytes(binary()) -> non_neg_integer() | undefined. parse_constrained_overage_bytes(ResponseBody) -> case decode_push_error_body(ResponseBody) of @@ -119,80 +101,66 @@ parse_non_neg_integer(Value) -> _ -> undefined end. --spec is_mozilla_push_endpoint(binary()) -> boolean(). -is_mozilla_push_endpoint(Endpoint) -> - LowerEndpoint = lowercase_binary(Endpoint), - case binary:match(LowerEndpoint, <<"push.services.mozilla.com">>) of - nomatch -> false; - _ -> true - end. - --spec lowercase_binary(binary()) -> binary(). -lowercase_binary(Value) -> - iolist_to_binary(string:lowercase(Value)). - -ifdef(TEST). -include_lib("eunit/include/eunit.hrl"). -is_mozilla_push_endpoint_test() -> - Endpoint = <<"https://updates.push.services.mozilla.com/wpush/v2/token">>, - ?assertEqual(true, is_mozilla_push_endpoint(Endpoint)), - ?assertEqual( - true, - is_mozilla_push_endpoint(<<"https://push.services.mozilla.com/wpush/x">>) - ), - ?assertEqual( - false, - is_mozilla_push_endpoint(<<"https://fcm.googleapis.com/fcm/send">>) - ). - -initial_record_size_for_endpoint_test() -> - MozUrl = <<"https://updates.push.services.mozilla.com/wpush/v2/x">>, - ?assertEqual(?MOZILLA_COMPAT_PUSH_RECORD_SIZE, initial_record_size_for_endpoint(MozUrl)), - FcmUrl = <<"https://fcm.googleapis.com/fcm/send">>, - ?assertEqual(?STANDARD_PUSH_RECORD_SIZE, initial_record_size_for_endpoint(FcmUrl)). +initial_record_size_is_shared_by_every_endpoint_test() -> + ?assertEqual(2816, initial_record_size()). next_record_size_for_payload_too_large_overage_test() -> ResponseBody = << "{\"code\":413,\"errno\":104,\"error\":\"Payload Too Large\"," "\"message\":\"This message is intended for a constrained device and is limited in size. " - "Converted buffer is too long by 1441 bytes\"}" + "Converted buffer is too long by 441 bytes\"}" >>, ?assertEqual( - 2655, - next_record_size_for_payload_too_large( - ?STANDARD_PUSH_RECORD_SIZE, - <<"https://updates.push.services.mozilla.com/wpush/v2/x">>, - ResponseBody - ) + 2375, + next_record_size_for_payload_too_large(?PUSH_RECORD_SIZE, ResponseBody) ). next_record_size_for_payload_too_large_fallback_test() -> ResponseBody = <<"{\"code\":413,\"errno\":104,\"error\":\"Payload Too Large\"}">>, - MozillaEndpoint = <<"https://updates.push.services.mozilla.com/wpush/v2/x">>, ?assertEqual( - ?MOZILLA_COMPAT_PUSH_RECORD_SIZE, - next_record_size_for_payload_too_large( - ?STANDARD_PUSH_RECORD_SIZE, MozillaEndpoint, ResponseBody - ) - ), - ?assertEqual( - ?MOZILLA_CONSTRAINED_PUSH_RECORD_SIZE, - next_record_size_for_payload_too_large( - ?MOZILLA_COMPAT_PUSH_RECORD_SIZE, MozillaEndpoint, ResponseBody - ) + ?CONSTRAINED_PUSH_RECORD_SIZE, + next_record_size_for_payload_too_large(?PUSH_RECORD_SIZE, ResponseBody) ), ?assertEqual( undefined, next_record_size_for_payload_too_large( - ?MOZILLA_CONSTRAINED_PUSH_RECORD_SIZE, MozillaEndpoint, ResponseBody - ) - ), - ?assertEqual( - undefined, - next_record_size_for_payload_too_large( - ?STANDARD_PUSH_RECORD_SIZE, <<"https://fcm.googleapis.com/fcm/send">>, ResponseBody + ?CONSTRAINED_PUSH_RECORD_SIZE, ResponseBody ) ). +next_record_size_is_clamped_to_the_minimum_test() -> + ResponseBody = << + "{\"code\":413,\"message\":\"Converted buffer is too long by 2000 bytes\"}" + >>, + ?assertEqual( + ?MIN_PUSH_RECORD_SIZE, + next_record_size_for_payload_too_large(?PUSH_RECORD_SIZE, ResponseBody) + ). + +maybe_retry_stops_after_the_attempt_cap_test() -> + Response = {ok, 413, [], <<>>}, + ?assertEqual( + {retry, ?CONSTRAINED_PUSH_RECORD_SIZE}, + maybe_retry_with_smaller_record_size(Response, ?PUSH_RECORD_SIZE, 0) + ), + ?assertEqual( + no_retry, + maybe_retry_with_smaller_record_size( + Response, ?PUSH_RECORD_SIZE, ?MAX_PAYLOAD_RETRY_ATTEMPTS + ) + ). + +maybe_retry_ignores_non_413_responses_test() -> + ?assertEqual( + no_retry, + maybe_retry_with_smaller_record_size({ok, 400, [], <<>>}, ?PUSH_RECORD_SIZE, 0) + ), + ?assertEqual( + no_retry, + maybe_retry_with_smaller_record_size({error, timeout}, ?PUSH_RECORD_SIZE, 0) + ). + -endif. diff --git a/fluxer_gateway/src/push/push_subscriptions.erl b/fluxer_gateway/src/push/push_subscriptions.erl index 873bcab9a..5f1411e8a 100644 --- a/fluxer_gateway/src/push/push_subscriptions.erl +++ b/fluxer_gateway/src/push/push_subscriptions.erl @@ -141,35 +141,35 @@ abandon_subscription_batches(Remaining, {Tasks, FailedBatches, FailedUsers, Cons -spec fetch_next_subscription_batch([[integer()]], subscription_batch_acc(), integer()) -> subscription_batch_acc(). -fetch_next_subscription_batch( - [Batch | Rest], {Tasks, FailedBatches, FailedUsers, Consecutive}, Deadline -) -> +fetch_next_subscription_batch([Batch | Rest], Acc, Deadline) -> + fetch_subscription_batches(Rest, fetch_subscription_batch(Batch, Acc), Deadline). + +-spec fetch_subscription_batch([integer()], subscription_batch_acc()) -> + subscription_batch_acc(). +fetch_subscription_batch(Batch, {Tasks, FailedBatches, FailedUsers, Consecutive}) -> Req = #{ <<"type">> => <<"get_push_subscriptions">>, <<"user_ids">> => [integer_to_binary(UserId) || UserId <- Batch] }, - case rpc_client:call(Req) of + Fill = push_ets_cache:reserve_subscriptions(Batch), + try rpc_client:call(Req) of {ok, BatchData} -> BatchTasks = lists:foldl( fun(UserId, Acc) -> - add_fetched_user_subscription_task(UserId, BatchData, Acc) + add_fetched_user_subscription_task(UserId, BatchData, Fill, Acc) end, Tasks, Batch ), - fetch_subscription_batches( - Rest, {BatchTasks, FailedBatches, FailedUsers, 0}, Deadline - ); + {BatchTasks, FailedBatches, FailedUsers, 0}; {error, Reason} -> logger:debug( "Push: RPC failed to fetch subscriptions", #{user_count => length(Batch), reason => Reason} ), - fetch_subscription_batches( - Rest, - {Tasks, FailedBatches + 1, FailedUsers + length(Batch), Consecutive + 1}, - Deadline - ) + {Tasks, FailedBatches + 1, FailedUsers + length(Batch), Consecutive + 1} + after + push_ets_cache:release(Fill) end. -spec batched_user_count([[integer()]], non_neg_integer()) -> non_neg_integer(). @@ -213,17 +213,19 @@ count_subscription_fetch_loss(UserCount) -> bump_counter(subscription_fetch_calls_failed, 1), bump_counter(subscription_fetch_users_dropped, UserCount). --spec add_fetched_user_subscription_task(integer(), map(), [{integer(), list()}]) -> +-spec add_fetched_user_subscription_task( + integer(), map(), push_ets_cache:fill(), [{integer(), list()}] +) -> [{integer(), list()}]. -add_fetched_user_subscription_task(UserId, SubscriptionsData, Acc) -> +add_fetched_user_subscription_task(UserId, SubscriptionsData, Fill, Acc) -> UserIdBin = integer_to_binary(UserId), case maps:get(UserIdBin, SubscriptionsData, []) of [] -> - push_ets_cache:put_subscriptions(UserId, []), + push_ets_cache:put_subscriptions(UserId, [], Fill), logger:debug("Push: no subscriptions for user", #{user_id => UserId}), Acc; Subscriptions -> - push_ets_cache:put_subscriptions(UserId, Subscriptions), + push_ets_cache:put_subscriptions(UserId, Subscriptions, Fill), logger:debug( "Push: found subscriptions for user", #{user_id => UserId, count => length(Subscriptions)} @@ -259,18 +261,32 @@ fetch_and_send_clear_notification_from_rpc(UserId, ChannelId, MessageId, BadgeCo "Push: fetching subscriptions for notification clear", #{user_id => UserId, channel_id => ChannelId, message_id => MessageId} ), - Result = rpc_client:call(SubscriptionsReq), - send_clear_rpc_result(UserId, ChannelId, MessageId, BadgeCount, Result). + Fill = push_ets_cache:reserve_subscriptions([UserId]), + try rpc_client:call(SubscriptionsReq) of + Result -> send_clear_rpc_result(UserId, ChannelId, MessageId, BadgeCount, Fill, Result) + after + push_ets_cache:release(Fill) + end. -spec send_clear_rpc_result( - integer(), integer(), integer(), non_neg_integer(), {ok, map()} | {error, term()} + integer(), + integer(), + integer(), + non_neg_integer(), + push_ets_cache:fill(), + {ok, map()} | {error, term()} ) -> ok. -send_clear_rpc_result(UserId, ChannelId, MessageId, BadgeCount, {ok, SubscriptionsData}) -> +send_clear_rpc_result(UserId, ChannelId, MessageId, BadgeCount, Fill, {ok, SubscriptionsData}) -> UserIdBin = integer_to_binary(UserId), send_clear_fetched_subscriptions( - UserId, ChannelId, MessageId, BadgeCount, maps:get(UserIdBin, SubscriptionsData, []) + UserId, + ChannelId, + MessageId, + BadgeCount, + Fill, + maps:get(UserIdBin, SubscriptionsData, []) ); -send_clear_rpc_result(UserId, _ChannelId, _MessageId, _BadgeCount, {error, Reason}) -> +send_clear_rpc_result(UserId, _ChannelId, _MessageId, _BadgeCount, _Fill, {error, Reason}) -> count_subscription_fetch_loss(1), logger:debug( "Push: RPC failed to fetch subscriptions for notification clear", @@ -279,13 +295,13 @@ send_clear_rpc_result(UserId, _ChannelId, _MessageId, _BadgeCount, {error, Reaso ok. -spec send_clear_fetched_subscriptions( - integer(), integer(), integer(), non_neg_integer(), list() + integer(), integer(), integer(), non_neg_integer(), push_ets_cache:fill(), list() ) -> ok. -send_clear_fetched_subscriptions(UserId, _ChannelId, _MessageId, _BadgeCount, []) -> +send_clear_fetched_subscriptions(UserId, _ChannelId, _MessageId, _BadgeCount, _Fill, []) -> logger:debug("Push: no subscriptions for notification clear", #{user_id => UserId}), ok; -send_clear_fetched_subscriptions(UserId, ChannelId, MessageId, BadgeCount, Subscriptions) -> - push_ets_cache:put_subscriptions(UserId, Subscriptions), +send_clear_fetched_subscriptions(UserId, ChannelId, MessageId, BadgeCount, Fill, Subscriptions) -> + push_ets_cache:put_subscriptions(UserId, Subscriptions, Fill), push_sender:send_clear_to_user_subscriptions( UserId, Subscriptions, @@ -688,19 +704,22 @@ fetch_and_cache_user_guild_settings(UserId, GuildId) -> "Push: fetching user guild settings via RPC", #{user_id => UserId, guild_id => GuildId} ), - case rpc_client:call(Req) of + Fill = push_ets_cache:reserve_user_guild_settings([UserId], GuildId), + try rpc_client:call(Req) of {ok, Data} -> - cache_user_guild_settings(UserId, GuildId, Data); + cache_user_guild_settings(UserId, GuildId, Data, Fill); {error, Reason} -> logger:debug( "Push: RPC failed to fetch user guild settings", #{user_id => UserId, guild_id => GuildId, reason => Reason} ), null + after + push_ets_cache:release(Fill) end. --spec cache_user_guild_settings(integer(), integer(), map()) -> map(). -cache_user_guild_settings(UserId, GuildId, Data) -> +-spec cache_user_guild_settings(integer(), integer(), map(), push_ets_cache:fill()) -> map(). +cache_user_guild_settings(UserId, GuildId, Data, Fill) -> SettingsData = case maps:get(<<"user_guild_settings">>, Data, [null]) of [First | _] -> First; @@ -712,7 +731,7 @@ cache_user_guild_settings(UserId, GuildId, Data) -> "Push: user guild settings returned null; caching empty sentinel", #{user_id => UserId, guild_id => GuildId} ), - push_ets_cache:put_user_guild_settings(UserId, GuildId, #{}), + push_ets_cache:put_user_guild_settings(UserId, GuildId, #{}, Fill), #{}; Settings -> logger:debug( @@ -724,7 +743,7 @@ cache_user_guild_settings(UserId, GuildId, Data) -> mobile_push => maps:get(mobile_push, Settings, undefined) } ), - push_ets_cache:put_user_guild_settings(UserId, GuildId, Settings), + push_ets_cache:put_user_guild_settings(UserId, GuildId, Settings, Fill), Settings end. @@ -1076,7 +1095,7 @@ fetch_missing_in_batches_sends_only_successful_batches_test() -> ok = meck:new(push_ets_cache, [passthrough, no_link]), application:set_env(fluxer_gateway, push_subscription_fetch_batch_size, 2), try - ok = meck:expect(push_ets_cache, put_subscriptions, fun(_UserId, _Subs) -> ok end), + ok = meck:expect(push_ets_cache, put_subscriptions, fun(_UserId, _Subs, _Fill) -> ok end), ok = meck:expect(rpc_client, call, fun(#{<<"user_ids">> := Ids}) -> case Ids of [<<"1">>, <<"2">>] -> {ok, #{<<"1">> => [sub1], <<"2">> => [sub2]}}; diff --git a/fluxer_gateway/src/push/push_token_cache.erl b/fluxer_gateway/src/push/push_token_cache.erl deleted file mode 100644 index d6f6127d2..000000000 --- a/fluxer_gateway/src/push/push_token_cache.erl +++ /dev/null @@ -1,46 +0,0 @@ -%% SPDX-License-Identifier: AGPL-3.0-or-later - --module(push_token_cache). --typing([eqwalizer]). - --export([init/0, get/1, put/3]). - --define(TABLE, push_bearer_tokens). - --spec init() -> ok. -init() -> - case ets:whereis(?TABLE) of - undefined -> create_table(); - _ -> ok - end. - --spec create_table() -> ok. -create_table() -> - try - _ = ets:new(?TABLE, [ - named_table, public, set, {read_concurrency, true}, {write_concurrency, true} - ]), - ok - catch - error:badarg -> ok - end. - --spec get(term()) -> {ok, binary(), integer()} | undefined. -get(Key) -> - try ets:lookup(?TABLE, Key) of - [{_, Token, ExpiresAt}] when is_binary(Token), is_integer(ExpiresAt) -> - {ok, Token, ExpiresAt}; - _ -> - undefined - catch - error:badarg -> undefined - end. - --spec put(term(), binary(), integer()) -> ok. -put(Key, Token, ExpiresAt) when is_binary(Token), is_integer(ExpiresAt) -> - ok = init(), - try ets:insert(?TABLE, {Key, Token, ExpiresAt}) of - _ -> ok - catch - error:badarg -> ok - end. diff --git a/fluxer_gateway/src/push/push_utils.erl b/fluxer_gateway/src/push/push_utils.erl index 43f8c3bc8..474ddc071 100644 --- a/fluxer_gateway/src/push/push_utils.erl +++ b/fluxer_gateway/src/push/push_utils.erl @@ -14,16 +14,20 @@ base64url_encode/1, base64url_decode/1, encrypt_payload/4, + plaintext_budget/1, decode_subscription_key/1, hkdf_expand/4, hkdf_expand_loop/6, - parse_timestamp/1, normalize_binary/1, normalize_binary/2, avatar_index/1, wrap_avatar_index/1 ]). +-define(RECORD_HEADER_BYTES, 86). +-define(RECORD_TAG_BYTES, 16). +-define(RECORD_DELIMITER_BYTES, 1). + -spec construct_avatar_url(binary(), binary()) -> binary(). construct_avatar_url(UserId, Hash) -> MediaProxyBin = media_proxy_endpoint_binary(), @@ -242,6 +246,13 @@ base64url_decode(Data) -> error -> error end. +-spec plaintext_budget(pos_integer()) -> non_neg_integer(). +plaintext_budget(RecordSize) when is_integer(RecordSize) -> + case RecordSize - ?RECORD_HEADER_BYTES - ?RECORD_TAG_BYTES - ?RECORD_DELIMITER_BYTES of + Budget when Budget > 0 -> Budget; + _ -> 0 + end. + -spec encrypt_payload(binary(), binary(), binary(), non_neg_integer()) -> {ok, binary()} | {error, term()}. encrypt_payload(Message, PeerPubB64, AuthSecretB64, RecordSize0) -> @@ -250,7 +261,7 @@ encrypt_payload(Message, PeerPubB64, AuthSecretB64, RecordSize0) -> AuthSecret = decode_subscription_key(AuthSecretB64), RecordSize = case RecordSize0 of - 0 -> 4096; + 0 -> push_sender_retry:initial_record_size(); _ -> RecordSize0 end, Salt = crypto:strong_rand_bytes(16), @@ -354,16 +365,6 @@ hkdf_expand_loop(PRK, Info, Length, I, Tprev, Acc) -> T = crypto:mac(hmac, sha256, PRK, <>), hkdf_expand_loop(PRK, Info, Length, I + 1, T, <>). --spec parse_timestamp(binary() | term()) -> integer() | undefined. -parse_timestamp(Str) when is_binary(Str) -> - try - binary_to_integer(Str) - catch - _:_ -> undefined - end; -parse_timestamp(_) -> - undefined. - -spec normalize_binary(term()) -> binary() | undefined. normalize_binary(Value) when is_binary(Value) -> Value; normalize_binary(Value) when is_list(Value) -> type_conv:to_binary(Value); diff --git a/fluxer_gateway/src/session/session_monitor.erl b/fluxer_gateway/src/session/session_monitor.erl index df218b5c1..42c23eec8 100644 --- a/fluxer_gateway/src/session/session_monitor.erl +++ b/fluxer_gateway/src/session/session_monitor.erl @@ -26,7 +26,7 @@ handle_process_down(Ref, Reason, State) -> PresenceRef = maps:get(presence_mref, State, undefined), case Ref of SocketRef when Ref =:= SocketRef -> - handle_socket_down(State); + handle_socket_down(Reason, State); PresenceRef when Ref =:= PresenceRef -> handle_presence_down(State); _ -> @@ -63,8 +63,29 @@ handle_call_or_ignore(Ref, Reason, State, Calls) -> {noreply, State} end. --spec handle_socket_down(session_state()) -> {noreply, session_state()}. -handle_socket_down(State) -> +-spec handle_socket_down(term(), session_state()) -> + {noreply, session_state()} | {stop, normal, session_state()}. +handle_socket_down({shutdown, client_closed}, State) -> + end_or_hold_session(enrolled_in_push_delivery(State), State); +handle_socket_down(_Reason, State) -> + hold_session_for_resume(State). + +-spec end_or_hold_session(boolean(), session_state()) -> + {noreply, session_state()} | {stop, normal, session_state()}. +end_or_hold_session(true, State) -> + {stop, normal, State#{socket_pid => undefined, socket_mref => undefined}}; +end_or_hold_session(false, State) -> + hold_session_for_resume(State). + +-spec enrolled_in_push_delivery(session_state()) -> boolean(). +enrolled_in_push_delivery(State) -> + case maps:get(user_id, State, undefined) of + UserId when is_integer(UserId) -> push_delivery_config:is_enrolled(UserId); + _ -> false + end. + +-spec hold_session_for_resume(session_state()) -> {noreply, session_state()}. +hold_session_for_resume(State) -> ResumeToken = make_ref(), ResumeTimerRef = erlang:send_after( constants:resume_timeout(), self(), {resume_timeout, ResumeToken} diff --git a/fluxer_gateway/test/push_endpoint_guard_tests.erl b/fluxer_gateway/test/push_endpoint_guard_tests.erl new file mode 100644 index 000000000..94fcc70f9 --- /dev/null +++ b/fluxer_gateway/test/push_endpoint_guard_tests.erl @@ -0,0 +1,231 @@ +%% SPDX-License-Identifier: AGPL-3.0-or-later + +-module(push_endpoint_guard_tests). +-typing([eqwalizer]). + +-include_lib("eunit/include/eunit.hrl"). + +-define(ENDPOINT, <<"https://push.example.com/wpush/v2/abc">>). + +resolves_to(Addresses) -> + fun(_Host) -> {ok, Addresses} end. + +fails_with(Reason) -> + fun(_Host) -> {error, Reason} end. + +link_local_metadata_address_is_refused_test() -> + ?assertEqual( + {error, endpoint_blocked}, + push_endpoint_guard:check(?ENDPOINT, resolves_to([{169, 254, 169, 254}])) + ). + +private_address_is_refused_test() -> + ?assertEqual( + {error, endpoint_blocked}, + push_endpoint_guard:check(?ENDPOINT, resolves_to([{10, 0, 0, 1}])) + ). + +loopback_address_is_refused_test() -> + ?assertEqual( + {error, endpoint_blocked}, + push_endpoint_guard:check(?ENDPOINT, resolves_to([{127, 0, 0, 1}])) + ). + +every_reserved_ipv4_range_is_refused_test() -> + Blocked = [ + {0, 0, 0, 1}, + {10, 1, 2, 3}, + {100, 64, 0, 1}, + {127, 0, 0, 1}, + {169, 254, 169, 254}, + {172, 16, 0, 1}, + {172, 31, 255, 254}, + {192, 0, 0, 1}, + {192, 0, 2, 1}, + {192, 88, 99, 1}, + {192, 168, 1, 1}, + {198, 18, 0, 1}, + {198, 51, 100, 1}, + {203, 0, 113, 1}, + {224, 0, 0, 1}, + {240, 0, 0, 1}, + {255, 255, 255, 255} + ], + lists:foreach( + fun(Address) -> + ?assertEqual( + {error, endpoint_blocked}, + push_endpoint_guard:check(?ENDPOINT, resolves_to([Address])) + ) + end, + Blocked + ). + +every_reserved_ipv6_range_is_refused_test() -> + Blocked = [ + {0, 0, 0, 0, 0, 0, 0, 0}, + {0, 0, 0, 0, 0, 0, 0, 1}, + {16#2001, 16#0db8, 0, 0, 0, 0, 0, 1}, + {16#fd00, 0, 0, 0, 0, 0, 0, 1}, + {16#fe80, 0, 0, 0, 0, 0, 0, 1}, + {16#ff02, 0, 0, 0, 0, 0, 0, 1} + ], + lists:foreach( + fun(Address) -> + ?assertEqual( + {error, endpoint_blocked}, + push_endpoint_guard:check(?ENDPOINT, resolves_to([Address])) + ) + end, + Blocked + ). + +ipv4_mapped_form_of_a_private_address_is_refused_test() -> + Mapped = {0, 0, 0, 0, 0, 16#ffff, 16#0a00, 16#0001}, + ?assertEqual( + {error, endpoint_blocked}, + push_endpoint_guard:check(?ENDPOINT, resolves_to([Mapped])) + ). + +ipv4_compatible_and_nat64_and_sixtofour_forms_are_refused_test() -> + Compatible = {0, 0, 0, 0, 0, 0, 16#a9fe, 16#a9fe}, + Nat64 = {16#0064, 16#ff9b, 0, 0, 0, 0, 16#0a00, 16#0001}, + SixToFour = {16#2002, 16#0a00, 16#0001, 0, 0, 0, 0, 0}, + lists:foreach( + fun(Address) -> + ?assertEqual( + {error, endpoint_blocked}, + push_endpoint_guard:check(?ENDPOINT, resolves_to([Address])) + ) + end, + [Compatible, Nat64, SixToFour] + ). + +one_private_address_refuses_the_whole_set_test() -> + Mixed = [{93, 184, 216, 34}, {10, 0, 0, 1}], + ?assertEqual( + {error, endpoint_blocked}, + push_endpoint_guard:check(?ENDPOINT, resolves_to(Mixed)) + ), + ?assertEqual( + {error, endpoint_blocked}, + push_endpoint_guard:check(?ENDPOINT, resolves_to(lists:reverse(Mixed))) + ). + +one_private_ipv6_address_refuses_the_whole_set_test() -> + Mixed = [ + {93, 184, 216, 34}, + {16#2606, 16#4700, 16#4700, 0, 0, 0, 0, 16#1111}, + {16#fd00, 0, 0, 0, 0, 0, 0, 1} + ], + ?assertEqual( + {error, endpoint_blocked}, + push_endpoint_guard:check(?ENDPOINT, resolves_to(Mixed)) + ). + +public_addresses_are_allowed_test() -> + Public = [ + {93, 184, 216, 34}, + {16#2606, 16#4700, 16#4700, 0, 0, 0, 0, 16#1111}, + {0, 0, 0, 0, 0, 16#ffff, 16#5db8, 16#d822} + ], + ?assertEqual(ok, push_endpoint_guard:check(?ENDPOINT, resolves_to(Public))). + +a_host_that_resolves_to_nothing_is_refused_test() -> + ?assertEqual({error, nxdomain}, push_endpoint_guard:check(?ENDPOINT, resolves_to([]))). + +an_unresolvable_host_reports_the_resolver_error_test() -> + ?assertEqual({error, nxdomain}, push_endpoint_guard:check(?ENDPOINT, fails_with(nxdomain))), + ?assertEqual({error, timeout}, push_endpoint_guard:check(?ENDPOINT, fails_with(timeout))). + +an_unresolvable_host_fails_cleanly_against_the_real_resolver_test() -> + ?assertMatch({error, _}, push_endpoint_guard:check(<<"https://push.invalid/sub">>)). + +a_public_host_is_allowed_by_the_real_resolver_test() -> + case inet:getaddrs("one.one.one.one", inet, 3000) of + {ok, [_ | _]} -> + ?assertEqual(ok, push_endpoint_guard:check(<<"https://one.one.one.one/sub">>)); + _ -> + ok + end. + +plain_http_is_refused_test() -> + ?assertEqual( + {error, endpoint_rejected}, + push_endpoint_guard:check( + <<"http://push.example.com/sub">>, resolves_to([{1, 1, 1, 1}]) + ) + ). + +non_standard_ports_are_refused_test() -> + ?assertEqual( + {error, endpoint_rejected}, + push_endpoint_guard:check( + <<"https://push.example.com:8080/sub">>, resolves_to([{1, 1, 1, 1}]) + ) + ), + ?assertEqual( + ok, + push_endpoint_guard:check( + <<"https://push.example.com:443/sub">>, resolves_to([{1, 1, 1, 1}]) + ) + ). + +userinfo_is_refused_test() -> + ?assertEqual( + {error, endpoint_rejected}, + push_endpoint_guard:check( + <<"https://push.example.com@169.254.169.254/sub">>, resolves_to([{1, 1, 1, 1}]) + ) + ), + ?assertEqual( + {error, endpoint_rejected}, + push_endpoint_guard:check( + <<"https://user:secret@push.example.com/sub">>, resolves_to([{1, 1, 1, 1}]) + ) + ). + +ip_literals_skip_dns_and_are_screened_directly_test() -> + Never = fun(_Host) -> erlang:error(resolver_called) end, + ?assertEqual( + {error, endpoint_blocked}, + push_endpoint_guard:check(<<"https://169.254.169.254/latest">>, Never) + ), + ?assertEqual( + {error, endpoint_blocked}, push_endpoint_guard:check(<<"https://127.0.0.1/sub">>, Never) + ), + ?assertEqual( + {error, endpoint_blocked}, push_endpoint_guard:check(<<"https://[::1]/sub">>, Never) + ), + ?assertEqual( + {error, endpoint_blocked}, + push_endpoint_guard:check(<<"https://[::ffff:10.0.0.1]/sub">>, Never) + ), + ?assertEqual(ok, push_endpoint_guard:check(<<"https://93.184.216.34/sub">>, Never)). + +malformed_and_non_fqdn_hosts_are_refused_test() -> + Never = fun(_Host) -> erlang:error(resolver_called) end, + lists:foreach( + fun(Endpoint) -> + ?assertEqual( + {error, endpoint_rejected}, push_endpoint_guard:check(Endpoint, Never) + ) + end, + [ + <<"not-a-url">>, + <<>>, + <<"https://localhost/sub">>, + <<"https://metadata/sub">>, + <<"https://push.example.123/sub">>, + <<"https://-push.example.com/sub">>, + <<"ftp://push.example.com/sub">> + ] + ). + +uppercase_hosts_are_normalised_test() -> + ?assertEqual( + {error, endpoint_blocked}, + push_endpoint_guard:check( + <<"https://PUSH.EXAMPLE.COM/sub">>, resolves_to([{10, 0, 0, 1}]) + ) + ). diff --git a/fluxer_gateway/test/push_ets_cache_stress_tests.erl b/fluxer_gateway/test/push_ets_cache_stress_tests.erl index c1da4f02f..979d69fd8 100644 --- a/fluxer_gateway/test/push_ets_cache_stress_tests.erl +++ b/fluxer_gateway/test/push_ets_cache_stress_tests.erl @@ -50,7 +50,7 @@ concurrent_subscription_reads_and_writes_keep_cache_consistent() -> end. put_subscription_for_even_user(UserId) when UserId rem 2 =:= 0 -> - ok = push_ets_cache:put_subscriptions(UserId, [#{endpoint => endpoint(UserId)}]); + ok = put_subscriptions(UserId, [#{endpoint => endpoint(UserId)}]); put_subscription_for_even_user(_UserId) -> ok. @@ -68,7 +68,7 @@ writer_loop(WriterIndex, Count) -> End = Start + Count - 1, lists:foreach( fun(UserId) -> - ok = push_ets_cache:put_subscriptions(UserId, [#{endpoint => endpoint(UserId)}]) + ok = put_subscriptions(UserId, [#{endpoint => endpoint(UserId)}]) end, lists:seq(Start, End) ). @@ -103,6 +103,11 @@ collect_done(Ref, Message, Remaining) -> ?assert(false) end. +put_subscriptions(UserId, Subscriptions) -> + push_ets_cache:put_subscriptions( + UserId, Subscriptions, push_ets_cache:reserve_subscriptions([UserId]) + ). + endpoint(UserId) -> UserIdBin = integer_to_binary(UserId), <<"endpoint-", UserIdBin/binary>>. diff --git a/fluxer_gateway/test/push_ets_cache_tests.erl b/fluxer_gateway/test/push_ets_cache_tests.erl index b20e9cf7f..02590fc36 100644 --- a/fluxer_gateway/test/push_ets_cache_tests.erl +++ b/fluxer_gateway/test/push_ets_cache_tests.erl @@ -34,7 +34,7 @@ subscriptions_test() -> cleanup_tables(), ok = push_ets_cache:init(), ?assertEqual(undefined, push_ets_cache:get_subscriptions(1)), - ok = push_ets_cache:put_subscriptions(1, [sub1, sub2]), + ok = seed_subscriptions(1, [sub1, sub2]), ?assertEqual([sub1, sub2], push_ets_cache:get_subscriptions(1)), ok = push_ets_cache:delete_subscriptions(1), ?assertEqual(undefined, push_ets_cache:get_subscriptions(1)), @@ -52,7 +52,7 @@ badge_count_test() -> cleanup_tables(), ok = push_ets_cache:init(), ?assertEqual(undefined, push_ets_cache:get_badge_count(1)), - ok = push_ets_cache:put_badge_count(1, 5, 1000), + ok = seed_badge_count(1, 5, 1000), ?assertEqual({5, 1000}, push_ets_cache:get_badge_count(1)), ok = push_ets_cache:delete_badge_count(1), ?assertEqual(undefined, push_ets_cache:get_badge_count(1)), @@ -61,20 +61,21 @@ badge_count_test() -> badge_count_keeps_fresher_timestamp_test() -> cleanup_tables(), ok = push_ets_cache:init(), - ok = push_ets_cache:put_badge_count(1, 5, 2000), - ok = push_ets_cache:put_badge_count(1, 9, 1000), + First = push_ets_cache:reserve_badge_counts([1]), + ok = push_ets_cache:put_badge_count(1, 5, 2000, First), + ok = push_ets_cache:put_badge_count(1, 9, 1000, First), ?assertEqual({5, 2000}, push_ets_cache:get_badge_count(1)), - ok = push_ets_cache:put_badge_count(1, 7, 3000), + ok = seed_badge_count(1, 7, 3000), ?assertEqual({7, 3000}, push_ets_cache:get_badge_count(1)), - ok = push_ets_cache:put_badge_count(1, 8, 3000), + ok = seed_badge_count(1, 8, 3000), ?assertEqual({8, 3000}, push_ets_cache:get_badge_count(1)), cleanup_tables(). cache_stats_test() -> cleanup_tables(), ok = push_ets_cache:init(), - ok = push_ets_cache:put_subscriptions(1, []), - ok = push_ets_cache:put_subscriptions(2, []), + ok = seed_subscriptions(1, []), + ok = seed_subscriptions(2, []), Stats = push_ets_cache:cache_stats(), ?assertEqual(2, maps:get(push_subscriptions_size, Stats)), ?assertEqual(0, maps:get(user_guild_settings_size, Stats)), @@ -83,7 +84,7 @@ cache_stats_test() -> evict_tables_test() -> cleanup_tables(), ok = push_ets_cache:init(), - lists:foreach(fun(I) -> push_ets_cache:put_subscriptions(I, []) end, lists:seq(1, 10)), + lists:foreach(fun(I) -> ok = seed_subscriptions(I, []) end, lists:seq(1, 10)), ?assertEqual(10, push_ets_cache:table_size(push_subscriptions)), ok = push_ets_cache:evict_tables(#{subscriptions => 5}), ?assertEqual(5, push_ets_cache:table_size(push_subscriptions)), @@ -98,8 +99,8 @@ rebalance_evicts_remote_owned_entries_test() -> RoleMap = #{push => Members, all => Members}, persistent_term:put({gateway_cluster_membership, members}, Members), persistent_term:put({gateway_cluster_membership, members_by_role}, RoleMap), - ok = push_ets_cache:put_subscriptions(LocalUserId, [local]), - ok = push_ets_cache:put_subscriptions(RemoteUserId, [remote]), + ok = seed_subscriptions(LocalUserId, [local]), + ok = seed_subscriptions(RemoteUserId, [remote]), ok = push_ets_cache:put_user_guild_settings(LocalUserId, 10, #{local => true}), ok = push_ets_cache:put_user_guild_settings(RemoteUserId, 10, #{remote => true}), ok = push_ets_cache:rebalance(), @@ -111,6 +112,16 @@ rebalance_evicts_remote_owned_entries_test() -> persistent_term:erase({gateway_cluster_membership, members_by_role}), cleanup_tables(). +seed_subscriptions(UserId, Subscriptions) -> + push_ets_cache:put_subscriptions( + UserId, Subscriptions, push_ets_cache:reserve_subscriptions([UserId]) + ). + +seed_badge_count(UserId, Count, CachedAt) -> + push_ets_cache:put_badge_count( + UserId, Count, CachedAt, push_ets_cache:reserve_badge_counts([UserId]) + ). + find_split_user_ids(Members, RemoteNode) -> Local = hd([ @@ -131,6 +142,7 @@ cleanup_tables() -> delete_table(push_subscriptions), delete_table(push_blocked_ids), delete_table(push_badge_counts), + delete_table(push_bearer_tokens), ok. delete_table(Table) -> diff --git a/fluxer_gateway/test/push_utils_tests.erl b/fluxer_gateway/test/push_utils_tests.erl index a7dfec0b6..5370b2bcc 100644 --- a/fluxer_gateway/test/push_utils_tests.erl +++ b/fluxer_gateway/test/push_utils_tests.erl @@ -33,15 +33,6 @@ wrap_avatar_index_test() -> ?assertEqual(0, push_utils:wrap_avatar_index(6)), ?assertEqual(1, push_utils:wrap_avatar_index(7)). -parse_timestamp_valid_test() -> - ?assertEqual(123456789, push_utils:parse_timestamp(<<"123456789">>)), - ?assertEqual(0, push_utils:parse_timestamp(<<"0">>)). - -parse_timestamp_invalid_test() -> - ?assertEqual(undefined, push_utils:parse_timestamp(<<"not_a_number">>)), - ?assertEqual(undefined, push_utils:parse_timestamp(123)), - ?assertEqual(undefined, push_utils:parse_timestamp(undefined)). - base64url_encode_test() -> Encoded = push_utils:base64url_encode(<<"test">>), ?assert(is_binary(Encoded)). @@ -50,6 +41,28 @@ base64url_decode_test() -> Encoded = push_utils:base64url_encode(<<"test">>), ?assertEqual(<<"test">>, push_utils:base64url_decode(Encoded)). +plaintext_budget_test() -> + ?assertEqual(2713, push_utils:plaintext_budget(2816)), + ?assertEqual(3993, push_utils:plaintext_budget(4096)), + ?assertEqual(0, push_utils:plaintext_budget(1)). + +encrypt_payload_fills_the_record_at_the_budget_test() -> + RecordSize = push_sender_retry:initial_record_size(), + Budget = push_utils:plaintext_budget(RecordSize), + {PeerPub, _PeerPriv} = crypto:generate_key(ecdh, prime256v1), + P256dh = push_utils:base64url_encode(PeerPub), + Auth = push_utils:base64url_encode(crypto:strong_rand_bytes(16)), + {ok, Body} = push_utils:encrypt_payload( + binary:copy(<<"x">>, Budget), P256dh, Auth, RecordSize + ), + ?assertEqual(RecordSize, byte_size(Body)), + ?assertEqual( + {error, max_pad_exceeded}, + push_utils:encrypt_payload( + binary:copy(<<"x">>, Budget + 1), P256dh, Auth, RecordSize + ) + ). + hkdf_expand_test() -> IKM = crypto:strong_rand_bytes(32), Salt = crypto:strong_rand_bytes(16), diff --git a/fluxer_push/Cargo.toml b/fluxer_push/Cargo.toml new file mode 100644 index 000000000..dff23eba8 --- /dev/null +++ b/fluxer_push/Cargo.toml @@ -0,0 +1,37 @@ +# SPDX-License-Identifier: AGPL-3.0-or-later + +[package] +name = "fluxer-push" +version = "0.1.0" +edition.workspace = true +license.workspace = true +publish = false + +[lib] +name = "fluxer_push" +path = "src/lib.rs" + +[[bin]] +name = "fluxer-push" +path = "src/main.rs" + +[dependencies] +fluxer-svc = { path = "../fluxer_svc", default-features = false } +anyhow = "1.0.104" +axum = { version = "0.8.9", default-features = false, features = ["http1", "http2", "tokio"] } +base64 = "0.23.1" +clap = { version = "4.6.7", features = ["derive"] } +futures = "0.3.34" +hmac = "0.13.0" +p256 = { version = "0.13.2", default-features = false, features = ["ecdh", "ecdsa", "pkcs8", "pem", "std"] } +rand = "0.10.2" +reqwest = { version = "0.13.5", default-features = false, features = ["http2", "json", "rustls"] } +ring = "0.17.14" +serde = { version = "1.0.229", features = ["derive"] } +serde_json = "1.0.151" +sha2 = "0.11.0" +thiserror = "2.0.20" +tokio = { version = "1.53.1", features = ["macros", "net", "rt-multi-thread", "signal", "sync", "time"] } +tracing = "0.1.44" +tracing-subscriber = { version = "0.3.23", features = ["env-filter", "fmt", "json"] } +url = "2.5.8" diff --git a/fluxer_push/Dockerfile b/fluxer_push/Dockerfile new file mode 100644 index 000000000..bac309cb7 --- /dev/null +++ b/fluxer_push/Dockerfile @@ -0,0 +1,42 @@ +# SPDX-License-Identifier: AGPL-3.0-or-later + +FROM rust:1-trixie AS builder + +WORKDIR /usr/src/app + +COPY . . + +RUN cargo build --release -p fluxer-push + +FROM debian:trixie-slim + +ARG BUILD_VERSION="" +ARG SOURCE_SHA="" +ARG SOURCE_DATE="" + +LABEL org.opencontainers.image.title="fluxer-push" +LABEL org.opencontainers.image.description="Fluxer push notification delivery service and push relay" +LABEL org.opencontainers.image.licenses="AGPL-3.0-or-later" +LABEL org.opencontainers.image.vendor="Fluxer" +LABEL org.opencontainers.image.url="https://fluxer.app" +LABEL org.opencontainers.image.documentation="https://docs.fluxer.app" +LABEL org.opencontainers.image.source="https://github.com/fluxerapp/fluxer" +LABEL org.opencontainers.image.version="${BUILD_VERSION}" +LABEL org.opencontainers.image.revision="${SOURCE_SHA}" +LABEL org.opencontainers.image.created="${SOURCE_DATE}" +LABEL app.fluxer.build-version="${BUILD_VERSION}" + +WORKDIR /usr/local/bin + +RUN apt-get update && apt-get install -y --no-install-recommends ca-certificates && \ + rm -rf /var/lib/apt/lists/* + +COPY --from=builder /usr/src/app/target/release/fluxer-push /usr/local/bin/fluxer-push + +ENV BUILD_VERSION="${BUILD_VERSION}" + +USER 65532:65532 + +EXPOSE 8126 8127 + +CMD ["/usr/local/bin/fluxer-push"] diff --git a/fluxer_push/src/cli.rs b/fluxer_push/src/cli.rs new file mode 100644 index 000000000..ff7f94243 --- /dev/null +++ b/fluxer_push/src/cli.rs @@ -0,0 +1,59 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use crate::config::{Config, Mode}; +use anyhow::Context as _; +use clap::{Parser, Subcommand}; +use std::net::IpAddr; + +#[derive(Debug, Parser)] +#[command(name = "fluxer-push", disable_help_subcommand = true)] +pub struct Args { + #[arg(long = "mode", value_name = "MODE", value_enum, default_value_t = Mode::Delivery, global = true)] + pub mode: Mode, + + #[arg(long = "bind-host", value_name = "HOST")] + pub bind_host: Option, + + #[arg(long = "port", value_name = "PORT")] + pub port: Option, + + #[command(subcommand)] + pub command: Option, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq, Subcommand)] +pub enum Command { + Healthcheck, +} + +pub fn load_config(args: &Args) -> anyhow::Result { + load_config_from_iter(args, std::env::vars()) +} + +fn load_config_from_iter(args: &Args, vars: I) -> anyhow::Result +where + I: IntoIterator, + K: Into, + V: Into, +{ + let mut cfg = Config::load_from_iter(args.mode, vars)?; + apply_overrides(args, &mut cfg)?; + Ok(cfg) +} + +fn apply_overrides(args: &Args, cfg: &mut Config) -> anyhow::Result<()> { + let bind_addr = cfg.bind_addr_mut(); + if let Some(bind_host) = args.bind_host.as_deref() { + let bind_host = bind_host.trim(); + anyhow::ensure!(!bind_host.is_empty(), "--bind-host cannot be empty"); + bind_addr.set_ip( + bind_host + .parse::() + .with_context(|| format!("--bind-host is not an IP address: {bind_host}"))?, + ); + } + if let Some(port) = args.port { + bind_addr.set_port(port); + } + Ok(()) +} diff --git a/fluxer_push/src/config.rs b/fluxer_push/src/config.rs new file mode 100644 index 000000000..899de2727 --- /dev/null +++ b/fluxer_push/src/config.rs @@ -0,0 +1,670 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use crate::secret::SecretString; +use anyhow::Context as _; +use serde::Deserialize; +use std::net::{IpAddr, SocketAddr}; + +pub const RPC_AUTH_HEADER: &str = "x-fluxer-rpc-auth"; +pub const DEFAULT_APP_ID: &str = "stable"; + +const RPC_PATH: &str = "/internal/rpc"; +const DEFAULT_HOST: &str = "0.0.0.0"; +const DEFAULT_DELIVERY_PORT: u16 = 8126; +const DEFAULT_RELAY_PORT: u16 = 8127; +const DEFAULT_NATS_URL: &str = "nats://127.0.0.1:4222"; +const DEFAULT_QUEUE_CAPACITY: usize = 10_000; +const DEFAULT_SEND_CONCURRENCY: usize = 256; +const DEFAULT_RELAY_MAX_CONCURRENT: usize = 1_024; +const DEFAULT_RELAY_MAX_BODY_BYTES: usize = 2_816; +const DEFAULT_DEVICE_TOKEN_BUCKET_ENTRIES: usize = 1_000_000; +const DEFAULT_DEVICE_TOKEN_BUCKET_PER_MINUTE: u32 = 60; +const DEFAULT_DEVICE_TOKEN_BUCKET_BURST: u32 = 20; +const DEFAULT_SOURCE_BUCKET_ENTRIES: usize = 100_000; +const DEFAULT_SOURCE_BUCKET_PER_MINUTE: u32 = 600; +const DEFAULT_SOURCE_BUCKET_BURST: u32 = 200; +const DEFAULT_FCM_TOKEN_URI: &str = "https://oauth2.googleapis.com/token"; +const DEFAULT_FCM_BASE_URL: &str = "https://fcm.googleapis.com"; +const DEFAULT_CLIENT_IP_HEADER_NAME: &str = "x-forwarded-for"; +const APNS_PRODUCTION_BASE_URL: &str = "https://api.push.apple.com"; +const APNS_DEVELOPMENT_BASE_URL: &str = "https://api.sandbox.push.apple.com"; + +#[derive(Clone, Copy, Debug, Default, Eq, PartialEq, clap::ValueEnum)] +pub enum Mode { + #[default] + Delivery, + Relay, +} + +impl Mode { + pub fn default_port(self) -> u16 { + match self { + Self::Delivery => DEFAULT_DELIVERY_PORT, + Self::Relay => DEFAULT_RELAY_PORT, + } + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum ProviderEnvironment { + Production, + Development, +} + +impl ProviderEnvironment { + pub fn label(self) -> &'static str { + match self { + Self::Production => "production", + Self::Development => "development", + } + } + + pub fn from_label(raw: &str) -> Option { + match raw { + "production" => Some(Self::Production), + "development" => Some(Self::Development), + _ => None, + } + } +} + +#[derive(Clone, Debug)] +pub struct ProviderApp { + pub app_id: String, + pub topic: Option, + pub environment: Option, + pub project_id: Option, +} + +#[derive(Debug)] +pub struct NatsConfig { + pub url: String, + pub auth_token: Option, +} + +#[derive(Debug)] +pub struct RpcConfig { + pub url: String, + pub auth_token: SecretString, +} + +#[derive(Debug)] +pub struct VapidConfig { + pub email: String, + pub public_key: String, + pub private_key: SecretString, +} + +#[derive(Debug)] +pub struct ApnsConfig { + pub team_id: String, + pub key_id: String, + pub private_key: SecretString, + pub default_environment: ProviderEnvironment, + pub apps: Vec, + pub base_url_override: Option, +} + +impl ApnsConfig { + pub fn topic_for(&self, app_id: &str, environment: ProviderEnvironment) -> Option<&str> { + let exact = self.apps.iter().find(|app| { + app.app_id == app_id && app.environment == Some(environment) && app.topic.is_some() + }); + exact + .or_else(|| { + self.apps + .iter() + .find(|app| app.app_id == app_id && app.topic.is_some()) + }) + .and_then(|app| app.topic.as_deref()) + } + + pub fn base_url(&self, environment: ProviderEnvironment) -> &str { + if let Some(base_url) = self.base_url_override.as_deref() { + return base_url; + } + match environment { + ProviderEnvironment::Production => APNS_PRODUCTION_BASE_URL, + ProviderEnvironment::Development => APNS_DEVELOPMENT_BASE_URL, + } + } +} + +#[derive(Debug)] +pub struct FcmConfig { + pub project_id: String, + pub client_email: String, + pub private_key: SecretString, + pub token_uri: String, + pub apps: Vec, + pub base_url: String, +} + +impl FcmConfig { + pub fn project_id_for(&self, app_id: &str) -> &str { + self.apps + .iter() + .find(|app| app.app_id == app_id && app.project_id.is_some()) + .and_then(|app| app.project_id.as_deref()) + .unwrap_or(self.project_id.as_str()) + } + + pub fn listed_project_id(&self, app_id: &str) -> Option<&str> { + self.apps + .iter() + .find(|app| app.app_id == app_id) + .map(|app| { + app.project_id + .as_deref() + .unwrap_or(self.project_id.as_str()) + }) + } +} + +#[derive(Debug)] +pub struct DeliveryConfig { + pub bind_addr: SocketAddr, + pub nats: NatsConfig, + pub rpc: RpcConfig, + pub queue_capacity: usize, + pub send_concurrency: usize, + pub vapid: VapidConfig, + pub apns: Option, + pub fcm: Option, +} + +#[derive(Clone, Copy, Debug)] +pub struct BucketConfig { + pub entries: usize, + pub per_minute: u32, + pub burst: u32, +} + +#[derive(Debug)] +pub struct RelayConfig { + pub bind_addr: SocketAddr, + pub max_concurrent: usize, + pub max_body_bytes: usize, + pub trust_client_ip_header: bool, + pub client_ip_header_name: String, + pub device_token_bucket: BucketConfig, + pub source_bucket: Option, + pub apns: Option, + pub fcm: Option, +} + +#[derive(Debug)] +pub enum Config { + Delivery(Box), + Relay(Box), +} + +impl Config { + pub fn load_from_iter(mode: Mode, vars: I) -> anyhow::Result + where + I: IntoIterator, + K: Into, + V: Into, + { + Ok(match mode { + Mode::Delivery => Self::Delivery(Box::new(DeliveryConfig::load_from_iter(vars)?)), + Mode::Relay => Self::Relay(Box::new(RelayConfig::load_from_iter(vars)?)), + }) + } + + pub fn bind_addr_mut(&mut self) -> &mut SocketAddr { + match self { + Self::Delivery(cfg) => &mut cfg.bind_addr, + Self::Relay(cfg) => &mut cfg.bind_addr, + } + } +} + +impl DeliveryConfig { + pub fn load_from_iter(vars: I) -> anyhow::Result + where + I: IntoIterator, + K: Into, + V: Into, + { + let env = Env::from_iter(vars); + Ok(Self { + bind_addr: bind_addr(&env, Mode::Delivery)?, + nats: nats_config(&env), + rpc: rpc_config(&env)?, + queue_capacity: parse_number( + "FLUXER_PUSH_SERVICE_QUEUE_CAPACITY", + env.get("FLUXER_PUSH_SERVICE_QUEUE_CAPACITY"), + DEFAULT_QUEUE_CAPACITY, + 1, + 1_000_000, + )?, + send_concurrency: parse_number( + "FLUXER_PUSH_SERVICE_SEND_CONCURRENCY", + env.get("FLUXER_PUSH_SERVICE_SEND_CONCURRENCY"), + DEFAULT_SEND_CONCURRENCY, + 1, + 65_536, + )?, + vapid: vapid_config(&env)?, + apns: apns_config(&env)?, + fcm: fcm_config(&env)?, + }) + } +} + +impl RelayConfig { + pub fn load_from_iter(vars: I) -> anyhow::Result + where + I: IntoIterator, + K: Into, + V: Into, + { + let env = Env::from_iter(vars); + let cfg = Self { + bind_addr: bind_addr(&env, Mode::Relay)?, + max_concurrent: parse_number( + "FLUXER_PUSH_RELAY_MAX_CONCURRENT", + env.get("FLUXER_PUSH_RELAY_MAX_CONCURRENT"), + DEFAULT_RELAY_MAX_CONCURRENT, + 1, + 1_000_000, + )?, + max_body_bytes: parse_number( + "FLUXER_PUSH_RELAY_MAX_BODY_BYTES", + env.get("FLUXER_PUSH_RELAY_MAX_BODY_BYTES"), + DEFAULT_RELAY_MAX_BODY_BYTES, + 1, + 1_048_576, + )?, + trust_client_ip_header: parse_bool( + "FLUXER_TRUST_CLIENT_IP_HEADER", + env.get("FLUXER_TRUST_CLIENT_IP_HEADER"), + )? + .unwrap_or(false), + client_ip_header_name: env + .get("FLUXER_CLIENT_IP_HEADER_NAME") + .unwrap_or(DEFAULT_CLIENT_IP_HEADER_NAME) + .to_ascii_lowercase(), + device_token_bucket: bucket_config( + &env, + "FLUXER_PUSH_RELAY_TOKEN_BUCKET", + BucketConfig { + entries: DEFAULT_DEVICE_TOKEN_BUCKET_ENTRIES, + per_minute: DEFAULT_DEVICE_TOKEN_BUCKET_PER_MINUTE, + burst: DEFAULT_DEVICE_TOKEN_BUCKET_BURST, + }, + )?, + source_bucket: source_bucket_config(&env)?, + apns: apns_config(&env)?, + fcm: fcm_config(&env)?, + }; + for (var_name, apps) in [ + ( + "FLUXER_PUSH_APNS_APPS", + cfg.apns.as_ref().map(|apns| &apns.apps), + ), + ( + "FLUXER_PUSH_FCM_APPS", + cfg.fcm.as_ref().map(|fcm| &fcm.apps), + ), + ] { + anyhow::ensure!( + apps.is_none_or(|apps| !apps.is_empty()), + "{var_name} must list at least one app" + ); + } + Ok(cfg) + } +} + +fn bind_addr(env: &Env, mode: Mode) -> anyhow::Result { + let host = env.get("FLUXER_PUSH_SERVICE_HOST").unwrap_or(DEFAULT_HOST); + let ip = host + .parse::() + .with_context(|| format!("FLUXER_PUSH_SERVICE_HOST is not an IP address: {host}"))?; + let port = parse_number( + "FLUXER_PUSH_SERVICE_PORT", + env.get("FLUXER_PUSH_SERVICE_PORT"), + mode.default_port(), + 1, + u16::MAX, + )?; + Ok(SocketAddr::new(ip, port)) +} + +fn nats_config(env: &Env) -> NatsConfig { + NatsConfig { + url: env + .get("FLUXER_SVC_NATS_URL") + .unwrap_or(DEFAULT_NATS_URL) + .to_owned(), + auth_token: env + .get("FLUXER_NATS_AUTH_TOKEN") + .map(|token| SecretString::new(token.to_owned())), + } +} + +fn rpc_config(env: &Env) -> anyhow::Result { + let api_endpoint = env + .require("FLUXER_INTERNAL_API_ENDPOINT")? + .trim_end_matches('/') + .to_owned(); + anyhow::ensure!( + !api_endpoint.is_empty(), + "FLUXER_INTERNAL_API_ENDPOINT must not be only slashes" + ); + Ok(RpcConfig { + url: format!("{api_endpoint}{RPC_PATH}"), + auth_token: SecretString::new(env.require("FLUXER_GATEWAY_RPC_AUTH_TOKEN")?.to_owned()), + }) +} + +fn vapid_config(env: &Env) -> anyhow::Result { + Ok(VapidConfig { + email: env.require("FLUXER_VAPID_EMAIL")?.to_owned(), + public_key: env.require("FLUXER_VAPID_PUBLIC_KEY")?.to_owned(), + private_key: SecretString::new(env.require("FLUXER_VAPID_PRIVATE_KEY")?.to_owned()), + }) +} + +fn bucket_config(env: &Env, prefix: &str, defaults: BucketConfig) -> anyhow::Result { + let entries_var = format!("{prefix}_ENTRIES"); + let per_minute_var = format!("{prefix}_PER_MINUTE"); + let burst_var = format!("{prefix}_BURST"); + Ok(BucketConfig { + entries: parse_number( + &entries_var, + env.get(&entries_var), + defaults.entries, + 1, + 100_000_000, + )?, + per_minute: parse_number( + &per_minute_var, + env.get(&per_minute_var), + defaults.per_minute, + 1, + 1_000_000, + )?, + burst: parse_number( + &burst_var, + env.get(&burst_var), + defaults.burst, + 1, + 1_000_000, + )?, + }) +} + +fn source_bucket_config(env: &Env) -> anyhow::Result> { + if !parse_bool( + "FLUXER_PUSH_RELAY_SOURCE_BUCKET_ENABLED", + env.get("FLUXER_PUSH_RELAY_SOURCE_BUCKET_ENABLED"), + )? + .unwrap_or(false) + { + return Ok(None); + } + bucket_config( + env, + "FLUXER_PUSH_RELAY_SOURCE_BUCKET", + BucketConfig { + entries: DEFAULT_SOURCE_BUCKET_ENTRIES, + per_minute: DEFAULT_SOURCE_BUCKET_PER_MINUTE, + burst: DEFAULT_SOURCE_BUCKET_BURST, + }, + ) + .map(Some) +} + +fn apns_config(env: &Env) -> anyhow::Result> { + if !parse_bool( + "FLUXER_PUSH_APNS_ENABLED", + env.get("FLUXER_PUSH_APNS_ENABLED"), + )? + .unwrap_or(false) + { + return Ok(None); + } + let default_environment = match env.get("FLUXER_PUSH_APNS_DEFAULT_ENVIRONMENT") { + Some(raw) => parse_environment("FLUXER_PUSH_APNS_DEFAULT_ENVIRONMENT", raw)?, + None => ProviderEnvironment::Production, + }; + Ok(Some(ApnsConfig { + team_id: required_field( + env.get("FLUXER_PUSH_APNS_TEAM_ID"), + "FLUXER_PUSH_APNS_TEAM_ID", + "APNs push", + )? + .to_owned(), + key_id: required_field( + env.get("FLUXER_PUSH_APNS_KEY_ID"), + "FLUXER_PUSH_APNS_KEY_ID", + "APNs push", + )? + .to_owned(), + private_key: SecretString::new(read_key( + env, + "FLUXER_PUSH_APNS_PRIVATE_KEY", + "FLUXER_PUSH_APNS_PRIVATE_KEY_PATH", + "APNs push", + )?), + default_environment, + apps: parse_apps("FLUXER_PUSH_APNS_APPS", env.get("FLUXER_PUSH_APNS_APPS"))?, + base_url_override: env + .get("FLUXER_PUSH_SERVICE_APNS_BASE_URL") + .map(ToOwned::to_owned), + })) +} + +fn fcm_config(env: &Env) -> anyhow::Result> { + if !parse_bool( + "FLUXER_PUSH_FCM_ENABLED", + env.get("FLUXER_PUSH_FCM_ENABLED"), + )? + .unwrap_or(false) + { + return Ok(None); + } + let service_account = match env.get("FLUXER_PUSH_FCM_SERVICE_ACCOUNT_JSON_PATH") { + Some(path) => Some(read_service_account(path)?), + None => None, + }; + let project_id = env + .get("FLUXER_PUSH_FCM_PROJECT_ID") + .map(ToOwned::to_owned) + .or_else(|| service_account.as_ref().and_then(|a| a.project_id.clone())); + let client_email = env + .get("FLUXER_PUSH_FCM_CLIENT_EMAIL") + .map(ToOwned::to_owned) + .or_else(|| { + service_account + .as_ref() + .and_then(|a| a.client_email.clone()) + }); + let private_key = match env.get("FLUXER_PUSH_FCM_PRIVATE_KEY") { + Some(key) => key.to_owned(), + None => match env.get("FLUXER_PUSH_FCM_PRIVATE_KEY_PATH") { + Some(path) => read_file("FLUXER_PUSH_FCM_PRIVATE_KEY_PATH", path)?, + None => service_account + .as_ref() + .and_then(|a| a.private_key.clone()) + .ok_or_else(|| { + anyhow::anyhow!( + "FLUXER_PUSH_FCM_PRIVATE_KEY or FLUXER_PUSH_FCM_PRIVATE_KEY_PATH is required when FCM push is enabled" + ) + })?, + }, + }; + + Ok(Some(FcmConfig { + project_id: required_field( + project_id.as_deref(), + "FLUXER_PUSH_FCM_PROJECT_ID", + "FCM push", + )? + .to_owned(), + client_email: required_field( + client_email.as_deref(), + "FLUXER_PUSH_FCM_CLIENT_EMAIL", + "FCM push", + )? + .to_owned(), + private_key: SecretString::new(private_key), + token_uri: env + .get("FLUXER_PUSH_FCM_TOKEN_URI") + .unwrap_or(DEFAULT_FCM_TOKEN_URI) + .to_owned(), + apps: parse_apps("FLUXER_PUSH_FCM_APPS", env.get("FLUXER_PUSH_FCM_APPS"))?, + base_url: env + .get("FLUXER_PUSH_SERVICE_FCM_BASE_URL") + .unwrap_or(DEFAULT_FCM_BASE_URL) + .to_owned(), + })) +} + +#[derive(Default, Deserialize)] +struct ServiceAccount { + project_id: Option, + client_email: Option, + private_key: Option, +} + +fn read_service_account(path: &str) -> anyhow::Result { + let raw = read_file("FLUXER_PUSH_FCM_SERVICE_ACCOUNT_JSON_PATH", path)?; + serde_json::from_str(&raw).with_context(|| { + format!("FLUXER_PUSH_FCM_SERVICE_ACCOUNT_JSON_PATH is not a service account JSON: {path}") + }) +} + +fn read_key(env: &Env, key_var: &str, path_var: &str, provider: &str) -> anyhow::Result { + if let Some(key) = env.get(key_var) { + return Ok(key.to_owned()); + } + match env.get(path_var) { + Some(path) => read_file(path_var, path), + None => Err(anyhow::anyhow!( + "{key_var} or {path_var} is required when {provider} is enabled" + )), + } +} + +fn read_file(var_name: &str, path: &str) -> anyhow::Result { + std::fs::read_to_string(path).with_context(|| format!("{var_name} could not be read: {path}")) +} + +fn required_field<'a>( + value: Option<&'a str>, + var_name: &str, + provider: &str, +) -> anyhow::Result<&'a str> { + value.ok_or_else(|| anyhow::anyhow!("{var_name} is required when {provider} is configured")) +} + +#[derive(Deserialize)] +struct RawProviderApp { + app_id: Option, + topic: Option, + environment: Option, + project_id: Option, +} + +fn parse_apps(var_name: &str, raw: Option<&str>) -> anyhow::Result> { + let Some(raw) = raw else { + return Ok(Vec::new()); + }; + let entries: Vec = serde_json::from_str(raw) + .with_context(|| format!("{var_name} must be a JSON array of push app objects"))?; + entries + .into_iter() + .map(|entry| { + let app_id = entry + .app_id + .filter(|app_id| !app_id.trim().is_empty()) + .ok_or_else(|| anyhow::anyhow!("{var_name} contains an entry with no app_id"))?; + let environment = match entry.environment.as_deref() { + Some(value) => Some(parse_environment(var_name, value)?), + None => None, + }; + Ok(ProviderApp { + app_id, + topic: entry.topic.filter(|topic| !topic.trim().is_empty()), + environment, + project_id: entry + .project_id + .filter(|project_id| !project_id.trim().is_empty()), + }) + }) + .collect() +} + +fn parse_environment(var_name: &str, raw: &str) -> anyhow::Result { + ProviderEnvironment::from_label(&raw.trim().to_ascii_lowercase()) + .ok_or_else(|| anyhow::anyhow!("{var_name} must be one of: production, development")) +} + +fn parse_bool(var_name: &str, raw: Option<&str>) -> anyhow::Result> { + let Some(raw) = raw else { + return Ok(None); + }; + match raw.to_ascii_lowercase().as_str() { + "true" | "1" | "yes" => Ok(Some(true)), + "false" | "0" | "no" => Ok(Some(false)), + _ => Err(anyhow::anyhow!( + "{var_name} must be a boolean: true, false, 1, 0, yes, or no" + )), + } +} + +fn parse_number( + var_name: &str, + raw: Option<&str>, + default_value: T, + min_value: T, + max_value: T, +) -> anyhow::Result +where + T: std::str::FromStr + PartialOrd + std::fmt::Display + Copy, +{ + let Some(raw) = raw else { + return Ok(default_value); + }; + let parsed = raw + .parse::() + .map_err(|_| anyhow::anyhow!("{var_name} must be a number"))?; + anyhow::ensure!( + parsed >= min_value && parsed <= max_value, + "{var_name} must be between {min_value} and {max_value}" + ); + Ok(parsed) +} + +struct Env(Vec<(String, String)>); + +impl Env { + fn from_iter(vars: I) -> Self + where + I: IntoIterator, + K: Into, + V: Into, + { + Self( + vars.into_iter() + .map(|(key, value)| (key.into(), value.into())) + .collect(), + ) + } + + fn get(&self, key: &str) -> Option<&str> { + self.0 + .iter() + .find_map(|(name, value)| (name == key).then_some(value.trim())) + .filter(|value| !value.is_empty()) + } + + fn require(&self, key: &str) -> anyhow::Result<&str> { + self.get(key) + .ok_or_else(|| anyhow::anyhow!("{key} is required")) + } +} diff --git a/fluxer_push/src/crypto.rs b/fluxer_push/src/crypto.rs new file mode 100644 index 000000000..1622fc0b2 --- /dev/null +++ b/fluxer_push/src/crypto.rs @@ -0,0 +1,233 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use base64::prelude::*; +use hmac::{Hmac, KeyInit, Mac}; +use p256::PublicKey; +use p256::ecdsa::signature::Signer as _; +use p256::ecdsa::{Signature, SigningKey}; +use p256::elliptic_curve::sec1::ToEncodedPoint as _; +use p256::pkcs8::DecodePrivateKey as _; +use rand::Rng as _; +use ring::aead::{AES_128_GCM, Aad, LessSafeKey, Nonce, UnboundKey}; +use ring::rand::SystemRandom; +use ring::signature::{RSA_PKCS1_SHA256, RsaKeyPair}; +use serde_json::Value; +use sha2::Sha256; +use thiserror::Error; + +type HmacSha256 = Hmac; + +pub const ES256: &str = "ES256"; + +const SALT_LEN: usize = 16; +const TAG_LEN: usize = 16; +const RECORD_SIZE_LEN: usize = 4; +const KEY_ID_LEN_LEN: usize = 1; +const IKM_LEN: usize = 32; +const CEK_LEN: usize = 16; +const NONCE_LEN: usize = 12; +const SCALAR_LEN: usize = 32; +const UNCOMPRESSED_POINT_LEN: usize = 65; +const UNCOMPRESSED_POINT_TAG: u8 = 0x04; +const PADDING_DELIMITER: u8 = 0x02; +const EPHEMERAL_KEY_TRIES: usize = 4; + +#[derive(Debug, Error)] +pub enum CryptoError { + #[error("subscription key is not base64")] + InvalidSubscriptionKey, + #[error("subscription public key is not an uncompressed P-256 point")] + InvalidPeerKey, + #[error("private key is not a P-256 key")] + InvalidP256Key, + #[error("private key is not an RSA key")] + InvalidRsaKey, + #[error("no ephemeral P-256 key could be generated")] + EphemeralKey, + #[error("record size {0} is too small to frame a payload")] + RecordSizeTooSmall(usize), + #[error("payload does not fit in a {0} byte record")] + MaxPadExceeded(usize), + #[error("sealing the payload failed")] + SealFailed, + #[error("signing failed")] + SigningFailed, +} + +pub fn decode_subscription_key(value: &str) -> Result, CryptoError> { + let trimmed = value.trim().trim_end_matches('='); + BASE64_URL_SAFE_NO_PAD + .decode(trimmed) + .or_else(|_| BASE64_STANDARD_NO_PAD.decode(trimmed)) + .map_err(|_| CryptoError::InvalidSubscriptionKey) +} + +pub fn parse_p256_private_scalar(value: &str) -> Result { + let raw = decode_subscription_key(value)?; + SigningKey::from_slice(&raw).map_err(|_| CryptoError::InvalidP256Key) +} + +pub fn parse_p256_pkcs8_pem(pem: &str) -> Result { + let secret = p256::SecretKey::from_pkcs8_pem(&expand_escaped_newlines(pem)) + .map_err(|_| CryptoError::InvalidP256Key)?; + Ok(SigningKey::from(&secret)) +} + +pub fn parse_rsa_pkcs8_pem(pem: &str) -> Result { + let der = pem_to_der(pem)?; + RsaKeyPair::from_pkcs8(&der).map_err(|_| CryptoError::InvalidRsaKey) +} + +pub fn es256_jwt(header: &Value, claims: &Value, key: &SigningKey) -> Result { + let signing_input = signing_input(header, claims); + let signature: Signature = key + .try_sign(signing_input.as_bytes()) + .map_err(|_| CryptoError::SigningFailed)?; + Ok(format!( + "{signing_input}.{}", + BASE64_URL_SAFE_NO_PAD.encode(signature.to_bytes()) + )) +} + +pub fn rs256_jwt(header: &Value, claims: &Value, key: &RsaKeyPair) -> Result { + let signing_input = signing_input(header, claims); + let mut signature = vec![0u8; key.public().modulus_len()]; + key.sign( + &RSA_PKCS1_SHA256, + &SystemRandom::new(), + signing_input.as_bytes(), + &mut signature, + ) + .map_err(|_| CryptoError::SigningFailed)?; + Ok(format!( + "{signing_input}.{}", + BASE64_URL_SAFE_NO_PAD.encode(&signature) + )) +} + +pub fn encrypt_aes128gcm( + message: &[u8], + p256dh: &[u8], + auth: &[u8], + record_size: usize, +) -> Result, CryptoError> { + let mut rng = rand::rng(); + let mut salt = [0u8; SALT_LEN]; + rng.fill_bytes(&mut salt); + let local_secret = ephemeral_secret(&mut rng)?; + seal(message, p256dh, auth, &local_secret, salt, record_size) +} + +fn seal( + message: &[u8], + p256dh: &[u8], + auth: &[u8], + local_secret: &p256::SecretKey, + salt: [u8; SALT_LEN], + record_size: usize, +) -> Result, CryptoError> { + if p256dh.len() != UNCOMPRESSED_POINT_LEN || p256dh[0] != UNCOMPRESSED_POINT_TAG { + return Err(CryptoError::InvalidPeerKey); + } + let peer_public = + PublicKey::from_sec1_bytes(p256dh).map_err(|_| CryptoError::InvalidPeerKey)?; + let local_point = local_secret.public_key().to_encoded_point(false); + let local_pub = local_point.as_bytes(); + + let shared = + p256::ecdh::diffie_hellman(local_secret.to_nonzero_scalar(), peer_public.as_affine()); + let mut info = Vec::with_capacity(b"WebPush: info\0".len() + p256dh.len() + local_pub.len()); + info.extend_from_slice(b"WebPush: info\0"); + info.extend_from_slice(p256dh); + info.extend_from_slice(local_pub); + let ikm = hkdf_sha256(shared.raw_secret_bytes(), auth, &info, IKM_LEN); + let cek = hkdf_sha256(&ikm, &salt, b"Content-Encoding: aes128gcm\0", CEK_LEN); + let nonce = hkdf_sha256(&ikm, &salt, b"Content-Encoding: nonce\0", NONCE_LEN); + + let header_len = SALT_LEN + RECORD_SIZE_LEN + KEY_ID_LEN_LEN + local_pub.len(); + let required = record_size + .checked_sub(TAG_LEN) + .and_then(|record_len| record_len.checked_sub(header_len)) + .ok_or(CryptoError::RecordSizeTooSmall(record_size))?; + let record_size_be = u32::try_from(record_size) + .map_err(|_| CryptoError::RecordSizeTooSmall(record_size))? + .to_be_bytes(); + + let mut data = Vec::with_capacity(required + TAG_LEN); + data.extend_from_slice(message); + data.push(PADDING_DELIMITER); + if data.len() > required { + return Err(CryptoError::MaxPadExceeded(record_size)); + } + data.resize(required, 0); + + let key = + LessSafeKey::new(UnboundKey::new(&AES_128_GCM, &cek).map_err(|_| CryptoError::SealFailed)?); + let nonce = Nonce::try_assume_unique_for_key(&nonce).map_err(|_| CryptoError::SealFailed)?; + key.seal_in_place_append_tag(nonce, Aad::empty(), &mut data) + .map_err(|_| CryptoError::SealFailed)?; + + let mut body = Vec::with_capacity(header_len + data.len()); + body.extend_from_slice(&salt); + body.extend_from_slice(&record_size_be); + body.push(u8::try_from(local_pub.len()).map_err(|_| CryptoError::InvalidP256Key)?); + body.extend_from_slice(local_pub); + body.append(&mut data); + Ok(body) +} + +fn ephemeral_secret(rng: &mut impl rand::Rng) -> Result { + let mut scalar = [0u8; SCALAR_LEN]; + for _ in 0..EPHEMERAL_KEY_TRIES { + rng.fill_bytes(&mut scalar); + if let Ok(secret) = p256::SecretKey::from_slice(&scalar) { + return Ok(secret); + } + } + Err(CryptoError::EphemeralKey) +} + +fn hkdf_sha256(ikm: &[u8], salt: &[u8], info: &[u8], len: usize) -> Vec { + let mut extract = HmacSha256::new_from_slice(salt).expect("hmac accepts any key length"); + extract.update(ikm); + let prk = extract.finalize().into_bytes(); + + let mut out = Vec::with_capacity(len + IKM_LEN); + let mut block = Vec::new(); + let mut counter: u8 = 1; + while out.len() < len { + let mut expand = HmacSha256::new_from_slice(&prk).expect("hmac accepts any key length"); + expand.update(&block); + expand.update(info); + expand.update(&[counter]); + block = expand.finalize().into_bytes().to_vec(); + out.extend_from_slice(&block); + counter = counter.wrapping_add(1); + } + out.truncate(len); + out +} + +fn signing_input(header: &Value, claims: &Value) -> String { + format!("{}.{}", encode_segment(header), encode_segment(claims)) +} + +fn encode_segment(value: &Value) -> String { + BASE64_URL_SAFE_NO_PAD.encode(serde_json::to_vec(value).expect("a json value serialises")) +} + +fn expand_escaped_newlines(value: &str) -> String { + value.replace("\\n", "\n") +} + +fn pem_to_der(pem: &str) -> Result, CryptoError> { + let normalized = expand_escaped_newlines(pem); + let body: String = normalized + .lines() + .filter(|line| !line.trim().starts_with("-----")) + .flat_map(|line| line.chars().filter(|c| !c.is_ascii_whitespace())) + .collect(); + BASE64_STANDARD + .decode(body) + .map_err(|_| CryptoError::InvalidRsaKey) +} diff --git a/fluxer_push/src/dedupe.rs b/fluxer_push/src/dedupe.rs new file mode 100644 index 000000000..3493b0522 --- /dev/null +++ b/fluxer_push/src/dedupe.rs @@ -0,0 +1,123 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use crate::job::{ClearJob, MessageJob}; +use sha2::{Digest, Sha256}; +use std::collections::{HashMap, HashSet, VecDeque}; +use std::sync::{Arc, Mutex}; +use std::time::{Duration, Instant}; + +const DONE_TTL: Duration = Duration::from_secs(300); +const MAX_DONE: usize = 200_000; +const KEY_BYTES: usize = 16; + +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct JobKey([u8; KEY_BYTES]); + +impl JobKey { + pub fn of_message(job: &MessageJob, user_id: &str) -> Self { + Self::digest(&["message", &job.message_id, &job.channel_id, user_id]) + } + + pub fn of_clear(job: &ClearJob) -> Self { + Self::digest(&["clear", &job.user_id, &job.channel_id, &job.message_id]) + } + + fn digest(parts: &[&str]) -> Self { + let mut hasher = Sha256::new(); + for part in parts { + hasher.update(part.as_bytes()); + hasher.update([0]); + } + let mut key = [0u8; KEY_BYTES]; + key.copy_from_slice(&hasher.finalize()[..KEY_BYTES]); + Self(key) + } +} + +pub enum Seen { + New(Claim), + Done, + Running, +} + +#[derive(Clone, Default)] +pub struct DoneJobs { + inner: Arc>, +} + +#[derive(Default)] +struct Entries { + running: HashSet, + done: HashMap, + done_order: VecDeque<(Instant, JobKey)>, +} + +impl Entries { + fn prune(&mut self, now: Instant) { + while let Some(&(at, key)) = self.done_order.front() { + if now.duration_since(at) < DONE_TTL && self.done_order.len() <= MAX_DONE { + break; + } + self.done_order.pop_front(); + if self.done.get(&key) == Some(&at) { + self.done.remove(&key); + } + } + } +} + +impl DoneJobs { + pub fn claim(&self, key: JobKey) -> Seen { + let mut entries = self + .inner + .lock() + .expect("the done job set is never poisoned"); + entries.prune(Instant::now()); + if entries.done.contains_key(&key) { + return Seen::Done; + } + if !entries.running.insert(key) { + return Seen::Running; + } + Seen::New(Claim { + key, + jobs: self.clone(), + finished: false, + }) + } + + fn finish(&self, key: JobKey, done: bool) { + let mut entries = self + .inner + .lock() + .expect("the done job set is never poisoned"); + entries.running.remove(&key); + if done { + let now = Instant::now(); + entries.done.insert(key, now); + entries.done_order.push_back((now, key)); + entries.prune(now); + } + } +} + +pub struct Claim { + key: JobKey, + jobs: DoneJobs, + finished: bool, +} + +impl Claim { + pub fn done(mut self) { + self.jobs.finish(self.key, true); + self.finished = true; + } +} + +impl Drop for Claim { + fn drop(&mut self) { + if !self.finished { + self.jobs.finish(self.key, false); + } + } +} diff --git a/fluxer_push/src/delivery.rs b/fluxer_push/src/delivery.rs new file mode 100644 index 000000000..669676260 --- /dev/null +++ b/fluxer_push/src/delivery.rs @@ -0,0 +1,555 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use crate::dedupe::{Claim, DoneJobs, JobKey, Seen}; +use crate::job::{ + self, ClearJob, JobError, MessageJob, QUEUE_GROUP, SUBJECT_CLEAR, SUBJECT_MESSAGE, +}; +use crate::metrics::{JobKind, JobRejection, Provider, elapsed_ms}; +use crate::payload; +use crate::providers::{self, SendOutcome}; +use crate::retry::{self, RETRY_DEADLINE}; +use crate::rpc::RpcError; +use crate::server::AppState; +use crate::subscription::Subscription; +use fluxer_svc::metrics::now_ms; +use fluxer_svc::transport::{Transport, TransportMessage, TransportSubscriber}; +use futures::future::join_all; +use serde_json::Value; +use std::collections::HashMap; +use std::sync::Arc; +use std::time::{Duration, Instant}; +use tokio::sync::{Semaphore, TryAcquireError, watch}; +use tokio::task::JoinSet; +use tracing::{error, info, warn}; + +const JOB_REPLY_DEADLINE: Duration = Duration::from_secs(90); + +const UNKNOWN_PROVIDER: &str = "unknown"; +const OVERLOADED: &str = "overloaded"; +const INVALID_JOB: &str = "invalid_job"; +const RPC_UNAVAILABLE: &str = "rpc_unavailable"; +const DEADLINE_EXCEEDED: &str = "deadline_exceeded"; +const RUNNING: &str = "in_flight"; + +enum Job { + Message(Box), + Clear(ClearJob), +} + +impl Job { + fn kind(&self) -> JobKind { + match self { + Self::Message(_) => JobKind::Message, + Self::Clear(_) => JobKind::Clear, + } + } +} + +struct Claimed { + job: Job, + claims: Vec, + recipients_running: bool, +} + +fn claim_recipients(done: &DoneJobs, job: Job) -> Claimed { + match job { + Job::Message(mut job) => { + let mut claims = Vec::new(); + let mut recipients = Vec::new(); + let mut recipients_running = false; + for user_id in std::mem::take(&mut job.user_ids) { + match done.claim(JobKey::of_message(&job, &user_id)) { + Seen::New(claim) => { + claims.push(claim); + recipients.push(user_id); + } + Seen::Done => {} + Seen::Running => recipients_running = true, + } + } + job.user_ids = recipients; + Claimed { + job: Job::Message(job), + claims, + recipients_running, + } + } + Job::Clear(job) => { + let (claims, recipients_running) = match done.claim(JobKey::of_clear(&job)) { + Seen::New(claim) => (vec![claim], false), + Seen::Done => (Vec::new(), false), + Seen::Running => (Vec::new(), true), + }; + Claimed { + job: Job::Clear(job), + claims, + recipients_running, + } + } + } +} + +enum Answer { + Done, + NotDone(&'static str), +} + +impl Answer { + fn body(&self) -> String { + match self { + Self::Done => r#"{"ok":true}"#.to_owned(), + Self::NotDone(error) => format!(r#"{{"ok":false,"error":"{error}"}}"#), + } + } +} + +struct Summary { + subscriptions: u64, + accepted: u64, + deleted: u64, +} + +struct Delivered { + accepted: bool, + deletion: Option<(String, String)>, +} + +pub async fn run_job_subscribers(transport: T, state: Arc) { + let admission = Arc::new(Semaphore::new(state.cfg.queue_capacity)); + let sends = Arc::new(Semaphore::new(state.cfg.send_concurrency)); + let done = DoneJobs::default(); + tokio::join!( + run_subject( + transport.clone(), + Arc::clone(&state), + Arc::clone(&admission), + Arc::clone(&sends), + done.clone(), + JobKind::Message, + ), + run_subject(transport, state, admission, sends, done, JobKind::Clear), + ); +} + +async fn run_subject( + transport: T, + state: Arc, + admission: Arc, + sends: Arc, + done: DoneJobs, + kind: JobKind, +) { + let subject = subject_of(kind); + let capacity = state.cfg.queue_capacity; + let mut draining = state.draining.subscribe(); + let mut running = JoinSet::new(); + 'serving: while !is_draining(&draining) { + let mut subscriber = match transport.subscribe_queue(subject, QUEUE_GROUP).await { + Ok(subscriber) => subscriber, + Err(error) => { + warn!( + error = %error, + subject, + "push job subscribe failed" + ); + tokio::select! { + () = drained(&mut draining) => {} + () = transport.wait_for_reconnect() => {} + } + continue; + } + }; + info!( + subject, + queue_group = QUEUE_GROUP, + capacity, + send_concurrency = state.cfg.send_concurrency, + "listening for push jobs" + ); + + loop { + let message = tokio::select! { + biased; + () = drained(&mut draining) => break 'serving, + finished = running.join_next(), if !running.is_empty() => { + reap(finished.expect("a non-empty push job set yields a result")); + record_depth(&state, &admission, capacity); + continue; + } + message = subscriber.next() => { + let Some(message) = message else { + warn!(subject, "push job subscription ended, will re-subscribe"); + break; + }; + message + } + }; + + while let Some(finished) = running.try_join_next() { + reap(finished); + } + state.metrics.record_job_received(kind); + let reply_to = message.reply_subject().map(str::to_owned); + + let job = match decode(kind, message.payload()) { + Ok(job) => job, + Err(error) => { + state.metrics.record_job_rejected(JobRejection::Decode); + warn!(error = %error, subject, "push job rejected"); + answer(&transport, reply_to, Answer::NotDone(INVALID_JOB)).await; + continue; + } + }; + let permit = match Arc::clone(&admission).try_acquire_owned() { + Ok(permit) => permit, + Err(TryAcquireError::NoPermits) => { + state.metrics.record_job_rejected(JobRejection::QueueFull); + warn!(subject, capacity, "push job refused, no capacity left"); + answer(&transport, reply_to, Answer::NotDone(OVERLOADED)).await; + continue; + } + Err(TryAcquireError::Closed) => return, + }; + let Claimed { + job, + claims, + recipients_running, + } = claim_recipients(&done, job); + if claims.is_empty() { + let settled = if recipients_running { + Answer::NotDone(RUNNING) + } else { + Answer::Done + }; + answer(&transport, reply_to, settled).await; + continue; + } + record_depth(&state, &admission, capacity); + + let job_state = Arc::clone(&state); + let job_sends = Arc::clone(&sends); + let job_transport = transport.clone(); + running.spawn(async move { + let _permit = permit; + let mut result = run_job(&job_state, &job_sends, job).await; + if matches!(result, Answer::Done) { + for claim in claims { + claim.done(); + } + if recipients_running { + result = Answer::NotDone(RUNNING); + } + } + answer(&job_transport, reply_to, result).await; + }); + } + } + + info!( + subject, + running = running.len(), + "push job subscription stopped" + ); + while let Some(finished) = running.join_next().await { + reap(finished); + record_depth(&state, &admission, capacity); + } +} + +fn is_draining(draining: &watch::Receiver) -> bool { + *draining.borrow() +} + +async fn drained(draining: &mut watch::Receiver) { + let _ = draining.wait_for(|draining| *draining).await; +} + +async fn run_job(state: &AppState, sends: &Semaphore, job: Job) -> Answer { + let kind = job.kind(); + let work = async { + match job { + Job::Message(job) => run_message_job(state, sends, *job).await, + Job::Clear(job) => run_clear_job(state, sends, job).await, + } + }; + match tokio::time::timeout(JOB_REPLY_DEADLINE, work).await { + Ok(Ok(())) => Answer::Done, + Ok(Err(error)) => { + error!(kind = kind.label(), error = %error, "push job failed"); + Answer::NotDone(RPC_UNAVAILABLE) + } + Err(_) => { + error!( + kind = kind.label(), + deadline_ms = JOB_REPLY_DEADLINE.as_millis() as u64, + "push job ran past its reply deadline" + ); + Answer::NotDone(DEADLINE_EXCEEDED) + } + } +} + +async fn answer(transport: &T, reply_to: Option, answer: Answer) { + let Some(reply_to) = reply_to else { + return; + }; + if let Err(error) = transport.publish(&reply_to, answer.body().as_bytes()).await { + warn!(error = %error, "push job answer failed"); + } +} + +async fn run_message_job( + state: &AppState, + sends: &Semaphore, + job: MessageJob, +) -> anyhow::Result<()> { + let started_ms = now_ms(); + let deadline = Instant::now() + RETRY_DEADLINE; + let (badges, subscriptions) = lookup(deadline, || async { + tokio::try_join!( + state.rpc.badge_counts(&job.user_ids), + state.rpc.push_subscriptions(&job.user_ids) + ) + }) + .await?; + state.metrics.record_recipients(job.user_ids.len() as u64); + + let envelopes = job + .user_ids + .iter() + .filter(|user_id| subscribed(&subscriptions, user_id)) + .map(|user_id| { + ( + user_id.as_str(), + payload::web_push_message(&job, user_id, badge_of(&badges, user_id)), + ) + }) + .collect(); + let summary = deliver(state, sends, &subscriptions, envelopes, deadline).await; + let duration_ms = elapsed_ms(started_ms); + + info!( + kind = JobKind::Message.label(), + guild_id = %job.guild_id, + channel_id = %job.channel_id, + message_id = %job.message_id, + config_version = job.config_version, + recipients = job.user_ids.len(), + subscriptions = summary.subscriptions, + accepted = summary.accepted, + deleted = summary.deleted, + duration_ms, + "push job" + ); + state + .metrics + .record_job_completed(JobKind::Message, duration_ms); + Ok(()) +} + +async fn run_clear_job(state: &AppState, sends: &Semaphore, job: ClearJob) -> anyhow::Result<()> { + let started_ms = now_ms(); + let deadline = Instant::now() + RETRY_DEADLINE; + let user_ids = std::slice::from_ref(&job.user_id); + let (badges, subscriptions) = lookup(deadline, || async { + tokio::try_join!( + state.rpc.badge_counts(user_ids), + state.rpc.push_subscriptions(user_ids) + ) + }) + .await?; + state.metrics.record_recipients(1); + + let envelope = payload::web_push_clear(&job, badge_of(&badges, &job.user_id)); + let envelopes = vec![(job.user_id.as_str(), envelope)]; + let summary = deliver(state, sends, &subscriptions, envelopes, deadline).await; + let duration_ms = elapsed_ms(started_ms); + + info!( + kind = JobKind::Clear.label(), + user_id = %job.user_id, + channel_id = %job.channel_id, + message_id = %job.message_id, + config_version = job.config_version, + recipients = 1, + subscriptions = summary.subscriptions, + accepted = summary.accepted, + deleted = summary.deleted, + duration_ms, + "push job" + ); + state + .metrics + .record_job_completed(JobKind::Clear, duration_ms); + Ok(()) +} + +async fn lookup(deadline: Instant, mut call: F) -> Result +where + F: FnMut() -> Fut, + Fut: Future>, +{ + let mut attempt = 0; + loop { + let error = match call().await { + Ok(found) => return Ok(found), + Err(error) if error.is_retryable() => error, + Err(error) => return Err(error), + }; + let Some(at) = retry::next_attempt(attempt, Instant::now(), deadline) else { + return Err(error); + }; + warn!(attempt, error = %error, "push job lookup failed, retrying"); + tokio::time::sleep_until(at.into()).await; + attempt += 1; + } +} + +async fn deliver( + state: &AppState, + sends: &Semaphore, + subscriptions: &HashMap>, + envelopes: Vec<(&str, Value)>, + deadline: Instant, +) -> Summary { + let mut pending = Vec::new(); + for (user_id, envelope) in envelopes { + let Some(subscriptions) = subscriptions.get(user_id) else { + continue; + }; + let envelope = Arc::new(envelope); + for subscription in subscriptions { + pending.push(send_one( + state, + sends, + user_id, + subscription, + Arc::clone(&envelope), + deadline, + )); + } + } + + let mut summary = Summary { + subscriptions: pending.len() as u64, + accepted: 0, + deleted: 0, + }; + state.metrics.record_subscriptions(summary.subscriptions); + + let mut deletions = Vec::new(); + for delivered in join_all(pending).await { + if delivered.accepted { + summary.accepted += 1; + } + if let Some(deletion) = delivered.deletion { + deletions.push(deletion); + } + } + summary.deleted = deletions.len() as u64; + + if !deletions.is_empty() + && let Err(error) = state.rpc.delete_push_subscriptions(&deletions).await + { + warn!( + error = %error, + count = deletions.len(), + "push subscription cleanup failed" + ); + } + summary +} + +async fn send_one( + state: &AppState, + sends: &Semaphore, + user_id: &str, + subscription: &Subscription, + envelope: Arc, + deadline: Instant, +) -> Delivered { + let provider = subscription.platform().map(providers::provider_of); + let mut attempt = 0; + let outcome = loop { + let outcome = { + let _permit = sends + .acquire() + .await + .expect("the push send semaphore is never closed"); + providers::send(state, subscription, &envelope).await + }; + if !matches!(outcome, SendOutcome::Transient { .. }) { + break outcome; + } + let now = Instant::now(); + let Some(at) = retry::next_attempt(attempt, now, deadline) else { + break outcome; + }; + warn!( + provider = provider.map_or(UNKNOWN_PROVIDER, Provider::label), + subscription_id = %subscription.subscription_id, + reason = outcome.reason(), + attempt, + "push delivery retrying" + ); + tokio::time::sleep_until(at.into()).await; + attempt += 1; + }; + + if !matches!(outcome, SendOutcome::Accepted) { + error!( + provider = provider.map_or(UNKNOWN_PROVIDER, Provider::label), + subscription_id = %subscription.subscription_id, + reason = outcome.reason(), + attempts = attempt + 1, + "push delivery failed" + ); + } + let mut deletion = None; + if outcome.deletes_token() { + if let Some(provider) = provider { + state.metrics.record_token_deletion(provider); + } + deletion = Some((user_id.to_owned(), subscription.subscription_id.clone())); + } + Delivered { + accepted: matches!(outcome, SendOutcome::Accepted), + deletion, + } +} + +fn decode(kind: JobKind, payload: &[u8]) -> Result { + match kind { + JobKind::Message => job::decode_message(payload).map(Box::new).map(Job::Message), + JobKind::Clear => job::decode_clear(payload).map(Job::Clear), + } +} + +fn subject_of(kind: JobKind) -> &'static str { + match kind { + JobKind::Message => SUBJECT_MESSAGE, + JobKind::Clear => SUBJECT_CLEAR, + } +} + +fn subscribed(subscriptions: &HashMap>, user_id: &str) -> bool { + subscriptions + .get(user_id) + .is_some_and(|subscriptions| !subscriptions.is_empty()) +} + +fn badge_of(badges: &HashMap, user_id: &str) -> u32 { + badges.get(user_id).copied().unwrap_or(0) +} + +fn record_depth(state: &AppState, admission: &Semaphore, capacity: usize) { + state + .metrics + .set_queue_depth(capacity.saturating_sub(admission.available_permits()) as u64); +} + +fn reap(result: Result<(), tokio::task::JoinError>) { + if let Err(error) = result { + warn!(error = %error, "push job task failed"); + } +} diff --git a/fluxer_push/src/healthcheck.rs b/fluxer_push/src/healthcheck.rs new file mode 100644 index 000000000..03970b5b7 --- /dev/null +++ b/fluxer_push/src/healthcheck.rs @@ -0,0 +1,59 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use crate::config::Mode; +use anyhow::Context as _; +use std::{ + env, + net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}, + time::Duration, +}; + +pub async fn run(mode: Mode) -> anyhow::Result<()> { + let addr = target( + env::var("FLUXER_PUSH_SERVICE_HOST").ok().as_deref(), + env::var("FLUXER_PUSH_SERVICE_PORT").ok().as_deref(), + mode, + )?; + probe(addr).await +} + +fn target(host: Option<&str>, port: Option<&str>, mode: Mode) -> anyhow::Result { + let host = host + .map(str::trim) + .filter(|host| !host.is_empty()) + .unwrap_or("127.0.0.1"); + let ip = host + .parse::() + .with_context(|| format!("FLUXER_PUSH_SERVICE_HOST is not an IP address: {host}"))?; + let ip = match ip { + IpAddr::V4(ip) if ip.is_unspecified() => IpAddr::V4(Ipv4Addr::LOCALHOST), + IpAddr::V6(ip) if ip.is_unspecified() => IpAddr::V6(Ipv6Addr::LOCALHOST), + ip => ip, + }; + let port = match port.map(str::trim).filter(|port| !port.is_empty()) { + Some(port) => port + .parse::() + .with_context(|| format!("FLUXER_PUSH_SERVICE_PORT is not a port number: {port}"))?, + None => mode.default_port(), + }; + Ok(SocketAddr::new(ip, port)) +} + +async fn probe(addr: SocketAddr) -> anyhow::Result<()> { + let client = reqwest::Client::builder() + .connect_timeout(Duration::from_millis(500)) + .timeout(Duration::from_millis(2_000)) + .no_proxy() + .build()?; + let status = client + .get(format!("http://{addr}/_health")) + .send() + .await + .with_context(|| format!("health request to {addr} failed"))? + .status(); + anyhow::ensure!( + status == reqwest::StatusCode::OK, + "health returned {status}" + ); + Ok(()) +} diff --git a/fluxer_push/src/job.rs b/fluxer_push/src/job.rs new file mode 100644 index 000000000..d87210bb1 --- /dev/null +++ b/fluxer_push/src/job.rs @@ -0,0 +1,77 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use serde::Deserialize; +use thiserror::Error; + +pub const SUBJECT_MESSAGE: &str = "push.job.message"; +pub const SUBJECT_CLEAR: &str = "push.job.clear"; +pub const QUEUE_GROUP: &str = "fluxer-push"; + +const SUPPORTED_VERSION: u8 = 1; +const DIRECT_MESSAGE_GUILD_ID: &str = "0"; + +#[derive(Clone, Debug, Deserialize, Eq, PartialEq)] +pub struct MessageJob { + pub v: u8, + pub config_version: u64, + pub guild_id: String, + pub channel_id: String, + pub message_id: String, + pub notification: NotificationFields, + pub user_ids: Vec, +} + +#[derive(Clone, Debug, Deserialize, Eq, PartialEq)] +pub struct NotificationFields { + pub title: String, + pub body: String, + pub icon: String, + pub badge: String, + pub tag: String, + pub notification_tag: String, + pub url: String, + #[serde(default)] + pub image_url: Option, +} + +#[derive(Clone, Debug, Deserialize, Eq, PartialEq)] +pub struct ClearJob { + pub v: u8, + pub config_version: u64, + pub user_id: String, + pub channel_id: String, + pub message_id: String, +} + +#[derive(Debug, Error)] +pub enum JobError { + #[error("push job version {0} is not supported")] + UnsupportedVersion(u8), + #[error("push job is not a valid job document: {0}")] + Decode(#[from] serde_json::Error), +} + +impl MessageJob { + pub fn is_direct_message(&self) -> bool { + self.guild_id == DIRECT_MESSAGE_GUILD_ID + } +} + +pub fn decode_message(bytes: &[u8]) -> Result { + let job: MessageJob = serde_json::from_slice(bytes)?; + supported(job.v)?; + Ok(job) +} + +pub fn decode_clear(bytes: &[u8]) -> Result { + let job: ClearJob = serde_json::from_slice(bytes)?; + supported(job.v)?; + Ok(job) +} + +fn supported(version: u8) -> Result<(), JobError> { + if version == SUPPORTED_VERSION { + return Ok(()); + } + Err(JobError::UnsupportedVersion(version)) +} diff --git a/fluxer_push/src/lib.rs b/fluxer_push/src/lib.rs new file mode 100644 index 000000000..351ed6f62 --- /dev/null +++ b/fluxer_push/src/lib.rs @@ -0,0 +1,30 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +pub mod cli; +pub mod config; +mod crypto; +mod dedupe; +mod delivery; +pub mod healthcheck; +mod job; +mod metrics; +mod payload; +mod providers; +mod relay; +mod resolver; +mod retry; +mod rollout; +mod rpc; +mod secret; +pub mod server; +mod subscription; +mod tokens; +mod vendor; + +pub use server::run; + +fn unix_seconds() -> i64 { + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .map_or(0, |since| since.as_secs() as i64) +} diff --git a/fluxer_push/src/main.rs b/fluxer_push/src/main.rs new file mode 100644 index 000000000..ebe611745 --- /dev/null +++ b/fluxer_push/src/main.rs @@ -0,0 +1,21 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use clap::Parser; +use fluxer_push::{cli, healthcheck, run}; +use tracing_subscriber::{EnvFilter, layer::SubscriberExt, util::SubscriberInitExt}; + +#[tokio::main(flavor = "multi_thread")] +async fn main() -> anyhow::Result<()> { + let args = cli::Args::parse(); + if matches!(args.command, Some(cli::Command::Healthcheck)) { + return healthcheck::run(args.mode).await; + } + + tracing_subscriber::registry() + .with(EnvFilter::try_from_default_env().unwrap_or_else(|_| EnvFilter::new("info"))) + .with(tracing_subscriber::fmt::layer().json()) + .init(); + + let cfg = cli::load_config(&args)?; + run(cfg).await +} diff --git a/fluxer_push/src/metrics.rs b/fluxer_push/src/metrics.rs new file mode 100644 index 000000000..be77bdb75 --- /dev/null +++ b/fluxer_push/src/metrics.rs @@ -0,0 +1,698 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use crate::relay::reject::{REASON_COUNT, Reason}; +use crate::rollout::{RolloutOutcome, RolloutSnapshot}; +use fluxer_svc::metrics::now_ms; +use std::fmt::{self, Write as _}; +use std::sync::atomic::{AtomicU64, Ordering}; + +const ORDERING: Ordering = Ordering::Relaxed; + +pub fn elapsed_ms(started_ms: i64) -> u64 { + u64::try_from(now_ms() - started_ms).unwrap_or(0) +} + +const HISTOGRAM_BUCKETS_MS: &[u64] = &[ + 1, 5, 10, 25, 50, 100, 250, 500, 1000, 2500, 5000, 10000, 30000, +]; + +struct Histogram { + buckets: [AtomicU64; 13], + inf: AtomicU64, + sum_ms: AtomicU64, + count: AtomicU64, +} + +impl Histogram { + const fn new() -> Self { + Self { + buckets: [const { AtomicU64::new(0) }; 13], + inf: AtomicU64::new(0), + sum_ms: AtomicU64::new(0), + count: AtomicU64::new(0), + } + } + + fn observe(&self, ms: u64) { + for (index, upper) in HISTOGRAM_BUCKETS_MS.iter().copied().enumerate() { + if ms <= upper { + self.buckets[index].fetch_add(1, ORDERING); + break; + } + } + self.inf.fetch_add(1, ORDERING); + self.sum_ms.fetch_add(ms, ORDERING); + self.count.fetch_add(1, ORDERING); + } + + fn render_series(&self, out: &mut String, name: &str, labels: &str) -> fmt::Result { + let mut cumulative = 0; + for (index, upper) in HISTOGRAM_BUCKETS_MS.iter().copied().enumerate() { + cumulative += self.buckets[index].load(ORDERING); + writeln!(out, "{name}_bucket{{{labels},le=\"{upper}\"}} {cumulative}")?; + } + writeln!( + out, + "{name}_bucket{{{labels},le=\"+Inf\"}} {}", + self.inf.load(ORDERING) + )?; + writeln!(out, "{name}_sum{{{labels}}} {}", self.sum_ms.load(ORDERING))?; + writeln!( + out, + "{name}_count{{{labels}}} {}", + self.count.load(ORDERING) + ) + } +} + +impl Default for Histogram { + fn default() -> Self { + Self::new() + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +#[repr(usize)] +pub enum JobKind { + Message, + Clear, +} + +impl JobKind { + pub const ALL: [Self; 2] = [Self::Message, Self::Clear]; + + pub fn label(self) -> &'static str { + match self { + Self::Message => "message", + Self::Clear => "clear", + } + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +#[repr(usize)] +pub enum JobRejection { + Decode, + QueueFull, +} + +impl JobRejection { + pub const ALL: [Self; 2] = [Self::Decode, Self::QueueFull]; + + pub fn label(self) -> &'static str { + match self { + Self::Decode => "decode", + Self::QueueFull => "queue_full", + } + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +#[repr(usize)] +pub enum Provider { + WebPush, + UnifiedPush, + Fcm, + Apns, +} + +impl Provider { + pub const ALL: [Self; 4] = [Self::WebPush, Self::UnifiedPush, Self::Fcm, Self::Apns]; + + pub fn label(self) -> &'static str { + match self { + Self::WebPush => "web_push", + Self::UnifiedPush => "unified_push", + Self::Fcm => "fcm", + Self::Apns => "apns", + } + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +#[repr(usize)] +pub enum SendResult { + Accepted, + TokenInvalid, + Permanent, + Transient, +} + +impl SendResult { + pub const ALL: [Self; 4] = [ + Self::Accepted, + Self::TokenInvalid, + Self::Permanent, + Self::Transient, + ]; + + pub fn label(self) -> &'static str { + match self { + Self::Accepted => "accepted", + Self::TokenInvalid => "token_invalid", + Self::Permanent => "permanent", + Self::Transient => "transient", + } + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +#[repr(usize)] +pub enum RelayLeg { + Apns, + Fcm, +} + +impl RelayLeg { + pub const ALL: [Self; 2] = [Self::Apns, Self::Fcm]; + + pub fn label(self) -> &'static str { + match self { + Self::Apns => "apns", + Self::Fcm => "fcm", + } + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +#[repr(usize)] +pub enum RelayResult { + Accepted, + Rejected, + Failed, +} + +impl RelayResult { + pub const ALL: [Self; 3] = [Self::Accepted, Self::Rejected, Self::Failed]; + + pub fn label(self) -> &'static str { + match self { + Self::Accepted => "accepted", + Self::Rejected => "rejected", + Self::Failed => "failed", + } + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +#[repr(usize)] +pub enum AuthProvider { + Vapid, + Apns, + Fcm, +} + +impl AuthProvider { + pub const ALL: [Self; 3] = [Self::Vapid, Self::Apns, Self::Fcm]; + + pub fn label(self) -> &'static str { + match self { + Self::Vapid => "vapid", + Self::Apns => "apns", + Self::Fcm => "fcm", + } + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +#[repr(usize)] +pub enum RpcMethod { + GetPushSubscriptions, + GetBadgeCounts, + DeletePushSubscriptions, + GetPushServiceDeliveryConfig, +} + +impl RpcMethod { + pub const ALL: [Self; 4] = [ + Self::GetPushSubscriptions, + Self::GetBadgeCounts, + Self::DeletePushSubscriptions, + Self::GetPushServiceDeliveryConfig, + ]; + + pub fn label(self) -> &'static str { + match self { + Self::GetPushSubscriptions => "get_push_subscriptions", + Self::GetBadgeCounts => "get_badge_counts", + Self::DeletePushSubscriptions => "delete_push_subscriptions", + Self::GetPushServiceDeliveryConfig => "get_push_service_delivery_config", + } + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +#[repr(usize)] +pub enum DeliveryRoute { + WebPush, + LegacyApns, + LegacyFcm, +} + +impl DeliveryRoute { + pub const ALL: [Self; 3] = [Self::WebPush, Self::LegacyApns, Self::LegacyFcm]; + + pub fn label(self) -> &'static str { + match self { + Self::WebPush => "web_push", + Self::LegacyApns => "legacy_apns", + Self::LegacyFcm => "legacy_fcm", + } + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +#[repr(usize)] +pub enum BucketKey { + DeviceToken, + Source, +} + +impl BucketKey { + pub const ALL: [Self; 2] = [Self::DeviceToken, Self::Source]; + + pub fn label(self) -> &'static str { + match self { + Self::DeviceToken => "device_token", + Self::Source => "source", + } + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +#[repr(usize)] +pub enum PayloadShrink { + Media, + Icons, + Body, + Minimal, +} + +impl PayloadShrink { + pub const ALL: [Self; 4] = [Self::Media, Self::Icons, Self::Body, Self::Minimal]; + + pub fn label(self) -> &'static str { + match self { + Self::Media => "media", + Self::Icons => "icons", + Self::Body => "body", + Self::Minimal => "minimal", + } + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +#[repr(usize)] +pub enum RpcOutcome { + Ok, + Error, +} + +impl RpcOutcome { + pub const ALL: [Self; 2] = [Self::Ok, Self::Error]; + + pub fn label(self) -> &'static str { + match self { + Self::Ok => "ok", + Self::Error => "error", + } + } +} + +const JOB_KIND_COUNT: usize = JobKind::ALL.len(); +const JOB_REJECTION_COUNT: usize = JobRejection::ALL.len(); +const PROVIDER_COUNT: usize = Provider::ALL.len(); +const SEND_RESULT_COUNT: usize = SendResult::ALL.len(); +const AUTH_PROVIDER_COUNT: usize = AuthProvider::ALL.len(); +const RELAY_LEG_COUNT: usize = RelayLeg::ALL.len(); +const RELAY_RESULT_COUNT: usize = RelayResult::ALL.len(); +const RPC_METHOD_COUNT: usize = RpcMethod::ALL.len(); +const BUCKET_KEY_COUNT: usize = BucketKey::ALL.len(); +const PAYLOAD_SHRINK_COUNT: usize = PayloadShrink::ALL.len(); +const DELIVERY_ROUTE_COUNT: usize = DeliveryRoute::ALL.len(); +const RPC_OUTCOME_COUNT: usize = RpcOutcome::ALL.len(); +const ROLLOUT_OUTCOME_COUNT: usize = RolloutOutcome::ALL.len(); + +pub struct Metrics { + jobs_received: [AtomicU64; JOB_KIND_COUNT], + jobs_rejected: [AtomicU64; JOB_REJECTION_COUNT], + jobs_completed: [AtomicU64; JOB_KIND_COUNT], + recipients: AtomicU64, + subscriptions: AtomicU64, + sends: [[AtomicU64; SEND_RESULT_COUNT]; PROVIDER_COUNT], + token_deletions: [AtomicU64; PROVIDER_COUNT], + payload_shrinks: [AtomicU64; PAYLOAD_SHRINK_COUNT], + auth_tokens_minted: [AtomicU64; AUTH_PROVIDER_COUNT], + rpc_requests: [[AtomicU64; RPC_OUTCOME_COUNT]; RPC_METHOD_COUNT], + rollout_updates: [AtomicU64; ROLLOUT_OUTCOME_COUNT], + delivery_routes: [[AtomicU64; SEND_RESULT_COUNT]; DELIVERY_ROUTE_COUNT], + relay_served: [[AtomicU64; RELAY_RESULT_COUNT]; RELAY_LEG_COUNT], + relay_vendor_requests: [[AtomicU64; RELAY_RESULT_COUNT]; RELAY_LEG_COUNT], + relay_rejected: [AtomicU64; REASON_COUNT], + relay_bucket_drops: [AtomicU64; BUCKET_KEY_COUNT], + job_duration: [Histogram; JOB_KIND_COUNT], + rpc_duration: [Histogram; RPC_METHOD_COUNT], + send_duration: [Histogram; PROVIDER_COUNT], + queue_depth: AtomicU64, + rollout_enabled: AtomicU64, + rollout_basis_points: AtomicU64, + rollout_config_version: AtomicU64, + start_ms: i64, +} + +impl Metrics { + pub fn new() -> Self { + Self { + jobs_received: [const { AtomicU64::new(0) }; JOB_KIND_COUNT], + jobs_rejected: [const { AtomicU64::new(0) }; JOB_REJECTION_COUNT], + jobs_completed: [const { AtomicU64::new(0) }; JOB_KIND_COUNT], + recipients: AtomicU64::new(0), + subscriptions: AtomicU64::new(0), + sends: [const { [const { AtomicU64::new(0) }; SEND_RESULT_COUNT] }; PROVIDER_COUNT], + token_deletions: [const { AtomicU64::new(0) }; PROVIDER_COUNT], + payload_shrinks: [const { AtomicU64::new(0) }; PAYLOAD_SHRINK_COUNT], + auth_tokens_minted: [const { AtomicU64::new(0) }; AUTH_PROVIDER_COUNT], + rpc_requests: [const { [const { AtomicU64::new(0) }; RPC_OUTCOME_COUNT] }; + RPC_METHOD_COUNT], + rollout_updates: [const { AtomicU64::new(0) }; ROLLOUT_OUTCOME_COUNT], + delivery_routes: [const { [const { AtomicU64::new(0) }; SEND_RESULT_COUNT] }; + DELIVERY_ROUTE_COUNT], + relay_served: [const { [const { AtomicU64::new(0) }; RELAY_RESULT_COUNT] }; + RELAY_LEG_COUNT], + relay_vendor_requests: [const { [const { AtomicU64::new(0) }; RELAY_RESULT_COUNT] }; + RELAY_LEG_COUNT], + relay_rejected: [const { AtomicU64::new(0) }; REASON_COUNT], + relay_bucket_drops: [const { AtomicU64::new(0) }; BUCKET_KEY_COUNT], + job_duration: [const { Histogram::new() }; JOB_KIND_COUNT], + rpc_duration: [const { Histogram::new() }; RPC_METHOD_COUNT], + send_duration: [const { Histogram::new() }; PROVIDER_COUNT], + queue_depth: AtomicU64::new(0), + rollout_enabled: AtomicU64::new(0), + rollout_basis_points: AtomicU64::new(0), + rollout_config_version: AtomicU64::new(0), + start_ms: now_ms(), + } + } + + pub fn record_job_received(&self, kind: JobKind) { + self.jobs_received[kind as usize].fetch_add(1, ORDERING); + } + + pub fn record_job_rejected(&self, reason: JobRejection) { + self.jobs_rejected[reason as usize].fetch_add(1, ORDERING); + } + + pub fn record_job_completed(&self, kind: JobKind, duration_ms: u64) { + self.jobs_completed[kind as usize].fetch_add(1, ORDERING); + self.job_duration[kind as usize].observe(duration_ms); + } + + pub fn record_recipients(&self, count: u64) { + self.recipients.fetch_add(count, ORDERING); + } + + pub fn record_subscriptions(&self, count: u64) { + self.subscriptions.fetch_add(count, ORDERING); + } + + pub fn record_send(&self, provider: Provider, result: SendResult, duration_ms: u64) { + self.sends[provider as usize][result as usize].fetch_add(1, ORDERING); + self.send_duration[provider as usize].observe(duration_ms); + } + + pub fn record_token_deletion(&self, provider: Provider) { + self.token_deletions[provider as usize].fetch_add(1, ORDERING); + } + + pub fn record_payload_shrink(&self, step: PayloadShrink) { + self.payload_shrinks[step as usize].fetch_add(1, ORDERING); + } + + pub fn record_auth_token_minted(&self, provider: AuthProvider) { + self.auth_tokens_minted[provider as usize].fetch_add(1, ORDERING); + } + + pub fn record_rpc(&self, method: RpcMethod, outcome: RpcOutcome, duration_ms: u64) { + self.rpc_requests[method as usize][outcome as usize].fetch_add(1, ORDERING); + self.rpc_duration[method as usize].observe(duration_ms); + } + + pub fn record_rollout_update(&self, outcome: RolloutOutcome) { + self.rollout_updates[outcome as usize].fetch_add(1, ORDERING); + } + + pub fn record_rollout_snapshot(&self, snapshot: &RolloutSnapshot) { + self.rollout_enabled + .store(u64::from(snapshot.enabled), ORDERING); + self.rollout_basis_points + .store(u64::from(snapshot.rollout_basis_points), ORDERING); + self.rollout_config_version + .store(snapshot.config_version, ORDERING); + } + + pub fn record_delivery_route(&self, route: DeliveryRoute, result: SendResult) { + self.delivery_routes[route as usize][result as usize].fetch_add(1, ORDERING); + } + + pub fn record_relay_served(&self, leg: RelayLeg, result: RelayResult) { + self.relay_served[leg as usize][result as usize].fetch_add(1, ORDERING); + } + + pub fn record_relay_vendor_request(&self, leg: RelayLeg, result: RelayResult) { + self.relay_vendor_requests[leg as usize][result as usize].fetch_add(1, ORDERING); + } + + pub fn record_relay_rejected(&self, reason: Reason) { + self.relay_rejected[reason as usize].fetch_add(1, ORDERING); + } + + pub fn record_bucket_drop(&self, key: BucketKey) { + self.relay_bucket_drops[key as usize].fetch_add(1, ORDERING); + } + + pub fn set_queue_depth(&self, depth: u64) { + self.queue_depth.store(depth, ORDERING); + } + + pub fn render(&self) -> String { + let mut out = String::new(); + self.render_into(&mut out) + .expect("writing push metrics to a String cannot fail"); + out + } + + fn render_into(&self, out: &mut String) -> fmt::Result { + render_labelled_counter( + out, + "fluxer_push_jobs_received_total", + "kind", + JobKind::ALL.map(JobKind::label), + &self.jobs_received, + )?; + render_labelled_counter( + out, + "fluxer_push_jobs_rejected_total", + "reason", + JobRejection::ALL.map(JobRejection::label), + &self.jobs_rejected, + )?; + render_labelled_counter( + out, + "fluxer_push_jobs_completed_total", + "kind", + JobKind::ALL.map(JobKind::label), + &self.jobs_completed, + )?; + render_counter(out, "fluxer_push_recipients_total", &self.recipients)?; + render_counter(out, "fluxer_push_subscriptions_total", &self.subscriptions)?; + + render_labelled_grid( + out, + "fluxer_push_sends_total", + ("provider", Provider::ALL.map(Provider::label)), + ("result", SendResult::ALL.map(SendResult::label)), + &self.sends, + )?; + + render_labelled_counter( + out, + "fluxer_push_token_deletions_total", + "provider", + Provider::ALL.map(Provider::label), + &self.token_deletions, + )?; + render_labelled_counter( + out, + "fluxer_push_payload_shrinks_total", + "step", + PayloadShrink::ALL.map(PayloadShrink::label), + &self.payload_shrinks, + )?; + render_labelled_counter( + out, + "fluxer_push_auth_tokens_minted_total", + "provider", + AuthProvider::ALL.map(AuthProvider::label), + &self.auth_tokens_minted, + )?; + + render_labelled_grid( + out, + "fluxer_push_rpc_requests_total", + ("method", RpcMethod::ALL.map(RpcMethod::label)), + ("result", RpcOutcome::ALL.map(RpcOutcome::label)), + &self.rpc_requests, + )?; + + render_labelled_counter( + out, + "fluxer_push_rollout_updates_total", + "result", + RolloutOutcome::ALL.map(RolloutOutcome::label), + &self.rollout_updates, + )?; + render_labelled_histogram( + out, + "fluxer_push_job_duration_ms", + "kind", + JobKind::ALL.map(JobKind::label), + &self.job_duration, + )?; + render_labelled_histogram( + out, + "fluxer_push_rpc_duration_ms", + "method", + RpcMethod::ALL.map(RpcMethod::label), + &self.rpc_duration, + )?; + render_labelled_histogram( + out, + "fluxer_push_send_duration_ms", + "provider", + Provider::ALL.map(Provider::label), + &self.send_duration, + )?; + render_labelled_grid( + out, + "fluxer_push_delivery_routes_total", + ("route", DeliveryRoute::ALL.map(DeliveryRoute::label)), + ("result", SendResult::ALL.map(SendResult::label)), + &self.delivery_routes, + )?; + render_labelled_grid( + out, + "fluxer_push_relay_served_total", + ("leg", RelayLeg::ALL.map(RelayLeg::label)), + ("result", RelayResult::ALL.map(RelayResult::label)), + &self.relay_served, + )?; + render_labelled_grid( + out, + "fluxer_push_relay_vendor_requests_total", + ("leg", RelayLeg::ALL.map(RelayLeg::label)), + ("result", RelayResult::ALL.map(RelayResult::label)), + &self.relay_vendor_requests, + )?; + render_labelled_counter( + out, + "fluxer_push_relay_rejected_total", + "reason", + Reason::ALL.map(Reason::label), + &self.relay_rejected, + )?; + render_labelled_counter( + out, + "fluxer_push_relay_token_bucket_drops_total", + "key", + BucketKey::ALL.map(BucketKey::label), + &self.relay_bucket_drops, + )?; + render_gauge(out, "fluxer_push_queue_depth", &self.queue_depth)?; + render_gauge(out, "fluxer_push_rollout_enabled", &self.rollout_enabled)?; + render_gauge( + out, + "fluxer_push_rollout_basis_points", + &self.rollout_basis_points, + )?; + render_gauge( + out, + "fluxer_push_rollout_config_version", + &self.rollout_config_version, + )?; + + writeln!(out, "# TYPE fluxer_push_uptime_seconds gauge")?; + writeln!( + out, + "fluxer_push_uptime_seconds {:.3}", + (now_ms() - self.start_ms) as f64 / 1000.0 + ) + } +} + +impl Default for Metrics { + fn default() -> Self { + Self::new() + } +} + +fn render_counter(out: &mut String, name: &str, counter: &AtomicU64) -> fmt::Result { + writeln!(out, "# TYPE {name} counter")?; + writeln!(out, "{name} {}", counter.load(ORDERING)) +} + +fn render_gauge(out: &mut String, name: &str, gauge: &AtomicU64) -> fmt::Result { + writeln!(out, "# TYPE {name} gauge")?; + writeln!(out, "{name} {}", gauge.load(ORDERING)) +} + +fn render_labelled_counter( + out: &mut String, + name: &str, + label: &str, + values: [&'static str; N], + counters: &[AtomicU64; N], +) -> fmt::Result { + writeln!(out, "# TYPE {name} counter")?; + for (index, value) in values.into_iter().enumerate() { + writeln!( + out, + "{name}{{{label}=\"{value}\"}} {}", + counters[index].load(ORDERING) + )?; + } + Ok(()) +} + +fn render_labelled_grid( + out: &mut String, + name: &str, + (row_label, rows): (&str, [&'static str; R]), + (column_label, columns): (&str, [&'static str; C]), + counters: &[[AtomicU64; C]; R], +) -> fmt::Result { + writeln!(out, "# TYPE {name} counter")?; + for (row_index, row) in rows.into_iter().enumerate() { + for (column_index, column) in columns.into_iter().enumerate() { + writeln!( + out, + "{name}{{{row_label}=\"{row}\",{column_label}=\"{column}\"}} {}", + counters[row_index][column_index].load(ORDERING) + )?; + } + } + Ok(()) +} + +fn render_labelled_histogram( + out: &mut String, + name: &str, + label: &str, + values: [&'static str; N], + histograms: &[Histogram; N], +) -> fmt::Result { + writeln!(out, "# TYPE {name} histogram")?; + for (index, value) in values.into_iter().enumerate() { + histograms[index].render_series(out, name, &format!("{label}=\"{value}\""))?; + } + Ok(()) +} diff --git a/fluxer_push/src/payload.rs b/fluxer_push/src/payload.rs new file mode 100644 index 000000000..78eeab337 --- /dev/null +++ b/fluxer_push/src/payload.rs @@ -0,0 +1,521 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use crate::job::{ClearJob, MessageJob}; +use crate::metrics::PayloadShrink; +use crate::unix_seconds; +use serde_json::{Map, Value, json}; + +const WEB_PUSH_MARKER: u64 = 8030; +const CLEAR_TYPE: &str = "notification_clear"; +const CLEAR_ACTION: &str = "clear_channel"; +const FALLBACK_TAG: &str = "fluxer-message"; +const FALLBACK_TITLE: &str = "Fluxer"; +const APNS_CATEGORY: &str = "FLUXER_MESSAGE"; +const APNS_SOUND: &str = "default"; +const APNS_ALERT_EXPIRATION_SECONDS: i64 = 86_400; +const APNS_BACKGROUND_EXPIRATION_SECONDS: i64 = 3_600; +const APNS_COLLAPSE_ID_MAX_BYTES: usize = 64; +const SHRUNK_BODY_MAX_BYTES: usize = 40; +const MINIMAL_TITLE_MAX_BYTES: usize = 120; +const MEDIA_KEYS: [&str; 2] = ["image_url", "image"]; +const ICON_KEYS: [&str; 3] = ["icon", "badge", "author_avatar_url"]; +const MINIMAL_DATA_KEYS: [&str; 7] = [ + "channel_id", + "message_id", + "guild_id", + "target_user_id", + "notification_tag", + "url", + "badge_count", +]; + +pub fn web_push_message(job: &MessageJob, target_user_id: &str, badge_count: u32) -> Value { + let fields = &job.notification; + let image_url = fields.image_url.as_deref(); + let guild_id = if job.is_direct_message() { + Value::Null + } else { + Value::String(job.guild_id.clone()) + }; + let mut data = json!({ + "channel_id": job.channel_id, + "author_avatar_url": fields.icon, + "message_id": job.message_id, + "notification_tag": fields.notification_tag, + "guild_id": guild_id, + "url": fields.url, + "badge_count": badge_count, + "target_user_id": target_user_id, + "has_media": image_url.is_some(), + }); + merge_image_fields(&mut data, image_url); + let mut notification = json!({ + "title": fields.title, + "body": fields.body, + "icon": fields.icon, + "badge": fields.badge, + "tag": fields.tag, + "navigate": fields.url, + "app_badge": badge_count.to_string(), + "data": data, + }); + merge_image_fields(&mut notification, image_url); + let mut envelope = json!({ + "web_push": WEB_PUSH_MARKER, + "notification": notification, + "title": fields.title, + "body": fields.body, + "icon": fields.icon, + "badge": fields.badge, + "tag": fields.tag, + "data": data, + }); + merge_image_fields(&mut envelope, image_url); + envelope +} + +pub fn web_push_clear(job: &ClearJob, badge_count: u32) -> Value { + let tag = format!("channel:{}", job.channel_id); + let data = json!({ + "type": CLEAR_TYPE, + "action": CLEAR_ACTION, + "channel_id": job.channel_id, + "message_id": job.message_id, + "target_user_id": job.user_id, + "notification_tag": tag, + "tag": tag, + "badge_count": badge_count, + }); + json!({ + "type": CLEAR_TYPE, + "action": CLEAR_ACTION, + "silent": true, + "tag": tag, + "notification_tag": tag, + "data": data, + "badge_count": badge_count, + "web_push": WEB_PUSH_MARKER, + "notification": { + "tag": tag, + "data": data, + "silent": true, + "close": true, + }, + }) +} + +pub fn fcm_message(device_token: &str, envelope: &Value) -> Value { + if is_clear(envelope) { + return fcm_clear_message(device_token, envelope); + } + fcm_notification_message(device_token, envelope) +} + +fn fcm_clear_message(device_token: &str, envelope: &Value) -> Value { + let tag = match envelope.get("notification_tag") { + Some(value) => value.as_str().unwrap_or(FALLBACK_TAG), + None => envelope + .get("tag") + .and_then(Value::as_str) + .unwrap_or(FALLBACK_TAG), + }; + let mut data = stringified_data(envelope.get("data")); + data.insert("type".to_owned(), CLEAR_TYPE.into()); + data.insert("action".to_owned(), CLEAR_ACTION.into()); + data.insert("notification_tag".to_owned(), tag.into()); + json!({ + "message": { + "token": device_token, + "data": data, + "android": { + "priority": "NORMAL", + "ttl": "3600s", + "collapse_key": format!("clear:{tag}"), + }, + "fcm_options": {"analytics_label": CLEAR_TYPE}, + }, + }) +} + +fn fcm_notification_message(device_token: &str, envelope: &Value) -> Value { + let notification = envelope.get("notification"); + let title = sanitized_text(fcm_text(notification, envelope, "title", FALLBACK_TITLE)); + let body = sanitized_text(fcm_text(notification, envelope, "body", "")); + let tag = envelope + .get("tag") + .and_then(Value::as_str) + .unwrap_or(FALLBACK_TAG); + let image_url = first_non_empty([ + field(envelope, "image_url"), + notification.and_then(|value| field(value, "image")), + notification.and_then(|value| field(value, "image_url")), + ]); + let mut notification_body = json!({"title": title, "body": body}); + put_image(&mut notification_body, image_url); + let mut data = stringified_data(envelope.get("data")); + data.insert("title".to_owned(), title.as_str().into()); + data.insert("body".to_owned(), body.as_str().into()); + data.insert("tag".to_owned(), tag.into()); + if let Some(url) = image_url { + data.insert("image_url".to_owned(), url.into()); + } + let mut android_notification = json!({ + "channel_id": "fluxer_default_push", + "tag": tag, + "click_action": APNS_CATEGORY, + }); + put_image(&mut android_notification, image_url); + json!({ + "message": { + "token": device_token, + "notification": notification_body, + "data": data, + "android": { + "priority": "HIGH", + "ttl": "86400s", + "notification": android_notification, + }, + "fcm_options": {"analytics_label": "message_create"}, + }, + }) +} + +pub fn apns_payload(envelope: &Value) -> Value { + let data = envelope.get("data"); + let mut payload = data.and_then(Value::as_object).cloned().unwrap_or_default(); + if is_clear(envelope) { + payload.insert("type".to_owned(), CLEAR_TYPE.into()); + payload.insert("action".to_owned(), CLEAR_ACTION.into()); + payload.insert("aps".to_owned(), json!({"content-available": 1})); + return Value::Object(payload); + } + let notification = envelope.get("notification"); + let title = notification + .and_then(|value| field(value, "title")) + .or_else(|| field(envelope, "title")) + .unwrap_or(FALLBACK_TITLE); + let body = notification + .and_then(|value| field(value, "body")) + .or_else(|| field(envelope, "body")) + .unwrap_or(""); + let url = data + .and_then(|data| field(data, "url")) + .or_else(|| notification.and_then(|value| field(value, "navigate"))); + let image_url = first_non_empty([ + field(envelope, "image_url"), + notification.and_then(|value| field(value, "image")), + ]); + let thread_id = data + .and_then(|data| field(data, "notification_tag")) + .map(str::to_owned) + .or_else(|| { + data.and_then(|data| field(data, "channel_id")) + .map(|channel_id| format!("channel:{channel_id}")) + }) + .unwrap_or_else(|| FALLBACK_TAG.to_owned()); + let mut aps = Map::new(); + aps.insert("alert".to_owned(), json!({"title": title, "body": body})); + aps.insert("sound".to_owned(), APNS_SOUND.into()); + aps.insert("thread-id".to_owned(), thread_id.into()); + aps.insert("category".to_owned(), APNS_CATEGORY.into()); + aps.insert("interruption-level".to_owned(), "active".into()); + aps.insert("relevance-score".to_owned(), Value::from(0.5)); + if let Some(badge) = badge_number(data.and_then(|data| data.get("badge_count"))) { + aps.insert("badge".to_owned(), badge.into()); + } + if image_url.is_some() { + aps.insert("mutable-content".to_owned(), 1.into()); + } + payload.insert("title".to_owned(), title.into()); + payload.insert("body".to_owned(), body.into()); + payload.remove("url"); + if let Some(url) = url { + payload.insert("url".to_owned(), url.into()); + } + payload.remove("image_url"); + if let Some(image_url) = image_url { + payload.insert("image_url".to_owned(), image_url.into()); + } + payload.insert("aps".to_owned(), Value::Object(aps)); + Value::Object(payload) +} + +pub fn apns_delivery_headers(envelope: &Value) -> Vec<(String, String)> { + let clear = is_clear(envelope); + let expiration = unix_seconds() + + if clear { + APNS_BACKGROUND_EXPIRATION_SECONDS + } else { + APNS_ALERT_EXPIRATION_SECONDS + }; + let mut headers = vec![ + ( + "apns-push-type".to_owned(), + if clear { "background" } else { "alert" }.to_owned(), + ), + ( + "apns-priority".to_owned(), + if clear { "5" } else { "10" }.to_owned(), + ), + ("apns-expiration".to_owned(), expiration.to_string()), + ("content-type".to_owned(), "application/json".to_owned()), + ]; + let collapse_id = field(envelope, "tag").or_else(|| { + envelope + .get("data") + .and_then(|data| field(data, "message_id")) + }); + if let Some(collapse_id) = collapse_id.filter(|id| id.len() <= APNS_COLLAPSE_ID_MAX_BYTES) { + headers.push(("apns-collapse-id".to_owned(), collapse_id.to_owned())); + } + headers +} + +fn put_image(target: &mut Value, image_url: Option<&str>) { + let (Some(image_url), Some(object)) = (image_url, target.as_object_mut()) else { + return; + }; + object.insert("image".to_owned(), image_url.into()); +} + +fn fcm_text<'a>( + notification: Option<&'a Value>, + envelope: &'a Value, + key: &str, + fallback: &'a str, +) -> &'a str { + notification + .and_then(|value| value.get(key)) + .or_else(|| envelope.get(key)) + .and_then(Value::as_str) + .unwrap_or(fallback) +} + +fn stringified_data(data: Option<&Value>) -> Map { + data.and_then(Value::as_object) + .map(|data| { + data.iter() + .map(|(key, value)| (key.clone(), stringified(value))) + .collect() + }) + .unwrap_or_default() +} + +fn stringified(value: &Value) -> Value { + match value { + Value::String(text) => Value::String(text.clone()), + Value::Null => Value::String("null".to_owned()), + Value::Bool(flag) => Value::String(flag.to_string()), + Value::Number(number) => Value::String(number.to_string()), + other => Value::String(other.to_string()), + } +} + +fn sanitized_text(value: &str) -> String { + value + .chars() + .filter(|character| displayable(*character)) + .collect() +} + +fn displayable(character: char) -> bool { + match character { + '\n' | '\t' => true, + '\u{0}'..='\u{1f}' + | '\u{7f}'..='\u{9f}' + | '\u{200e}'..='\u{200f}' + | '\u{202a}'..='\u{202e}' + | '\u{2066}'..='\u{2069}' => false, + _ => true, + } +} + +fn field<'a>(value: &'a Value, key: &str) -> Option<&'a str> { + value.get(key).and_then(non_empty) +} + +fn first_non_empty(candidates: [Option<&str>; N]) -> Option<&str> { + candidates.into_iter().flatten().next() +} + +fn badge_number(value: Option<&Value>) -> Option { + match value? { + Value::Number(number) => { + let badge = number.as_f64()?; + badge.is_finite().then(|| badge.max(0.0).trunc() as u64) + } + Value::String(text) => text + .parse::() + .ok() + .map(|badge| badge.max(0).unsigned_abs()), + _ => None, + } +} + +pub fn is_clear(envelope: &Value) -> bool { + envelope.get("type").and_then(Value::as_str) == Some(CLEAR_TYPE) + || envelope.get("action").and_then(Value::as_str) == Some(CLEAR_ACTION) +} + +pub fn fit(envelope: &Value, budget: usize) -> (Vec, Option) { + let serialized = serialize(envelope); + if serialized.len() <= budget { + return (serialized, None); + } + let mut working = envelope.clone(); + for step in PayloadShrink::ALL { + working = shrink(&working, step, budget); + let serialized = serialize(&working); + if serialized.len() <= budget { + return (serialized, Some(step)); + } + } + (serialize(&working), Some(PayloadShrink::Minimal)) +} + +fn shrink(envelope: &Value, step: PayloadShrink, budget: usize) -> Value { + let mut working = envelope.clone(); + match step { + PayloadShrink::Media => drop_keys(&mut working, &MEDIA_KEYS), + PayloadShrink::Icons => drop_keys(&mut working, &ICON_KEYS), + PayloadShrink::Body => truncate_body(&mut working), + PayloadShrink::Minimal => working = minimal(envelope, budget), + } + working +} + +fn minimal(envelope: &Value, budget: usize) -> Value { + let title = truncate_bytes( + first_text(envelope, "title").unwrap_or(FALLBACK_TITLE), + MINIMAL_TITLE_MAX_BYTES, + ); + let tag = first_text(envelope, "tag") + .unwrap_or(FALLBACK_TAG) + .to_owned(); + let url = first_text(envelope, "navigate") + .or_else(|| { + envelope + .get("data") + .and_then(|data| data.get("url")) + .and_then(non_empty) + }) + .unwrap_or_default() + .to_owned(); + let data = minimal_data(envelope.get("data")); + + let candidates = [ + minimal_envelope(&title, &tag, &url, data.clone()), + minimal_envelope(&title, &tag, "", data), + minimal_envelope(&title, &tag, "", Map::new()), + minimal_envelope(&title, "", "", Map::new()), + ]; + for candidate in candidates { + if serialize(&candidate).len() <= budget { + return candidate; + } + } + title_only(&title, budget) +} + +fn minimal_envelope(title: &str, tag: &str, url: &str, data: Map) -> Value { + json!({ + "web_push": WEB_PUSH_MARKER, + "title": title, + "tag": tag, + "data": data, + "notification": { + "title": title, + "tag": tag, + "navigate": url, + "data": data, + }, + }) +} + +fn title_only(title: &str, budget: usize) -> Value { + let mut allowance = title.len(); + loop { + let candidate = json!({ + "web_push": WEB_PUSH_MARKER, + "title": truncate_bytes(title, allowance), + }); + if allowance == 0 || serialize(&candidate).len() <= budget { + return candidate; + } + allowance -= 1; + } +} + +fn minimal_data(data: Option<&Value>) -> Map { + let Some(data) = data.and_then(Value::as_object) else { + return Map::new(); + }; + MINIMAL_DATA_KEYS + .into_iter() + .filter_map(|key| data.get(key).map(|value| (key.to_owned(), value.clone()))) + .collect() +} + +fn drop_keys(envelope: &mut Value, keys: &[&str]) { + for_each_block(envelope, &mut |block| { + for key in keys { + block.remove(*key); + } + }); +} + +fn truncate_body(envelope: &mut Value) { + for_each_block(envelope, &mut |block| { + let Some(body) = block.get("body").and_then(Value::as_str) else { + return; + }; + let shortened = truncate_bytes(body, SHRUNK_BODY_MAX_BYTES); + block.insert("body".to_owned(), shortened.into()); + }); +} + +fn for_each_block(envelope: &mut Value, apply: &mut impl FnMut(&mut Map)) { + let Some(root) = envelope.as_object_mut() else { + return; + }; + for key in ["notification", "data"] { + if let Some(nested) = root.get_mut(key) { + for_each_block(nested, apply); + } + } + apply(root); +} + +fn truncate_bytes(text: &str, max_bytes: usize) -> String { + if text.len() <= max_bytes { + return text.to_owned(); + } + let mut end = max_bytes; + while end > 0 && !text.is_char_boundary(end) { + end -= 1; + } + text[..end].to_owned() +} + +fn first_text<'a>(envelope: &'a Value, key: &str) -> Option<&'a str> { + envelope + .get(key) + .and_then(non_empty) + .or_else(|| envelope.get("notification")?.get(key).and_then(non_empty)) +} + +fn merge_image_fields(target: &mut Value, image_url: Option<&str>) { + let (Some(image_url), Some(object)) = (image_url, target.as_object_mut()) else { + return; + }; + object.insert("image_url".to_owned(), image_url.into()); + object.insert("image".to_owned(), image_url.into()); +} + +fn non_empty(value: &Value) -> Option<&str> { + value.as_str().filter(|text| !text.is_empty()) +} + +fn serialize(value: &Value) -> Vec { + serde_json::to_vec(value).expect("a json value serialises") +} diff --git a/fluxer_push/src/providers/apns.rs b/fluxer_push/src/providers/apns.rs new file mode 100644 index 000000000..307cbd28a --- /dev/null +++ b/fluxer_push/src/providers/apns.rs @@ -0,0 +1,43 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use crate::payload; +use crate::providers::SendOutcome; +use crate::server::AppState; +use crate::subscription::Subscription; +use crate::vendor::{self, ApnsRequest}; +use serde_json::Value; + +pub async fn send(state: &AppState, sub: &Subscription, envelope: &Value) -> SendOutcome { + let Some(cfg) = state.cfg.apns.as_ref() else { + return SendOutcome::permanent("apns_unavailable"); + }; + if sub.endpoint.is_empty() { + return SendOutcome::permanent("missing_device_token"); + } + let environment = sub.environment(cfg.default_environment); + let Some(topic) = cfg.topic_for(sub.app_id(), environment) else { + return SendOutcome::permanent("apns_topic_missing"); + }; + + let headers = payload::apns_delivery_headers(envelope); + let request = ApnsRequest { + environment, + topic, + device_token: &sub.endpoint, + headers: &headers, + body: serde_json::to_vec(&payload::apns_payload(envelope)) + .expect("a json value serialises"), + }; + match vendor::send_apns( + &state.apns_http, + &state.tokens, + &state.metrics, + cfg, + request, + ) + .await + { + Ok(outcome) => SendOutcome::from(outcome), + Err(error) => SendOutcome::permanent(format!("apns_auth: {error}")), + } +} diff --git a/fluxer_push/src/providers/fcm.rs b/fluxer_push/src/providers/fcm.rs new file mode 100644 index 000000000..086d8238d --- /dev/null +++ b/fluxer_push/src/providers/fcm.rs @@ -0,0 +1,33 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use crate::payload; +use crate::providers::SendOutcome; +use crate::server::AppState; +use crate::subscription::Subscription; +use crate::vendor; +use serde_json::Value; + +pub async fn send(state: &AppState, sub: &Subscription, envelope: &Value) -> SendOutcome { + let Some(cfg) = state.cfg.fcm.as_ref() else { + return SendOutcome::permanent("fcm_unavailable"); + }; + if sub.endpoint.is_empty() { + return SendOutcome::permanent("missing_device_token"); + } + + let body = serde_json::to_vec(&payload::fcm_message(&sub.endpoint, envelope)) + .expect("a json value serialises"); + match vendor::send_fcm( + &state.http, + &state.tokens, + &state.metrics, + cfg, + cfg.project_id_for(sub.app_id()), + body, + ) + .await + { + Ok(outcome) => SendOutcome::from(outcome), + Err(error) => SendOutcome::transient(format!("fcm_auth: {error}")), + } +} diff --git a/fluxer_push/src/providers/mod.rs b/fluxer_push/src/providers/mod.rs new file mode 100644 index 000000000..6e83674dc --- /dev/null +++ b/fluxer_push/src/providers/mod.rs @@ -0,0 +1,128 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +pub mod apns; +pub mod fcm; +pub mod web_push; + +use crate::metrics::{DeliveryRoute, Provider, SendResult, elapsed_ms}; +use crate::server::AppState; +use crate::subscription::{Platform, Subscription}; +use crate::vendor::VendorOutcome; +use fluxer_svc::metrics::now_ms; +use serde_json::Value; + +#[derive(Clone, Debug, Eq, PartialEq)] +pub enum SendOutcome { + Accepted, + TokenInvalid { reason: &'static str }, + Permanent { reason: String }, + Transient { reason: String }, +} + +impl SendOutcome { + pub fn deletes_token(&self) -> bool { + matches!(self, Self::TokenInvalid { .. }) + } + + pub fn reason(&self) -> &str { + match self { + Self::Accepted => "accepted", + Self::TokenInvalid { reason } => reason, + Self::Permanent { reason } | Self::Transient { reason } => reason, + } + } + + fn permanent(reason: impl Into) -> Self { + Self::Permanent { + reason: reason.into(), + } + } + + fn transient(reason: impl Into) -> Self { + Self::Transient { + reason: reason.into(), + } + } +} + +impl From for SendOutcome { + fn from(outcome: VendorOutcome) -> Self { + match outcome { + VendorOutcome::Accepted => Self::Accepted, + VendorOutcome::Unreachable => Self::transient("transport"), + VendorOutcome::Refused(refusal) => match refusal.dead_token { + Some(dead_token) => Self::TokenInvalid { + reason: dead_token.label(), + }, + None if refusal.is_transient() => { + Self::transient(format!("http_{}_{}", refusal.status, refusal.reason)) + } + None => Self::permanent(format!("http_{}_{}", refusal.status, refusal.reason)), + }, + } + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum Route { + WebPush, + LegacyApns, + LegacyFcm, +} + +pub fn route_of(sub: &Subscription) -> Option { + match sub.platform()? { + Platform::WebPush | Platform::AndroidUnifiedPush => Some(Route::WebPush), + Platform::IosApns if sub.is_web_push_registration() => Some(Route::WebPush), + Platform::AndroidFcm if sub.is_web_push_registration() => Some(Route::WebPush), + Platform::IosApns => Some(Route::LegacyApns), + Platform::AndroidFcm => Some(Route::LegacyFcm), + } +} + +pub async fn send(state: &AppState, sub: &Subscription, envelope: &Value) -> SendOutcome { + let (Some(platform), Some(route)) = (sub.platform(), route_of(sub)) else { + return SendOutcome::permanent("unsupported_platform"); + }; + let started_ms = now_ms(); + let outcome = match route { + Route::WebPush => web_push::send(state, sub, envelope).await, + Route::LegacyApns => apns::send(state, sub, envelope).await, + Route::LegacyFcm => fcm::send(state, sub, envelope).await, + }; + state.metrics.record_send( + provider_of(platform), + result_of(&outcome), + elapsed_ms(started_ms), + ); + state + .metrics + .record_delivery_route(route_label(route), result_of(&outcome)); + outcome +} + +fn route_label(route: Route) -> DeliveryRoute { + match route { + Route::WebPush => DeliveryRoute::WebPush, + Route::LegacyApns => DeliveryRoute::LegacyApns, + Route::LegacyFcm => DeliveryRoute::LegacyFcm, + } +} + +pub fn provider_of(platform: Platform) -> Provider { + match platform { + Platform::WebPush => Provider::WebPush, + Platform::AndroidUnifiedPush => Provider::UnifiedPush, + Platform::AndroidFcm => Provider::Fcm, + Platform::IosApns => Provider::Apns, + } +} + +fn result_of(outcome: &SendOutcome) -> SendResult { + match outcome { + SendOutcome::Accepted => SendResult::Accepted, + SendOutcome::TokenInvalid { .. } => SendResult::TokenInvalid, + SendOutcome::Permanent { .. } => SendResult::Permanent, + SendOutcome::Transient { .. } => SendResult::Transient, + } +} diff --git a/fluxer_push/src/providers/web_push.rs b/fluxer_push/src/providers/web_push.rs new file mode 100644 index 000000000..412da9793 --- /dev/null +++ b/fluxer_push/src/providers/web_push.rs @@ -0,0 +1,200 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use crate::crypto; +use crate::payload; +use crate::providers::SendOutcome; +use crate::resolver; +use crate::server::AppState; +use crate::subscription::Subscription; +use crate::vendor::is_transient_status; +use rand::RngExt as _; +use reqwest::header::{AUTHORIZATION, CONTENT_ENCODING, CONTENT_TYPE}; +use serde_json::Value; +use std::net::IpAddr; +use std::time::Duration; +use url::{Host, Url}; + +pub const RECORD_SIZE: usize = 2816; + +const HEADER_BYTES: usize = 86; +const TAG_BYTES: usize = 16; +const PADDING_DELIMITER_BYTES: usize = 1; +pub const PLAINTEXT_BUDGET: usize = + RECORD_SIZE - HEADER_BYTES - TAG_BYTES - PADDING_DELIMITER_BYTES; + +const MAX_TRANSIENT_RETRIES: u32 = 2; +const BASE_RETRY_DELAY_MS: u64 = 200; +const MAX_RETRY_DELAY_MS: u64 = 2_000; +const ALERT_TTL_SECONDS: &str = "86400"; +const CLEAR_TTL_SECONDS: &str = "3600"; +const TTL_HEADER: &str = "TTL"; +const URGENCY_HEADER: &str = "Urgency"; +const ALERT_URGENCY: &str = "high"; +const CLEAR_URGENCY: &str = "low"; +const OCTET_STREAM: &str = "application/octet-stream"; +const AES128GCM: &str = "aes128gcm"; +const NOT_FOUND: u16 = 404; +const GONE: u16 = 410; +const MAX_HOSTNAME_BYTES: usize = 253; +const MAX_LABEL_BYTES: usize = 63; + +pub async fn send(state: &AppState, sub: &Subscription, envelope: &Value) -> SendOutcome { + if !endpoint_is_allowed(&sub.endpoint) { + return SendOutcome::permanent("endpoint_rejected"); + } + let (Some(p256dh), Some(auth)) = (sub.p256dh_key.as_deref(), sub.auth_key.as_deref()) else { + return SendOutcome::permanent("missing_keys"); + }; + let (Ok(p256dh), Ok(auth)) = ( + crypto::decode_subscription_key(p256dh), + crypto::decode_subscription_key(auth), + ) else { + return SendOutcome::permanent("invalid_keys"); + }; + + let vapid = &state.cfg.vapid; + let token = match state + .tokens + .vapid(&origin_of(&sub.endpoint), vapid, &state.metrics) + .await + { + Ok(token) => token, + Err(error) => return SendOutcome::permanent(format!("vapid_token: {error}")), + }; + let authorization = format!("vapid t={token}, k={}", vapid.public_key); + + let (plaintext, shrunk) = payload::fit(envelope, PLAINTEXT_BUDGET); + if let Some(step) = shrunk { + state.metrics.record_payload_shrink(step); + } + let body = match crypto::encrypt_aes128gcm(&plaintext, &p256dh, &auth, RECORD_SIZE) { + Ok(body) => body, + Err(error) => return SendOutcome::permanent(format!("encrypt: {error}")), + }; + let clear = payload::is_clear(envelope); + + let mut attempt: u32 = 0; + loop { + let response = state + .web_push_http + .post(&sub.endpoint) + .header( + TTL_HEADER, + if clear { + CLEAR_TTL_SECONDS + } else { + ALERT_TTL_SECONDS + }, + ) + .header( + URGENCY_HEADER, + if clear { CLEAR_URGENCY } else { ALERT_URGENCY }, + ) + .header(CONTENT_TYPE, OCTET_STREAM) + .header(CONTENT_ENCODING, AES128GCM) + .header(AUTHORIZATION, &authorization) + .body(body.clone()) + .send() + .await; + + let status = match response { + Ok(response) => response.status().as_u16(), + Err(_) if attempt >= MAX_TRANSIENT_RETRIES => { + return SendOutcome::transient("transport"); + } + Err(_) => { + tokio::time::sleep(retry_delay(attempt)).await; + attempt += 1; + continue; + } + }; + if is_transient_status(status) && attempt < MAX_TRANSIENT_RETRIES { + tokio::time::sleep(retry_delay(attempt)).await; + attempt += 1; + continue; + } + return classify(status); + } +} + +fn classify(status: u16) -> SendOutcome { + match status { + 200..=299 => SendOutcome::Accepted, + GONE => SendOutcome::TokenInvalid { reason: "expired" }, + NOT_FOUND => SendOutcome::TokenInvalid { + reason: "not_found", + }, + _ if is_transient_status(status) => SendOutcome::transient(format!("http_{status}")), + _ => SendOutcome::permanent(format!("http_{status}")), + } +} + +fn endpoint_is_allowed(endpoint: &str) -> bool { + let Ok(url) = Url::parse(endpoint) else { + return false; + }; + if url.scheme() != "https" { + return false; + } + if !url.username().is_empty() || url.password().is_some() { + return false; + } + if !matches!(url.port_or_known_default(), Some(80 | 443)) { + return false; + } + match url.host() { + Some(Host::Domain(host)) => is_public_hostname(host), + Some(Host::Ipv4(ip)) => !resolver::is_blocked(IpAddr::V4(ip)), + Some(Host::Ipv6(ip)) => !resolver::is_blocked(IpAddr::V6(ip)), + None => false, + } +} + +fn is_public_hostname(host: &str) -> bool { + let host = host.strip_suffix('.').unwrap_or(host); + if host.is_empty() || host.len() > MAX_HOSTNAME_BYTES || !host.contains('.') { + return false; + } + let mut labels = host.split('.'); + let mut top_level = ""; + for label in &mut labels { + if !is_hostname_label(label) { + return false; + } + top_level = label; + } + !top_level.bytes().all(|byte| byte.is_ascii_digit()) +} + +fn is_hostname_label(label: &str) -> bool { + let bytes = label.as_bytes(); + let (Some(first), Some(last)) = (bytes.first(), bytes.last()) else { + return false; + }; + bytes.len() <= MAX_LABEL_BYTES + && first.is_ascii_alphanumeric() + && last.is_ascii_alphanumeric() + && bytes + .iter() + .all(|byte| byte.is_ascii_alphanumeric() || *byte == b'-') +} + +fn origin_of(endpoint: &str) -> String { + match endpoint.split_once("://") { + Some((scheme, rest)) => { + let host = rest.split('/').next().unwrap_or(rest); + format!("{scheme}://{host}") + } + None => endpoint.to_owned(), + } +} + +fn retry_delay(attempt: u32) -> Duration { + let base = MAX_RETRY_DELAY_MS.min( + BASE_RETRY_DELAY_MS + .checked_shl(attempt) + .unwrap_or(MAX_RETRY_DELAY_MS), + ); + let jitter = rand::rng().random_range(1..=(base / 4).max(1)); + Duration::from_millis(MAX_RETRY_DELAY_MS.min(base + jitter - 1)) +} diff --git a/fluxer_push/src/relay/client_ip.rs b/fluxer_push/src/relay/client_ip.rs new file mode 100644 index 000000000..49b1d7904 --- /dev/null +++ b/fluxer_push/src/relay/client_ip.rs @@ -0,0 +1,27 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use crate::config::RelayConfig; +use axum::http::HeaderMap; +use std::net::{IpAddr, SocketAddr}; + +fn resolve(cfg: &RelayConfig, peer: SocketAddr, headers: &HeaderMap) -> Option { + if !cfg.trust_client_ip_header { + return Some(peer.ip().to_canonical()); + } + headers + .get(&cfg.client_ip_header_name) + .and_then(|value| value.to_str().ok()) + .and_then(nearest_entry) + .map(|ip| ip.to_canonical()) +} + +pub fn for_rate_limit(cfg: &RelayConfig, peer: SocketAddr, headers: &HeaderMap) -> IpAddr { + resolve(cfg, peer, headers).unwrap_or_else(|| peer.ip().to_canonical()) +} + +fn nearest_entry(value: &str) -> Option { + value + .rsplit(',') + .filter_map(|entry| entry.trim().trim_matches(['[', ']']).parse().ok()) + .next() +} diff --git a/fluxer_push/src/relay/envelope.rs b/fluxer_push/src/relay/envelope.rs new file mode 100644 index 000000000..07f120acb --- /dev/null +++ b/fluxer_push/src/relay/envelope.rs @@ -0,0 +1,119 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use super::reject::{Reason, Rejection}; +use base64::prelude::*; +use serde_json::{Value, json}; + +pub const FORMAT_VERSION: u64 = 1; +pub const ALERT_TTL_CAP_SECONDS: i64 = 86_400; +pub const BACKGROUND_TTL_CAP_SECONDS: i64 = 3_600; +pub const APNS_BODY_MAX_BYTES: usize = 4_096; +pub const FCM_DATA_MAX_BYTES: usize = 4_096; + +const ALERT_LOC_KEY: &str = "PUSH_NEW_MESSAGE"; +const APNS_PUSH_TYPE_HEADER: &str = "apns-push-type"; +const APNS_PRIORITY_HEADER: &str = "apns-priority"; +const APNS_EXPIRATION_HEADER: &str = "apns-expiration"; +const CONTENT_TYPE_HEADER: &str = "content-type"; +const JSON_CONTENT_TYPE: &str = "application/json"; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum Urgency { + Alert, + Background, +} + +impl Urgency { + pub fn from_header(raw: Option<&str>) -> Option { + match raw.map(str::trim) { + None => Some(Self::Alert), + Some(value) => match value.to_ascii_lowercase().as_str() { + "high" | "normal" => Some(Self::Alert), + "low" | "very-low" => Some(Self::Background), + _ => None, + }, + } + } + + pub fn ttl_cap_seconds(self) -> i64 { + match self { + Self::Alert => ALERT_TTL_CAP_SECONDS, + Self::Background => BACKGROUND_TTL_CAP_SECONDS, + } + } +} + +pub fn encode_payload(body: &[u8]) -> String { + BASE64_URL_SAFE_NO_PAD.encode(body) +} + +pub fn apns_body(payload: &str, urgency: Urgency) -> Result, Rejection> { + let aps = match urgency { + Urgency::Alert => json!({ + "alert": {"loc-key": ALERT_LOC_KEY, "loc-args": []}, + "mutable-content": 1, + "interruption-level": "active", + }), + Urgency::Background => json!({"content-available": 1}), + }; + let body = serialize(&json!({"aps": aps, "v": FORMAT_VERSION, "p": payload})); + if body.len() > APNS_BODY_MAX_BYTES { + return Err(Rejection::new(Reason::PayloadTooLarge)); + } + Ok(body) +} + +pub fn apns_headers(urgency: Urgency, now_unix: i64, ttl_seconds: i64) -> Vec<(String, String)> { + let (push_type, priority) = match urgency { + Urgency::Alert => ("alert", "10"), + Urgency::Background => ("background", "5"), + }; + vec![ + (APNS_PUSH_TYPE_HEADER.to_owned(), push_type.to_owned()), + (APNS_PRIORITY_HEADER.to_owned(), priority.to_owned()), + ( + APNS_EXPIRATION_HEADER.to_owned(), + apns_expiration(now_unix, ttl_seconds, urgency).to_string(), + ), + (CONTENT_TYPE_HEADER.to_owned(), JSON_CONTENT_TYPE.to_owned()), + ] +} + +pub fn fcm_body( + device_token: &str, + payload: &str, + urgency: Urgency, + ttl_seconds: i64, +) -> Result, Rejection> { + let data = json!({"v": FORMAT_VERSION.to_string(), "p": payload}); + if serialize(&data).len() > FCM_DATA_MAX_BYTES { + return Err(Rejection::new(Reason::PayloadTooLarge)); + } + let priority = match urgency { + Urgency::Alert => "HIGH", + Urgency::Background => "NORMAL", + }; + Ok(serialize(&json!({ + "message": { + "token": device_token, + "data": data, + "android": { + "priority": priority, + "ttl": format!("{}s", capped_ttl(ttl_seconds, urgency)), + }, + }, + }))) +} + +fn apns_expiration(now_unix: i64, ttl_seconds: i64, urgency: Urgency) -> i64 { + let ttl = capped_ttl(ttl_seconds, urgency); + if ttl == 0 { 0 } else { now_unix + ttl } +} + +fn capped_ttl(ttl_seconds: i64, urgency: Urgency) -> i64 { + ttl_seconds.clamp(0, urgency.ttl_cap_seconds()) +} + +fn serialize(value: &Value) -> Vec { + serde_json::to_vec(value).expect("a json value serialises") +} diff --git a/fluxer_push/src/relay/mod.rs b/fluxer_push/src/relay/mod.rs new file mode 100644 index 000000000..737f3927e --- /dev/null +++ b/fluxer_push/src/relay/mod.rs @@ -0,0 +1,453 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +mod client_ip; +pub mod envelope; +mod quota; +pub mod reject; + +use crate::config::{ProviderEnvironment, RelayConfig}; +use crate::metrics::{Metrics, RelayLeg, RelayResult}; +use crate::server::{Sidecar, serve, sidecar_router}; +use crate::tokens::{TokenCache, TokenError}; +use crate::unix_seconds; +use crate::vendor::{self, ApnsRequest, DeadToken, Refusal, VendorOutcome}; +use axum::Router; +use axum::body::Body; +use axum::extract::{ConnectInfo, Path, State}; +use axum::http::{HeaderMap, HeaderValue, StatusCode, header}; +use axum::response::{IntoResponse, Response}; +use axum::routing::post; +use base64::prelude::*; +use envelope::Urgency; +use fluxer_svc::shutdown::wait_for_shutdown; +use quota::Quota; +use reject::{Reason, Rejection}; +use sha2::{Digest as _, Sha256}; +use std::net::SocketAddr; +use std::sync::Arc; +use std::time::{Duration, Instant}; +use tracing::{info, warn}; + +pub const APNS_ROUTE: &str = "/relay/v1/apns/{app_id}/{environment}/{device_token}"; +pub const FCM_ROUTE: &str = "/relay/v1/fcm/{app_id}/{device_token}"; + +const AES128GCM: &str = "aes128gcm"; +const TTL_HEADER: &str = "ttl"; +const URGENCY_HEADER: &str = "urgency"; +const JSON_CONTENT_TYPE: &str = "application/json"; +const DIGEST_BYTES: usize = 8; +const APNS_DEVICE_TOKEN_LEN: usize = 64; +const MAX_FCM_DEVICE_TOKEN_LEN: usize = 512; +const MAX_TTL_SECONDS: i64 = 86_400; +const BODY_READ_TIMEOUT: Duration = Duration::from_secs(15); + +pub struct AppState { + pub(crate) cfg: RelayConfig, + pub(crate) metrics: Arc, + pub(crate) sidecar: Arc, + quota: Quota, + http: reqwest::Client, + apns_http: reqwest::Client, + tokens: TokenCache, +} + +impl AppState { + pub(crate) fn try_new(cfg: RelayConfig) -> anyhow::Result { + let metrics = Arc::new(Metrics::new()); + Ok(Self { + sidecar: Arc::new(Sidecar::new(Arc::clone(&metrics))), + quota: Quota::new( + cfg.max_concurrent, + &cfg.device_token_bucket, + cfg.source_bucket.as_ref(), + ), + http: vendor::http_client()?, + apns_http: vendor::apns_http_client()?, + tokens: TokenCache::new(), + metrics, + cfg, + }) + } +} + +pub async fn run(cfg: RelayConfig) -> anyhow::Result<()> { + let state = Arc::new(AppState::try_new(cfg)?); + let addr = state.cfg.bind_addr; + let serving = serve(addr, router(Arc::clone(&state))).await?; + state.sidecar.set_serving(true); + info!( + %addr, + apns = state.cfg.apns.is_some(), + fcm = state.cfg.fcm.is_some(), + max_body_bytes = state.cfg.max_body_bytes, + source_rate_limit = source_rate_limit_owner(&state.cfg), + trust_client_ip_header = state.cfg.trust_client_ip_header, + "push relay listening" + ); + + wait_for_shutdown().await; + state.sidecar.set_serving(false); + serving.stop().await; + Ok(()) +} + +fn source_rate_limit_owner(cfg: &RelayConfig) -> &'static str { + if cfg.source_bucket.is_some() { + "relay" + } else { + "edge" + } +} + +pub fn router(state: Arc) -> Router { + let sidecar = Arc::clone(&state.sidecar); + Router::new() + .route(APNS_ROUTE, post(apns_route)) + .route(FCM_ROUTE, post(fcm_route)) + .fallback(unmatched_route) + .with_state(state) + .merge(sidecar_router(sidecar)) +} + +async fn unmatched_route(State(state): State>) -> Response { + let rejection = Rejection::new(Reason::BadRequest); + state.metrics.record_relay_rejected(rejection.reason); + respond(rejection.reason.status(), Some(rejection)) +} + +async fn apns_route( + State(state): State>, + ConnectInfo(peer): ConnectInfo, + Path((app_id, environment, device_token)): Path<(String, String, String)>, + headers: HeaderMap, + body: Body, +) -> Response { + let environment = ProviderEnvironment::from_label(&environment); + relay( + &state, + Incoming { + leg: RelayLeg::Apns, + app_id, + environment, + device_token, + peer, + }, + headers, + body, + ) + .await +} + +async fn fcm_route( + State(state): State>, + ConnectInfo(peer): ConnectInfo, + Path((app_id, device_token)): Path<(String, String)>, + headers: HeaderMap, + body: Body, +) -> Response { + relay( + &state, + Incoming { + leg: RelayLeg::Fcm, + app_id, + environment: None, + device_token, + peer, + }, + headers, + body, + ) + .await +} + +struct Incoming { + leg: RelayLeg, + app_id: String, + environment: Option, + device_token: String, + peer: SocketAddr, +} + +struct Delivery { + urgency: Urgency, + ttl_seconds: i64, +} + +enum Target<'a> { + Apns { + environment: ProviderEnvironment, + topic: &'a str, + }, + Fcm { + project_id: &'a str, + }, +} + +async fn relay(state: &AppState, incoming: Incoming, headers: HeaderMap, body: Body) -> Response { + let started = Instant::now(); + let mut bytes = 0; + let outcome = forward(state, &incoming, &headers, body, &mut bytes).await; + + let rejection = outcome.err(); + let status = rejection.map_or(StatusCode::OK, |rejection| rejection.reason.status()); + state.metrics.record_relay_served( + incoming.leg, + if status.is_success() { + RelayResult::Accepted + } else if status.is_server_error() { + RelayResult::Failed + } else { + RelayResult::Rejected + }, + ); + if let Some(rejection) = rejection { + state.metrics.record_relay_rejected(rejection.reason); + } + + info!( + leg = incoming.leg.label(), + app_id = incoming.app_id, + environment = incoming.environment.map_or("-", ProviderEnvironment::label), + device_token = digest(&incoming.device_token), + status = status.as_u16(), + reason = rejection.map_or("accepted", |rejection| rejection.reason.label()), + bytes, + duration_ms = started.elapsed().as_millis(), + "relay request" + ); + respond(status, rejection) +} + +async fn forward( + state: &AppState, + incoming: &Incoming, + headers: &HeaderMap, + body: Body, + bytes: &mut usize, +) -> Result<(), Rejection> { + let _permit = state.quota.admit()?; + let target = resolve(state, incoming)?; + if !device_token_is_shaped(incoming.leg, &incoming.device_token) { + return Err(Rejection::new(Reason::DeviceTokenInvalid)); + } + if !is_aes128gcm(headers) { + return Err(Rejection::new(Reason::BadRequest)); + } + let delivery = Delivery { + urgency: Urgency::from_header(header(headers, URGENCY_HEADER)) + .ok_or(Rejection::new(Reason::BadRequest))?, + ttl_seconds: ttl_seconds(headers)?, + }; + state.quota.take( + &state.metrics, + &incoming.device_token, + client_ip::for_rate_limit(&state.cfg, incoming.peer, headers), + Instant::now(), + )?; + + let body = tokio::time::timeout( + BODY_READ_TIMEOUT, + axum::body::to_bytes(body, state.cfg.max_body_bytes), + ) + .await + .map_err(|_| Rejection::new(Reason::BadRequest))? + .map_err(|_| Rejection::new(Reason::PayloadTooLarge))?; + *bytes = body.len(); + if body.is_empty() { + return Err(Rejection::new(Reason::BadRequest)); + } + let payload = envelope::encode_payload(&body); + + match target { + Target::Apns { environment, topic } => { + send_apns(state, incoming, &delivery, &payload, environment, topic).await + } + Target::Fcm { project_id } => { + send_fcm(state, incoming, &delivery, &payload, project_id).await + } + } +} + +fn resolve<'a>(state: &'a AppState, incoming: &Incoming) -> Result, Rejection> { + let unknown = Rejection::new(Reason::AppUnknown); + match incoming.leg { + RelayLeg::Apns => { + let cfg = state.cfg.apns.as_ref().ok_or(unknown)?; + let environment = incoming.environment.ok_or(unknown)?; + Ok(Target::Apns { + environment, + topic: cfg + .topic_for(&incoming.app_id, environment) + .ok_or(unknown)?, + }) + } + RelayLeg::Fcm => { + let cfg = state.cfg.fcm.as_ref().ok_or(unknown)?; + Ok(Target::Fcm { + project_id: cfg.listed_project_id(&incoming.app_id).ok_or(unknown)?, + }) + } + } +} + +async fn send_apns( + state: &AppState, + incoming: &Incoming, + delivery: &Delivery, + payload: &str, + environment: ProviderEnvironment, + topic: &str, +) -> Result<(), Rejection> { + let cfg = state + .cfg + .apns + .as_ref() + .ok_or(Rejection::new(Reason::AppUnknown))?; + let request = ApnsRequest { + environment, + topic, + device_token: &incoming.device_token, + headers: &envelope::apns_headers(delivery.urgency, unix_seconds(), delivery.ttl_seconds), + body: envelope::apns_body(payload, delivery.urgency)?, + }; + let outcome = vendor::send_apns( + &state.apns_http, + &state.tokens, + &state.metrics, + cfg, + request, + ) + .await; + finish(state, RelayLeg::Apns, outcome) +} + +async fn send_fcm( + state: &AppState, + incoming: &Incoming, + delivery: &Delivery, + payload: &str, + project_id: &str, +) -> Result<(), Rejection> { + let cfg = state + .cfg + .fcm + .as_ref() + .ok_or(Rejection::new(Reason::AppUnknown))?; + let body = envelope::fcm_body( + &incoming.device_token, + payload, + delivery.urgency, + delivery.ttl_seconds, + )?; + let outcome = vendor::send_fcm( + &state.http, + &state.tokens, + &state.metrics, + cfg, + project_id, + body, + ) + .await; + finish(state, RelayLeg::Fcm, outcome) +} + +fn finish( + state: &AppState, + leg: RelayLeg, + outcome: Result, +) -> Result<(), Rejection> { + let outcome = outcome.map_err(|error| { + warn!(%error, leg = leg.label(), "the relay could not mint its own vendor credential"); + Rejection::new(Reason::Internal) + })?; + let (result, verdict) = match outcome { + VendorOutcome::Accepted => (RelayResult::Accepted, Ok(())), + VendorOutcome::Unreachable => ( + RelayResult::Failed, + Err(Rejection::new(Reason::ProviderUnavailable)), + ), + VendorOutcome::Refused(refusal) => ( + if refusal.is_transient() { + RelayResult::Failed + } else { + RelayResult::Rejected + }, + Err(Rejection::new(refusal_reason(&refusal))), + ), + }; + state.metrics.record_relay_vendor_request(leg, result); + verdict +} + +fn refusal_reason(refusal: &Refusal) -> Reason { + match refusal.dead_token { + Some(DeadToken::Gone(_)) => Reason::DeviceTokenGone, + Some(DeadToken::Invalid(_)) => Reason::DeviceTokenInvalid, + None if refusal.is_transient() => Reason::ProviderUnavailable, + None => Reason::BadRequest, + } +} + +fn device_token_is_shaped(leg: RelayLeg, device_token: &str) -> bool { + match leg { + RelayLeg::Apns => { + device_token.len() == APNS_DEVICE_TOKEN_LEN + && device_token.bytes().all(|byte| byte.is_ascii_hexdigit()) + } + RelayLeg::Fcm => { + !device_token.is_empty() + && device_token.len() <= MAX_FCM_DEVICE_TOKEN_LEN + && device_token.bytes().all(is_fcm_token_byte) + } + } +} + +fn is_fcm_token_byte(byte: u8) -> bool { + byte.is_ascii_alphanumeric() || matches!(byte, b':' | b'-' | b'_' | b'.') +} + +fn is_aes128gcm(headers: &HeaderMap) -> bool { + header(headers, header::CONTENT_ENCODING.as_str()) + .is_some_and(|value| value.eq_ignore_ascii_case(AES128GCM)) +} + +fn ttl_seconds(headers: &HeaderMap) -> Result { + let raw = header(headers, TTL_HEADER).ok_or(Rejection::new(Reason::BadRequest))?; + let parsed = raw + .parse::() + .map_err(|_| Rejection::new(Reason::BadRequest))?; + Ok(parsed.clamp(0, MAX_TTL_SECONDS)) +} + +fn header<'a>(headers: &'a HeaderMap, name: &str) -> Option<&'a str> { + headers + .get(name) + .and_then(|value| value.to_str().ok()) + .map(str::trim) + .filter(|value| !value.is_empty()) +} + +fn respond(status: StatusCode, rejection: Option) -> Response { + let mut response = match rejection { + None => status.into_response(), + Some(rejection) => ( + status, + [(header::CONTENT_TYPE, JSON_CONTENT_TYPE)], + rejection.body(), + ) + .into_response(), + }; + if let Some(seconds) = rejection.and_then(|rejection| rejection.retry_after) + && let Ok(value) = HeaderValue::from_str(&seconds.to_string()) + { + response.headers_mut().insert(header::RETRY_AFTER, value); + } + response +} + +fn digest(value: &str) -> String { + if value.is_empty() { + return "-".to_owned(); + } + BASE64_URL_SAFE_NO_PAD.encode(&Sha256::digest(value.as_bytes())[..DIGEST_BYTES]) +} diff --git a/fluxer_push/src/relay/quota.rs b/fluxer_push/src/relay/quota.rs new file mode 100644 index 000000000..be6da86ce --- /dev/null +++ b/fluxer_push/src/relay/quota.rs @@ -0,0 +1,172 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use super::reject::{Reason, Rejection}; +use crate::config::BucketConfig; +use crate::metrics::{BucketKey, Metrics}; +use sha2::{Digest as _, Sha256}; +use std::collections::HashMap; +use std::net::IpAddr; +use std::sync::Mutex; +use std::time::Instant; +use tokio::sync::{Semaphore, SemaphorePermit}; + +const BUSY_RETRY_AFTER_SECONDS: i64 = 1; +const RATE_LIMIT_RETRY_AFTER_SECONDS: i64 = 1; +const SECONDS_PER_MINUTE: f64 = 60.0; +const KEY_BYTES: usize = 16; +const IPV4_PREFIX_BYTES: usize = 3; +const IPV6_PREFIX_BYTES: usize = 8; + +type Key = [u8; KEY_BYTES]; + +pub struct Quota { + admissions: Semaphore, + device_tokens: Buckets, + sources: Option, +} + +impl Quota { + pub fn new( + max_concurrent: usize, + device_tokens: &BucketConfig, + sources: Option<&BucketConfig>, + ) -> Self { + Self { + admissions: Semaphore::new(max_concurrent), + device_tokens: Buckets::new(device_tokens), + sources: sources.map(Buckets::new), + } + } + + pub fn admit(&self) -> Result, Rejection> { + self.admissions + .try_acquire() + .map_err(|_| Rejection::after(Reason::RelayUnavailable, BUSY_RETRY_AFTER_SECONDS)) + } + + pub fn take( + &self, + metrics: &Metrics, + device_token: &str, + client_ip: IpAddr, + now: Instant, + ) -> Result<(), Rejection> { + self.check( + metrics, + BucketKey::DeviceToken, + &self.device_tokens, + device_token_key(device_token), + now, + )?; + let Some(sources) = self.sources.as_ref() else { + return Ok(()); + }; + self.check( + metrics, + BucketKey::Source, + sources, + source_key(client_ip), + now, + ) + } + + fn check( + &self, + metrics: &Metrics, + which: BucketKey, + buckets: &Buckets, + key: Key, + now: Instant, + ) -> Result<(), Rejection> { + if buckets.take(key, now) { + return Ok(()); + } + metrics.record_bucket_drop(which); + Err(Rejection::after( + Reason::RateLimited, + RATE_LIMIT_RETRY_AFTER_SECONDS, + )) + } +} + +pub fn device_token_key(device_token: &str) -> Key { + let digest = Sha256::digest(device_token.as_bytes()); + let mut key = [0u8; KEY_BYTES]; + key.copy_from_slice(&digest[..KEY_BYTES]); + key +} + +pub fn source_key(ip: IpAddr) -> Key { + let mut key = [0u8; KEY_BYTES]; + match ip.to_canonical() { + IpAddr::V4(ip) => { + key[..IPV4_PREFIX_BYTES].copy_from_slice(&ip.octets()[..IPV4_PREFIX_BYTES]) + } + IpAddr::V6(ip) => { + key[..IPV6_PREFIX_BYTES].copy_from_slice(&ip.octets()[..IPV6_PREFIX_BYTES]) + } + } + key +} + +struct Bucket { + tokens: f64, + refilled_at: Instant, +} + +struct Held { + live: HashMap, + aged: HashMap, +} + +pub struct Buckets { + refill_per_second: f64, + burst: f64, + entries: usize, + held: Mutex, +} + +impl Buckets { + pub fn new(cfg: &BucketConfig) -> Self { + Self { + refill_per_second: f64::from(cfg.per_minute) / SECONDS_PER_MINUTE, + burst: f64::from(cfg.burst), + entries: cfg.entries, + held: Mutex::new(Held { + live: HashMap::new(), + aged: HashMap::new(), + }), + } + } + + pub fn take(&self, key: Key, now: Instant) -> bool { + let mut held = self + .held + .lock() + .expect("the relay quota lock is not poisoned"); + let mut bucket = held + .live + .remove(&key) + .or_else(|| held.aged.remove(&key)) + .unwrap_or(Bucket { + tokens: self.burst, + refilled_at: now, + }); + + let elapsed = now + .saturating_duration_since(bucket.refilled_at) + .as_secs_f64(); + bucket.refilled_at = now; + bucket.tokens = (bucket.tokens + elapsed * self.refill_per_second).min(self.burst); + let allowed = bucket.tokens >= 1.0; + if allowed { + bucket.tokens -= 1.0; + } + + if held.live.len() >= self.entries { + held.aged = std::mem::take(&mut held.live); + } + held.live.insert(key, bucket); + allowed + } +} diff --git a/fluxer_push/src/relay/reject.rs b/fluxer_push/src/relay/reject.rs new file mode 100644 index 000000000..821eb99bf --- /dev/null +++ b/fluxer_push/src/relay/reject.rs @@ -0,0 +1,87 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use axum::http::StatusCode; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +#[repr(usize)] +pub enum Reason { + BadRequest, + PayloadTooLarge, + DeviceTokenInvalid, + AppUnknown, + DeviceTokenGone, + RateLimited, + ProviderUnavailable, + RelayUnavailable, + Internal, +} + +impl Reason { + pub const ALL: [Self; 9] = [ + Self::BadRequest, + Self::PayloadTooLarge, + Self::DeviceTokenInvalid, + Self::AppUnknown, + Self::DeviceTokenGone, + Self::RateLimited, + Self::ProviderUnavailable, + Self::RelayUnavailable, + Self::Internal, + ]; + + pub fn label(self) -> &'static str { + match self { + Self::BadRequest => "bad_request", + Self::PayloadTooLarge => "payload_too_large", + Self::DeviceTokenInvalid => "device_token_invalid", + Self::AppUnknown => "app_unknown", + Self::DeviceTokenGone => "device_token_gone", + Self::RateLimited => "rate_limited", + Self::ProviderUnavailable => "provider_unavailable", + Self::RelayUnavailable => "relay_unavailable", + Self::Internal => "internal", + } + } + + pub fn status(self) -> StatusCode { + match self { + Self::BadRequest + | Self::PayloadTooLarge + | Self::DeviceTokenInvalid + | Self::AppUnknown => StatusCode::BAD_REQUEST, + Self::DeviceTokenGone => StatusCode::GONE, + Self::RateLimited => StatusCode::TOO_MANY_REQUESTS, + Self::ProviderUnavailable => StatusCode::BAD_GATEWAY, + Self::RelayUnavailable => StatusCode::SERVICE_UNAVAILABLE, + Self::Internal => StatusCode::INTERNAL_SERVER_ERROR, + } + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub struct Rejection { + pub reason: Reason, + pub retry_after: Option, +} + +impl Rejection { + pub fn new(reason: Reason) -> Self { + Self { + reason, + retry_after: None, + } + } + + pub fn after(reason: Reason, seconds: i64) -> Self { + Self { + reason, + retry_after: Some(seconds), + } + } + + pub fn body(self) -> String { + format!("{{\"reason\":\"{}\"}}", self.reason.label()) + } +} + +pub const REASON_COUNT: usize = Reason::ALL.len(); diff --git a/fluxer_push/src/resolver.rs b/fluxer_push/src/resolver.rs new file mode 100644 index 000000000..eec0c8622 --- /dev/null +++ b/fluxer_push/src/resolver.rs @@ -0,0 +1,124 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use reqwest::dns::{Addrs, Name, Resolve, Resolving}; +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}; +use tokio::net::lookup_host; + +type ResolveError = Box; + +const BLOCKED_V4: &[(Ipv4Addr, u32)] = &[ + (Ipv4Addr::new(0, 0, 0, 0), 8), + (Ipv4Addr::new(10, 0, 0, 0), 8), + (Ipv4Addr::new(100, 64, 0, 0), 10), + (Ipv4Addr::new(127, 0, 0, 0), 8), + (Ipv4Addr::new(169, 254, 0, 0), 16), + (Ipv4Addr::new(172, 16, 0, 0), 12), + (Ipv4Addr::new(192, 0, 0, 0), 24), + (Ipv4Addr::new(192, 0, 2, 0), 24), + (Ipv4Addr::new(192, 88, 99, 0), 24), + (Ipv4Addr::new(192, 168, 0, 0), 16), + (Ipv4Addr::new(198, 18, 0, 0), 15), + (Ipv4Addr::new(198, 51, 100, 0), 24), + (Ipv4Addr::new(203, 0, 113, 0), 24), + (Ipv4Addr::new(224, 0, 0, 0), 4), + (Ipv4Addr::new(240, 0, 0, 0), 4), +]; + +const BLOCKED_V6: &[(Ipv6Addr, u32)] = &[ + (Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, 0), 128), + (Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, 1), 128), + (Ipv6Addr::new(0x2001, 0x0db8, 0, 0, 0, 0, 0, 0), 32), + (Ipv6Addr::new(0xfc00, 0, 0, 0, 0, 0, 0, 0), 7), + (Ipv6Addr::new(0xfe80, 0, 0, 0, 0, 0, 0, 0), 10), + (Ipv6Addr::new(0xff00, 0, 0, 0, 0, 0, 0, 0), 8), +]; + +const NAT64_PREFIX: [u8; 4] = [0x00, 0x64, 0xff, 0x9b]; +const SIXTOFOUR_PREFIX: [u8; 2] = [0x20, 0x02]; + +pub struct PublicOnlyResolver; + +impl Resolve for PublicOnlyResolver { + fn resolve(&self, name: Name) -> Resolving { + let host = name.as_str().to_owned(); + Box::pin(async move { resolve_public(&host).await }) + } +} + +async fn resolve_public(host: &str) -> Result { + let resolved: Vec = lookup_host((host, 0)).await?.collect(); + screen(resolved) +} + +fn screen(resolved: Vec) -> Result { + if resolved.is_empty() { + return Err("host resolved to no addresses".into()); + } + if resolved.iter().any(|addr| is_blocked(addr.ip())) { + return Err("host resolved into blocked address space".into()); + } + Ok(Box::new(resolved.into_iter())) +} + +pub fn is_blocked(ip: IpAddr) -> bool { + match ip { + IpAddr::V4(ip) => is_blocked_v4(ip), + IpAddr::V6(ip) => match embedded_v4(ip) { + Some(embedded) => is_blocked_v4(embedded), + None => is_blocked_v6(ip), + }, + } +} + +fn is_blocked_v4(ip: Ipv4Addr) -> bool { + let value = ip.to_bits(); + BLOCKED_V4 + .iter() + .any(|(network, prefix)| masked_v4(value, *prefix) == masked_v4(network.to_bits(), *prefix)) +} + +fn is_blocked_v6(ip: Ipv6Addr) -> bool { + let value = ip.to_bits(); + BLOCKED_V6 + .iter() + .any(|(network, prefix)| masked_v6(value, *prefix) == masked_v6(network.to_bits(), *prefix)) +} + +fn masked_v4(value: u32, prefix: u32) -> u32 { + match prefix { + 0 => 0, + _ => value & (u32::MAX << (u32::BITS - prefix)), + } +} + +fn masked_v6(value: u128, prefix: u32) -> u128 { + match prefix { + 0 => 0, + _ => value & (u128::MAX << (u128::BITS - prefix)), + } +} + +fn embedded_v4(ip: Ipv6Addr) -> Option { + let octets = ip.octets(); + let quad = |start: usize| { + Ipv4Addr::new( + octets[start], + octets[start + 1], + octets[start + 2], + octets[start + 3], + ) + }; + if octets[..10] == [0; 10] && octets[10] == 0xff && octets[11] == 0xff { + return Some(quad(12)); + } + if octets[..4] == NAT64_PREFIX && octets[4..12] == [0; 8] { + return Some(quad(12)); + } + if octets[..12] == [0; 12] { + return Some(quad(12)); + } + if octets[..2] == SIXTOFOUR_PREFIX { + return Some(quad(2)); + } + None +} diff --git a/fluxer_push/src/retry.rs b/fluxer_push/src/retry.rs new file mode 100644 index 000000000..0276031e4 --- /dev/null +++ b/fluxer_push/src/retry.rs @@ -0,0 +1,18 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use rand::RngExt as _; +use std::time::{Duration, Instant}; + +pub const RETRY_DEADLINE: Duration = Duration::from_secs(60); +const BASE_DELAY_MS: u64 = 500; +const MAX_DELAY_MS: u64 = 10_000; + +pub fn next_attempt(attempt: u32, now: Instant, deadline: Instant) -> Option { + let at = now + backoff(attempt); + (at <= deadline).then_some(at) +} + +fn backoff(attempt: u32) -> Duration { + let ceiling = MAX_DELAY_MS.min(BASE_DELAY_MS << attempt.min(8)); + Duration::from_millis(rand::rng().random_range(ceiling / 2..=ceiling)) +} diff --git a/fluxer_push/src/rollout.rs b/fluxer_push/src/rollout.rs new file mode 100644 index 000000000..020d4de77 --- /dev/null +++ b/fluxer_push/src/rollout.rs @@ -0,0 +1,290 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use crate::metrics::{Metrics, RpcMethod}; +use crate::rpc::RpcClient; +use fluxer_svc::transport::{Transport, TransportMessage, TransportSubscriber}; +use serde_json::Value; +use std::fmt::Debug; +use std::sync::{Arc, RwLock}; +use std::time::Duration; +use tokio::time::{Instant, MissedTickBehavior}; +use tracing::{info, warn}; + +pub const RECONCILE_INTERVAL: Duration = Duration::from_secs(30); + +const MAX_BASIS_POINTS: u64 = 10_000; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +#[repr(usize)] +pub enum RolloutOutcome { + Updated, + Unchanged, + Stale, + Rejected, +} + +impl RolloutOutcome { + pub const ALL: [Self; 4] = [Self::Updated, Self::Unchanged, Self::Stale, Self::Rejected]; + + pub fn label(self) -> &'static str { + match self { + Self::Updated => "updated", + Self::Unchanged => "unchanged", + Self::Stale => "stale", + Self::Rejected => "rejected", + } + } +} + +pub trait RolloutConfig: Debug + Eq + Send + Sync + Sized + 'static { + const NAME: &'static str; + const SUBJECT: &'static str; + const MESSAGE_TYPE: &'static str; + const RPC_METHOD: RpcMethod; + + fn parse(config: &Value) -> Option; + fn enabled(&self) -> bool; + fn config_version(&self) -> u64; + fn record_update(metrics: &Metrics, outcome: RolloutOutcome); + fn record_held(&self, metrics: &Metrics); +} + +#[derive(Clone, Debug, Default, Eq, PartialEq)] +pub struct RolloutSnapshot { + pub enabled: bool, + pub config_version: u64, + pub rollout_basis_points: u32, +} + +impl RolloutConfig for RolloutSnapshot { + const NAME: &'static str = "push delivery"; + const SUBJECT: &'static str = "config.push.delivery"; + const MESSAGE_TYPE: &'static str = "push_service_delivery_config"; + const RPC_METHOD: RpcMethod = RpcMethod::GetPushServiceDeliveryConfig; + + fn parse(config: &Value) -> Option { + let rollout_basis_points = match config.get("rollout_basis_points") { + None | Some(Value::Null) => 0, + Some(value) => { + let points = value.as_u64()?; + if points > MAX_BASIS_POINTS { + return None; + } + u32::try_from(points).ok()? + } + }; + Some(Self { + enabled: parse_enabled(config)?, + config_version: parse_config_version(config)?, + rollout_basis_points, + }) + } + + fn enabled(&self) -> bool { + self.enabled + } + + fn config_version(&self) -> u64 { + self.config_version + } + + fn record_update(metrics: &Metrics, outcome: RolloutOutcome) { + metrics.record_rollout_update(outcome); + } + + fn record_held(&self, metrics: &Metrics) { + metrics.record_rollout_snapshot(self); + } +} + +pub fn parse_enabled(config: &Value) -> Option { + match config.get("enabled") { + None | Some(Value::Null) => Some(false), + Some(Value::Bool(enabled)) => Some(*enabled), + Some(_) => None, + } +} + +pub fn parse_config_version(config: &Value) -> Option { + match config.get("config_version") { + None | Some(Value::Null) => Some(0), + Some(value) => value.as_u64(), + } +} + +struct Current { + held: Option>, + highest_version: u64, +} + +pub struct RolloutStore { + current: RwLock>, +} + +impl Default for RolloutStore { + fn default() -> Self { + Self { + current: RwLock::new(Current { + held: None, + highest_version: 0, + }), + } + } +} + +impl RolloutStore { + pub fn new() -> Self { + Self::default() + } + + pub fn snapshot(&self) -> Option> { + self.current + .read() + .expect("rollout config lock poisoned") + .held + .clone() + } + + fn apply(&self, payload: &[u8]) -> RolloutOutcome { + let Ok(value) = serde_json::from_slice::(payload) else { + warn!(config = C::NAME, "rollout config payload is not JSON"); + return RolloutOutcome::Rejected; + }; + let Some(config) = config_object(&value, C::MESSAGE_TYPE) else { + warn!( + config = C::NAME, + "rollout config payload has no config object of its type" + ); + return RolloutOutcome::Rejected; + }; + self.update(config) + } + + fn update(&self, config: &Value) -> RolloutOutcome { + let Some(offered) = C::parse(config) else { + warn!(config = C::NAME, "rollout config rejected as invalid"); + return RolloutOutcome::Rejected; + }; + let mut current = self.current.write().expect("rollout config lock poisoned"); + match current.held.as_deref() { + Some(_) if offered.enabled() && offered.config_version() < current.highest_version => { + warn!( + config = C::NAME, + highest = current.highest_version, + offered = offered.config_version(), + "rollout config ignored a lower config_version" + ); + return RolloutOutcome::Stale; + } + Some(held) if *held == offered => return RolloutOutcome::Unchanged, + _ => {} + } + info!(config = C::NAME, held = ?offered, "rollout config updated"); + current.highest_version = current.highest_version.max(offered.config_version()); + current.held = Some(Arc::new(offered)); + RolloutOutcome::Updated + } +} + +pub async fn run_rollout_subscriber( + transport: T, + rpc: &RpcClient, + store: &RolloutStore, + metrics: &Metrics, + reconcile_every: Duration, +) { + loop { + let mut subscriber = match transport.subscribe(C::SUBJECT).await { + Ok(subscriber) => subscriber, + Err(error) => { + warn!( + config = C::NAME, + error = %error, + subject = C::SUBJECT, + "rollout config subscribe failed" + ); + transport.wait_for_reconnect().await; + continue; + } + }; + info!( + config = C::NAME, + subject = C::SUBJECT, + "listening for rollout config updates" + ); + let outcome = fetch_config(rpc, store, metrics).await; + info!( + config = C::NAME, + outcome = outcome.label(), + "rollout config read after subscribing" + ); + let mut reconcile = + tokio::time::interval_at(Instant::now() + reconcile_every, reconcile_every); + reconcile.set_missed_tick_behavior(MissedTickBehavior::Delay); + loop { + tokio::select! { + message = subscriber.next() => { + let Some(message) = message else { + break; + }; + let outcome = store.apply(message.payload()); + record_outcome(metrics, store, outcome); + } + _ = reconcile.tick() => { + fetch_config(rpc, store, metrics).await; + } + } + } + warn!( + config = C::NAME, + subject = C::SUBJECT, + "rollout config subscription ended, will re-subscribe" + ); + } +} + +async fn fetch_config( + rpc: &RpcClient, + store: &RolloutStore, + metrics: &Metrics, +) -> RolloutOutcome { + match rpc.rollout_config(C::RPC_METHOD).await { + Ok(config) => { + let outcome = store.update(&config); + record_outcome(metrics, store, outcome); + outcome + } + Err(error) => { + warn!(config = C::NAME, error = %error, "rollout config read failed"); + C::record_update(metrics, RolloutOutcome::Rejected); + RolloutOutcome::Rejected + } + } +} + +fn record_outcome( + metrics: &Metrics, + store: &RolloutStore, + outcome: RolloutOutcome, +) { + C::record_update(metrics, outcome); + if outcome == RolloutOutcome::Rejected { + return; + } + if let Some(held) = store.snapshot() { + held.record_held(metrics); + } +} + +fn config_object<'a>(value: &'a Value, message_type: &str) -> Option<&'a Value> { + if value + .get("type") + .is_some_and(|found| found.as_str() != Some(message_type)) + { + return None; + } + if let Some(config) = value.get("config").filter(|config| config.is_object()) { + return Some(config); + } + value.is_object().then_some(value) +} diff --git a/fluxer_push/src/rpc.rs b/fluxer_push/src/rpc.rs new file mode 100644 index 000000000..09bd94a8f --- /dev/null +++ b/fluxer_push/src/rpc.rs @@ -0,0 +1,220 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use crate::config::{RPC_AUTH_HEADER, RpcConfig}; +use crate::metrics::{Metrics, RpcMethod, RpcOutcome, elapsed_ms}; +use crate::secret::SecretString; +use crate::subscription::Subscription; +use fluxer_svc::metrics::now_ms; +use rand::RngExt as _; +use serde::Deserialize; +use serde::de::{DeserializeOwned, IgnoredAny}; +use serde_json::{Value, json}; +use std::collections::HashMap; +use std::sync::Arc; +use std::time::Duration; +use thiserror::Error; +use tracing::warn; + +const USER_BATCH_MAX: usize = 2_000; +const DELETION_BATCH_MAX: usize = 100; + +const REQUEST_TIMEOUT: Duration = Duration::from_secs(5); +const MAX_ATTEMPTS: u32 = 3; +const BASE_BACKOFF_MS: u64 = 200; +const MAX_BACKOFF_MS: u64 = 2_000; +const JITTER_MS: u64 = 200; +const ERROR_MESSAGE_MAX: usize = 512; + +#[derive(Debug, Error)] +pub enum RpcError { + #[error("internal rpc transport failed: {0}")] + Transport(#[from] reqwest::Error), + #[error("internal rpc returned {code}: {message}")] + Status { code: u16, message: String }, + #[error("internal rpc response could not be decoded: {0}")] + Decode(#[from] serde_json::Error), +} + +impl RpcError { + pub(crate) fn is_retryable(&self) -> bool { + match self { + Self::Transport(_) => true, + Self::Status { code, .. } => *code >= 500, + Self::Decode(_) => false, + } + } +} + +pub struct RpcClient { + http: reqwest::Client, + url: String, + auth: SecretString, + metrics: Arc, +} + +impl RpcClient { + pub fn new(cfg: &RpcConfig, http: reqwest::Client, metrics: Arc) -> Self { + Self { + http, + url: cfg.url.clone(), + auth: cfg.auth_token.clone(), + metrics, + } + } + + pub async fn rollout_config(&self, method: RpcMethod) -> Result { + let data: RolloutConfigData = self.call(method, &json!({"type": method.label()})).await?; + Ok(data.config) + } + + pub async fn badge_counts( + &self, + user_ids: &[String], + ) -> Result, RpcError> { + let mut counts = HashMap::new(); + for batch in user_ids.chunks(USER_BATCH_MAX) { + let data: BadgeCountsData = self + .call( + RpcMethod::GetBadgeCounts, + &json!({"type": "get_badge_counts", "user_ids": batch}), + ) + .await?; + counts.extend(data.badge_counts); + } + Ok(counts) + } + + pub async fn push_subscriptions( + &self, + user_ids: &[String], + ) -> Result>, RpcError> { + let mut subscriptions = HashMap::new(); + for batch in user_ids.chunks(USER_BATCH_MAX) { + let data: HashMap> = self + .call( + RpcMethod::GetPushSubscriptions, + &json!({"type": "get_push_subscriptions", "user_ids": batch}), + ) + .await?; + subscriptions.extend(data); + } + Ok(subscriptions) + } + + pub async fn delete_push_subscriptions( + &self, + subscriptions: &[(String, String)], + ) -> Result<(), RpcError> { + for batch in subscriptions.chunks(DELETION_BATCH_MAX) { + let entries = batch + .iter() + .map(|(user_id, subscription_id)| { + json!({"user_id": user_id, "subscription_id": subscription_id}) + }) + .collect::>(); + self.call::( + RpcMethod::DeletePushSubscriptions, + &json!({"type": "delete_push_subscriptions", "subscriptions": entries}), + ) + .await?; + } + Ok(()) + } + + async fn call( + &self, + method: RpcMethod, + body: &Value, + ) -> Result { + let mut attempt = 1; + loop { + let error = match self.attempt(method, body).await { + Ok(data) => return Ok(data), + Err(error) => error, + }; + if !error.is_retryable() || attempt >= MAX_ATTEMPTS { + return Err(error); + } + let delay_ms = backoff_delay_ms(attempt); + warn!( + method = method.label(), + attempt, + max_attempts = MAX_ATTEMPTS, + delay_ms, + error = %error, + "internal rpc retrying" + ); + tokio::time::sleep(Duration::from_millis(delay_ms)).await; + attempt += 1; + } + } + + async fn attempt( + &self, + method: RpcMethod, + body: &Value, + ) -> Result { + let started_ms = now_ms(); + let result = self.request(body).await; + let duration_ms = elapsed_ms(started_ms); + let outcome = if result.is_ok() { + RpcOutcome::Ok + } else { + RpcOutcome::Error + }; + self.metrics.record_rpc(method, outcome, duration_ms); + result + } + + async fn request(&self, body: &Value) -> Result { + let response = self + .http + .post(&self.url) + .header(RPC_AUTH_HEADER, self.auth.expose()) + .timeout(REQUEST_TIMEOUT) + .json(body) + .send() + .await?; + let code = response.status().as_u16(); + let payload = response.bytes().await?; + if !(200..300).contains(&code) { + return Err(RpcError::Status { + code, + message: error_message(&payload), + }); + } + let mut envelope: HashMap = serde_json::from_slice(&payload)?; + let data = envelope.remove("data").unwrap_or_else(|| json!({})); + Ok(serde_json::from_value(data)?) + } +} + +#[derive(Deserialize)] +struct RolloutConfigData { + config: Value, +} + +#[derive(Deserialize)] +struct BadgeCountsData { + #[serde(default)] + badge_counts: HashMap, +} + +fn backoff_delay_ms(attempt: u32) -> u64 { + let exponential = BASE_BACKOFF_MS.saturating_mul(1u64 << (attempt - 1).min(16)); + exponential.min(MAX_BACKOFF_MS) + rand::rng().random_range(0..JITTER_MS) +} + +fn error_message(payload: &[u8]) -> String { + serde_json::from_slice::(payload) + .ok() + .and_then(|body| { + body.get("message") + .and_then(Value::as_str) + .map(ToOwned::to_owned) + }) + .unwrap_or_else(|| String::from_utf8_lossy(payload).into_owned()) + .chars() + .take(ERROR_MESSAGE_MAX) + .collect() +} diff --git a/fluxer_push/src/secret.rs b/fluxer_push/src/secret.rs new file mode 100644 index 000000000..b0b6645b0 --- /dev/null +++ b/fluxer_push/src/secret.rs @@ -0,0 +1,38 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use std::fmt; +use std::sync::atomic::{Ordering, compiler_fence}; + +const REDACTED: &str = "[REDACTED]"; + +fn zero(bytes: &mut [u8]) { + for byte in bytes.iter_mut() { + unsafe { std::ptr::write_volatile(byte, 0) }; + } + compiler_fence(Ordering::SeqCst); +} + +#[derive(Clone)] +pub struct SecretString(String); + +impl SecretString { + pub fn new(value: String) -> Self { + Self(value) + } + + pub fn expose(&self) -> &str { + self.0.as_str() + } +} + +impl Drop for SecretString { + fn drop(&mut self) { + zero(unsafe { self.0.as_mut_vec() }); + } +} + +impl fmt::Debug for SecretString { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str(REDACTED) + } +} diff --git a/fluxer_push/src/server.rs b/fluxer_push/src/server.rs new file mode 100644 index 000000000..55b23965d --- /dev/null +++ b/fluxer_push/src/server.rs @@ -0,0 +1,207 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use crate::config::{Config, DeliveryConfig}; +use crate::delivery; +use crate::metrics::Metrics; +use crate::relay; +use crate::rollout::{self, RolloutSnapshot, RolloutStore}; +use crate::rpc::RpcClient; +use crate::secret::SecretString; +use crate::tokens::TokenCache; +use crate::vendor; +use axum::Router; +use axum::extract::{ConnectInfo, State}; +use axum::http::{HeaderValue, StatusCode, header}; +use axum::response::{IntoResponse, Response}; +use axum::routing::get; +use fluxer_svc::shutdown::{DEFAULT_DRAIN_TIMEOUT, drain_with_timeout, wait_for_shutdown}; +use fluxer_svc::transport::NatsTransport; +use std::net::SocketAddr; +use std::sync::Arc; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::time::Duration; +use tokio::net::TcpListener; +use tokio::sync::{oneshot, watch}; +use tokio::task::JoinHandle; +use tracing::info; + +const JOB_DRAIN_TIMEOUT: Duration = Duration::from_secs(25); + +pub struct Sidecar { + metrics: Arc, + serving: AtomicBool, +} + +impl Sidecar { + pub fn new(metrics: Arc) -> Self { + Self { + metrics, + serving: AtomicBool::new(false), + } + } + + pub fn set_serving(&self, serving: bool) { + self.serving.store(serving, Ordering::SeqCst); + } +} + +pub fn sidecar_router(sidecar: Arc) -> Router { + Router::new() + .route("/_health", get(readiness)) + .route("/_healthz", get(async || "OK")) + .route("/_metrics", get(metrics_handler)) + .with_state(sidecar) +} + +pub struct Serving { + stop: oneshot::Sender<()>, + served: JoinHandle>, +} + +impl Serving { + pub async fn stop(self) { + let _ = self.stop.send(()); + drain_with_timeout( + async { + let _ = self.served.await; + }, + DEFAULT_DRAIN_TIMEOUT, + ) + .await; + } +} + +pub async fn serve(addr: SocketAddr, router: Router) -> anyhow::Result { + let listener = TcpListener::bind(addr).await?; + let (stop, stopped) = oneshot::channel::<()>(); + let served = tokio::spawn(async move { + axum::serve( + listener, + router.into_make_service_with_connect_info::(), + ) + .with_graceful_shutdown(async move { + let _ = stopped.await; + }) + .await + }); + Ok(Serving { stop, served }) +} + +pub struct AppState { + pub(crate) cfg: DeliveryConfig, + pub(crate) metrics: Arc, + pub(crate) sidecar: Arc, + pub(crate) rollout: RolloutStore, + pub(crate) rpc: RpcClient, + pub(crate) http: reqwest::Client, + pub(crate) web_push_http: reqwest::Client, + pub(crate) apns_http: reqwest::Client, + pub(crate) tokens: TokenCache, + pub(crate) draining: watch::Sender, +} + +impl AppState { + pub(crate) fn try_new(cfg: DeliveryConfig) -> anyhow::Result { + let metrics = Arc::new(Metrics::new()); + let http = vendor::http_client()?; + Ok(Self { + rpc: RpcClient::new(&cfg.rpc, http.clone(), Arc::clone(&metrics)), + rollout: RolloutStore::new(), + sidecar: Arc::new(Sidecar::new(Arc::clone(&metrics))), + web_push_http: vendor::web_push_http_client()?, + apns_http: vendor::apns_http_client()?, + tokens: TokenCache::new(), + draining: watch::Sender::new(false), + cfg, + metrics, + http, + }) + } +} + +pub async fn run(cfg: Config) -> anyhow::Result<()> { + match cfg { + Config::Delivery(cfg) => run_delivery(*cfg).await, + Config::Relay(cfg) => relay::run(*cfg).await, + } +} + +async fn run_delivery(cfg: DeliveryConfig) -> anyhow::Result<()> { + let state = Arc::new(AppState::try_new(cfg)?); + let addr = state.cfg.bind_addr; + let serving = serve(addr, sidecar_router(Arc::clone(&state.sidecar))).await?; + info!(%addr, vapid_subject = state.cfg.vapid.email, "push sidecar listening"); + + let transport = NatsTransport::connect( + &state.cfg.nats.url, + state.cfg.nats.auth_token.as_ref().map(SecretString::expose), + ) + .await?; + + state.sidecar.set_serving(true); + + let subscriber = tokio::spawn({ + let transport = transport.clone(); + let state = Arc::clone(&state); + async move { + rollout::run_rollout_subscriber( + transport, + &state.rpc, + &state.rollout, + &state.metrics, + rollout::RECONCILE_INTERVAL, + ) + .await + } + }); + let jobs = tokio::spawn(delivery::run_job_subscribers(transport, Arc::clone(&state))); + + wait_for_shutdown().await; + state.sidecar.set_serving(false); + subscriber.abort(); + drain_jobs(&state, jobs, JOB_DRAIN_TIMEOUT).await; + serving.stop().await; + Ok(()) +} + +async fn drain_jobs(state: &AppState, jobs: JoinHandle<()>, timeout: Duration) { + state.draining.send_replace(true); + let running = jobs.abort_handle(); + drain_with_timeout( + async { + let _ = jobs.await; + }, + timeout, + ) + .await; + running.abort(); +} + +async fn readiness(State(sidecar): State>) -> impl IntoResponse { + if sidecar.serving.load(Ordering::SeqCst) { + (StatusCode::OK, "OK") + } else { + (StatusCode::SERVICE_UNAVAILABLE, "NOT READY") + } +} + +async fn metrics_handler( + ConnectInfo(peer): ConnectInfo, + State(sidecar): State>, +) -> Response { + if !is_loopback_peer(&peer) { + return (StatusCode::FORBIDDEN, "FORBIDDEN").into_response(); + } + ( + [( + header::CONTENT_TYPE, + HeaderValue::from_static("text/plain; version=0.0.4; charset=utf-8"), + )], + sidecar.metrics.render(), + ) + .into_response() +} + +fn is_loopback_peer(peer: &SocketAddr) -> bool { + peer.ip().to_canonical().is_loopback() +} diff --git a/fluxer_push/src/subscription.rs b/fluxer_push/src/subscription.rs new file mode 100644 index 000000000..b67a9a143 --- /dev/null +++ b/fluxer_push/src/subscription.rs @@ -0,0 +1,64 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use crate::config::{DEFAULT_APP_ID, ProviderEnvironment}; +use serde::Deserialize; + +const DEFAULT_PLATFORM: &str = "web_push"; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum Platform { + WebPush, + AndroidUnifiedPush, + AndroidFcm, + IosApns, +} + +#[derive(Clone, Debug, Deserialize)] +pub struct Subscription { + pub subscription_id: String, + pub endpoint: String, + pub p256dh_key: Option, + pub auth_key: Option, + pub platform: Option, + pub app_id: Option, + pub provider_environment: Option, +} + +impl Subscription { + pub fn platform(&self) -> Option { + match self.platform.as_deref().unwrap_or(DEFAULT_PLATFORM) { + "web_push" => Some(Platform::WebPush), + "android_unified_push" => Some(Platform::AndroidUnifiedPush), + "android_fcm" => Some(Platform::AndroidFcm), + "ios_apns" => Some(Platform::IosApns), + _ => None, + } + } + + pub fn is_web_push_registration(&self) -> bool { + endpoint_is_url(&self.endpoint) && self.has_web_push_keys() + } + + pub fn app_id(&self) -> &str { + self.app_id.as_deref().unwrap_or(DEFAULT_APP_ID) + } + + pub fn environment(&self, default: ProviderEnvironment) -> ProviderEnvironment { + match self.provider_environment.as_deref() { + None => default, + Some("development" | "sandbox") => ProviderEnvironment::Development, + Some(_) => ProviderEnvironment::Production, + } + } + + fn has_web_push_keys(&self) -> bool { + [self.p256dh_key.as_deref(), self.auth_key.as_deref()] + .into_iter() + .all(|key| key.is_some_and(|key| !key.trim().is_empty())) + } +} + +fn endpoint_is_url(endpoint: &str) -> bool { + let endpoint = endpoint.trim(); + endpoint.starts_with("https://") || endpoint.starts_with("http://") +} diff --git a/fluxer_push/src/tokens.rs b/fluxer_push/src/tokens.rs new file mode 100644 index 000000000..b7e065358 --- /dev/null +++ b/fluxer_push/src/tokens.rs @@ -0,0 +1,180 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use crate::config::{ApnsConfig, FcmConfig, VapidConfig}; +use crate::crypto::{self, CryptoError, ES256}; +use crate::metrics::{AuthProvider, Metrics}; +use crate::unix_seconds; +use serde_json::{Value, json}; +use std::collections::HashMap; +use thiserror::Error; +use tokio::sync::Mutex; + +const VAPID_TOKEN_TTL_SECONDS: i64 = 43_200; +const VAPID_TOKEN_SKEW_SECONDS: i64 = 60; +const MAX_VAPID_AUDIENCES: usize = 10_000; +const APNS_TOKEN_TTL_SECONDS: i64 = 50 * 60; +const FCM_ASSERTION_TTL_SECONDS: i64 = 3_600; +const FCM_TOKEN_SKEW_SECONDS: i64 = 60; +const FCM_DEFAULT_EXPIRES_IN_SECONDS: i64 = 3_600; +const FCM_SCOPE: &str = "https://www.googleapis.com/auth/firebase.messaging"; +const FCM_GRANT_TYPE: &str = "urn%3Aietf%3Aparams%3Aoauth%3Agrant-type%3Ajwt-bearer"; +const FORM_CONTENT_TYPE: &str = "application/x-www-form-urlencoded"; + +#[derive(Debug, Error)] +pub enum TokenError { + #[error(transparent)] + Crypto(#[from] CryptoError), + #[error("the token request failed")] + Request(#[from] reqwest::Error), + #[error("the token endpoint returned status {0}")] + Status(u16), + #[error("the token response has no access_token")] + Malformed, +} + +struct Cached { + token: String, + expires_at: i64, +} + +#[derive(Default)] +pub struct TokenCache { + vapid: Mutex>, + apns: Mutex>, + fcm: Mutex>, +} + +impl TokenCache { + pub fn new() -> Self { + Self::default() + } + + pub async fn vapid( + &self, + audience: &str, + cfg: &VapidConfig, + metrics: &Metrics, + ) -> Result { + let now = unix_seconds(); + let mut cache = self.vapid.lock().await; + if let Some(cached) = cache.get(audience) + && cached.expires_at - VAPID_TOKEN_SKEW_SECONDS > now + { + return Ok(cached.token.clone()); + } + + let expires_at = now + VAPID_TOKEN_TTL_SECONDS; + let key = crypto::parse_p256_private_scalar(cfg.private_key.expose())?; + let token = crypto::es256_jwt( + &json!({"alg": ES256, "typ": "JWT"}), + &json!({ + "sub": format!("mailto:{}", cfg.email), + "aud": audience, + "exp": expires_at, + }), + &key, + )?; + metrics.record_auth_token_minted(AuthProvider::Vapid); + if cache.len() >= MAX_VAPID_AUDIENCES { + cache.retain(|_, cached| cached.expires_at - VAPID_TOKEN_SKEW_SECONDS > now); + } + if cache.len() >= MAX_VAPID_AUDIENCES { + cache.clear(); + } + cache.insert( + audience.to_owned(), + Cached { + token: token.clone(), + expires_at, + }, + ); + Ok(token) + } + + pub async fn apns(&self, cfg: &ApnsConfig, metrics: &Metrics) -> Result { + let now = unix_seconds(); + let mut cache = self.apns.lock().await; + if let Some(cached) = cache.as_ref() + && cached.expires_at > now + { + return Ok(cached.token.clone()); + } + + let key = crypto::parse_p256_pkcs8_pem(cfg.private_key.expose())?; + let token = crypto::es256_jwt( + &json!({"alg": ES256, "kid": cfg.key_id}), + &json!({"iss": cfg.team_id, "iat": now}), + &key, + )?; + metrics.record_auth_token_minted(AuthProvider::Apns); + *cache = Some(Cached { + token: token.clone(), + expires_at: now + APNS_TOKEN_TTL_SECONDS, + }); + Ok(token) + } + + pub async fn fcm( + &self, + cfg: &FcmConfig, + http: &reqwest::Client, + metrics: &Metrics, + ) -> Result { + let now = unix_seconds(); + let mut cache = self.fcm.lock().await; + if let Some(cached) = cache.as_ref() + && cached.expires_at - FCM_TOKEN_SKEW_SECONDS > now + { + return Ok(cached.token.clone()); + } + + let key = crypto::parse_rsa_pkcs8_pem(cfg.private_key.expose())?; + let assertion = crypto::rs256_jwt( + &json!({"alg": "RS256", "typ": "JWT"}), + &json!({ + "iss": cfg.client_email, + "scope": FCM_SCOPE, + "aud": cfg.token_uri, + "iat": now, + "exp": now + FCM_ASSERTION_TTL_SECONDS, + }), + &key, + )?; + + let response = http + .post(&cfg.token_uri) + .header(reqwest::header::CONTENT_TYPE, FORM_CONTENT_TYPE) + .body(format!("grant_type={FCM_GRANT_TYPE}&assertion={assertion}")) + .send() + .await?; + let status = response.status(); + if !status.is_success() { + return Err(TokenError::Status(status.as_u16())); + } + let body: Value = + serde_json::from_slice(&response.bytes().await?).map_err(|_| TokenError::Malformed)?; + let token = body + .get("access_token") + .and_then(Value::as_str) + .ok_or(TokenError::Malformed)? + .to_owned(); + + metrics.record_auth_token_minted(AuthProvider::Fcm); + *cache = Some(Cached { + token: token.clone(), + expires_at: now + normalize_expires_in(body.get("expires_in")), + }); + Ok(token) + } +} + +fn normalize_expires_in(value: Option<&Value>) -> i64 { + let parsed = match value { + Some(Value::Number(number)) => number.as_i64(), + Some(Value::String(text)) => text.trim().parse::().ok(), + _ => None, + }; + parsed + .filter(|seconds| *seconds > 0) + .unwrap_or(FCM_DEFAULT_EXPIRES_IN_SECONDS) +} diff --git a/fluxer_push/src/vendor.rs b/fluxer_push/src/vendor.rs new file mode 100644 index 000000000..9716f2e59 --- /dev/null +++ b/fluxer_push/src/vendor.rs @@ -0,0 +1,226 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +use crate::config::{ApnsConfig, FcmConfig, ProviderEnvironment}; +use crate::metrics::Metrics; +use crate::resolver::PublicOnlyResolver; +use crate::tokens::{TokenCache, TokenError}; +use reqwest::header::{AUTHORIZATION, CONTENT_TYPE}; +use reqwest::redirect::Policy; +use serde_json::Value; +use std::time::Duration; + +const HTTP_TIMEOUT: Duration = Duration::from_secs(10); +const APNS_TIMEOUT: Duration = Duration::from_secs(5); +const CONNECT_TIMEOUT: Duration = Duration::from_secs(5); +const APNS_TOPIC_HEADER: &str = "apns-topic"; +pub const FCM_CONTENT_TYPE: &str = "application/json; charset=UTF-8"; +const TOO_MANY_REQUESTS: u16 = 429; +const MAX_ERROR_BODY_BYTES: usize = 8_192; +const HTTP_ERROR: &str = "http_error"; +const UNREGISTERED: &str = "UNREGISTERED"; +const INVALID_ARGUMENT: &str = "INVALID_ARGUMENT"; + +pub fn http_client() -> reqwest::Result { + reqwest::Client::builder() + .redirect(Policy::none()) + .connect_timeout(CONNECT_TIMEOUT) + .timeout(HTTP_TIMEOUT) + .build() +} + +pub fn web_push_http_client() -> reqwest::Result { + reqwest::Client::builder() + .redirect(Policy::none()) + .dns_resolver(PublicOnlyResolver) + .connect_timeout(CONNECT_TIMEOUT) + .timeout(HTTP_TIMEOUT) + .build() +} + +pub fn apns_http_client() -> reqwest::Result { + reqwest::Client::builder() + .redirect(Policy::none()) + .http2_prior_knowledge() + .connect_timeout(CONNECT_TIMEOUT) + .timeout(APNS_TIMEOUT) + .build() +} + +pub struct ApnsRequest<'a> { + pub environment: ProviderEnvironment, + pub topic: &'a str, + pub device_token: &'a str, + pub headers: &'a [(String, String)], + pub body: Vec, +} + +#[derive(Debug, Eq, PartialEq)] +pub enum VendorOutcome { + Accepted, + Refused(Refusal), + Unreachable, +} + +#[derive(Debug, Eq, PartialEq)] +pub struct Refusal { + pub status: u16, + pub reason: String, + pub dead_token: Option, +} + +impl Refusal { + pub fn is_transient(&self) -> bool { + is_transient_status(self.status) + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum DeadToken { + Gone(&'static str), + Invalid(&'static str), +} + +impl DeadToken { + pub fn label(self) -> &'static str { + match self { + Self::Gone(label) | Self::Invalid(label) => label, + } + } +} + +pub async fn send_apns( + http: &reqwest::Client, + tokens: &TokenCache, + metrics: &Metrics, + cfg: &ApnsConfig, + request: ApnsRequest<'_>, +) -> Result { + let token = tokens.apns(cfg, metrics).await?; + let url = format!( + "{}/3/device/{}", + cfg.base_url(request.environment).trim_end_matches('/'), + request.device_token + ); + let mut builder = http + .post(&url) + .header(AUTHORIZATION, format!("bearer {token}")) + .header(APNS_TOPIC_HEADER, request.topic) + .body(request.body); + for (name, value) in request.headers { + builder = builder.header(name, value); + } + Ok(outcome(builder.send().await, apns_refusal).await) +} + +pub async fn send_fcm( + http: &reqwest::Client, + tokens: &TokenCache, + metrics: &Metrics, + cfg: &FcmConfig, + project_id: &str, + body: Vec, +) -> Result { + let token = tokens.fcm(cfg, http, metrics).await?; + let url = format!( + "{}/v1/projects/{project_id}/messages:send", + cfg.base_url.trim_end_matches('/') + ); + let response = http + .post(&url) + .header(AUTHORIZATION, format!("Bearer {token}")) + .header(CONTENT_TYPE, FCM_CONTENT_TYPE) + .body(body) + .send() + .await; + Ok(outcome(response, fcm_refusal).await) +} + +async fn outcome( + response: reqwest::Result, + refusal: fn(u16, &[u8]) -> Refusal, +) -> VendorOutcome { + let Ok(response) = response else { + return VendorOutcome::Unreachable; + }; + if response.status().is_success() { + return VendorOutcome::Accepted; + } + let status = response.status().as_u16(); + VendorOutcome::Refused(refusal(status, &read_error_body(response).await)) +} + +fn apns_refusal(status: u16, body: &[u8]) -> Refusal { + let reason = reason_field(body).unwrap_or_else(|| format!("http_{status}")); + Refusal { + status, + dead_token: apns_dead_token(status, &reason), + reason, + } +} + +fn fcm_refusal(status: u16, body: &[u8]) -> Refusal { + let reason = fcm_error_code(body); + Refusal { + status, + dead_token: fcm_dead_token(&reason), + reason, + } +} + +fn apns_dead_token(status: u16, reason: &str) -> Option { + match (status, reason) { + (_, "Unregistered") => Some(DeadToken::Gone("unregistered")), + (410, _) => Some(DeadToken::Gone("gone")), + (400, "BadDeviceToken") => Some(DeadToken::Invalid("bad_device_token")), + (400, "DeviceTokenNotForTopic") => Some(DeadToken::Invalid("device_token_not_for_topic")), + _ => None, + } +} + +fn fcm_dead_token(code: &str) -> Option { + match code { + UNREGISTERED => Some(DeadToken::Gone("unregistered")), + INVALID_ARGUMENT => Some(DeadToken::Invalid("invalid_argument")), + _ => None, + } +} + +pub fn reason_field(body: &[u8]) -> Option { + let parsed: Value = serde_json::from_slice(body).ok()?; + Some(parsed.get("reason")?.as_str()?.to_owned()) +} + +fn fcm_error_code(body: &[u8]) -> String { + let Ok(parsed) = serde_json::from_slice::(body) else { + return HTTP_ERROR.to_owned(); + }; + let Some(error) = parsed.get("error") else { + return HTTP_ERROR.to_owned(); + }; + if let Some(details) = error.get("details").and_then(Value::as_array) { + return details + .iter() + .find_map(|detail| detail.get("errorCode").and_then(Value::as_str)) + .unwrap_or(HTTP_ERROR) + .to_owned(); + } + error + .get("status") + .and_then(Value::as_str) + .unwrap_or(HTTP_ERROR) + .to_owned() +} + +pub fn is_transient_status(status: u16) -> bool { + status >= 500 || status == TOO_MANY_REQUESTS +} + +pub async fn read_error_body(response: reqwest::Response) -> Vec { + let mut body = response + .bytes() + .await + .map(|bytes| bytes.to_vec()) + .unwrap_or_default(); + body.truncate(MAX_ERROR_BODY_BYTES); + body +} diff --git a/knip.json b/knip.json index a6073ff18..1c6f90b20 100644 --- a/knip.json +++ b/knip.json @@ -30,6 +30,7 @@ "fluxer_desktop/src/main/LinuxDesktopEntry.ts": ["exports"], "fluxer_desktop/src/main/NotificationState.ts": ["exports", "types"], "packages/schema/src/domains/admin/AdminUserSchemas.ts": ["exports"], + "packages/schema/src/domains/admin/PushServiceDeliverySchemas.ts": ["exports"], "packages/schema/src/domains/admin/ScreenShareDeliverySchemas.ts": ["exports"], "packages/schema/src/domains/admin/VoiceNoiseSuppressionSchemas.ts": ["exports"], "packages/schema/src/domains/download/DownloadSchemas.ts": ["exports"], diff --git a/packages/schema/src/domains/admin/AdminSchemas.ts b/packages/schema/src/domains/admin/AdminSchemas.ts index 59cf581eb..c1dd1c573 100644 --- a/packages/schema/src/domains/admin/AdminSchemas.ts +++ b/packages/schema/src/domains/admin/AdminSchemas.ts @@ -15,6 +15,10 @@ import { GatewayRolloutConfigResponse, GatewayRolloutConfigUpdateRequest, } from '@fluxer/schema/src/domains/admin/GatewayRolloutSchemas'; +import { + PushServiceDeliveryConfigResponse, + PushServiceDeliveryConfigUpdateRequest, +} from '@fluxer/schema/src/domains/admin/PushServiceDeliverySchemas'; import { ScreenShareDeliveryConfigResponse, ScreenShareDeliveryConfigUpdateRequest, @@ -649,6 +653,7 @@ export const InstanceConfigResponse = z.object({ gateway_rollout: GatewayRolloutConfigResponse, voice_noise_suppression: VoiceNoiseSuppressionConfigResponse, screen_share_delivery: ScreenShareDeliveryConfigResponse, + push_service_delivery: PushServiceDeliveryConfigResponse, experiment_delivery: ExperimentDeliveryConfigResponse, registration: InstanceRegistrationResponse, self_hosted: z.boolean(), @@ -686,6 +691,7 @@ export const InstanceConfigUpdateRequest = z.object({ gateway_rollout: GatewayRolloutConfigUpdateRequest.nullish(), voice_noise_suppression: VoiceNoiseSuppressionConfigUpdateRequest.nullish(), screen_share_delivery: ScreenShareDeliveryConfigUpdateRequest.nullish(), + push_service_delivery: PushServiceDeliveryConfigUpdateRequest.nullish(), experiment_delivery: ExperimentDeliveryConfigUpdateRequest.nullish(), registration: z .object({ diff --git a/packages/schema/src/domains/admin/PushServiceDeliverySchemas.ts b/packages/schema/src/domains/admin/PushServiceDeliverySchemas.ts new file mode 100644 index 000000000..7951ce0fd --- /dev/null +++ b/packages/schema/src/domains/admin/PushServiceDeliverySchemas.ts @@ -0,0 +1,57 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +import {EXPERIMENT_BUCKET_RESOLUTION, experimentBucket} from '@fluxer/schema/src/domains/experiment/ExperimentBucket'; +import {z} from 'zod'; + +const PUSH_SERVICE_DELIVERY_ROLLOUT_BASIS_POINTS_MAX = EXPERIMENT_BUCKET_RESOLUTION; +const PUSH_SERVICE_DELIVERY_MAX_TARGETED_USERS = 1000; +const DEFAULT_PUSH_SERVICE_DELIVERY_SALT = 'push-service-delivery-v1'; + +const PUSH_SERVICE_DELIVERY_SALT_PATTERN = /^[\x20-\x7e]+$/u; + +const PushServiceDeliveryTargetIdSchema = z.string().regex(/^\d{1,20}$/u); +const PushServiceDeliveryTargetedUserIdsSchema = z + .array(PushServiceDeliveryTargetIdSchema) + .max(PUSH_SERVICE_DELIVERY_MAX_TARGETED_USERS); + +const pushServiceDeliveryConfigFields = { + enabled: z.boolean(), + config_version: z.number().int().min(0), + rollout_basis_points: z.number().int().min(0).max(PUSH_SERVICE_DELIVERY_ROLLOUT_BASIS_POINTS_MAX), + rollout_salt: z.string().trim().min(1).max(64).regex(PUSH_SERVICE_DELIVERY_SALT_PATTERN), + included_user_ids: PushServiceDeliveryTargetedUserIdsSchema, + excluded_user_ids: PushServiceDeliveryTargetedUserIdsSchema, +}; + +export const PushServiceDeliveryConfigSchema = z.object({ + enabled: pushServiceDeliveryConfigFields.enabled.default(false), + config_version: pushServiceDeliveryConfigFields.config_version.default(0), + rollout_basis_points: pushServiceDeliveryConfigFields.rollout_basis_points.default(0), + rollout_salt: pushServiceDeliveryConfigFields.rollout_salt.default(DEFAULT_PUSH_SERVICE_DELIVERY_SALT), + included_user_ids: pushServiceDeliveryConfigFields.included_user_ids.default([]), + excluded_user_ids: pushServiceDeliveryConfigFields.excluded_user_ids.default([]), +}); + +export type PushServiceDeliveryConfig = z.infer; + +export const DEFAULT_PUSH_SERVICE_DELIVERY_CONFIG: PushServiceDeliveryConfig = PushServiceDeliveryConfigSchema.parse( + {}, +); + +export const PushServiceDeliveryConfigUpdateRequest = z + .object(pushServiceDeliveryConfigFields) + .omit({config_version: true}) + .partial(); + +export type PushServiceDeliveryConfigUpdateRequest = z.infer; + +export const PushServiceDeliveryConfigResponse = PushServiceDeliveryConfigSchema; + +export type PushServiceDeliveryConfigResponse = z.infer; + +export function pushServiceDeliveryEnrols(config: PushServiceDeliveryConfig, userId: string): boolean { + if (!config.enabled) return false; + if (config.excluded_user_ids.includes(userId)) return false; + if (config.included_user_ids.includes(userId)) return true; + return experimentBucket(userId, config.rollout_salt) < config.rollout_basis_points; +} diff --git a/packages/schema/src/domains/rpc/RpcSchemas.ts b/packages/schema/src/domains/rpc/RpcSchemas.ts index e19311bec..eea84de01 100644 --- a/packages/schema/src/domains/rpc/RpcSchemas.ts +++ b/packages/schema/src/domains/rpc/RpcSchemas.ts @@ -2,6 +2,7 @@ import {RTC_REGION_ID_MAX_LENGTH, RTC_REGION_ID_MIN_LENGTH} from '@fluxer/constants/src/LimitConstants'; import {GatewayRolloutConfigResponse} from '@fluxer/schema/src/domains/admin/GatewayRolloutSchemas'; +import {PushServiceDeliveryConfigResponse} from '@fluxer/schema/src/domains/admin/PushServiceDeliverySchemas'; import {WebAuthnCredentialResponse} from '@fluxer/schema/src/domains/auth/AuthSchemas'; import {ChannelResponse, RtcRegionResponse} from '@fluxer/schema/src/domains/channel/ChannelSchemas'; import {VoiceStateResponse} from '@fluxer/schema/src/domains/gateway/GatewaySchemas'; @@ -212,6 +213,11 @@ export const RpcRequest = z.discriminatedUnion('type', [ z.object({ type: z.literal('get_gateway_rollout_config').describe('Request type for fetching gateway rollout configuration'), }), + z.object({ + type: z + .literal('get_push_service_delivery_config') + .describe('Request type for fetching push service delivery configuration'), + }), ]); export type RpcRequest = z.infer; @@ -517,6 +523,16 @@ export const RpcResponse = z.discriminatedUnion('type', [ }) .describe('Gateway rollout config result'), }), + z.object({ + type: z + .literal('get_push_service_delivery_config') + .describe('Response type for push service delivery configuration'), + data: z + .object({ + config: PushServiceDeliveryConfigResponse.describe('Push service delivery configuration'), + }) + .describe('Push service delivery config result'), + }), ]); export type RpcResponse = z.infer; diff --git a/packages/schema/src/domains/user/UserRequestSchemas.ts b/packages/schema/src/domains/user/UserRequestSchemas.ts index 925627985..3811fa7cc 100644 --- a/packages/schema/src/domains/user/UserRequestSchemas.ts +++ b/packages/schema/src/domains/user/UserRequestSchemas.ts @@ -481,7 +481,9 @@ const MobilePushProviderEnvironmentSchema = createNamedStringLiteralUnion( export const RegisterMobileDeviceRequest = z .object({ platform: MobilePushPlatformSchema.describe('The mobile push notification platform'), - token: createStringType(1, 4096).describe('The platform-specific push notification token or endpoint URL'), + token: createStringType(1, 4096).describe( + 'The Web Push endpoint URL when encryption keys are supplied, otherwise the raw platform push token', + ), user_agent: createStringType(1, 1024).optional().describe('The user agent string identifying the device'), app_id: createStringType(1, 128) .optional() @@ -491,32 +493,44 @@ export const RegisterMobileDeviceRequest = z ), encryption_key: createStringType(1, 1024) .optional() - .describe('The P-256 ECDH public key for UnifiedPush encryption (base64url)'), + .describe('The P-256 ECDH public key for Web Push encryption (base64url)'), auth_secret: createStringType(1, 1024) .optional() - .describe('The authentication secret for UnifiedPush encryption (base64url)'), + .describe('The authentication secret for Web Push encryption (base64url)'), }) .superRefine((value, ctx) => { - if (value.platform !== 'android_unified_push') return; - if (!URLType.safeParse(value.token).success) { + const tokenIsUrl = URLType.safeParse(value.token).success; + const isWebPushRegistration = + value.platform === 'android_unified_push' || value.encryption_key != null || value.auth_secret != null; + if (!isWebPushRegistration) { + if (tokenIsUrl) { + ctx.addIssue({ + code: 'custom', + path: ['token'], + message: 'Endpoint URL registrations require encryption_key and auth_secret', + }); + } + return; + } + if (!tokenIsUrl) { ctx.addIssue({ code: 'custom', path: ['token'], - message: 'UnifiedPush registrations require a valid endpoint URL', + message: 'Web Push registrations require a valid endpoint URL', }); } if (!value.encryption_key) { ctx.addIssue({ code: 'custom', path: ['encryption_key'], - message: 'UnifiedPush registrations require encryption_key', + message: 'Web Push registrations require encryption_key', }); } if (!value.auth_secret) { ctx.addIssue({ code: 'custom', path: ['auth_secret'], - message: 'UnifiedPush registrations require auth_secret', + message: 'Web Push registrations require auth_secret', }); } }); @@ -525,7 +539,9 @@ export type RegisterMobileDeviceRequest = z.infer = services.iter().copied().collect(); assert_eq!(services.len(), unique.len()); - assert_eq!(services.len(), 17); + assert_eq!(services.len(), 18); } #[test] @@ -1016,7 +1020,7 @@ mod tests { .lines() .filter(|line| line.starts_with(" image: ")) .count(), - 17 + 18 ); let api = manifest @@ -1068,7 +1072,7 @@ mod tests { let mut sorted = services.clone(); sorted.sort_unstable(); assert_eq!(services, sorted); - assert_eq!(services.len(), 17); + assert_eq!(services.len(), 18); } #[test] diff --git a/tools/dev/src/dev.rs b/tools/dev/src/dev.rs index e4c9fdf68..b0cc1b172 100644 --- a/tools/dev/src/dev.rs +++ b/tools/dev/src/dev.rs @@ -3,7 +3,7 @@ use crate::gateway::{build_gateway_cluster_nodes, setup_gateway_config}; use crate::manifest::{ ADMIN_PORT, ANY_HOST, API_PORT, APP_PORT, APP_PROXY_PORT, DEV_PROXY_GATEWAY_PORTS_ENV, - DEV_PROXY_PORT, GATEWAY_PORT, LOOPBACK_HOST, MEDIA_PROXY_PORT, rust_services, + DEV_PROXY_PORT, GATEWAY_PORT, LOOPBACK_HOST, MEDIA_PROXY_PORT, PUSH_PORT, rust_services, }; use crate::object_store::s3_endpoint; use crate::paths::{DESKTOP_DIR, ROOT}; @@ -31,6 +31,7 @@ const DEFAULT_TASKS: &[&str] = &[ "proxy", "services", "media", + "push", "admin", "api", "gateway-single", @@ -462,6 +463,24 @@ pub fn task_table() -> Result> { cwd: ROOT.clone(), env: Vec::new(), }); + insert(DevTask { + name: "push", + args: strings(&[ + "cargo", + "run", + "-p", + "fluxer-push", + "--bin", + "fluxer-push", + "--", + "--bind-host", + ANY_HOST, + "--port", + &PUSH_PORT.to_string(), + ]), + cwd: ROOT.clone(), + env: Vec::new(), + }); insert(DevTask { name: "services", args: tool_args(&self_tool, &["rust-services"]), @@ -906,6 +925,7 @@ mod tests { "proxy", "services", "media", + "push", "admin", "api", "gateway-single", @@ -1096,6 +1116,17 @@ mod tests { "-p".to_owned(), "fluxer-media-proxy".to_owned() ])); + assert!(tasks["push"].args.starts_with(&[ + "cargo".to_owned(), + "run".to_owned(), + "-p".to_owned(), + "fluxer-push".to_owned() + ])); + assert!( + tasks["push"] + .args + .ends_with(&["--port".to_owned(), PUSH_PORT.to_string()]) + ); assert_eq!( tasks["admin"].args, vec![ diff --git a/tools/dev/src/manifest.rs b/tools/dev/src/manifest.rs index 73f064816..4417ac824 100644 --- a/tools/dev/src/manifest.rs +++ b/tools/dev/src/manifest.rs @@ -15,6 +15,7 @@ pub const API_PORT: u16 = 8080; pub const GATEWAY_PORT: u16 = 8771; pub const GATEWAY_WEBSOCKET_PORTS: &[u16] = &[8771, 8772, 8774]; pub const MEDIA_PROXY_PORT: u16 = 8082; +pub const PUSH_PORT: u16 = 8126; pub const LIVEKIT_PORT: u16 = 7880; pub const DEVMAIL_PORT: u16 = 8025;