Compare commits

...
90 changed files with 2001 additions and 10082 deletions
+6 -59
View File
@@ -10524,7 +10524,7 @@
},
"gateway_rollout": {"$ref": "#/components/schemas/GatewayRolloutConfigResponse"},
"voice_noise_suppression": {"$ref": "#/components/schemas/VoiceNoiseSuppressionConfigResponse"},
"push_service_delivery": {"$ref": "#/components/schemas/PushServiceDeliveryConfigResponse"},
"push_relay": {"$ref": "#/components/schemas/PushRelayConfigResponse"},
"domain_migration": {"$ref": "#/components/schemas/DomainMigrationConfigResponse"},
"altcha_captcha": {"$ref": "#/components/schemas/AltchaCaptchaConfigResponse"},
"experiment_delivery": {"$ref": "#/components/schemas/ExperimentDeliveryConfigResponse"},
@@ -10954,7 +10954,7 @@
"sso",
"gateway_rollout",
"voice_noise_suppression",
"push_service_delivery",
"push_relay",
"domain_migration",
"altcha_captcha",
"experiment_delivery",
@@ -11091,10 +11091,7 @@
"nullable": true,
"allOf": [{"$ref": "#/components/schemas/VoiceNoiseSuppressionConfigUpdateRequest"}]
},
"push_service_delivery": {
"nullable": true,
"allOf": [{"$ref": "#/components/schemas/PushServiceDeliveryConfigUpdateRequest"}]
},
"push_relay": {"nullable": true, "allOf": [{"$ref": "#/components/schemas/PushRelayConfigUpdateRequest"}]},
"domain_migration": {
"nullable": true,
"allOf": [{"$ref": "#/components/schemas/DomainMigrationConfigUpdateRequest"}]
@@ -15237,25 +15234,7 @@
"standalone_forwarding": {"type": "boolean"}
}
},
"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}$"}
},
"relay_consent_accepted": {"type": "boolean"}
}
},
"PushRelayConfigUpdateRequest": {"type": "object", "properties": {"relay_consent_accepted": {"type": "boolean"}}},
"VoiceNoiseSuppressionConfigUpdateRequest": {
"type": "object",
"properties": {
@@ -15403,31 +15382,9 @@
],
"additionalProperties": false
},
"PushServiceDeliveryConfigResponse": {
"PushRelayConfigResponse": {
"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}$"}
},
"relay_consent_accepted": {"default": false, "type": "boolean"},
"relay_consent_accepted_at": {
"default": null,
@@ -15438,17 +15395,7 @@
},
"relay_consent_accepted_by": {"default": null, "nullable": true, "type": "string", "pattern": "^\\d{1,20}$"}
},
"required": [
"enabled",
"config_version",
"rollout_basis_points",
"rollout_salt",
"included_user_ids",
"excluded_user_ids",
"relay_consent_accepted",
"relay_consent_accepted_at",
"relay_consent_accepted_by"
],
"required": ["relay_consent_accepted", "relay_consent_accepted_at", "relay_consent_accepted_by"],
"additionalProperties": false
},
"VoiceNoiseSuppressionConfigResponse": {
+5 -38
View File
@@ -23,7 +23,7 @@ pub struct InstanceConfigResponse {
#[serde(default)]
pub voice_noise_suppression: VoiceNoiseSuppressionConfigResponse,
#[serde(default)]
pub push_service_delivery: PushServiceDeliveryConfigResponse,
pub push_relay: PushRelayConfigResponse,
#[serde(default)]
pub domain_migration: DomainMigrationConfigResponse,
#[serde(default)]
@@ -453,7 +453,6 @@ 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 DOMAIN_MIGRATION_DEFAULT_SALT: &str = "domain-migration-v1";
pub const ALTCHA_CAPTCHA_DEFAULT_SALT: &str = "altcha-captcha-v1";
pub const ALTCHA_CAPTCHA_COST_RANGE: std::ops::RangeInclusive<u32> = 1_000..=100_000;
@@ -548,48 +547,16 @@ pub struct VoiceNoiseSuppressionConfigUpdateRequest {
pub suppression_strength: Option<u32>,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
#[derive(Clone, Debug, Default, 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<String>,
pub excluded_user_ids: Vec<String>,
pub struct PushRelayConfigResponse {
pub relay_consent_accepted: bool,
pub relay_consent_accepted_at: Option<String>,
pub relay_consent_accepted_by: Option<String>,
}
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(),
relay_consent_accepted: false,
relay_consent_accepted_at: None,
relay_consent_accepted_by: None,
}
}
}
#[derive(Clone, Debug, Default, Serialize)]
pub struct PushServiceDeliveryConfigUpdateRequest {
#[serde(skip_serializing_if = "Option::is_none")]
pub enabled: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub rollout_basis_points: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub rollout_salt: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub included_user_ids: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub excluded_user_ids: Option<Vec<String>>,
pub struct PushRelayConfigUpdateRequest {
#[serde(skip_serializing_if = "Option::is_none")]
pub relay_consent_accepted: Option<bool>,
}
@@ -806,7 +773,7 @@ pub struct InstanceConfigUpdateRequest {
#[serde(skip_serializing_if = "Option::is_none")]
pub voice_noise_suppression: Option<VoiceNoiseSuppressionConfigUpdateRequest>,
#[serde(skip_serializing_if = "Option::is_none")]
pub push_service_delivery: Option<PushServiceDeliveryConfigUpdateRequest>,
pub push_relay: Option<PushRelayConfigUpdateRequest>,
#[serde(skip_serializing_if = "Option::is_none")]
pub domain_migration: Option<DomainMigrationConfigUpdateRequest>,
#[serde(skip_serializing_if = "Option::is_none")]
+21 -54
View File
@@ -20,9 +20,9 @@ use crate::{
InstancePolicyUpdateRequest, InstanceRegistrationConfigUpdateRequest,
InstanceServicesUpdateRequest, InstanceYoutubeIntegrationUpdateRequest,
LimitConfigUpdateRequest, LimitRule, LimitRuleFilters, NoiseSuppressionBackend,
PremiumMode, PushServiceDeliveryConfigUpdateRequest, RegistrationMode,
SsoConfigUpdateRequest, VOICE_NS_MAX_GUILD_OVERRIDES, VoiceE2eeScope,
VoiceNoiseSuppressionConfigUpdateRequest, VoiceNoiseSuppressionGuildOverride,
PremiumMode, PushRelayConfigUpdateRequest, RegistrationMode, SsoConfigUpdateRequest,
VOICE_NS_MAX_GUILD_OVERRIDES, VoiceE2eeScope, VoiceNoiseSuppressionConfigUpdateRequest,
VoiceNoiseSuppressionGuildOverride,
},
},
config::AdminConfig,
@@ -209,10 +209,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_push_relay" => {
let update = build_push_relay_update(&form);
instance_config_result(client.update_instance_config(&update).await)
}
"update_domain_migration" => match build_domain_migration_update(&form) {
Ok(update) => instance_config_result(client.update_instance_config(&update).await),
Err(message) => FlashData::error(message),
@@ -658,39 +658,13 @@ fn build_voice_noise_suppression_update(
})
}
fn build_push_service_delivery_update(
form: &MultiValueForm,
) -> Result<InstanceConfigUpdateRequest, String> {
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_ascii_experiment_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",
)?),
relay_consent_accepted: Some(
form.bool_value("push_service_delivery_relay_consent_accepted"),
),
fn build_push_relay_update(form: &MultiValueForm) -> InstanceConfigUpdateRequest {
InstanceConfigUpdateRequest {
push_relay: Some(PushRelayConfigUpdateRequest {
relay_consent_accepted: Some(form.bool_value("push_relay_consent_accepted")),
}),
..Default::default()
})
}
}
fn build_domain_migration_update(
@@ -1784,26 +1758,19 @@ mod tests {
}
#[test]
fn build_push_service_delivery_update_reads_the_relay_consent_checkbox() {
let unchecked = MultiValueForm::parse(b"_csrf=token");
fn build_push_relay_update_reads_the_consent_checkbox() {
let unchecked = build_push_relay_update(&MultiValueForm::parse(b"_csrf=token"));
assert_eq!(
build_push_service_delivery_update(&unchecked)
.expect("valid form")
.push_service_delivery
.expect("push service delivery update")
.relay_consent_accepted,
Some(false)
serde_json::to_value(&unchecked).expect("serialize update"),
serde_json::json!({"push_relay": {"relay_consent_accepted": false}})
);
let checked =
MultiValueForm::parse(b"_csrf=token&push_service_delivery_relay_consent_accepted=true");
let checked = build_push_relay_update(&MultiValueForm::parse(
b"_csrf=token&push_relay_consent_accepted=true",
));
assert_eq!(
build_push_service_delivery_update(&checked)
.expect("valid form")
.push_service_delivery
.expect("push service delivery update")
.relay_consent_accepted,
Some(true)
serde_json::to_value(&checked).expect("serialize update"),
serde_json::json!({"push_relay": {"relay_consent_accepted": true}})
);
}
@@ -8,9 +8,8 @@ use crate::{
ExperimentDeliveryConfigResponse, GatewayRolloutConfigResponse, InstanceConfigResponse,
InstanceIntegrationsResponse, InstanceMediaResponse, InstancePolicyResponse,
InstanceRegistrationResponse, LimitConfigResponse, NoiseSuppressionBackend,
PUSH_SERVICE_DELIVERY_DEFAULT_SALT, PendingRegistrationResponse,
PushServiceDeliveryConfigResponse, RegistrationUrlResponse, SsoConfigResponse,
VOICE_NS_MAX_GUILD_OVERRIDES, VoiceNoiseSuppressionConfigResponse,
PendingRegistrationResponse, PushRelayConfigResponse, RegistrationUrlResponse,
SsoConfigResponse, VOICE_NS_MAX_GUILD_OVERRIDES, VoiceNoiseSuppressionConfigResponse,
},
config::AdminConfig,
middleware::auth::AuthContext,
@@ -138,6 +137,13 @@ pub fn instance_config_page(
(integrations_config_section(base, csrf_token, &instance_config.integrations))
},
))
(config_group(
"Push notifications",
"Consent for the relay that delivers official mobile app notifications.",
html! {
(push_relay_section(base, csrf_token, &instance_config.push_relay))
},
))
(config_group(
"Media & retention",
"Attachment expiry rules that can be changed without editing environment variables.",
@@ -151,7 +157,6 @@ pub fn instance_config_page(
html! {
(gateway_rollout_section(base, csrf_token, &instance_config.gateway_rollout))
(voice_noise_suppression_section(base, csrf_token, &instance_config.voice_noise_suppression))
(push_service_delivery_section(base, csrf_token, &instance_config.push_service_delivery))
(domain_migration_section(base, csrf_token, &instance_config.domain_migration))
(altcha_captcha_section(base, csrf_token, &instance_config.altcha_captcha))
(experiment_delivery_section(base, csrf_token, &instance_config.experiment_delivery))
@@ -1182,138 +1187,69 @@ fn voice_noise_suppression_section(
)
}
fn push_service_delivery_section(
fn push_relay_section(
base: &str,
csrf_token: &str,
push_service_delivery: &PushServiceDeliveryConfigResponse,
push_relay: &PushRelayConfigResponse,
) -> Markup {
let status = if push_service_delivery.enabled {
("Live", BadgeVariant::Success)
let status = if push_relay.relay_consent_accepted {
("Accepted", 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");
let relay_consent_stamp = match (
push_service_delivery.relay_consent_accepted_at.as_deref(),
push_service_delivery.relay_consent_accepted_by.as_deref(),
) {
(Some(at), Some(by)) => Some(format!("Accepted {at} by user {by}")),
(Some(at), None) => Some(format!("Accepted {at}")),
_ => None,
("Not accepted", BadgeVariant::Default)
};
let accepted_at =
format_optional_admin_timestamp(push_relay.relay_consent_accepted_at.as_deref(), "Never");
let accepted_by = push_relay
.relay_consent_accepted_by
.as_deref()
.unwrap_or("Nobody");
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.",
"Push Relay",
"Official mobile app notifications travel through Fluxer's relay to Apple and Google. \
The relay delivers them only after an operator accepts its privacy notice.",
html! {
form method="post" action={(base) "/instance-config?action=update_push_service_delivery"} {
form method="post" action={(base) "/instance-config?action=update_push_relay"} {
(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" }
h3 class="text-sm font-semibold text-neutral-900" { "Relay consent" }
(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" { "Managed relay consent" }
(checkbox(
"push_service_delivery_relay_consent_accepted",
"push_relay_consent_accepted",
"true",
"Accept the push relay supplemental privacy notice",
push_service_delivery.relay_consent_accepted,
push_relay.relay_consent_accepted,
true,
))
p class="text-xs text-neutral-500" {
"Required only for the official mobile apps, whose notifications travel \
through Fluxer's relay to Apple and Google. Until this is accepted those \
notifications are dropped. Self-hosted UnifiedPush and ntfy endpoints \
never reach the relay and are unaffected. "
"Until this is accepted official mobile app notifications are dropped. \
Self-hosted UnifiedPush and ntfy endpoints never reach the relay and are \
unaffected. "
a href="https://fluxer.com/push-relay" target="_blank" rel="noreferrer"
class="text-neutral-900 underline decoration-neutral-300 hover:text-neutral-600 hover:decoration-neutral-500" {
"Read the notice"
}
}
@if let Some(stamp) = relay_consent_stamp {
p class="text-xs text-neutral-500" { (stamp) }
}
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,
div class="grid grid-cols-1 gap-4 sm:grid-cols-2" {
(form_field_group("Accepted at", "push_relay_consent_accepted_at", false, None, None,
html! {
input type="text" id="push_relay_consent_accepted_at"
value=(accepted_at)
disabled class=(FORM_INPUT_CLASS);
},
))
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,
(form_field_group("Accepted by user ID", "push_relay_consent_accepted_by", false, None, None,
html! {
input type="text" id="push_relay_consent_accepted_by"
value=(accepted_by)
disabled class=(FORM_INPUT_CLASS);
},
))
(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"))
(submit_button("Save Push Relay Settings"))
}))
}
}
@@ -2259,26 +2195,28 @@ mod tests {
}
#[test]
fn push_service_delivery_section_shows_the_relay_consent_toggle() {
let accepted = PushServiceDeliveryConfigResponse {
fn push_relay_section_shows_the_consent_toggle() {
let accepted = PushRelayConfigResponse {
relay_consent_accepted: true,
relay_consent_accepted_at: Some("2026-09-27T10:11:12.000Z".to_owned()),
relay_consent_accepted_by: Some("1130650140672000000".to_owned()),
..PushServiceDeliveryConfigResponse::default()
};
let markup = push_service_delivery_section("/admin", "csrf", &accepted).into_string();
assert!(markup.contains("name=\"push_service_delivery_relay_consent_accepted\""));
let markup = push_relay_section("/admin", "csrf", &accepted).into_string();
assert!(markup.contains("action=update_push_relay"));
assert!(markup.contains("name=\"push_relay_consent_accepted\""));
assert!(markup.contains("https://fluxer.com/push-relay"));
assert!(markup.contains("Accepted 2026-09-27T10:11:12.000Z by user 1130650140672000000"));
assert!(markup.contains("value=\"Sep 27, 2026, 10:11 AM UTC\""));
assert!(markup.contains("value=\"1130650140672000000\""));
assert!(!markup.contains("name=\"push_relay_consent_accepted_at\""));
assert!(!markup.contains("name=\"push_relay_consent_accepted_by\""));
assert!(!markup.to_lowercase().contains("rollout"));
let unaccepted = push_service_delivery_section(
"/admin",
"csrf",
&PushServiceDeliveryConfigResponse::default(),
)
.into_string();
assert!(unaccepted.contains("name=\"push_service_delivery_relay_consent_accepted\""));
assert!(!unaccepted.contains("Accepted "));
let unaccepted =
push_relay_section("/admin", "csrf", &PushRelayConfigResponse::default()).into_string();
assert!(unaccepted.contains("name=\"push_relay_consent_accepted\""));
assert!(unaccepted.contains("Not accepted"));
assert!(unaccepted.contains("value=\"Never\""));
assert!(unaccepted.contains("value=\"Nobody\""));
}
#[test]
+13 -39
View File
@@ -409,13 +409,7 @@ fn deserialize_instance_config_response_with_unknown_keys() {
"future_object_knob": {"nested": true},
"future_list_knob": ["a", "b"]
},
"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": [],
"push_relay": {
"relay_consent_accepted": true,
"relay_consent_accepted_at": "2026-09-27T10:11:12.000Z",
"relay_consent_accepted_by": "1130650140672000000"
@@ -579,7 +573,7 @@ fn deserialize_instance_config_response_with_unknown_keys() {
assert_eq!(resp.domain_migration.included_user_ids.len(), 1);
assert_eq!(resp.domain_migration.anonymous_rollout_basis_points, 100);
assert!(resp.domain_migration.standalone_forwarding);
assert!(resp.push_service_delivery.relay_consent_accepted);
assert!(resp.push_relay.relay_consent_accepted);
assert!(resp.altcha_captcha.enabled);
assert_eq!(resp.altcha_captcha.config_version, 3);
assert!(resp.altcha_captcha.anonymous_enabled);
@@ -618,15 +612,9 @@ fn deserialize_instance_config_response_with_unknown_keys() {
}
#[test]
fn deserialize_push_service_delivery_relay_consent() {
let accepted: types::PushServiceDeliveryConfigResponse = serde_json::from_str(
fn deserialize_push_relay_config() {
let accepted: types::PushRelayConfigResponse = serde_json::from_str(
r#"{
"enabled": true,
"config_version": 3,
"rollout_basis_points": 5000,
"rollout_salt": "push-service-delivery-v1",
"included_user_ids": [],
"excluded_user_ids": [],
"relay_consent_accepted": true,
"relay_consent_accepted_at": "2026-09-27T10:11:12.000Z",
"relay_consent_accepted_by": "1130650140672000000"
@@ -644,37 +632,23 @@ fn deserialize_push_service_delivery_relay_consent() {
Some("1130650140672000000")
);
let legacy: types::PushServiceDeliveryConfigResponse = serde_json::from_str(
r#"{
"enabled": true,
"config_version": 3,
"rollout_basis_points": 5000,
"rollout_salt": "push-service-delivery-v1",
"included_user_ids": [],
"excluded_user_ids": []
}"#,
)
.expect("a response written before relay consent must still deserialize");
let empty: types::PushRelayConfigResponse =
serde_json::from_str("{}").expect("an empty push relay config must deserialize");
assert!(!legacy.relay_consent_accepted);
assert!(legacy.relay_consent_accepted_at.is_none());
assert!(legacy.relay_consent_accepted_by.is_none());
assert!(!empty.relay_consent_accepted);
assert!(empty.relay_consent_accepted_at.is_none());
assert!(empty.relay_consent_accepted_by.is_none());
}
#[test]
fn serialize_push_service_delivery_update_omits_an_unset_relay_consent() {
let without = types::PushServiceDeliveryConfigUpdateRequest {
enabled: Some(true),
..Default::default()
};
fn serialize_push_relay_update_omits_an_unset_consent() {
assert_eq!(
serde_json::to_value(&without).unwrap(),
serde_json::json!({"enabled": true})
serde_json::to_value(types::PushRelayConfigUpdateRequest::default()).unwrap(),
serde_json::json!({})
);
let with = types::PushServiceDeliveryConfigUpdateRequest {
let with = types::PushRelayConfigUpdateRequest {
relay_consent_accepted: Some(true),
..Default::default()
};
assert_eq!(
serde_json::to_value(&with).unwrap(),
@@ -16,7 +16,7 @@ import {OpenAPI} from '@app/api/middleware/ResponseTypeMiddleware';
import {
getGatewayRolloutConfigPublisher,
getInstanceConfigRepository,
getPushServiceDeliveryConfigPublisher,
getPushRelayConfigPublisher,
} from '@app/api/middleware/ServiceSingletons';
import {RateLimitConfigs} from '@app/api/RateLimitConfig';
import type {HonoApp, HonoEnv} from '@app/api/types/HonoEnv';
@@ -37,11 +37,7 @@ import {
import {AltchaCaptchaConfigSchema} from '@fluxer/schema/src/domains/admin/AltchaCaptchaSchemas';
import {DomainMigrationConfigSchema} from '@fluxer/schema/src/domains/admin/DomainMigrationSchemas';
import {GatewayRolloutConfigSchema} from '@fluxer/schema/src/domains/admin/GatewayRolloutSchemas';
import {
type PushServiceDeliveryConfig,
PushServiceDeliveryConfigSchema,
type PushServiceDeliveryConfigUpdateRequest,
} from '@fluxer/schema/src/domains/admin/PushServiceDeliverySchemas';
import type {PushRelayConfig, PushRelayConfigUpdateRequest} from '@fluxer/schema/src/domains/admin/PushRelaySchemas';
import {VoiceNoiseSuppressionConfigSchema} from '@fluxer/schema/src/domains/admin/VoiceNoiseSuppressionSchemas';
import {UserIdParam} from '@fluxer/schema/src/domains/common/CommonParamSchemas';
import {ExperimentDeliveryConfigSchema} from '@fluxer/schema/src/domains/experiment/ExperimentSchemas';
@@ -70,7 +66,7 @@ async function buildInstanceConfigResponse(): Promise<InstanceConfigResponse> {
ssoConfig,
gatewayRollout,
voiceNoiseSuppression,
pushServiceDelivery,
pushRelay,
domainMigration,
altchaCaptcha,
experimentDelivery,
@@ -81,7 +77,7 @@ async function buildInstanceConfigResponse(): Promise<InstanceConfigResponse> {
instanceConfigRepository.getSsoConfig(),
instanceConfigRepository.getGatewayRolloutConfig(),
instanceConfigRepository.getVoiceNoiseSuppressionConfig(),
instanceConfigRepository.getPushServiceDeliveryConfig(),
instanceConfigRepository.getPushRelayConfig(),
instanceConfigRepository.getDomainMigrationConfig(),
instanceConfigRepository.getAltchaCaptchaConfig(),
instanceConfigRepository.getExperimentDeliveryConfig(),
@@ -115,7 +111,7 @@ async function buildInstanceConfigResponse(): Promise<InstanceConfigResponse> {
},
gateway_rollout: gatewayRollout,
voice_noise_suppression: voiceNoiseSuppression,
push_service_delivery: pushServiceDelivery,
push_relay: pushRelay,
domain_migration: domainMigration,
altcha_captcha: altchaCaptcha,
experiment_delivery: experimentDelivery,
@@ -203,10 +199,10 @@ async function grantSetupCompleterAdminACL(ctx: Context<HonoEnv>): Promise<boole
}
function relayConsentStamp(
current: PushServiceDeliveryConfig,
patch: Partial<PushServiceDeliveryConfigUpdateRequest>,
current: PushRelayConfig,
patch: PushRelayConfigUpdateRequest,
adminUserId: string,
): Partial<PushServiceDeliveryConfig> {
): Partial<PushRelayConfig> {
const accepted = patch.relay_consent_accepted;
if (accepted === undefined || accepted === current.relay_consent_accepted) {
return {};
@@ -295,19 +291,16 @@ export function InstanceConfigAdminController(app: HonoApp) {
);
}
}
if (data.push_service_delivery) {
const patch = omitUndefinedFields(data.push_service_delivery);
if (data.push_relay) {
const patch = omitUndefinedFields(data.push_relay);
if (Object.keys(patch).length > 0) {
const adminUserId = ctx.get('adminUserId').toString();
const landed = await instanceConfigRepository.updatePushServiceDeliveryConfig((current) =>
PushServiceDeliveryConfigSchema.parse({
...current,
...patch,
...relayConsentStamp(current, patch, adminUserId),
config_version: current.config_version + 1,
}),
);
await getPushServiceDeliveryConfigPublisher().publish(landed);
const landed = await instanceConfigRepository.updatePushRelayConfig((current) => ({
...current,
...patch,
...relayConsentStamp(current, patch, adminUserId),
}));
await getPushRelayConfigPublisher().publish(landed);
}
}
if (data.domain_migration) {
@@ -4,7 +4,7 @@ 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 {PushRelayConfigPublisher} from '@app/api/instance/PushRelayConfigPublisher';
import {InstanceConfigWriteRaceExecutor} from '@app/api/instance/tests/InstanceConfigWriteRaceExecutor';
import {getAdminRepository} from '@app/api/middleware/ServiceSingletons';
import type {ApiTestHarness} from '@app/api/test/ApiTestHarness';
@@ -16,12 +16,12 @@ 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';
type LegacyPushServiceDeliveryWire,
toLegacyPushServiceDeliveryWire,
} from '@fluxer/schema/src/domains/admin/PushRelaySchemas';
import {afterAll, afterEach, beforeAll, beforeEach, describe, expect, it, vi} from 'vitest';
const PUSH_SERVICE_DELIVERY_CONFIG_KEY = 'push_service_delivery_config';
const PUSH_RELAY_CONFIG_KEY = 'push_service_delivery_config';
describe('instance config admin PATCH under concurrent writes', () => {
let harness: ApiTestHarness;
@@ -55,13 +55,13 @@ describe('instance config admin PATCH under concurrent writes', () => {
const patchConfig = (admin: TestAccount, body: Record<string, unknown>) =>
createBuilder<InstanceConfigResponse>(harness, admin.token).patch('/admin/instance/config').body(body);
const spyOnPushDeliveryPublishes = () =>
vi.spyOn(PushServiceDeliveryConfigPublisher.prototype, 'publish').mockResolvedValue(undefined);
const spyOnPushRelayPublishes = () =>
vi.spyOn(PushRelayConfigPublisher.prototype, 'publish').mockResolvedValue(undefined);
async function readStoredPushServiceDelivery(): Promise<PushServiceDeliveryConfig> {
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 readStoredPushRelay(): Promise<LegacyPushServiceDeliveryWire> {
const raw = await executor.readDirectly(PUSH_RELAY_CONFIG_KEY);
if (raw === null) throw new Error('push relay config was never stored');
return JSON.parse(raw) as LegacyPushServiceDeliveryWire;
}
async function listConfigUpdateAudits(): Promise<Array<AdminAuditLog>> {
@@ -84,37 +84,32 @@ describe('instance config admin PATCH under concurrent writes', () => {
});
it('answers with a conflict and neither writes, publishes nor audits once every attempt has lost the race', async () => {
const publish = spyOnPushDeliveryPublishes();
const publish = spyOnPushRelayPublishes();
const admin = await createAdmin();
await patchConfig(admin, {push_service_delivery: {enabled: true, rollout_basis_points: 1000}}).execute();
await patchConfig(admin, {push_relay: {relay_consent_accepted: false}}).execute();
publish.mockClear();
const auditsBefore = await listConfigUpdateAudits();
executor.watch(PUSH_SERVICE_DELIVERY_CONFIG_KEY);
executor.watch(PUSH_RELAY_CONFIG_KEY);
const unaccepted = {
relay_consent_accepted: false,
relay_consent_accepted_at: null,
relay_consent_accepted_by: null,
};
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,
}),
PUSH_RELAY_CONFIG_KEY,
JSON.stringify(toLegacyPushServiceDeliveryWire(unaccepted, 100 + competingWrites)),
);
});
await patchConfig(admin, {push_service_delivery: {rollout_basis_points: 5000}})
await patchConfig(admin, {push_relay: {relay_consent_accepted: true}})
.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(await readStoredPushRelay()).toEqual(toLegacyPushServiceDeliveryWire(unaccepted, 100 + competingWrites));
expect(publish).not.toHaveBeenCalled();
expect(await listConfigUpdateAudits()).toHaveLength(auditsBefore.length);
});
@@ -2,24 +2,54 @@
import type {TestAccount} from '@app/api/auth/tests/AuthTestUtils';
import {createTestAccount, setUserACLs} from '@app/api/auth/tests/AuthTestUtils';
import {PushServiceDeliveryConfigPublisher} from '@app/api/instance/PushServiceDeliveryConfigPublisher';
import {setCassandraQueryExecutorForTesting} from '@app/api/database/CassandraQueryExecution';
import {PushRelayConfigPublisher} from '@app/api/instance/PushRelayConfigPublisher';
import {InstanceConfigWriteRaceExecutor} from '@app/api/instance/tests/InstanceConfigWriteRaceExecutor';
import {getInstanceConfigRepository} 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 type {InstanceConfigResponse} from '@fluxer/schema/src/domains/admin/AdminSchemas';
import type {LegacyPushServiceDeliveryWire} from '@fluxer/schema/src/domains/admin/PushRelaySchemas';
import {afterAll, afterEach, beforeAll, beforeEach, describe, expect, it, vi} from 'vitest';
const PUSH_RELAY_CONFIG_KEY = 'push_service_delivery_config';
const ACCEPTED_AT = '2026-09-20T08:00:00.000Z';
const ACCEPTED_BY = '1500000000000000007';
const PROD_ROW = {
enabled: true,
config_version: 41,
rollout_basis_points: 10000,
rollout_salt: 'push-service-delivery-v1',
included_user_ids: [],
excluded_user_ids: [],
relay_consent_accepted: true,
relay_consent_accepted_at: ACCEPTED_AT,
relay_consent_accepted_by: ACCEPTED_BY,
};
interface PushServiceDeliveryRpcResponse {
type: 'get_push_service_delivery_config';
data: {config: LegacyPushServiceDeliveryWire};
}
describe('push relay supplemental notice consent', () => {
let harness: ApiTestHarness;
let executor: InstanceConfigWriteRaceExecutor;
beforeAll(async () => {
harness = await createApiTestHarness();
executor = new InstanceConfigWriteRaceExecutor(new InMemoryCassandraQueryExecutor());
setCassandraQueryExecutorForTesting(executor);
});
beforeEach(async () => {
await harness.reset();
vi.spyOn(PushServiceDeliveryConfigPublisher.prototype, 'publish').mockResolvedValue(undefined);
vi.spyOn(PushRelayConfigPublisher.prototype, 'publish').mockResolvedValue(undefined);
});
afterEach(() => {
@@ -43,63 +73,88 @@ describe('push relay supplemental notice consent', () => {
const readConfig = (admin: TestAccount) =>
createBuilder<InstanceConfigResponse>(harness, admin.token).get('/admin/instance/config');
const readRpcConfig = async (): Promise<LegacyPushServiceDeliveryWire> => {
const response = await createBuilder<PushServiceDeliveryRpcResponse>(harness, '')
.post('/test/rpc-session-init')
.body({type: 'get_push_service_delivery_config'})
.expect(HTTP_STATUS.OK)
.execute();
expect(response.type).toBe('get_push_service_delivery_config');
return response.data.config;
};
async function storeRow(row: Record<string, unknown>): Promise<void> {
await executor.writeDirectly(PUSH_RELAY_CONFIG_KEY, JSON.stringify(row));
getInstanceConfigRepository().clearCacheForTesting();
}
async function readStoredRow(): Promise<unknown> {
const raw = await executor.readDirectly(PUSH_RELAY_CONFIG_KEY);
if (raw === null) throw new Error('push relay config was never stored');
return JSON.parse(raw);
}
it('reads back as unaccepted before an operator agrees', async () => {
const admin = await createAdmin();
const config = await readConfig(admin).execute();
expect(config.push_service_delivery).toMatchObject({
expect(config.push_relay).toEqual({
relay_consent_accepted: false,
relay_consent_accepted_at: null,
relay_consent_accepted_by: null,
});
});
it('keeps the consent of a stored push service delivery row', async () => {
const admin = await createAdmin();
await storeRow(PROD_ROW);
const config = await readConfig(admin).execute();
expect(config.push_relay).toEqual({
relay_consent_accepted: true,
relay_consent_accepted_at: ACCEPTED_AT,
relay_consent_accepted_by: ACCEPTED_BY,
});
});
it('reads a stored row without consent fields as unaccepted', async () => {
const admin = await createAdmin();
await storeRow({enabled: true, config_version: 3, rollout_basis_points: 10000});
const config = await readConfig(admin).execute();
expect(config.push_relay.relay_consent_accepted).toBe(false);
expect(await readRpcConfig()).toMatchObject({config_version: 3, relay_consent_accepted: false});
});
it('stamps the acting admin and the acceptance time when consent is given', async () => {
const admin = await createAdmin();
const updated = await patchConfig(admin, {push_service_delivery: {relay_consent_accepted: true}}).execute();
const updated = await patchConfig(admin, {push_relay: {relay_consent_accepted: true}}).execute();
expect(updated.push_service_delivery.relay_consent_accepted).toBe(true);
expect(updated.push_service_delivery.relay_consent_accepted_by).toBe(admin.userId);
expect(Date.parse(updated.push_service_delivery.relay_consent_accepted_at ?? '')).not.toBeNaN();
});
it('keeps the first acceptance stamp when a later patch changes only the rollout', async () => {
const admin = await createAdmin();
const accepted = await patchConfig(admin, {push_service_delivery: {relay_consent_accepted: true}}).execute();
const rolledOut = await patchConfig(admin, {
push_service_delivery: {enabled: true, rollout_basis_points: 2500},
}).execute();
expect(rolledOut.push_service_delivery).toMatchObject({
enabled: true,
rollout_basis_points: 2500,
relay_consent_accepted: true,
relay_consent_accepted_at: accepted.push_service_delivery.relay_consent_accepted_at,
relay_consent_accepted_by: admin.userId,
});
expect(updated.push_relay.relay_consent_accepted).toBe(true);
expect(updated.push_relay.relay_consent_accepted_by).toBe(admin.userId);
expect(Date.parse(updated.push_relay.relay_consent_accepted_at ?? '')).not.toBeNaN();
});
it('keeps the stamp untouched when consent is re-sent unchanged', async () => {
const admin = await createAdmin();
const accepted = await patchConfig(admin, {push_service_delivery: {relay_consent_accepted: true}}).execute();
const accepted = await patchConfig(admin, {push_relay: {relay_consent_accepted: true}}).execute();
const resent = await patchConfig(admin, {push_service_delivery: {relay_consent_accepted: true}}).execute();
const resent = await patchConfig(admin, {push_relay: {relay_consent_accepted: true}}).execute();
expect(resent.push_service_delivery.relay_consent_accepted_at).toBe(
accepted.push_service_delivery.relay_consent_accepted_at,
);
expect(resent.push_relay).toEqual(accepted.push_relay);
});
it('clears the stamp when an operator withdraws consent', async () => {
const admin = await createAdmin();
await patchConfig(admin, {push_service_delivery: {relay_consent_accepted: true}}).execute();
await patchConfig(admin, {push_relay: {relay_consent_accepted: true}}).execute();
const withdrawn = await patchConfig(admin, {push_service_delivery: {relay_consent_accepted: false}}).execute();
const withdrawn = await patchConfig(admin, {push_relay: {relay_consent_accepted: false}}).execute();
expect(withdrawn.push_service_delivery).toMatchObject({
expect(withdrawn.push_relay).toEqual({
relay_consent_accepted: false,
relay_consent_accepted_at: null,
relay_consent_accepted_by: null,
@@ -110,23 +165,146 @@ describe('push relay supplemental notice consent', () => {
const admin = await createAdmin();
const updated = await patchConfig(admin, {
push_service_delivery: {
push_relay: {
relay_consent_accepted: true,
relay_consent_accepted_at: '2020-01-01T00:00:00.000Z',
relay_consent_accepted_by: '1500000000000000009',
},
}).execute();
expect(updated.push_service_delivery.relay_consent_accepted_at).not.toBe('2020-01-01T00:00:00.000Z');
expect(updated.push_service_delivery.relay_consent_accepted_by).toBe(admin.userId);
expect(updated.push_relay.relay_consent_accepted_at).not.toBe('2020-01-01T00:00:00.000Z');
expect(updated.push_relay.relay_consent_accepted_by).toBe(admin.userId);
});
it('publishes the consent to the delivery services', async () => {
it('writes the full legacy document and bumps the stored config version', async () => {
const admin = await createAdmin();
const publish = vi.mocked(PushServiceDeliveryConfigPublisher.prototype.publish);
await storeRow({
...PROD_ROW,
relay_consent_accepted: false,
relay_consent_accepted_at: null,
relay_consent_accepted_by: null,
});
await patchConfig(admin, {push_service_delivery: {relay_consent_accepted: true}}).execute();
const updated = await patchConfig(admin, {push_relay: {relay_consent_accepted: true}}).execute();
expect(publish).toHaveBeenCalledWith(expect.objectContaining({relay_consent_accepted: true}));
expect(await readStoredRow()).toEqual({
enabled: true,
config_version: 42,
rollout_basis_points: 10000,
rollout_salt: 'push-service-delivery-v1',
included_user_ids: [],
excluded_user_ids: [],
relay_consent_accepted: true,
relay_consent_accepted_at: updated.push_relay.relay_consent_accepted_at,
relay_consent_accepted_by: admin.userId,
});
await patchConfig(admin, {push_relay: {relay_consent_accepted: false}}).execute();
expect(await readStoredRow()).toMatchObject({
enabled: true,
config_version: 43,
rollout_basis_points: 10000,
relay_consent_accepted: false,
relay_consent_accepted_at: null,
relay_consent_accepted_by: null,
});
});
it('rewrites a partially enrolled stored row as full enrolment', async () => {
const admin = await createAdmin();
await storeRow({
...PROD_ROW,
enabled: false,
rollout_basis_points: 250,
rollout_salt: 'custom-salt',
included_user_ids: ['1500000000000000003'],
excluded_user_ids: ['1500000000000000004'],
});
await patchConfig(admin, {push_relay: {relay_consent_accepted: false}}).execute();
expect(await readStoredRow()).toMatchObject({
enabled: true,
config_version: 42,
rollout_basis_points: 10000,
rollout_salt: 'push-service-delivery-v1',
included_user_ids: [],
excluded_user_ids: [],
});
});
it('publishes the legacy delivery document with the consent', async () => {
const admin = await createAdmin();
const publish = vi.mocked(PushRelayConfigPublisher.prototype.publish);
const updated = await patchConfig(admin, {push_relay: {relay_consent_accepted: true}}).execute();
expect(publish).toHaveBeenCalledTimes(1);
expect(publish).toHaveBeenCalledWith({
enabled: true,
config_version: 1,
rollout_basis_points: 10000,
rollout_salt: 'push-service-delivery-v1',
included_user_ids: [],
excluded_user_ids: [],
relay_consent_accepted: true,
relay_consent_accepted_at: updated.push_relay.relay_consent_accepted_at,
relay_consent_accepted_by: admin.userId,
});
});
it('does not write or publish for an empty push relay patch', async () => {
const admin = await createAdmin();
const publish = vi.mocked(PushRelayConfigPublisher.prototype.publish);
await patchConfig(admin, {push_relay: {}}).execute();
expect(publish).not.toHaveBeenCalled();
expect(await executor.readDirectly(PUSH_RELAY_CONFIG_KEY)).toBeNull();
});
it('ignores the retired push_service_delivery section', async () => {
const admin = await createAdmin();
const publish = vi.mocked(PushRelayConfigPublisher.prototype.publish);
const updated = await patchConfig(admin, {push_service_delivery: {relay_consent_accepted: true}}).execute();
expect(updated.push_relay.relay_consent_accepted).toBe(false);
expect(publish).not.toHaveBeenCalled();
});
it('answers the legacy delivery RPC with full enrolment and the stored consent', async () => {
await storeRow(PROD_ROW);
expect(await readRpcConfig()).toEqual(PROD_ROW);
});
it('answers the legacy delivery RPC with defaults before anything is stored', async () => {
expect(await readRpcConfig()).toEqual({
enabled: true,
config_version: 0,
rollout_basis_points: 10000,
rollout_salt: 'push-service-delivery-v1',
included_user_ids: [],
excluded_user_ids: [],
relay_consent_accepted: false,
relay_consent_accepted_at: null,
relay_consent_accepted_by: null,
});
});
it('answers the legacy delivery RPC with consent given through the admin API', async () => {
const admin = await createAdmin();
const updated = await patchConfig(admin, {push_relay: {relay_consent_accepted: true}}).execute();
expect(await readRpcConfig()).toMatchObject({
enabled: true,
config_version: 1,
rollout_basis_points: 10000,
relay_consent_accepted: true,
relay_consent_accepted_at: updated.push_relay.relay_consent_accepted_at,
relay_consent_accepted_by: admin.userId,
});
});
});
@@ -43,7 +43,6 @@ async function revokeSessionTargets(
scope === 'all'
? users.deleteAllPushSubscriptions(userId)
: users.deletePushSubscriptionsForAuthSessions(userId, sessionIdHashes, {deleteUnboundSubscriptions: true}),
() => gateway.invalidatePushSubscriptions({userId}),
];
if (scope === 'selected' || targets.length > 0) {
steps.push(
@@ -37,7 +37,6 @@ import type {GuildMemberResponse} from '@fluxer/schema/src/domains/guild/GuildMe
import type {GuildResponse} from '@fluxer/schema/src/domains/guild/GuildResponseSchemas';
import {ms} from 'itty-time';
const PUSH_BADGE_COUNT_BATCH_SIZE = 100;
const USER_PERMISSIONS_BATCH_SIZE = 100;
const GATEWAY_ERROR_TO_DOMAIN_ERROR: Record<string, () => Error> = {
@@ -67,18 +66,6 @@ interface DispatchPresenceParams {
data: unknown;
}
interface InvalidatePushBadgeCountParams {
userId: UserID;
}
interface InvalidatePushBadgeCountsParams {
userIds: Array<UserID>;
}
interface InvalidatePushSubscriptionsParams {
userId: UserID;
}
interface ClearPushChannelNotificationsParams {
userId: UserID;
channelId: ChannelID;
@@ -287,8 +274,6 @@ export class GatewayService {
private readonly MAX_BATCH_CONCURRENCY = 50;
private readonly PENDING_REQUEST_TIMEOUT_MS = ms('30 seconds');
private readonly AUTH_CONTEXT_FALLBACK_MS = ms('5 minutes');
private readonly BADGE_COUNTS_FALLBACK_MS = ms('5 minutes');
private badgeCountsUnsupportedUntil = 0;
constructor() {
this.rpcClient = GatewayRpcClient.getInstance();
@@ -704,48 +689,6 @@ export class GatewayService {
});
}
async invalidatePushBadgeCount({userId}: InvalidatePushBadgeCountParams): Promise<void> {
await this.call('push.invalidate_badge_count', {
user_id: userId.toString(),
});
}
async invalidatePushBadgeCounts({userIds}: InvalidatePushBadgeCountsParams): Promise<void> {
if (Date.now() < this.badgeCountsUnsupportedUntil) {
await this.invalidatePushBadgeCountsIndividually(userIds);
return;
}
const batches: Array<Array<UserID>> = [];
for (let index = 0; index < userIds.length; index += PUSH_BADGE_COUNT_BATCH_SIZE) {
batches.push(userIds.slice(index, index + PUSH_BADGE_COUNT_BATCH_SIZE));
}
try {
await Promise.all(
batches.map((batch) =>
this.call('push.invalidate_badge_counts', {user_ids: batch.map((userId) => userId.toString())}),
),
);
} catch (error) {
const transformedError = this.transformGatewayError(error);
if (!this.isAuthContextUnsupportedError(transformedError)) {
throw transformedError;
}
this.badgeCountsUnsupportedUntil = Date.now() + this.BADGE_COUNTS_FALLBACK_MS;
Logger.warn({error}, '[gateway-rpc] push.invalidate_badge_counts unavailable, falling back to per-user calls');
await this.invalidatePushBadgeCountsIndividually(userIds);
}
}
private async invalidatePushBadgeCountsIndividually(userIds: ReadonlyArray<UserID>): Promise<void> {
await Promise.all(userIds.map((userId) => this.invalidatePushBadgeCount({userId})));
}
async invalidatePushSubscriptions({userId}: InvalidatePushSubscriptionsParams): Promise<void> {
await this.call('push.invalidate_subscriptions', {
user_id: userId.toString(),
});
}
async clearPushChannelNotifications({
userId,
channelId,
@@ -292,12 +292,6 @@ export abstract class IGatewayService {
abstract dispatchPresence(params: {userId: UserID; event: GatewayDispatchEvent; data: unknown}): Promise<void>;
abstract invalidatePushBadgeCount(params: {userId: UserID}): Promise<void>;
abstract invalidatePushBadgeCounts(params: {userIds: Array<UserID>}): Promise<void>;
abstract invalidatePushSubscriptions(params: {userId: UserID}): Promise<void>;
abstract clearPushChannelNotifications(params: {
userId: UserID;
channelId: ChannelID;
@@ -41,9 +41,11 @@ import {
GatewayRolloutConfigSchema,
} from '@fluxer/schema/src/domains/admin/GatewayRolloutSchemas';
import {
type PushServiceDeliveryConfig,
PushServiceDeliveryConfigSchema,
} from '@fluxer/schema/src/domains/admin/PushServiceDeliverySchemas';
type LegacyPushServiceDeliveryWire,
type PushRelayConfig,
PushRelayConfigSchema,
toLegacyPushServiceDeliveryWire,
} from '@fluxer/schema/src/domains/admin/PushRelaySchemas';
import {
type VoiceNoiseSuppressionConfig,
VoiceNoiseSuppressionConfigSchema,
@@ -70,7 +72,7 @@ import {z} from 'zod';
const GATEWAY_ROLLOUT_CONFIG_KEY = 'gateway_rollout_config';
const VOICE_NOISE_SUPPRESSION_CONFIG_KEY = 'voice_noise_suppression_config';
const PUSH_SERVICE_DELIVERY_CONFIG_KEY = 'push_service_delivery_config';
const PUSH_RELAY_CONFIG_KEY = 'push_service_delivery_config';
const DOMAIN_MIGRATION_CONFIG_KEY = 'domain_migration_config';
const ALTCHA_CAPTCHA_CONFIG_KEY = 'altcha_captcha_config';
const EXPERIMENT_DELIVERY_CONFIG_KEY = 'experiment_delivery_config';
@@ -379,7 +381,7 @@ type StoredConfigSection =
| 'app public'
| 'gateway rollout'
| 'voice noise suppression'
| 'push service delivery'
| 'push relay'
| 'domain migration'
| 'altcha captcha'
| 'experiment delivery'
@@ -520,8 +522,25 @@ function parseStoredVoiceNoiseSuppressionConfig(raw: string | null): VoiceNoiseS
return parseStoredConfigOrDefault(VoiceNoiseSuppressionConfigSchema, raw, 'voice noise suppression');
}
function parseStoredPushServiceDeliveryConfig(raw: string | null): PushServiceDeliveryConfig {
return parseStoredConfigOrDefault(PushServiceDeliveryConfigSchema, raw, 'push service delivery');
const StoredPushRelayConfigSchema = PushRelayConfigSchema.extend({
config_version: z.number().int().min(0).default(0),
});
function parseStoredPushRelayConfig(raw: string | null): LegacyPushServiceDeliveryWire {
const {config_version, ...config} = salvageStoredConfig(
StoredPushRelayConfigSchema,
readStoredConfigValue(raw, 'push relay'),
'push relay',
);
return toLegacyPushServiceDeliveryWire(config, config_version);
}
function toPushRelayConfig(wire: LegacyPushServiceDeliveryWire): PushRelayConfig {
return {
relay_consent_accepted: wire.relay_consent_accepted,
relay_consent_accepted_at: wire.relay_consent_accepted_at,
relay_consent_accepted_by: wire.relay_consent_accepted_by,
};
}
function parseStoredDomainMigrationConfig(raw: string | null): DomainMigrationConfig {
@@ -1179,7 +1198,7 @@ export class InstanceConfigRepository {
parseStoredGatewayRolloutConfig(snapshot.get(GATEWAY_ROLLOUT_CONFIG_KEY) ?? null),
);
parseStoredVoiceNoiseSuppressionConfig(snapshot.get(VOICE_NOISE_SUPPRESSION_CONFIG_KEY) ?? null);
parseStoredPushServiceDeliveryConfig(snapshot.get(PUSH_SERVICE_DELIVERY_CONFIG_KEY) ?? null);
parseStoredPushRelayConfig(snapshot.get(PUSH_RELAY_CONFIG_KEY) ?? null);
parseStoredDomainMigrationConfig(snapshot.get(DOMAIN_MIGRATION_CONFIG_KEY) ?? null);
parseStoredAltchaCaptchaConfig(snapshot.get(ALTCHA_CAPTCHA_CONFIG_KEY) ?? null);
parseStoredExperimentDeliveryConfig(snapshot.get(EXPERIMENT_DELIVERY_CONFIG_KEY) ?? null);
@@ -1278,21 +1297,21 @@ export class InstanceConfigRepository {
);
}
async getPushServiceDeliveryConfig(): Promise<PushServiceDeliveryConfig> {
const raw = await this.getConfig(PUSH_SERVICE_DELIVERY_CONFIG_KEY);
return parseStoredPushServiceDeliveryConfig(raw);
async getLegacyPushServiceDeliveryWire(): Promise<LegacyPushServiceDeliveryWire> {
const raw = await this.getConfig(PUSH_RELAY_CONFIG_KEY);
return parseStoredPushRelayConfig(raw);
}
updatePushServiceDeliveryConfig(
update: (current: PushServiceDeliveryConfig) => PushServiceDeliveryConfig,
): Promise<PushServiceDeliveryConfig> {
return this.updateStoredConfig(PUSH_SERVICE_DELIVERY_CONFIG_KEY, (raw) =>
validateStoredConfig(
PushServiceDeliveryConfigSchema,
update(parseStoredPushServiceDeliveryConfig(raw)),
'push service delivery',
),
);
async getPushRelayConfig(): Promise<PushRelayConfig> {
return toPushRelayConfig(await this.getLegacyPushServiceDeliveryWire());
}
updatePushRelayConfig(update: (current: PushRelayConfig) => PushRelayConfig): Promise<LegacyPushServiceDeliveryWire> {
return this.updateStoredConfig(PUSH_RELAY_CONFIG_KEY, (raw) => {
const current = parseStoredPushRelayConfig(raw);
const next = validateStoredConfig(PushRelayConfigSchema, update(toPushRelayConfig(current)), 'push relay');
return toLegacyPushServiceDeliveryWire(next, current.config_version + 1);
});
}
async getDomainMigrationConfig(): Promise<DomainMigrationConfig> {
@@ -1,21 +1,21 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import type {PushServiceDeliveryConfig} from '@fluxer/schema/src/domains/admin/PushServiceDeliverySchemas';
import type {LegacyPushServiceDeliveryWire} from '@fluxer/schema/src/domains/admin/PushRelaySchemas';
import type {INatsConnectionManager} from '@pkgs/nats/src/INatsConnectionManager';
const textEncoder = new TextEncoder();
export const PUSH_SERVICE_DELIVERY_CONFIG_NATS_SUBJECT = 'config.push.delivery';
const PUSH_SERVICE_DELIVERY_CONFIG_NATS_SUBJECT = 'config.push.delivery';
interface PushServiceDeliveryConfigNatsMessage {
type: 'push_service_delivery_config';
config: PushServiceDeliveryConfig;
config: LegacyPushServiceDeliveryWire;
}
export class PushServiceDeliveryConfigPublisher {
export class PushRelayConfigPublisher {
constructor(private readonly connectionManager: INatsConnectionManager) {}
async publish(config: PushServiceDeliveryConfig): Promise<void> {
async publish(config: LegacyPushServiceDeliveryWire): Promise<void> {
if (this.connectionManager.isClosed()) {
await this.connectionManager.connect();
}
@@ -51,7 +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 {PushRelayConfigPublisher} from '@app/api/instance/PushRelayConfigPublisher';
import {InviteRepository} from '@app/api/invite/InviteRepository';
import {Logger} from '@app/api/Logger';
import {LimitConfigService} from '@app/api/limits/LimitConfigService';
@@ -157,13 +157,13 @@ export const getGatewayRolloutConfigPublisher = singleton(
),
);
export const getPushServiceDeliveryConfigPublisher = singleton(
export const getPushRelayConfigPublisher = singleton(
() =>
new PushServiceDeliveryConfigPublisher(
new PushRelayConfigPublisher(
new NatsConnectionManager({
url: Config.nats.coreUrl,
token: Config.nats.authToken || undefined,
name: 'fluxer-api-push-service-delivery-config',
name: 'fluxer-api-push-relay-config',
}),
),
);
@@ -15,53 +15,6 @@ import {ReadStateService} from '@app/api/read_state/ReadStateService';
import {BadGatewayError} from '@fluxer/errors/src/domains/core/BadGatewayError';
import {describe, expect, it, vi} from 'vitest';
describe('ReadStateService.bulkIncrementMentionCounts', () => {
it('invalidates badge counts for touched users in a single bulk call', async () => {
const channelId = createChannelID(2n);
const messageId = createMessageID(3n);
const touched: Array<{userId: UserID; channelId: ChannelID}> = [
{userId: createUserID(10n), channelId},
{userId: createUserID(11n), channelId},
{userId: createUserID(10n), channelId: createChannelID(4n)},
];
const repository = {
bulkIncrementMentionCounts: vi.fn().mockResolvedValue(touched),
} as unknown as IReadStateRepository;
const invalidatePushBadgeCounts = vi.fn().mockResolvedValue(undefined);
const invalidatePushBadgeCount = vi.fn().mockResolvedValue(undefined);
const gatewayService = {
invalidatePushBadgeCounts,
invalidatePushBadgeCount,
} as unknown as IGatewayService;
const service = new ReadStateService(repository, gatewayService);
await service.bulkIncrementMentionCounts([
{userId: createUserID(10n), channelId, messageId},
{userId: createUserID(11n), channelId, messageId},
{userId: createUserID(12n), channelId, messageId},
]);
expect(invalidatePushBadgeCount).not.toHaveBeenCalled();
expect(invalidatePushBadgeCounts).toHaveBeenCalledTimes(1);
expect(invalidatePushBadgeCounts).toHaveBeenCalledWith({userIds: [createUserID(10n), createUserID(11n)]});
});
it('skips the bulk call when no read state was touched', async () => {
const repository = {
bulkIncrementMentionCounts: vi.fn().mockResolvedValue([]),
} as unknown as IReadStateRepository;
const invalidatePushBadgeCounts = vi.fn().mockResolvedValue(undefined);
const gatewayService = {invalidatePushBadgeCounts} as unknown as IGatewayService;
const service = new ReadStateService(repository, gatewayService);
await service.bulkIncrementMentionCounts([
{userId: createUserID(10n), channelId: createChannelID(2n), messageId: createMessageID(3n)},
]);
expect(invalidatePushBadgeCounts).not.toHaveBeenCalled();
});
});
const USER_ID = createUserID(20n);
const CHANNEL_ID = createChannelID(21n);
const MESSAGE_ID = createMessageID(22n);
@@ -78,7 +31,7 @@ function makeReadState(channelId: ChannelID, messageId: MessageID, mentionCount
}
describe('ReadStateService gateway side effects after the write', () => {
it('returns the committed read state when the badge invalidation fails', async () => {
it('returns the committed read state when clearing push notifications fails', async () => {
const stored: Array<{channelId: ChannelID; messageId: MessageID}> = [];
const repository = {
upsertReadState: vi.fn(async (_userId: UserID, channelId: ChannelID, messageId: MessageID) => {
@@ -87,8 +40,7 @@ describe('ReadStateService gateway side effects after the write', () => {
}),
} as unknown as IReadStateRepository;
const gatewayService = {
invalidatePushBadgeCount: vi.fn().mockRejectedValue(new BadGatewayError()),
clearPushChannelNotifications: vi.fn().mockResolvedValue(undefined),
clearPushChannelNotifications: vi.fn().mockRejectedValue(new BadGatewayError()),
dispatchPresence: vi.fn().mockResolvedValue(undefined),
} as unknown as IGatewayService;
const service = new ReadStateService(repository, gatewayService);
@@ -113,7 +65,6 @@ describe('ReadStateService gateway side effects after the write', () => {
),
} as unknown as IReadStateRepository;
const gatewayService = {
invalidatePushBadgeCount: vi.fn().mockResolvedValue(undefined),
clearPushChannelNotifications: vi.fn().mockResolvedValue(undefined),
dispatchPresence: vi.fn().mockRejectedValue(new BadGatewayError()),
} as unknown as IGatewayService;
@@ -138,7 +89,6 @@ describe('ReadStateService gateway side effects after the write', () => {
}),
} as unknown as IReadStateRepository;
const gatewayService = {
invalidatePushBadgeCount: vi.fn().mockResolvedValue(undefined),
clearPushChannelNotifications: vi.fn().mockResolvedValue(undefined),
dispatchPresence: vi.fn().mockRejectedValue(new BadGatewayError()),
} as unknown as IGatewayService;
@@ -156,42 +106,13 @@ describe('ReadStateService gateway side effects after the write', () => {
expect(stored).toEqual(['21', '23']);
});
it('deletes the read state when the badge invalidation fails', async () => {
const deleteReadState = vi.fn().mockResolvedValue(undefined);
const repository = {deleteReadState} as unknown as IReadStateRepository;
const gatewayService = {
invalidatePushBadgeCount: vi.fn().mockRejectedValue(new BadGatewayError()),
} as unknown as IGatewayService;
const service = new ReadStateService(repository, gatewayService);
await expect(service.deleteReadState({userId: USER_ID, channelId: CHANNEL_ID})).resolves.toBeUndefined();
expect(deleteReadState).toHaveBeenCalledWith(USER_ID, CHANNEL_ID);
});
it('increments the mention count when the badge invalidation fails', async () => {
const incrementReadStateMentions = vi.fn().mockResolvedValue(makeReadState(CHANNEL_ID, MESSAGE_ID, 1));
const repository = {incrementReadStateMentions} as unknown as IReadStateRepository;
const gatewayService = {
invalidatePushBadgeCount: vi.fn().mockRejectedValue(new BadGatewayError()),
} as unknown as IGatewayService;
const service = new ReadStateService(repository, gatewayService);
await expect(
service.incrementMentionCount({userId: USER_ID, channelId: CHANNEL_ID, messageId: MESSAGE_ID}),
).resolves.toBeUndefined();
expect(incrementReadStateMentions).toHaveBeenCalledTimes(1);
});
it('returns the bulk acknowledged states when the badge invalidation fails', async () => {
it('returns the bulk acknowledged states when clearing push notifications fails', async () => {
const updated = [makeReadState(CHANNEL_ID, MESSAGE_ID)];
const repository = {
bulkAckMessages: vi.fn().mockResolvedValue(updated),
} as unknown as IReadStateRepository;
const gatewayService = {
invalidatePushBadgeCount: vi.fn().mockRejectedValue(new BadGatewayError()),
clearPushChannelNotifications: vi.fn().mockResolvedValue(undefined),
clearPushChannelNotifications: vi.fn().mockRejectedValue(new BadGatewayError()),
dispatchPresence: vi.fn().mockResolvedValue(undefined),
} as unknown as IGatewayService;
const service = new ReadStateService(repository, gatewayService);
@@ -34,7 +34,6 @@ export class ReadStateService {
undefined,
manual ?? false,
);
await this.invalidatePushBadgeCount(userId);
if (!silent) {
await this.clearPushChannelNotifications({userId, channelId, messageId});
}
@@ -115,7 +114,6 @@ export class ReadStateService {
try {
const updatedReadStates = await this.repository.bulkAckMessages(userId, readStates);
const readStatesByChannel = new Map(updatedReadStates.map((readState) => [readState.channelId, readState]));
await this.invalidatePushBadgeCount(userId);
await Promise.all(
readStates.map(({channelId, messageId}) =>
Promise.all([
@@ -145,7 +143,6 @@ export class ReadStateService {
async deleteReadState({userId, channelId}: {userId: UserID; channelId: ChannelID}): Promise<void> {
await this.repository.deleteReadState(userId, channelId);
await this.invalidatePushBadgeCount(userId);
}
async incrementMentionCount({
@@ -157,11 +154,7 @@ export class ReadStateService {
channelId: ChannelID;
messageId: MessageID;
}): Promise<void> {
const readState = await this.repository.incrementReadStateMentions(userId, channelId, messageId, 1);
if (readState == null) {
return;
}
await this.invalidatePushBadgeCount(userId);
await this.repository.incrementReadStateMentions(userId, channelId, messageId, 1);
}
async bulkIncrementMentionCounts(
@@ -175,15 +168,7 @@ export class ReadStateService {
return;
}
try {
const appliedUpdates = await this.repository.bulkIncrementMentionCounts(updates);
const uniqueUserIds = Array.from(new Set(appliedUpdates.map((update) => update.userId)));
if (uniqueUserIds.length === 0) {
return;
}
await this.gatewayService.invalidatePushBadgeCounts({userIds: uniqueUserIds}).catch((error) => {
Logger.error({userCount: uniqueUserIds.length, error}, 'Failed to invalidate push badge counts');
return null;
});
await this.repository.bulkIncrementMentionCounts(updates);
} catch (error) {
Logger.error({error}, 'Bulk increment mention counts failed');
throw error;
@@ -196,13 +181,6 @@ export class ReadStateService {
await this.dispatchPinsAck({userId, channelId, timestamp});
}
private async invalidatePushBadgeCount(userId: UserID): Promise<void> {
await this.gatewayService.invalidatePushBadgeCount({userId}).catch((error) => {
Logger.error({userId: userId.toString(), error}, 'Failed to invalidate push badge count');
return null;
});
}
private async dispatchMessageAck(params: {
userId: UserID;
channelId: ChannelID;
+1 -9
View File
@@ -99,7 +99,6 @@ 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';
@@ -433,13 +432,6 @@ 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,
@@ -646,7 +638,7 @@ export class RpcService {
};
}
case 'get_push_service_delivery_config': {
const config = await this.instanceConfigRepository.getPushServiceDeliveryConfig();
const config = await this.instanceConfigRepository.getLegacyPushServiceDeliveryWire();
return {
type: 'get_push_service_delivery_config',
data: {config},
@@ -795,12 +795,6 @@ export class NoopGatewayService extends IGatewayService {
async dispatchPresence(_params: {userId: UserID; event: GatewayDispatchEvent; data: unknown}): Promise<void> {}
async invalidatePushBadgeCount(_params: {userId: UserID}): Promise<void> {}
async invalidatePushBadgeCounts(_params: {userIds: Array<UserID>}): Promise<void> {}
async invalidatePushSubscriptions(_params: {userId: UserID}): Promise<void> {}
async clearPushChannelNotifications(_params: {
userId: UserID;
channelId: ChannelID;
@@ -392,7 +392,6 @@ export class UserContentService {
provider_environment: null,
};
const subscription = await this.storeWebPushSubscription(data, originKind ?? null, installedApp === true);
await this.gatewayService.invalidatePushSubscriptions({userId});
return subscription;
}
@@ -468,7 +467,6 @@ export class UserContentService {
async deletePushSubscription(userId: UserID, subscriptionId: string): Promise<void> {
await this.userRepository.deletePushSubscription(userId, subscriptionId);
await this.gatewayService.invalidatePushSubscriptions({userId});
}
async rotatePushSubscription(params: {
@@ -504,7 +502,6 @@ export class UserContentService {
provider_environment: null,
};
const subscription = await this.storeWebPushSubscription(data, originKind ?? null, installedApp === true);
await this.gatewayService.invalidatePushSubscriptions({userId});
return subscription;
}
@@ -530,7 +527,6 @@ export class UserContentService {
provider_environment: providerEnvironment,
};
const subscription = await this.userRepository.createPushSubscription(data);
await this.gatewayService.invalidatePushSubscriptions({userId});
return subscription;
}
@@ -572,7 +572,7 @@ export const SelfHostedSetupWizardGate = observer(() => {
setSingleCommunityEnabled(next.policy.single_community_enabled);
setDirectMessagesDisabled(next.policy.direct_messages_disabled);
setPremiumMode(next.policy.premium_mode);
setPushRelayConsentAccepted(next.push_service_delivery.relay_consent_accepted);
setPushRelayConsentAccepted(next.push_relay.relay_consent_accepted);
setServiceSelection({
gif: next.policy.services_resolved.gif_enabled,
youtube: next.policy.services_resolved.youtube_enabled,
@@ -771,8 +771,8 @@ export const SelfHostedSetupWizardGate = observer(() => {
const nextConfig = await updateInstanceConfig({
integrations: buildIntegrationsPatch(integrationDraft),
media: buildMediaPatch(mediaExpiryDraft),
push_service_delivery:
config.push_service_delivery.relay_consent_accepted === pushRelayConsentAccepted
push_relay:
config.push_relay.relay_consent_accepted === pushRelayConsentAccepted
? undefined
: {relay_consent_accepted: pushRelayConsentAccepted},
registration: {mode: registrationMode},
@@ -57,6 +57,7 @@ import type {GuildMember} from '@app/features/member/models/GuildMember';
import GuildMembers from '@app/features/member/state/GuildMembers';
import type {SearchContext} from '@app/features/member/state/MemberSearch';
import * as HighlightCommands from '@app/features/messaging/commands/HighlightCommands';
import * as MessageCommands from '@app/features/messaging/commands/MessageCommands';
import * as ReactionCommands from '@app/features/messaging/commands/ReactionCommands';
import Messages from '@app/features/messaging/state/MessagingMessages';
import {
@@ -77,6 +78,7 @@ import {
} from '@app/features/messaging/utils/AutocompleteOptionBuilders';
import {isAutocompleteTriggerAllowed, type TriggerType} from '@app/features/messaging/utils/AutocompleteTriggerPolicy';
import {toReactionEmoji} from '@app/features/messaging/utils/MessageReactionUtils';
import {getReactionShorthandTargetId} from '@app/features/messaging/utils/ReactionShorthandUtils';
import {
type AutocompleteTrigger,
detectAutocompleteTrigger,
@@ -705,10 +707,10 @@ export function useLexicalAutocomplete({
const matchStart = getComposerAutocompleteReplacementStart(currentTextUpToCursor, trigger.type, trigger.match);
if (trigger.type === 'emojiReaction' && isEmoji(option)) {
if (channel != null) {
const messages = Messages.getMessages(channel.id).toArray();
const mostRecent = messages[messages.length - 1];
if (mostRecent != null) {
ReactionCommands.addReaction(i18n, channel.id, mostRecent.id, toReactionEmoji(option.emoji));
const targetId = getReactionShorthandTargetId(channel.id);
if (targetId !== null) {
ReactionCommands.addReaction(i18n, channel.id, targetId, toReactionEmoji(option.emoji));
MessageCommands.stopReply(channel.id);
}
}
handle.clear();
@@ -20,6 +20,7 @@ import GuildMembers from '@app/features/member/state/GuildMembers';
import MemberSidebar from '@app/features/member/state/MemberSidebar';
import * as DraftCommands from '@app/features/messaging/commands/DraftCommands';
import * as MessageCommands from '@app/features/messaging/commands/MessageCommands';
import * as ReactionCommands from '@app/features/messaging/commands/ReactionCommands';
import type {Message} from '@app/features/messaging/models/MessagingMessage';
import type {MentionConfirmationInfo, MentionType} from '@app/features/messaging/state/MentionConfirmationStateMachine';
import Messages from '@app/features/messaging/state/MessagingMessages';
@@ -29,6 +30,10 @@ import {
isAttachmentOnlyMessage,
} from '@app/features/messaging/utils/MessageEditContentUtils';
import {canSubmitMessage, hasVisibleMessageContent} from '@app/features/messaging/utils/MessageRequestUtils';
import {
getReactionShorthandTargetId,
parseReactionShorthand,
} from '@app/features/messaging/utils/ReactionShorthandUtils';
import * as ReplaceCommandUtils from '@app/features/messaging/utils/ReplaceCommandUtils';
import {resolveTypedEmojiShortcodes} from '@app/features/messaging/utils/TypedEmojiShortcodeUtils';
import Permission from '@app/features/permissions/state/Permission';
@@ -458,6 +463,20 @@ export const useTextareaSubmit = ({
parsedCommand = lexicalCommand.command;
}
const replaceCommand = ReplaceCommandUtils.parseReplaceCommand(actualContent);
const reactionShorthand =
editingMessage === null && uploadAttachmentsLength === 0 && !hasPendingSticker
? parseReactionShorthand(resolvedContent)
: null;
const reactionTargetId = reactionShorthand === null ? null : getReactionShorthandTargetId(channelId);
if (reactionShorthand !== null && reactionTargetId !== null) {
ReactionCommands.addReaction(i18n, channelId, reactionTargetId, reactionShorthand);
setValue('');
clearSegments();
DraftCommands.deleteDraft(channelId);
TypingUtils.clear(channelId);
MessageCommands.stopReply(channelId);
return;
}
if (
shouldBlockSubmissionForSlowmode(
isSlowmodeActive,
@@ -0,0 +1,41 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import UnicodeEmojis from '@app/features/expressions/utils/UnicodeEmojis';
import MessageReply from '@app/features/messaging/state/MessageReply';
import Messages from '@app/features/messaging/state/MessagingMessages';
import {type ReactionEmoji, toReactionEmoji} from '@app/features/messaging/utils/MessageReactionUtils';
const REACTION_SHORTHAND_PATTERN = /^\+(\S+)$/u;
const CUSTOM_EMOJI_MARKDOWN_PATTERN = /^<(a)?:([a-zA-Z0-9_+-]{2,}):(\d+)>$/;
const SHORTCODE_PATTERN = /^:([^\s:]+):$/;
export function parseReactionShorthand(content: string): ReactionEmoji | null {
const match = REACTION_SHORTHAND_PATTERN.exec(content.trim());
if (match === null) {
return null;
}
const token = match[1];
const customMatch = CUSTOM_EMOJI_MARKDOWN_PATTERN.exec(token);
if (customMatch !== null) {
return {id: customMatch[3], name: customMatch[2], animated: customMatch[1] === 'a'};
}
const shortcodeMatch = SHORTCODE_PATTERN.exec(token);
if (shortcodeMatch !== null) {
const emoji = UnicodeEmojis.findEmojiByShortcodeName(shortcodeMatch[1]);
return emoji === null ? null : toReactionEmoji(emoji);
}
const name = UnicodeEmojis.nameForSurrogate(token, false);
if (name === '') {
return null;
}
return {name: UnicodeEmojis.surrogateForName(name, token)};
}
export function getReactionShorthandTargetId(channelId: string): string | null {
const reply = MessageReply.getReplyingMessage(channelId);
if (reply !== null) {
return reply.messageId;
}
const messages = Messages.getMessages(channelId).toArray();
return messages[messages.length - 1]?.id ?? null;
}
@@ -35,7 +35,7 @@ Missing settings use the defaults documented below. Invalid stored configuration
| sso | [SSO configuration](#sso-configuration-object) object | Single sign-on settings |
| gateway_rollout | [Gateway rollout configuration](#gateway-rollout-configuration-object) object | Gateway admission and dispatch tuning |
| voice_noise_suppression | [voice noise suppression configuration](#voice-noise-suppression-configuration-object) object | Client-side noise suppression rollout |
| push_service_delivery | [push service delivery configuration](#push-service-delivery-configuration-object) object | Push service delivery rollout |
| push_relay | [push relay configuration](#push-relay-configuration-object) object | Operator consent to the Fluxer-run push relay |
| domain_migration | [domain migration configuration](#domain-migration-configuration-object) object | Web domain migration rollout |
| altcha_captcha | [ALTCHA captcha configuration](#altcha-captcha-configuration-object) object | ALTCHA proof-of-work captcha rollout |
| experiment_delivery | [experiment delivery configuration](#experiment-delivery-configuration-object) object | Cadence every client polls the experiments route on |
@@ -123,23 +123,17 @@ Every field is present on read. An absent document or missing field uses the def
How often a client revalidates this rollout is not set here. It is set once for every experiment in the [experiment delivery configuration](#experiment-delivery-configuration-object) below.
## Push service delivery configuration object
## Push relay configuration object
The instance rollout of push service delivery.
The operator's consent to the push relay supplemental privacy notice.
### Structure
| Field | Type | Description |
| --- | --- | --- |
| enabled | boolean | Whether the rollout runs at all (default false) |
| config_version | integer | Revision counter, raised by Fluxer and never accepted from a request |
| rollout_basis_points | integer | Share of accounts the rollout selects, in basis points (0-10000, default 0) |
| rollout_salt | string | Salt of the sampling hash (1-64 printable ASCII characters, default `push-service-delivery-v1`) |
| included_user_ids | array[snowflake] | Accounts the rollout always selects, up to 1000 entries (default empty) |
| excluded_user_ids | array[snowflake] | Accounts the rollout never selects, up to 1000 entries (default empty) |
| relay_consent_accepted | boolean | Whether the operator accepted the push relay supplemental privacy notice (default false) |
| relay_consent_accepted_at | ?string | ISO 8601 timestamp of that acceptance, or null when the notice stands unaccepted (default null) |
| relay_consent_accepted_by | ?snowflake | Admin account that accepted the notice, or null when the notice stands unaccepted (default null) |
| relay_consent_accepted_by | ?snowflake | Account that accepted the notice, or null when the notice stands unaccepted (default null) |
Every field is present on read. An absent document or missing field uses the defaults above.
@@ -608,7 +602,7 @@ The body has one optional object for each section. Fluxer leaves an absent secti
| sso?<sup>1</sup> | object | Every [SSO configuration](#sso-configuration-object) field except `client_secret_set` and `redirect_uri`, plus `client_secret` |
| gateway_rollout? | object | Any subset of the [Gateway rollout configuration](#gateway-rollout-configuration-object) fields, each bound as documented there |
| voice_noise_suppression? | object | Any subset of the [noise suppression](#voice-noise-suppression-configuration-object) fields |
| push_service_delivery? | object | Any subset of the [push service delivery](#push-service-delivery-configuration-object) fields |
| push_relay? | object | `relay_consent_accepted` from the [push relay configuration](#push-relay-configuration-object) |
| domain_migration? | object | Any subset of the [domain migration](#domain-migration-configuration-object) fields |
| altcha_captcha? | object | Any subset of the [ALTCHA captcha](#altcha-captcha-configuration-object) fields |
| experiment_delivery? | object | Any subset of the [experiment delivery](#experiment-delivery-configuration-object) fields |
@@ -624,9 +618,9 @@ The body has one optional object for each section. Fluxer leaves an absent secti
`voice_noise_suppression` takes every [voice noise suppression configuration](#voice-noise-suppression-configuration-object) field except `config_version`, each bound as documented there. Fluxer raises `config_version` by one on each request that supplies at least one of them. A section that is absent, or present with no field set, writes nothing and leaves `config_version` alone.
`push_service_delivery` works the same way, over the [push service delivery configuration](#push-service-delivery-configuration-object) fields and its own `config_version`. It also takes `relay_consent_accepted`. Fluxer sets `relay_consent_accepted_at` and `relay_consent_accepted_by` itself on the request that changes that flag, and accepts neither from a request: turning the flag on stamps the current time and the acting Admin account, and turning it off clears both to null. A request that repeats the flag it already holds leaves the stamp alone.
`push_relay` takes only `relay_consent_accepted`. Fluxer sets `relay_consent_accepted_at` and `relay_consent_accepted_by` itself on the request that changes that flag, and accepts neither from a request. Turning the flag on stamps the current time and the acting account, and turning it off clears both to null. A request that repeats the flag it already holds leaves the stamp alone.
`domain_migration` works the same way, over the [domain migration configuration](#domain-migration-configuration-object) fields and its own `config_version`.
`domain_migration` works the same way as `voice_noise_suppression`, over the [domain migration configuration](#domain-migration-configuration-object) fields and its own `config_version`.
`altcha_captcha` works the same way, over the [ALTCHA captcha configuration](#altcha-captcha-configuration-object) fields and its own `config_version`.
@@ -666,12 +660,12 @@ 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`, `push_service_delivery`, `domain_migration`, `altcha_captcha`, `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`, `push_relay`, `domain_migration`, `altcha_captcha`, `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
Fluxer publishes a `gateway_rollout` change to the Gateway cluster. Premium mode changes affect the limits in force without replacing the saved limit configuration. On self-hosted deployments, `everyone` hides premium-filtered rules. Switching back to `mirror` restores them unless an Admin has replaced the limit configuration in the meantime. Enabling single community mode creates the community when none is designated, with the acting Admin as owner.
Fluxer publishes a `gateway_rollout` change to the Gateway cluster. A `push_relay` change reaches the push service without a restart. Premium mode changes affect the limits in force without replacing the saved limit configuration. On self-hosted deployments, `everyone` hides premium-filtered rules. Switching back to `mirror` restores them unless an Admin has replaced the limit configuration in the meantime. Enabling single community mode creates the community when none is designated, with the acting Admin as owner.
Initial setup completes on the first update that sets `app_public.setup.configured` to true from a session credential whose account holds neither `admin:authenticate` nor the wildcard. That update grants the account the wildcard Admin ACL and marks the deployment as bootstrapped.
@@ -811,17 +811,15 @@ For environment-based configuration, use `FLUXER_AUTH_BLUESKY_ENABLED`, `FLUXER_
| FLUXER_VAPID_PRIVATE_KEY | `CHANGE_ME` | The VAPID private key. Base64url of the 32-byte scalar, and the matching half of the pair |
| FLUXER_VAPID_EMAIL | unset | The VAPID contact address. Compose derives `admin@` followed by `FLUXER_DOMAIN` when it is unset |
The Gateway reads the same names. A malformed pair, or a private key that does not derive the public point, does not stop it. It records the fault in its log at startup and then drops every web push notification.
`FLUXER_GATEWAY_PUSH_ENABLED` is read by the Gateway alone, defaults to `true`, and skips that startup check on the VAPID pair when it is `false`.
`FLUXER_GATEWAY_PUSH_ENABLED` is read by the Gateway alone and defaults to `true`. When it is `false`, the Gateway hands no message or clear notification to `push`.
## Mobile push
`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.
`api` 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 `api`, `gateway`, and `push`.
Default `false`. The APNs switch. Must be set on `api` and `push`.
#### `FLUXER_PUSH_APNS_TEAM_ID`
@@ -849,7 +847,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 `api`, `gateway`, and `push`.
Default `false`. The FCM switch. Must be set on `api` and `push`.
#### `FLUXER_PUSH_FCM_PROJECT_ID`
@@ -907,6 +905,10 @@ No default. Replaces the APNs host in both environments. Set it only for a test
Default `https://fcm.googleapis.com`. The FCM host. Set it only for a proxy or a test double.
#### `FLUXER_PUSH_SERVICE_RELAY_CONSENT_ACCEPTED`
Default `false`. Accepts the push relay supplemental privacy notice for this `push` process. `true`, `1`, or `yes` accepts it, and `false`, `0`, or `no` leaves the decision to the [push relay configuration](/admin-api/instance/#push-relay-configuration-object). Any other value fails startup. When it is `true`, `push` sends through the Fluxer-run relay even while the stored consent is off. Compose does not forward it.
## Payments
Stripe billing, which the shipped stack keeps off. All are optional.
@@ -1711,7 +1713,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. `api`, `gateway`, and `push` all read these names. An override has to reach all three.
Mobile push cannot be configured at all from the example. `api` and `push` both read these names. An override has to reach both.
#### `RUST_LOG`, `LOG_LEVEL` and `LOGGER_LEVEL`
@@ -156,6 +156,7 @@ The run ends by naming the record it wrote and the command that rolls back to it
| `--dry-run` | `-DryRun` | Print the plan. Change nothing |
| `--backup-dir <path>` | `-BackupDir` | Where records go. Default `<dir>/backups` |
| `--no-volume-backup` | `-NoVolumeBackup` | Take the database dump and skip the uploads copy |
| `--no-volume-compression` | `-NoVolumeCompression` | Copy the uploads as a plain `.tar`. Less downtime, more disk |
| `--skip-backup-accept-data-loss` | `-SkipBackupAcceptDataLoss` | Upgrade with no backup at all |
`--dir` and `--ref` mean what they mean on an install, and the installer derives an omitted `--ref` from the `FLUXER_IMAGE_TAG` line in `.env`. Every failure prints a sentence on standard error and exits non-zero, and the script header lists what each code means. A Postgres major version change is exit 3.
@@ -218,7 +219,7 @@ Each upgrade writes one record directory, named `record-` and a UTC stamp, under
| Artefact | What it covers |
| --- | --- |
| `fluxer.dump` | Every account, message, guild and configuration row, in the Postgres custom format |
| `seaweedfs-data.tgz` | Every upload, avatar, report and harvest |
| `seaweedfs-data.tgz` | Every upload, avatar, report and harvest, or `seaweedfs-data.tar` under `--no-volume-compression` |
| `.env` | Every secret the instance was built with |
The stack files sit in the record beside them, so a rollback puts back the exact files the instance was running. The record directory is created `0700` and the `.env` copy inside it is `0600`. Keep records wherever you already keep secrets.
@@ -306,7 +307,7 @@ docker run --rm -v fluxer_seaweedfs-data:/data \
docker compose up -d
```
`tar` extracts over whatever is already on the volume, so the `find` empties it first. In PowerShell write `${PWD}` instead of `$PWD`. Success is an existing attachment URL answering 200 again. The `fluxer_` prefix is the Compose project name, which `docker-compose.yml` sets to `fluxer`.
For a `seaweedfs-data.tar`, write `tar xf` in place of `tar xzf`. `tar` extracts over whatever is already on the volume, so the `find` empties it first. In PowerShell write `${PWD}` instead of `$PWD`. Success is an existing attachment URL answering 200 again. The `fluxer_` prefix is the Compose project name, which `docker-compose.yml` sets to `fluxer`.
## Move to a new Postgres major version
+22 -5
View File
@@ -52,6 +52,7 @@ param(
[switch]$Update,
[switch]$Rollback,
[switch]$NoVolumeBackup,
[switch]$NoVolumeCompression,
[switch]$SkipBackupAcceptDataLoss,
[switch]$Help,
[Parameter(ValueFromRemainingArguments = $true)]
@@ -86,9 +87,8 @@ $FluxerImagesFile = 'images'
$FluxerTagFile = 'image-tag'
$FluxerDumpFile = 'fluxer.dump'
# Free space demanded before a volume copy, as a percentage of the measured volume size. The
# tarball compresses, so this is generous on purpose. A backup that fills the disk it writes to
# takes the instance down with it.
# Free space demanded before a volume copy, as a percentage of the measured volume size. A backup
# that fills the disk it writes to takes the instance down with it.
$FluxerVolumeHeadroomPercent = 110
$FluxerExitUsage = 1
@@ -238,6 +238,7 @@ function Show-FluxerUsage {
Write-FluxerLine ' -Rollback Restore the images and stack files of the last record.'
Write-FluxerLine ' -BackupDir <path> Where records go. Default: the backups folder under -Dir.'
Write-FluxerLine ' -NoVolumeBackup Take the database dump and skip the uploads copy.'
Write-FluxerLine ' -NoVolumeCompression Copy the uploads as a plain .tar. Faster, larger.'
Write-FluxerLine ' -SkipBackupAcceptDataLoss'
Write-FluxerLine ' Upgrade with no backup at all. Losable data is lost.'
Write-FluxerLine ' -Help Print this text.'
@@ -1413,6 +1414,12 @@ function Copy-FluxerVolumes([string]$Record, [string]$Project, [string]$TargetDi
if ($present.Count -eq 0) {
return
}
$tarFlags = 'czf'
$tarExtension = 'tgz'
if ($NoVolumeCompression) {
$tarFlags = 'cf'
$tarExtension = 'tar'
}
Write-FluxerLine 'Stopping the stack for a consistent copy of the uploads.'
if ((Invoke-FluxerDocker @('compose', 'stop')) -ne 0) {
Stop-Fluxer 'docker compose stop failed.' $FluxerExitBackup
@@ -1420,7 +1427,7 @@ function Copy-FluxerVolumes([string]$Record, [string]$Project, [string]$TargetDi
foreach ($volume in $present) {
$full = "${Project}_$volume"
Write-FluxerLine "Copying $full."
$code = Invoke-FluxerDocker @('run', '--rm', '-v', "${full}:/data:ro", '-v', "${Record}:/backup", $FluxerHelperImage, 'tar', 'czf', "/backup/$volume.tgz", '-C', '/data', '.')
$code = Invoke-FluxerDocker @('run', '--rm', '-v', "${full}:/data:ro", '-v', "${Record}:/backup", $FluxerHelperImage, 'tar', $tarFlags, "/backup/$volume.$tarExtension", '-C', '/data', '.')
if ($code -ne 0) {
[void](Invoke-FluxerDocker @('compose', 'up', '-d', '--remove-orphans'))
Stop-Fluxer "Copying $full failed. The stack is started again on the images it was running." $FluxerExitBackup
@@ -1529,7 +1536,11 @@ function Show-FluxerUpdatePlan([string]$TargetDir, [string]$EnvPath, [string]$Ba
} elseif ($NoVolumeBackup) {
Write-FluxerLine ' Backup: the database dump, .env, and the stack files'
} else {
Write-FluxerLine ' Backup: the database dump, the uploads volume, .env, and the stack files'
if ($NoVolumeCompression) {
Write-FluxerLine ' Backup: the database dump, the uploads volume uncompressed, .env, and the stack files'
} else {
Write-FluxerLine ' Backup: the database dump, the uploads volume, .env, and the stack files'
}
Write-FluxerLine ' Downtime: the stack stops for the uploads copy, then again for the recreate'
}
# The dry run downloads into a temporary directory so it can name the files that actually
@@ -1925,6 +1936,12 @@ function Invoke-FluxerInstall {
if ($SkipBackupAcceptDataLoss -and $NoVolumeBackup) {
Stop-Fluxer '-SkipBackupAcceptDataLoss already skips the volume copy.' $FluxerExitUsage
}
if ($NoVolumeCompression -and -not $Update) {
Stop-Fluxer '-NoVolumeCompression belongs to -Update.' $FluxerExitUsage
}
if ($NoVolumeCompression -and ($SkipBackupAcceptDataLoss -or $NoVolumeBackup)) {
Stop-Fluxer '-NoVolumeCompression changes the volume copy, which this run skips.' $FluxerExitUsage
}
Invoke-FluxerPreflight
+27 -4
View File
@@ -82,8 +82,8 @@ FLUXER_TAG_FILE='image-tag'
FLUXER_DUMP_FILE='fluxer.dump'
# Free space demanded before a volume copy, as a percentage of the measured
# volume size. The tarball compresses, so this is generous on purpose. A backup
# that fills the disk it writes to takes the instance down with it.
# volume size. A backup that fills the disk it writes to takes the instance down
# with it.
FLUXER_VOLUME_HEADROOM=110
# The keys .env carries, in the order they are written. The installer iterates
@@ -226,6 +226,7 @@ Options:
--rollback Restore the images and stack files of the last record.
--backup-dir <path> Where records go. Default <dir>/backups.
--no-volume-backup Take the database dump and skip the uploads copy.
--no-volume-compression Copy the uploads as a plain .tar. Faster, larger.
--skip-backup-accept-data-loss
Upgrade with no backup at all. Losable data is lost.
--allow-root Permit running as root.
@@ -298,6 +299,7 @@ opt_update=0
opt_rollback=0
opt_backup_dir=''
opt_no_volume_backup=0
opt_no_volume_compression=0
opt_skip_backup=0
opt_allow_root=0
@@ -372,6 +374,10 @@ while [ $# -gt 0 ]; do
opt_no_volume_backup=1
shift
;;
--no-volume-compression)
opt_no_volume_compression=1
shift
;;
--skip-backup-accept-data-loss)
opt_skip_backup=1
shift
@@ -687,6 +693,12 @@ fluxer_validate_options() {
if [ "$opt_skip_backup" -eq 1 ] && [ "$opt_no_volume_backup" -eq 1 ]; then
fluxer_bad_usage '--skip-backup-accept-data-loss already skips the volume copy.'
fi
if [ "$opt_no_volume_compression" -eq 1 ] && [ "$opt_update" -eq 0 ]; then
fluxer_bad_usage '--no-volume-compression belongs to --update.'
fi
if [ "$opt_no_volume_compression" -eq 1 ] && { [ "$opt_skip_backup" -eq 1 ] || [ "$opt_no_volume_backup" -eq 1 ]; }; then
fluxer_bad_usage '--no-volume-compression changes the volume copy, which this run skips.'
fi
}
fluxer_ref_for_tag() {
@@ -1566,6 +1578,13 @@ $(fluxer_volume_error ' ')" ;;
if [ "$fluxer_copy_any" -eq 0 ]; then
return 0
fi
if [ "$opt_no_volume_compression" -eq 1 ]; then
fluxer_tar_flags='cf'
fluxer_tar_ext='tar'
else
fluxer_tar_flags='czf'
fluxer_tar_ext='tgz'
fi
fluxer_say 'Stopping the stack for a consistent copy of the uploads.'
if ! $fluxer_engine compose stop; then
fluxer_fail 7 "$fluxer_engine compose stop failed in $opt_dir."
@@ -1574,7 +1593,7 @@ $(fluxer_volume_error ' ')" ;;
[ -n "$fluxer_volume" ] || continue
fluxer_full="${fluxer_project}_${fluxer_volume}"
fluxer_say "Copying $fluxer_full."
if ! $fluxer_engine run --rm -v "$fluxer_full:/data:ro" -v "$fluxer_record:/backup" "$FLUXER_HELPER_IMAGE" tar czf "/backup/$fluxer_volume.tgz" -C /data .; then
if ! $fluxer_engine run --rm -v "$fluxer_full:/data:ro" -v "$fluxer_record:/backup" "$FLUXER_HELPER_IMAGE" tar "$fluxer_tar_flags" "/backup/$fluxer_volume.$fluxer_tar_ext" -C /data .; then
$fluxer_engine compose up -d --remove-orphans || true
fluxer_fail 7 "Copying $fluxer_full failed. The stack is started again on the images it was running."
fi
@@ -1761,7 +1780,11 @@ fluxer_plan_update() {
elif [ "$opt_no_volume_backup" -eq 1 ]; then
fluxer_say ' backup the database dump, .env, and the stack files'
else
fluxer_say ' backup the database dump, the uploads volume, .env, and the stack files'
if [ "$opt_no_volume_compression" -eq 1 ]; then
fluxer_say ' backup the database dump, the uploads volume uncompressed, .env, and the stack files'
else
fluxer_say ' backup the database dump, the uploads volume, .env, and the stack files'
fi
fluxer_say ' downtime the stack stops for the uploads copy, then again for the recreate'
fi
fluxer_fetch_stack
+1 -3
View File
@@ -2,7 +2,6 @@
{erl_opts, [debug_info, nowarn_deprecated_function]}.
{deps, [
{cowboy, "2.19.0"},
{jose, "1.11.12"},
{ezstd, "1.2.4"},
{enats, "1.2.0"},
{opentelemetry_api, "1.5.0"},
@@ -23,7 +22,6 @@
{override, cowboy, [{deps, [cowlib, ranch]}]},
{override, enats, [{plugins, []}, {project_plugins, []}]},
{add, enats, [{erl_opts, [nowarn_deprecated_catch]}]},
{add, jose, [{erl_opts, [nowarn_deprecated_catch]}]},
{override, enats_msg, [{plugins, []}, {project_plugins, []}]},
{override, ezstd, [{plugins, []}, {project_plugins, []}]}
]}.
@@ -88,7 +86,7 @@
]},
{plt_apps, all_deps},
{plt_extra_apps, [
crypto, enats, inets, jose, public_key, ranch, ssl
crypto, enats, inets, public_key, ranch, ssl
]},
{warnings_file, "dialyzer.ignore-warnings"}
]}.
-3
View File
@@ -9,7 +9,6 @@
"eqwalizer_support"},
0},
{<<"ezstd">>,{pkg,<<"ezstd">>,<<"1.2.4">>},0},
{<<"jose">>,{pkg,<<"jose">>,<<"1.11.12">>},0},
{<<"opentelemetry_api">>,{pkg,<<"opentelemetry_api">>,<<"1.5.0">>},0},
{<<"ranch">>,{pkg,<<"ranch">>,<<"2.3.0">>},1}]}.
[
@@ -19,7 +18,6 @@
{<<"enats">>, <<"D7459C804013CAFA4AF880B18D446C48890D28D372D62AD66C76187E5779248D">>},
{<<"enats_msg">>, <<"50631124F37D88BE76A91A5B96A6565C5981EBF917CD819F0175CA658A966F43">>},
{<<"ezstd">>, <<"7AB3ED4BF5ED93E249C936F457A060CF99487F392A41B2B3FD0F93E2C792A7D6">>},
{<<"jose">>, <<"06E62B467B61D3726CBC19E9B5489F7549C37993DE846DFB3EE8259F9ED208B3">>},
{<<"opentelemetry_api">>, <<"1A676F3E3340CAB81C763E939A42E11A70C22863F645AA06AAFEFC689B5550CF">>},
{<<"ranch">>, <<"7DE7B041A9A6A5091A3AA5898D66C0564BE671D87DB4F9D63B1B5EE775B097DF">>}]},
{pkg_hash_ext,[
@@ -28,7 +26,6 @@
{<<"enats">>, <<"20DEB3CB1D3E960194DF8B136C40D2DB085B485BBA5E493B340AB5F9FD2BED22">>},
{<<"enats_msg">>, <<"C4F2139E5144FABC99AFE01B8B016AD9DA278CDDC60857AC5D5AFB0AD1283534">>},
{<<"ezstd">>, <<"C79A63C8F1706CA5402D4D97347A1934F41FD8DD055C7AF4CA92B9BCBEF6D1C0">>},
{<<"jose">>, <<"31E92B653E9210B696765CDD885437457DE1ADD2A9011D92F8CF63E4641BAB7B">>},
{<<"opentelemetry_api">>, <<"F53EC8A1337AE4A487D43AC89DA4BD3A3C99DDF576655D071DEED8B56A2D5DDA">>},
{<<"ranch">>, <<"6168EC49409D982F7CFBD83DD083144F6CBE67CAA4036551D2F0A3AD67C9D023">>}]}
].
@@ -11,7 +11,6 @@
public_key,
ssl,
inets,
jose,
ranch,
cowboy,
ezstd,
@@ -8,26 +8,11 @@
-spec start(application:start_type(), term()) -> {ok, pid()} | {error, term()}.
start(_StartType, _StartArgs) ->
erlang:system_flag(fullsweep_after, 10),
init_jose(),
init_subsystems(),
{ok, Pid} = fluxer_gateway_sup:start_link(),
{ok, _} = start_cowboy(),
{ok, Pid}.
-spec init_jose() -> ok.
init_jose() ->
case code:ensure_loaded(jose_json_otp) of
{module, jose_json_otp} -> ok;
{error, EnsureErr} -> erlang:error({jose_json_otp_missing, EnsureErr})
end,
application:set_env(jose, json_module, jose_json_otp),
{ok, _} = application:ensure_all_started(jose),
_ = jose:json_module(jose_json_otp),
case jose:json_module() of
jose_json_otp -> ok;
Other -> erlang:error({jose_json_module_registration_failed, Other})
end.
-spec init_subsystems() -> ok.
init_subsystems() ->
_ = fluxer_gateway_env:load(),
@@ -19,7 +19,6 @@
-type gateway_role() :: websocket | sessions | presence | guilds | calls | push | all.
-define(MAX_CLUSTER_STATIC_PEERS, 256).
-define(DEFAULT_MANAGED_RELAY_HOSTS, <<"push.fluxer.com">>).
-spec load() -> config().
load() ->
@@ -33,8 +32,6 @@ env_config() ->
<<"public">> => env_public_config(),
<<"proxy">> => env_proxy_config(),
<<"services">> => env_services_config(),
<<"auth">> => env_auth_config(),
<<"integrations">> => env_integrations_config(),
<<"telemetry">> => env_telemetry_config()
}.
@@ -79,21 +76,9 @@ 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_enrolled_clear_notifications_enabled">> => env_bool(
"FLUXER_GATEWAY_PUSH_ENROLLED_CLEAR_NOTIFICATIONS_ENABLED", true
),
<<"push_endpoint_guard_enabled">> => env_bool(
"FLUXER_GATEWAY_PUSH_ENDPOINT_GUARD_ENABLED", true
),
<<"push_managed_relay_hosts">> => env_binary(
"FLUXER_GATEWAY_PUSH_MANAGED_RELAY_HOSTS", ?DEFAULT_MANAGED_RELAY_HOSTS
),
<<"push_relay_consent_accepted">> => env_bool(
"FLUXER_GATEWAY_PUSH_RELAY_CONSENT_ACCEPTED", false
),
<<"push_outbox_request_timeout_ms">> => env_int(
"FLUXER_GATEWAY_PUSH_OUTBOX_REQUEST_TIMEOUT_MS", 100000
),
@@ -146,56 +131,6 @@ env_nats_config() ->
<<"auth_token">> => env_string("FLUXER_NATS_AUTH_TOKEN", "")
}.
-spec env_auth_config() -> map().
env_auth_config() ->
#{
<<"vapid">> => #{
<<"email">> => env_binary("FLUXER_VAPID_EMAIL", <<>>),
<<"public_key">> => env_optional_binary("FLUXER_VAPID_PUBLIC_KEY"),
<<"private_key">> => env_optional_binary("FLUXER_VAPID_PRIVATE_KEY")
}
}.
-spec env_integrations_config() -> map().
env_integrations_config() ->
#{
<<"push">> => #{
<<"apns">> => env_apns_config(),
<<"fcm">> => env_fcm_config()
}
}.
-spec env_apns_config() -> map().
env_apns_config() ->
#{
<<"enabled">> => env_bool("FLUXER_PUSH_APNS_ENABLED", false),
<<"team_id">> => env_optional_binary("FLUXER_PUSH_APNS_TEAM_ID"),
<<"key_id">> => env_optional_binary("FLUXER_PUSH_APNS_KEY_ID"),
<<"private_key">> => env_optional_binary("FLUXER_PUSH_APNS_PRIVATE_KEY"),
<<"private_key_path">> => env_optional_binary("FLUXER_PUSH_APNS_PRIVATE_KEY_PATH"),
<<"default_environment">> => env_binary(
"FLUXER_PUSH_APNS_DEFAULT_ENVIRONMENT", <<"production">>
),
<<"apps">> => env_json_list("FLUXER_PUSH_APNS_APPS", [])
}.
-spec env_fcm_config() -> map().
env_fcm_config() ->
#{
<<"enabled">> => env_bool("FLUXER_PUSH_FCM_ENABLED", false),
<<"project_id">> => env_optional_binary("FLUXER_PUSH_FCM_PROJECT_ID"),
<<"client_email">> => env_optional_binary("FLUXER_PUSH_FCM_CLIENT_EMAIL"),
<<"private_key">> => env_optional_binary("FLUXER_PUSH_FCM_PRIVATE_KEY"),
<<"private_key_path">> => env_optional_binary("FLUXER_PUSH_FCM_PRIVATE_KEY_PATH"),
<<"service_account_json_path">> => env_optional_binary(
"FLUXER_PUSH_FCM_SERVICE_ACCOUNT_JSON_PATH"
),
<<"token_uri">> => env_binary(
"FLUXER_PUSH_FCM_TOKEN_URI", <<"https://oauth2.googleapis.com/token">>
),
<<"apps">> => env_json_list("FLUXER_PUSH_FCM_APPS", [])
}.
-spec env_telemetry_config() -> map().
env_telemetry_config() ->
#{
@@ -209,10 +144,6 @@ build_config(RawConfig) ->
Internal = get_map(RawConfig, [<<"internal">>]),
Nats = get_map(RawConfig, [<<"services">>, <<"nats">>]),
Telemetry = get_map(RawConfig, [<<"telemetry">>]),
Vapid = get_map(RawConfig, [<<"auth">>, <<"vapid">>]),
Push = get_map(RawConfig, [<<"integrations">>, <<"push">>]),
Apns = get_map(Push, [<<"apns">>]),
Fcm = get_map(Push, [<<"fcm">>]),
Proxy = get_map(RawConfig, [<<"proxy">>]),
Public = get_map(RawConfig, [<<"public">>]),
lists:foldl(fun maps:merge/2, #{}, [
@@ -221,9 +152,6 @@ build_config(RawConfig) ->
build_sharding_config(Service),
build_http_config(Service),
build_cluster_config(Service, Public),
build_vapid_config(Vapid),
build_apns_config(Apns),
build_fcm_config(Fcm),
build_misc_config(Service, Telemetry)
]).
@@ -252,35 +180,13 @@ build_push_config(Service, Public) ->
push_enabled => get_bool(Service, <<"push_enabled">>, true),
push_user_guild_settings_cache_mb =>
get_int(Service, <<"push_user_guild_settings_cache_mb">>, 1024),
push_subscriptions_cache_mb => get_int(
Service, <<"push_subscriptions_cache_mb">>, 1024
),
push_blocked_ids_cache_mb => get_int(Service, <<"push_blocked_ids_cache_mb">>, 1024),
push_badge_counts_cache_mb => get_int(Service, <<"push_badge_counts_cache_mb">>, 256),
push_badge_counts_cache_ttl_seconds =>
get_int(Service, <<"push_badge_counts_cache_ttl_seconds">>, 60),
static_cdn_endpoint => public_endpoint(
get_binary(Service, <<"static_cdn_endpoint">>, <<"http://localhost:8088">>), 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_clear_notifications_enabled => get_bool(
Service, <<"push_clear_notifications_enabled">>, true
),
push_enrolled_clear_notifications_enabled => get_bool(
Service, <<"push_enrolled_clear_notifications_enabled">>, true
),
push_endpoint_guard_enabled => get_bool(
Service, <<"push_endpoint_guard_enabled">>, true
),
push_managed_relay_hosts => parse_host_list(
get_binary(Service, <<"push_managed_relay_hosts">>, ?DEFAULT_MANAGED_RELAY_HOSTS)
),
push_relay_consent_accepted => get_bool(
Service, <<"push_relay_consent_accepted">>, false
),
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(
@@ -309,14 +215,8 @@ build_sharding_config(Service) ->
-spec build_http_config(map()) -> config().
build_http_config(Service) ->
#{
gateway_http_push_connect_timeout_ms =>
get_int(Service, <<"gateway_http_push_connect_timeout_ms">>, 3000),
gateway_http_push_recv_timeout_ms =>
get_int(Service, <<"gateway_http_push_recv_timeout_ms">>, 5000),
gateway_http_rpc_max_concurrency =>
get_int(Service, <<"gateway_http_rpc_max_concurrency">>, 512),
gateway_http_push_max_concurrency =>
get_int(Service, <<"gateway_http_push_max_concurrency">>, 256),
gateway_http_failure_threshold => get_int(
Service, <<"gateway_http_failure_threshold">>, 6
),
@@ -347,44 +247,6 @@ build_cluster_config(Service, Public) ->
)
}.
-spec build_vapid_config(map()) -> config().
build_vapid_config(Vapid) ->
#{
vapid_email => get_binary(Vapid, <<"email">>, <<>>),
vapid_public_key => get_optional_binary(Vapid, <<"public_key">>),
vapid_private_key => get_optional_binary(Vapid, <<"private_key">>)
}.
-spec build_apns_config(map()) -> config().
build_apns_config(Apns) ->
#{
apns_enabled => get_bool(Apns, <<"enabled">>, false),
apns_team_id => get_optional_binary(Apns, <<"team_id">>),
apns_key_id => get_optional_binary(Apns, <<"key_id">>),
apns_private_key => get_optional_binary(Apns, <<"private_key">>),
apns_private_key_path => get_optional_binary(Apns, <<"private_key_path">>),
apns_default_environment => get_binary(
Apns, <<"default_environment">>, <<"production">>
),
apns_apps => get_list(Apns, <<"apps">>, [])
}.
-spec build_fcm_config(map()) -> config().
build_fcm_config(Fcm) ->
#{
fcm_enabled => get_bool(Fcm, <<"enabled">>, false),
fcm_project_id => get_optional_binary(Fcm, <<"project_id">>),
fcm_client_email => get_optional_binary(Fcm, <<"client_email">>),
fcm_private_key => get_optional_binary(Fcm, <<"private_key">>),
fcm_private_key_path => get_optional_binary(Fcm, <<"private_key_path">>),
fcm_service_account_json_path => get_optional_binary(
Fcm, <<"service_account_json_path">>
),
fcm_token_uri =>
get_binary(Fcm, <<"token_uri">>, <<"https://oauth2.googleapis.com/token">>),
fcm_apps => get_list(Fcm, <<"apps">>, [])
}.
-spec build_misc_config(map(), map()) -> config().
build_misc_config(Service, Telemetry) ->
#{
@@ -474,41 +336,6 @@ env_bool(Name, Default) ->
_ -> Default
end.
-spec env_json_list(string(), list()) -> list().
env_json_list(Name, Default) ->
parse_env_json_list(os:getenv(Name), Default).
-spec parse_env_json_list(false | string(), list()) -> list().
parse_env_json_list(false, Default) ->
Default;
parse_env_json_list("", Default) ->
Default;
parse_env_json_list(Value, Default) ->
parse_json_list(Value, Default).
-spec parse_json_list(string(), list()) -> list().
parse_json_list(Value, Default) ->
case unicode:characters_to_binary(Value) of
Encoded when is_binary(Encoded) -> decode_json_list(Encoded, Default);
_ -> Default
end.
-spec decode_json_list(binary(), list()) -> list().
decode_json_list(Value, Default) ->
try json:decode(Value) of
Decoded when is_list(Decoded) -> Decoded;
_ -> Default
catch
_:_ -> Default
end.
-spec get_list(map(), binary(), list()) -> list().
get_list(Map, Key, Default) when is_list(Default) ->
case get_value(Map, Key) of
V when is_list(V) -> V;
_ -> Default
end.
-spec get_int(map(), binary(), integer()) -> integer().
get_int(Map, Key, Default) when is_integer(Default) -> to_int(get_value(Map, Key), Default).
-spec get_optional_int(map(), binary()) -> integer() | undefined.
@@ -603,21 +430,6 @@ to_binary(Str, _) when is_list(Str) -> list_to_binary(config_char_list(Str));
to_binary(Atom, _) when is_atom(Atom) -> list_to_binary(atom_to_list(Atom));
to_binary(_, Default) -> Default.
-spec parse_host_list(binary()) -> [binary()].
parse_host_list(Bin) ->
[
lower_ascii(Host)
|| Host <- binary:split(Bin, [<<",">>, <<" ">>, <<"\t">>], [global, trim_all])
].
-spec lower_ascii(binary()) -> binary().
lower_ascii(Bin) ->
<<<<(lower_byte(Byte))>> || <<Byte>> <= Bin>>.
-spec lower_byte(byte()) -> byte().
lower_byte(Byte) when Byte >= $A, Byte =< $Z -> Byte + 32;
lower_byte(Byte) -> Byte.
-spec parse_node_list(binary() | undefined) -> [node()].
parse_node_list(undefined) ->
[];
@@ -38,8 +38,7 @@ 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(push_delivery_config, push_delivery_config)
child_spec(gateway_rollout_config, gateway_rollout_config)
] ++ cluster_children() ++
[
child_spec(gateway_dispatch_relay, gateway_dispatch_relay),
@@ -99,7 +98,6 @@ role_specs(calls, Role) ->
[child_spec(call_manager, call_manager)] ++ calls_voice_state_counts_sync_children(Role);
role_specs(push, _Role) ->
[
child_spec(push_dispatcher, push_dispatcher),
child_spec(push_outbox, push_outbox),
child_spec(push, push)
].
@@ -7,7 +7,6 @@
-export([start_link/0, request/5, request/6]).
-export([init/1, handle_call/3, handle_cast/2, handle_info/2, terminate/2, code_change/3]).
-export([pick_sharded_profile/1, ensure_started/0, cleanup_max_age_ms/0]).
-export([push_max_concurrency/0]).
-define(SERVER, ?MODULE).
-define(CIRCUIT_TABLE, gateway_http_circuit_breaker).
@@ -16,11 +15,8 @@
-define(DEFAULT_RPC_CONNECT_TIMEOUT_MS, 5000).
-define(DEFAULT_RPC_RECV_TIMEOUT_MS, 30000).
-define(DEFAULT_PUSH_CONNECT_TIMEOUT_MS, 3000).
-define(DEFAULT_PUSH_RECV_TIMEOUT_MS, 5000).
-define(DEFAULT_RPC_MAX_CONCURRENCY, 512).
-define(DEFAULT_PUSH_MAX_CONCURRENCY, 256).
-define(DEFAULT_FAILURE_THRESHOLD, 500).
-define(DEFAULT_RECOVERY_TIMEOUT_MS, 5000).
@@ -28,9 +24,8 @@
-define(DEFAULT_CLEANUP_MAX_AGE_MS, 300000).
-define(RPC_PROFILE_SHARDS, 8).
-define(PUSH_PROFILE_SHARDS, 4).
-type workload() :: rpc | push.
-type workload() :: rpc.
-type method() :: get | post | put | patch | delete | head | options.
-type request_headers() :: [{binary() | string(), binary() | string()}].
-type request_options() :: #{
@@ -97,10 +92,7 @@ request(Workload, Method, Url, Headers, Body, Opts) when is_map(Opts) ->
-spec pick_sharded_profile(workload()) -> atom().
pick_sharded_profile(rpc) ->
Idx = erlang:phash2(self(), ?RPC_PROFILE_SHARDS),
sharded_profile_name(rpc, Idx);
pick_sharded_profile(push) ->
Idx = erlang:phash2(self(), ?PUSH_PROFILE_SHARDS),
sharded_profile_name(push, Idx).
sharded_profile_name(rpc, Idx).
-spec ensure_started() -> ok.
ensure_started() ->
@@ -113,10 +105,6 @@ ensure_started() ->
cleanup_max_age_ms() ->
get_int_or_default(gateway_http_cleanup_max_age_ms, ?DEFAULT_CLEANUP_MAX_AGE_MS).
-spec push_max_concurrency() -> pos_integer().
push_max_concurrency() ->
get_int_or_default(gateway_http_push_max_concurrency, ?DEFAULT_PUSH_MAX_CONCURRENCY).
-spec init([]) -> {ok, state()}.
init([]) ->
process_flag(trap_exit, true),
@@ -125,7 +113,6 @@ init([]) ->
ensure_table(?CIRCUIT_WINDOW_TABLE),
ensure_table(?INFLIGHT_TABLE),
ok = ensure_sharded_profiles(rpc, ?RPC_PROFILE_SHARDS),
ok = ensure_sharded_profiles(push, ?PUSH_PROFILE_SHARDS),
schedule_cleanup(),
{ok, #{}}.
@@ -315,9 +302,7 @@ ensure_sharded_profiles(Workload, ShardCount) ->
-spec sharded_profile_name(workload(), non_neg_integer()) -> atom().
sharded_profile_name(rpc, Idx) ->
list_to_atom("gateway_http_rpc_profile_" ++ integer_to_list(Idx));
sharded_profile_name(push, Idx) ->
list_to_atom("gateway_http_push_profile_" ++ integer_to_list(Idx)).
list_to_atom("gateway_http_rpc_profile_" ++ integer_to_list(Idx)).
-spec ensure_httpc_profile(atom(), workload()) -> ok.
ensure_httpc_profile(Profile, Workload) ->
@@ -339,9 +324,7 @@ workload_httpc_options(rpc) ->
{max_keep_alive_length, 128},
{max_pipeline_length, 0},
{keep_alive_timeout, 120000}
];
workload_httpc_options(push) ->
[{max_sessions, 512}, {max_keep_alive_length, 128}].
].
-spec merged_workload_options(workload(), request_options()) -> request_options().
merged_workload_options(Workload, Opts) ->
@@ -357,16 +340,6 @@ default_options(rpc) ->
gateway_http_rpc_max_concurrency,
?DEFAULT_RPC_MAX_CONCURRENCY,
<<"application/json">>
);
default_options(push) ->
default_options(
gateway_http_push_connect_timeout_ms,
?DEFAULT_PUSH_CONNECT_TIMEOUT_MS,
gateway_http_push_recv_timeout_ms,
?DEFAULT_PUSH_RECV_TIMEOUT_MS,
gateway_http_push_max_concurrency,
?DEFAULT_PUSH_MAX_CONCURRENCY,
<<"application/octet-stream">>
).
-spec default_options(atom(), integer(), atom(), integer(), atom(), integer(), binary()) ->
@@ -12,7 +12,7 @@
]).
-export_type([workload/0, method/0, request_headers/0, request_options/0, response/0]).
-type workload() :: rpc | push.
-type workload() :: rpc.
-type method() :: get | post | put | patch | delete | head | options.
-type request_headers() :: [{binary() | string(), binary() | string()}].
-type request_options() :: #{
@@ -187,18 +187,7 @@ retry_update_counter(Table, Key, Op) ->
-spec is_countable_circuit_failure({atom(), binary()}, response()) -> boolean().
is_countable_circuit_failure({rpc, _Host}, Result) ->
is_countable_failure_with_transport(Result);
is_countable_circuit_failure({_Workload, _Host}, Result) ->
is_countable_failure_without_transport(Result).
-spec is_countable_failure_without_transport(response()) -> boolean().
is_countable_failure_without_transport({error, nxdomain}) -> false;
is_countable_failure_without_transport({error, {failed_connect, _}}) -> false;
is_countable_failure_without_transport({error, timeout}) -> false;
is_countable_failure_without_transport({error, {timeout, _}}) -> false;
is_countable_failure_without_transport({error, _}) -> true;
is_countable_failure_without_transport({ok, StatusCode, _, _}) when StatusCode >= 500 -> true;
is_countable_failure_without_transport(_) -> false.
is_countable_failure_with_transport(Result).
-spec is_countable_failure_with_transport(response()) -> boolean().
is_countable_failure_with_transport({error, nxdomain}) -> true;
@@ -333,19 +322,6 @@ safe_delete(Table, Key) ->
-ifdef(TEST).
-include_lib("eunit/include/eunit.hrl").
circuit_ignores_transport_failures_for_non_rpc_workloads_test() ->
?assertEqual(false, is_countable_circuit_failure({push, <<"h">>}, {error, nxdomain})),
?assertEqual(
false, is_countable_circuit_failure({push, <<"h">>}, {error, {failed_connect, []}})
),
?assertEqual(false, is_countable_circuit_failure({push, <<"h">>}, {error, timeout})),
?assertEqual(
false, is_countable_circuit_failure({push, <<"h">>}, {error, {timeout, connect}})
),
?assertEqual(true, is_countable_circuit_failure({push, <<"h">>}, {error, closed})),
?assertEqual(true, is_countable_circuit_failure({push, <<"h">>}, {ok, 500, [], <<>>})),
?assertEqual(false, is_countable_circuit_failure({push, <<"h">>}, {ok, 200, [], <<>>})).
circuit_counts_transport_failures_for_rpc_test() ->
?assertEqual(true, is_countable_circuit_failure({rpc, <<"h">>}, {error, nxdomain})),
?assertEqual(
@@ -388,7 +388,6 @@ handle_msg(
-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.
@@ -36,15 +36,9 @@ handle_dispatch(#{<<"user_id">> := UserIdBin, <<"event">> := Event, <<"data">> :
-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) ->
handle_unreachable_dispatch(_EventAtom, _UserId, _Data) ->
gateway_rpc_error:raise(<<"presence_dispatch_error">>).
-spec dispatch_event_atom_or_error(term()) -> atom().
@@ -116,4 +116,31 @@ invalidate_badge_counts_rejects_invalid_snowflakes_test() ->
invalidate_badge_counts_accepts_an_empty_batch_test() ->
?assert(execute_method(<<"push.invalidate_badge_counts">>, #{<<"user_ids">> => []})).
invalidate_badge_counts_accepts_a_batch_test() ->
?assert(
execute_method(<<"push.invalidate_badge_counts">>, #{
<<"user_ids">> => [<<"1001">>, <<"1002">>]
})
).
invalidate_badge_count_accepts_a_user_test() ->
?assert(execute_method(<<"push.invalidate_badge_count">>, #{<<"user_id">> => <<"1001">>})).
invalidate_badge_count_rejects_an_invalid_snowflake_test() ->
?assertError(
{validation, _},
execute_method(<<"push.invalidate_badge_count">>, #{<<"user_id">> => <<"nope">>})
).
invalidate_subscriptions_accepts_a_user_test() ->
?assert(
execute_method(<<"push.invalidate_subscriptions">>, #{<<"user_id">> => <<"1001">>})
).
invalidate_subscriptions_rejects_an_invalid_snowflake_test() ->
?assertError(
{validation, _},
execute_method(<<"push.invalidate_subscriptions">>, #{<<"user_id">> => <<"nope">>})
).
-endif.
+49 -117
View File
@@ -44,8 +44,6 @@ render_metrics() ->
render_gateway_gauges(),
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()
].
@@ -254,97 +252,11 @@ count_registry_prefix(Prefix) ->
error:badarg -> 0
end.
-spec render_push_dispatcher_stats() -> iolist().
render_push_dispatcher_stats() ->
Stats = safe_apply_map(fun push_dispatcher:stats/0),
case map_size(Stats) of
0 ->
[];
_ ->
Queued = maps:get(queued, Stats, 0),
Inflight = maps:get(inflight, Stats, 0),
[
format_metric(
<<"fluxer_gateway_push_dispatcher_queued">>,
<<"gauge">>,
<<"Push dispatcher queued jobs">>,
integer_to_binary(Queued)
),
format_metric(
<<"fluxer_gateway_push_dispatcher_inflight">>,
<<"gauge">>,
<<"Push dispatcher in-flight jobs">>,
integer_to_binary(Inflight)
)
]
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)].
[render_push_outbox_queue_stats(Stats), render_push_outbox_dropped(Stats)].
-spec render_push_outbox_queue_stats(map()) -> iolist().
render_push_outbox_queue_stats(Stats) ->
@@ -393,34 +305,25 @@ render_push_outbox_queue_stats(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 render_push_outbox_dropped(map()) -> iolist().
render_push_outbox_dropped(Stats) ->
format_labeled_series(
<<"fluxer_gateway_push_outbox_dropped_total">>,
<<"counter">>,
<<"Push jobs dropped undelivered by kind and reason">>,
[
{push_outbox_dropped_label(Kind, Reason), integer_to_binary(Count)}
|| {{Kind, Reason}, Count} <- lists:sort(maps:to_list(maps:get(dropped, Stats, #{}))),
is_atom(Kind),
is_atom(Reason),
is_integer(Count)
]
).
-spec push_outbox_dropped_label(atom(), atom()) -> binary().
push_outbox_dropped_label(Kind, Reason) ->
<<"kind=\"", (atom_to_binary(Kind))/binary, "\",reason=\"", (atom_to_binary(Reason))/binary,
"\"">>.
-spec gate_counter(atom(), map()) -> binary().
gate_counter(Key, Counters) ->
@@ -519,3 +422,32 @@ format_labeled_series(Name, Type, Help, LabelValues) ->
|| {Label, Value} <- LabelValues
]
].
-ifdef(TEST).
-include_lib("eunit/include/eunit.hrl").
push_outbox_drops_render_one_series_per_kind_and_reason_test() ->
Rendered = iolist_to_binary(
render_push_outbox_dropped(#{
dropped => #{{message, expired} => 3, {clear, outbox_unavailable} => 1}
})
),
?assertNotEqual(
nomatch,
binary:match(
Rendered,
<<"fluxer_gateway_push_outbox_dropped_total{kind=\"message\",reason=\"expired\"} 3\n">>
)
),
?assertNotEqual(
nomatch,
binary:match(
Rendered,
<<"fluxer_gateway_push_outbox_dropped_total{kind=\"clear\",reason=\"outbox_unavailable\"} 1\n">>
)
).
push_outbox_without_drops_renders_no_dropped_series_test() ->
?assertEqual([], render_push_outbox_dropped(#{dropped => #{}})).
-endif.
@@ -1115,7 +1115,7 @@ presence_is_offline(UserId, Presences) ->
-spec grace_hold(map(), #{user_id() => boolean()}) -> grace_hold().
grace_hold(Sessions, SessionEligibility) ->
case suppressed_enrolled_sessions(Sessions, SessionEligibility) of
case suppressed_sessions(Sessions, SessionEligibility) of
Held when map_size(Held) =:= 0 ->
none;
Held ->
@@ -1126,8 +1126,8 @@ grace_hold(Sessions, SessionEligibility) ->
held_sessions(none) -> #{};
held_sessions({Held, _RecheckAt}) -> Held.
-spec suppressed_enrolled_sessions(map(), #{user_id() => boolean()}) -> grace_sessions().
suppressed_enrolled_sessions(Sessions, SessionEligibility) ->
-spec suppressed_sessions(map(), #{user_id() => boolean()}) -> grace_sessions().
suppressed_sessions(Sessions, SessionEligibility) ->
maps:fold(
fun(_Sid, Session, Acc) -> maybe_hold_session(Session, SessionEligibility, Acc) end,
#{},
@@ -1153,19 +1153,13 @@ hold_suppressed_session(UserId, Pid, SessionEligibility, Acc) when
->
case maps:get(UserId, SessionEligibility, true) of
false ->
hold_enrolled_session(push_delivery_config:is_enrolled(UserId), UserId, Pid, Acc);
Acc#{UserId => [Pid | maps:get(UserId, 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
+19 -39
View File
@@ -155,14 +155,8 @@ flush_push_buffer(#{push_buffer := Buffer} = State) ->
-spec maybe_update_push_eligibility(state()) -> state().
maybe_update_push_eligibility(State) ->
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).
Eligible = push_eligible(State),
flush_when_eligible(Eligible, record_push_eligibility(Eligible, State)).
-spec flush_when_eligible(boolean(), state()) -> state().
flush_when_eligible(Eligible, State) ->
@@ -171,13 +165,6 @@ flush_when_eligible(Eligible, 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};
@@ -253,13 +240,7 @@ route_push_notification(Params, State) ->
-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).
no_session_holds_push(maps:get(sessions, State, #{})).
-spec build_push_create_params(user_id(), map()) -> map() | undefined.
build_push_create_params(UserId, Data) ->
@@ -406,17 +387,6 @@ build_buffer_entry(ChannelId, MessageId, Params) when
build_buffer_entry(_, _, _) ->
undefined.
-spec is_push_eligible(map()) -> boolean().
is_push_eligible(Sessions) ->
case map_size(Sessions) of
0 -> true;
_ -> all_sessions_afk(Sessions)
end.
-spec all_sessions_afk(map()) -> boolean().
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)).
@@ -439,12 +409,22 @@ parse_snowflake(FieldName, Value) ->
-ifdef(TEST).
-include_lib("eunit/include/eunit.hrl").
is_push_eligible_test() ->
?assertEqual(true, is_push_eligible(#{})),
?assertEqual(false, is_push_eligible(#{<<"s1">> => #{mobile => true, afk => false}})),
?assertEqual(true, is_push_eligible(#{<<"s1">> => #{mobile => true, afk => true}})),
?assertEqual(true, is_push_eligible(#{<<"s1">> => #{mobile => false, afk => true}})),
?assertEqual(false, is_push_eligible(#{<<"s1">> => #{mobile => false, afk => false}})).
no_session_holds_push_test() ->
?assertEqual(true, no_session_holds_push(#{})),
?assertEqual(false, no_session_holds_push(#{<<"s1">> => #{mobile => true, afk => false}})),
?assertEqual(true, no_session_holds_push(#{<<"s1">> => #{mobile => true, afk => true}})),
?assertEqual(true, no_session_holds_push(#{<<"s1">> => #{mobile => false, afk => true}})),
?assertEqual(false, no_session_holds_push(#{<<"s1">> => #{mobile => false, afk => false}})),
?assertEqual(
true, no_session_holds_push(#{<<"s1">> => #{afk => false, status => offline}})
),
?assertEqual(
false,
no_session_holds_push(#{
<<"s1">> => #{afk => true},
<<"s2">> => #{afk => false, status => online}
})
).
custom_status_comparator_test() ->
Expected = #{
+114 -552
View File
@@ -12,14 +12,12 @@
sync_user_guild_settings_local/3,
sync_user_blocked_ids/2,
sync_user_blocked_ids_local/2,
invalidate_user_subscriptions/1,
invalidate_user_subscriptions_local/1,
invalidate_user_badge_count/1,
invalidate_user_badge_count_local/1,
invalidate_user_badge_counts_local/1,
clear_channel_notifications/3
]).
-export([get_cache_stats/0, delivery_gate_counters/0]).
-export([get_cache_stats/0]).
-export([push_owner_key/1]).
-define(EVICT_INTERVAL_MS, 60000).
@@ -31,27 +29,6 @@
-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_DISPATCH_DROPPED, push_loss_dispatch_dropped).
-define(CNT_DISPATCH_DROPPED_USERS, push_loss_dispatch_dropped_users).
-define(CNT_CLEAR_DROPPED, push_loss_clear_dropped).
-define(CNT_CLEAR_RETRIED, push_clear_retried).
-define(CLEAR_RETRY_ATTEMPTS, 3).
-define(CLEAR_RETRY_BASE_MS, 250).
-define(CNT_QUEUE_FULL, push_loss_queue_full).
-define(CNT_INVALID_JOB, push_loss_invalid_job).
-define(CNT_ENQUEUE_TIMEOUT, push_loss_enqueue_timeout).
-define(CNT_ENQUEUE_FAILED, push_loss_enqueue_failed).
-define(CNT_JOB_CRASHED, push_loss_job_crashed).
-define(CNT_WORKER_DIED, push_loss_worker_died).
-define(CNT_DISPATCHER_RESTARTS, push_dispatcher_restarts).
-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_FETCH_RPCS, 8).
-define(DEFAULT_FETCH_USERS, 2000).
-define(MAX_FETCH_USERS, 5000).
@@ -60,12 +37,8 @@
-define(MAX_FETCH_CHUNK, 1000).
-type state() :: #{
badge_counts_ttl_seconds := non_neg_integer(),
max_entries := non_neg_integer()
}.
-type worker_state() :: #{
badge_counts_ttl_seconds := non_neg_integer()
}.
-spec start_link() -> {ok, pid()} | {error, term()} | ignore.
start_link() ->
@@ -76,22 +49,15 @@ init([]) ->
erlang:process_flag(fullsweep_after, 10),
push_ets_cache:init(),
init_worker_counter(),
PushEnabled = env_boolean(push_enabled),
maybe_warn_vapid_misconfigured(PushEnabled),
case PushEnabled of
true ->
BcTtl = env_non_neg_integer(push_badge_counts_cache_ttl_seconds, 0),
schedule_eviction(),
{ok, #{
badge_counts_ttl_seconds => BcTtl,
max_entries => ?DEFAULT_MAX_ENTRIES
}};
false ->
{ok, #{
badge_counts_ttl_seconds => 0,
max_entries => ?DEFAULT_MAX_ENTRIES
}}
end.
maybe_schedule_eviction(env_boolean(push_enabled)),
{ok, #{max_entries => ?DEFAULT_MAX_ENTRIES}}.
-spec maybe_schedule_eviction(boolean()) -> ok.
maybe_schedule_eviction(true) ->
_ = schedule_eviction(),
ok;
maybe_schedule_eviction(false) ->
ok.
-spec handle_call(term(), gen_server:from(), state()) -> {reply, term(), state()}.
handle_call(get_cache_stats, _From, State) ->
@@ -109,17 +75,11 @@ handle_cast({sync_user_guild_settings, UserId, GuildId, UserGuildSettings}, Stat
{noreply, State};
handle_cast({sync_user_blocked_ids, UserId, BlockedIds}, State) when is_integer(UserId) ->
handle_sync_user_blocked_ids(UserId, BlockedIds, State);
handle_cast({invalidate_user_subscriptions, UserId}, State) when is_integer(UserId) ->
push_ets_cache:delete_subscriptions(UserId),
{noreply, State};
handle_cast({cache_user_guild_settings, UserId, GuildId, Settings}, State) when
is_integer(UserId), is_integer(GuildId), is_map(Settings)
->
push_ets_cache:put_user_guild_settings(UserId, GuildId, Settings),
{noreply, State};
handle_cast({invalidate_user_badge_count, UserId}, State) when is_integer(UserId) ->
push_ets_cache:delete_badge_count(UserId),
{noreply, State};
handle_cast({clear_channel_notifications, UserId, ChannelId, MessageId}, State) when
is_integer(UserId), is_integer(ChannelId), is_integer(MessageId)
->
@@ -132,14 +92,10 @@ handle_info(evict_caches, State) ->
MaxEntries = maps:get(max_entries, State),
push_ets_cache:evict_tables(#{
user_guild_settings => MaxEntries,
subscriptions => MaxEntries,
blocked_ids => MaxEntries,
badge_counts => MaxEntries
blocked_ids => MaxEntries
}),
schedule_eviction(),
{noreply, State};
handle_info({retry_clear_notifications, UserId, ChannelId, MessageId, Attempt}, State) ->
clear_via_dispatcher(UserId, ChannelId, MessageId, Attempt, State);
handle_info(_Info, State) ->
{noreply, State}.
@@ -182,41 +138,24 @@ sync_user_blocked_ids_local(UserId, BlockedIds) ->
ok
end.
-spec invalidate_user_subscriptions_local(integer()) -> ok.
invalidate_user_subscriptions_local(_UserId) ->
ok.
-spec invalidate_user_badge_count_local(integer()) -> ok.
invalidate_user_badge_count_local(_UserId) ->
ok.
-spec invalidate_user_badge_counts_local([integer()]) -> ok.
invalidate_user_badge_counts_local(_UserIds) ->
ok.
-spec put_blocked_ids_local(integer(), [integer()]) -> ok.
put_blocked_ids_local(UserId, TypedBlockedIds) ->
local_cache_mutation(fun() ->
push_ets_cache:put_blocked_ids(UserId, TypedBlockedIds)
end).
-spec invalidate_user_subscriptions(integer()) -> ok.
invalidate_user_subscriptions(UserId) ->
maybe_cast(UserId, {invalidate_user_subscriptions, UserId}).
-spec invalidate_user_subscriptions_local(integer()) -> ok.
invalidate_user_subscriptions_local(UserId) ->
local_cache_mutation(fun() ->
push_ets_cache:delete_subscriptions(UserId)
end).
-spec invalidate_user_badge_count(integer()) -> ok.
invalidate_user_badge_count(UserId) ->
maybe_cast(UserId, {invalidate_user_badge_count, UserId}).
-spec invalidate_user_badge_count_local(integer()) -> ok.
invalidate_user_badge_count_local(UserId) ->
local_cache_mutation(fun() ->
push_ets_cache:delete_badge_count(UserId)
end).
-spec invalidate_user_badge_counts_local(term()) -> ok.
invalidate_user_badge_counts_local(UserIds) ->
case push_normalize:integer_list(UserIds) of
{ok, TypedUserIds} ->
lists:foreach(fun invalidate_user_badge_count_local/1, TypedUserIds);
error ->
ok
end.
-spec maybe_cast(term(), term()) -> ok.
maybe_cast(Key, Msg) ->
case is_push_noop() of
@@ -236,42 +175,27 @@ is_push_noop() ->
clear_channel_notifications(UserId, ChannelId, MessageId) ->
case is_push_active() of
true ->
clear_for_enrolment(
push_delivery_config:is_enrolled(UserId), UserId, ChannelId, MessageId
ok = push_outbox:truncate_read(UserId, ChannelId, MessageId),
cast_clear_if_enabled(
clear_notifications_enabled(), 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(
enrolled_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 enrolled_clear_notifications_enabled() -> boolean().
enrolled_clear_notifications_enabled() ->
-spec clear_notifications_enabled() -> boolean().
clear_notifications_enabled() ->
case persistent_term:get(push_enrolled_clear_notifications_enabled, undefined) of
OperatorChoice when is_boolean(OperatorChoice) -> OperatorChoice;
_ -> env_boolean(push_enrolled_clear_notifications_enabled, true)
end.
-spec clear_notifications_enabled() -> boolean().
clear_notifications_enabled() ->
case persistent_term:get(push_clear_notifications_enabled, undefined) of
OperatorChoice when is_boolean(OperatorChoice) -> OperatorChoice;
_ -> env_boolean(push_clear_notifications_enabled, true)
end.
-spec get_cache_stats() -> {ok, map()}.
get_cache_stats() ->
gen_server:call(?MODULE, get_cache_stats, 5000).
@@ -320,18 +244,18 @@ resolve_push_owner(Key) ->
exit:_Reason -> unavailable
end.
-spec do_handle_message_create(map(), worker_state()) -> ok.
do_handle_message_create(Params, State) ->
-spec do_handle_message_create(map()) -> ok.
do_handle_message_create(Params) ->
case push_message_params:context(Params) of
{ok, Context} ->
do_handle_message_create_context(Context, State);
do_handle_message_create_context(Context);
{error, Reason} ->
logger:debug("Push: skipping malformed message create", #{reason => Reason}),
ok
end.
-spec do_handle_message_create_context(push_message_params:context(), worker_state()) -> ok.
do_handle_message_create_context(Context, State) ->
-spec do_handle_message_create_context(push_message_params:context()) -> ok.
do_handle_message_create_context(Context) ->
#{
message_data := MessageData,
user_ids := UserIds,
@@ -369,158 +293,59 @@ do_handle_message_create_context(Context, State) ->
channel_id => ChannelId,
eligible_count => length(EligibleUsers)
}),
route_partitioned_users(
split_for_delivery(EligibleUsers),
publish_eligible_users(
EligibleUsers,
MessageData,
MarkdownContext,
GuildId,
ChannelId,
MessageId,
GuildName,
ChannelName,
State
ChannelName
).
-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(
-spec publish_eligible_users(
[integer()],
non_neg_integer(),
map(),
map(),
integer(),
integer(),
integer(),
binary() | undefined,
binary() | undefined,
worker_state()
binary() | undefined
) -> ok.
route_service_users(
publish_eligible_users(
[],
_ConfigVersion,
_MessageData,
_MarkdownContext,
_GuildId,
_ChannelId,
_MessageId,
_GuildName,
_ChannelName,
_State
_ChannelName
) ->
ok;
route_service_users(
ServiceUsers,
ConfigVersion,
publish_eligible_users(
EligibleUsers,
MessageData,
MarkdownContext,
GuildId,
ChannelId,
MessageId,
GuildName,
ChannelName,
State
ChannelName
) ->
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).
_ = push_job_publisher:publish_message(
EligibleUsers,
MessageData,
MarkdownContext,
GuildId,
ChannelId,
MessageId,
GuildName,
ChannelName
),
ok.
-spec filter_eligible_users(
[integer()],
@@ -747,81 +572,6 @@ blocked_ids_counters() ->
blocked_ids_budget_exhausted => read_counter(?CNT_BUDGET_EXHAUSTED)
}.
-spec dispatch_if_eligible(
[integer()],
map(),
map(),
integer(),
integer(),
integer(),
binary() | undefined,
binary() | undefined,
worker_state()
) -> ok.
dispatch_if_eligible(
[],
_MessageData,
_MarkdownContext,
_GuildId,
_ChannelId,
_MessageId,
_GuildName,
_ChannelName,
_State
) ->
ok;
dispatch_if_eligible(
EligibleUsers,
MessageData,
MarkdownContext,
GuildId,
ChannelId,
MessageId,
GuildName,
ChannelName,
State
) ->
BadgeCountsTtl = maps:get(badge_counts_ttl_seconds, State),
case
push_dispatcher:enqueue_send_notifications(
EligibleUsers,
MessageData,
MarkdownContext,
GuildId,
ChannelId,
MessageId,
GuildName,
ChannelName,
BadgeCountsTtl
)
of
ok ->
ok;
dropped ->
EligibleCount = length(EligibleUsers),
count_dispatch_dropped(EligibleCount),
log_dispatch_drop(
loss_logging_enabled(), MessageId, ChannelId, GuildId, EligibleCount
),
ok
end.
-spec log_dispatch_drop(boolean(), integer(), integer(), integer(), non_neg_integer()) -> ok.
log_dispatch_drop(true, MessageId, ChannelId, GuildId, EligibleCount) ->
logger:error("Push: dispatcher saturated, dropping notification job", #{
message_id => MessageId,
channel_id => ChannelId,
guild_id => GuildId,
eligible_count => EligibleCount
});
log_dispatch_drop(false, MessageId, ChannelId, GuildId, EligibleCount) ->
logger:debug("Push: dispatcher saturated, dropping notification job", #{
message_id => MessageId,
channel_id => ChannelId,
guild_id => GuildId,
eligible_count => EligibleCount
}).
-spec handle_sync_user_blocked_ids(integer(), term(), state()) -> {noreply, state()}.
handle_sync_user_blocked_ids(UserId, BlockedIds, State) ->
case push_normalize:integer_list(BlockedIds) of
@@ -834,84 +584,16 @@ handle_sync_user_blocked_ids(UserId, BlockedIds, State) ->
-spec handle_message_create_cast(map(), state()) -> {noreply, state()}.
handle_message_create_cast(Params, State) ->
BadgeCountsTtl = maps:get(badge_counts_ttl_seconds, State),
WorkerState = #{badge_counts_ttl_seconds => BadgeCountsTtl},
SpawnResult = maybe_spawn_push_worker(fun() ->
do_handle_message_create(Params, WorkerState)
end),
SpawnResult = maybe_spawn_push_worker(fun() -> do_handle_message_create(Params) end),
log_message_worker_drop(SpawnResult, Params),
{noreply, 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) ->
clear_via_dispatcher(UserId, ChannelId, MessageId, 0, State).
-spec clear_via_dispatcher(
integer(), integer(), integer(), non_neg_integer(), state()
) -> {noreply, state()}.
clear_via_dispatcher(UserId, ChannelId, MessageId, Attempt, State) ->
BadgeCountsTtl = maps:get(badge_counts_ttl_seconds, State),
case
push_dispatcher:enqueue_clear_notifications(
UserId, ChannelId, MessageId, BadgeCountsTtl
)
of
ok ->
ok;
dropped ->
retry_or_drop_clear(UserId, ChannelId, MessageId, Attempt)
end,
_ = push_job_publisher:publish_clear(UserId, ChannelId, MessageId),
{noreply, State}.
-spec retry_or_drop_clear(integer(), integer(), integer(), non_neg_integer()) -> ok.
retry_or_drop_clear(UserId, ChannelId, MessageId, Attempt) when
Attempt < ?CLEAR_RETRY_ATTEMPTS
->
bump_counter(?CNT_CLEAR_RETRIED),
Delay = ?CLEAR_RETRY_BASE_MS bsl Attempt,
_ = erlang:send_after(
Delay, self(), {retry_clear_notifications, UserId, ChannelId, MessageId, Attempt + 1}
),
ok;
retry_or_drop_clear(UserId, ChannelId, MessageId, _Attempt) ->
count_clear_dropped(),
log_clear_drop(loss_logging_enabled(), UserId, ChannelId, MessageId).
-spec log_clear_drop(boolean(), integer(), integer(), integer()) -> ok.
log_clear_drop(true, UserId, ChannelId, MessageId) ->
logger:warning("Push: dispatcher saturated, dropping clear notification job", #{
user_id => UserId, channel_id => ChannelId, message_id => MessageId
});
log_clear_drop(false, UserId, ChannelId, MessageId) ->
logger:debug("Push: dispatcher saturated, dropping clear notification job", #{
user_id => UserId, channel_id => ChannelId, message_id => MessageId
}).
-spec log_message_worker_drop(ok | dropped, map()) -> ok.
log_message_worker_drop(ok, _Params) ->
ok;
@@ -937,60 +619,15 @@ 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()),
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.
maps:merge(Base, push_loss_counters()).
-spec push_loss_counters() -> map().
push_loss_counters() ->
#{
counters => counter_table_status(),
worker_pool_dropped => read_counter(?CNT_WORKER_POOL),
dispatch_dropped => read_counter(?CNT_DISPATCH_DROPPED),
dispatch_dropped_users => read_counter(?CNT_DISPATCH_DROPPED_USERS),
clear_dispatch_dropped => read_counter(?CNT_CLEAR_DROPPED),
dispatcher_queue_full => read_counter(?CNT_QUEUE_FULL),
dispatcher_invalid_job => read_counter(?CNT_INVALID_JOB),
dispatcher_enqueue_timeout => read_counter(?CNT_ENQUEUE_TIMEOUT),
dispatcher_enqueue_failed => read_counter(?CNT_ENQUEUE_FAILED),
dispatcher_job_crashed => read_counter(?CNT_JOB_CRASHED),
dispatcher_worker_died => read_counter(?CNT_WORKER_DIED),
dispatcher_restarts => read_counter(?CNT_DISPATCHER_RESTARTS),
dispatcher_restart_discarded => read_counter(?CNT_RESTART_DISCARDED),
dispatcher_queue_backlog => dispatcher_queue_backlog()
worker_pool_dropped => read_counter(?CNT_WORKER_POOL)
}.
-spec count_dispatch_dropped(non_neg_integer()) -> ok.
count_dispatch_dropped(EligibleCount) ->
bump_counter(?CNT_DISPATCH_DROPPED),
bump_counter(?CNT_DISPATCH_DROPPED_USERS, EligibleCount).
-spec count_clear_dropped() -> ok.
count_clear_dropped() ->
bump_counter(?CNT_CLEAR_DROPPED).
-spec counter_table_status() -> live | unavailable.
counter_table_status() ->
case ets:info(?PUSH_COUNTER_TABLE, size) of
@@ -998,19 +635,6 @@ counter_table_status() ->
_ -> unavailable
end.
-spec dispatcher_queue_backlog() -> non_neg_integer() | unavailable.
dispatcher_queue_backlog() ->
backlog(read_counter(?CNT_QUEUE_ENQUEUED), read_counter(?CNT_QUEUE_DEQUEUED)).
-spec backlog(non_neg_integer() | unavailable, non_neg_integer() | unavailable) ->
non_neg_integer() | unavailable.
backlog(Enqueued, Dequeued) when is_integer(Enqueued), is_integer(Dequeued) ->
max(0, Enqueued - Dequeued);
backlog(Enqueued, unavailable) when is_integer(Enqueued) ->
Enqueued;
backlog(_Enqueued, _Dequeued) ->
unavailable.
-spec loss_logging_enabled() -> boolean().
loss_logging_enabled() ->
application:get_env(fluxer_gateway, push_loss_logging, false) =:= true.
@@ -1082,43 +706,6 @@ maybe_spawn_push_worker(Fun) ->
schedule_eviction() ->
erlang:send_after(?EVICT_INTERVAL_MS, self(), evict_caches).
-spec maybe_warn_vapid_misconfigured(boolean()) -> ok.
maybe_warn_vapid_misconfigured(true) ->
Public = fluxer_gateway_env:get(vapid_public_key),
Private = fluxer_gateway_env:get(vapid_private_key),
case {Public, Private} of
{Public0, Private0} when
is_binary(Public0),
is_binary(Private0),
byte_size(Public0) > 0,
byte_size(Private0) > 0
->
warn_unless_vapid_pair_valid(Public0, Private0);
_ ->
logger:error(
"Push: push_enabled=true but VAPID keys are missing or empty; "
"all web push notifications will be silently dropped"
),
ok
end;
maybe_warn_vapid_misconfigured(_) ->
ok.
-spec warn_unless_vapid_pair_valid(binary(), binary()) -> ok.
warn_unless_vapid_pair_valid(Public, Private) ->
try
push_utils:assert_vapid_pair(Public, Private)
catch
_:Reason ->
logger:error(
"Push: FLUXER_VAPID_PUBLIC_KEY and FLUXER_VAPID_PRIVATE_KEY are not a "
"valid base64url P-256 pair; expected a 65-byte 0x04-prefixed point and "
"a 32-byte scalar; all web push notifications will be silently dropped",
#{reason => Reason}
),
ok
end.
-spec env_boolean(atom()) -> boolean().
env_boolean(Key) ->
case fluxer_gateway_env:get(Key) of
@@ -1133,13 +720,6 @@ env_boolean(Key, Default) ->
_ -> Default
end.
-spec env_non_neg_integer(atom(), non_neg_integer()) -> non_neg_integer().
env_non_neg_integer(Key, Default) ->
case fluxer_gateway_env:get(Key) of
Value when is_integer(Value), Value >= 0 -> Value;
_ -> Default
end.
-ifdef(TEST).
-include_lib("eunit/include/eunit.hrl").
@@ -1159,38 +739,59 @@ sync_user_blocked_ids_local_updates_local_cache_test() ->
?assertEqual([20, 30], push_ets_cache:get_blocked_ids(10))
end).
invalidate_user_badge_counts_local_deletes_every_cached_entry_test() ->
a_message_publishes_every_eligible_recipient_in_one_call_test() ->
push_ets_cache:init(),
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),
?assertEqual(undefined, push_ets_cache:get_badge_count(10)),
?assertEqual(undefined, push_ets_cache:get_badge_count(11)).
lists:foreach(fun(UserId) -> ok = push_ets_cache:put_blocked_ids(UserId, []) end, [51, 52]),
ok = push_ets_cache:put_blocked_ids(53, [7]),
Self = self(),
ok = meck:new(push_job_publisher, [passthrough, no_link]),
try
ok = meck:expect(
push_job_publisher,
publish_message,
fun(UserIds, _Data, _Markdown, GuildId, ChannelId, MessageId, _GName, _CName) ->
Self ! {published, UserIds, GuildId, ChannelId, MessageId},
ok
end
),
ok = do_handle_message_create(#{
message_data => #{
<<"channel_id">> => <<"123">>,
<<"id">> => <<"456">>,
<<"channel_type">> => 1
},
user_ids => [51, 52, 53],
guild_id => 0,
author_id => 7
}),
receive
{published, UserIds, GuildId, ChannelId, MessageId} ->
?assertEqual({[51, 52], 0, 123, 456}, {UserIds, GuildId, ChannelId, MessageId})
after 2000 -> erlang:error(nothing_published)
end
after
meck:unload(push_job_publisher)
end.
invalidate_user_badge_counts_local_ignores_untyped_ids_test() ->
a_message_without_eligible_recipients_publishes_nothing_test() ->
push_ets_cache:init(),
ok = seed_badge_count(12, 5, 1000),
with_registered_push(fun() ->
ok = invalidate_user_badge_counts_local([<<"12">>])
end),
?assertEqual({5, 1000}, push_ets_cache:get_badge_count(12)),
push_ets_cache:delete_badge_count(12).
invalidate_user_subscriptions_local_deletes_local_cache_test() ->
push_ets_cache:init(),
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))
end).
ok = push_ets_cache:put_blocked_ids(61, [7]),
ok = meck:new(push_job_publisher, [passthrough, no_link]),
try
ok = do_handle_message_create(#{
message_data => #{
<<"channel_id">> => <<"123">>,
<<"id">> => <<"456">>,
<<"channel_type">> => 1
},
user_ids => [61],
guild_id => 0,
author_id => 7
}),
?assertEqual(0, meck:num_calls(push_job_publisher, publish_message, '_'))
after
meck:unload(push_job_publisher)
end.
filter_eligible_users_fetches_large_metadata_once_test() ->
push_ets_cache:init(),
@@ -1308,41 +909,22 @@ assert_cache_stats_expose_blocked_ids() ->
push_loss_counters_expose_every_counter_push_writes_test() ->
with_counter_table(fun() ->
ok = log_message_worker_drop(dropped, #{}),
ok = count_dispatch_dropped(12),
ok = count_clear_dropped(),
Stats = push_loss_counters(),
?assertEqual(live, maps:get(counters, Stats)),
?assertEqual(1, maps:get(worker_pool_dropped, Stats)),
?assertEqual(1, maps:get(dispatch_dropped, Stats)),
?assertEqual(12, maps:get(dispatch_dropped_users, Stats)),
?assertEqual(1, maps:get(clear_dispatch_dropped, Stats))
?assertEqual(1, maps:get(worker_pool_dropped, Stats))
end).
push_loss_counters_are_unavailable_without_the_shared_table_test() ->
delete_counter_table(),
Stats = push_loss_counters(),
?assertEqual(unavailable, maps:get(counters, Stats)),
?assertEqual(unavailable, maps:get(worker_pool_dropped, Stats)),
?assertEqual(unavailable, maps:get(dispatch_dropped, Stats)),
?assertEqual(unavailable, maps:get(dispatcher_enqueue_timeout, Stats)),
?assertEqual(unavailable, maps:get(dispatcher_queue_backlog, Stats)).
?assertEqual(unavailable, maps:get(worker_pool_dropped, Stats)).
push_loss_counters_keep_a_genuine_zero_distinct_from_absent_test() ->
with_counter_table(fun() ->
ok = count_dispatch_dropped(0),
Stats = push_loss_counters(),
?assertEqual(1, maps:get(dispatch_dropped, Stats)),
?assertEqual(0, maps:get(dispatch_dropped_users, Stats)),
?assertEqual(unavailable, maps:get(clear_dispatch_dropped, Stats))
end).
dispatcher_queue_backlog_survives_an_untrappable_dispatcher_kill_test() ->
with_counter_table(fun() ->
?assertEqual(unavailable, dispatcher_queue_backlog()),
ok = bump_counter(?CNT_QUEUE_ENQUEUED, 9),
?assertEqual(9, dispatcher_queue_backlog()),
ok = bump_counter(?CNT_QUEUE_DEQUEUED, 4),
?assertEqual(5, dispatcher_queue_backlog())
?assertEqual(unavailable, maps:get(worker_pool_dropped, push_loss_counters())),
ok = bump_counter(?CNT_WORKER_POOL, 0),
?assertEqual(0, maps:get(worker_pool_dropped, push_loss_counters()))
end).
cache_stats_with_counters_carries_the_loss_surface_test() ->
@@ -1351,22 +933,7 @@ cache_stats_with_counters_carries_the_loss_surface_test() ->
Stats = cache_stats_with_counters(),
lists:foreach(
fun(Key) -> ?assertEqual(true, maps:is_key(Key, Stats)) end,
[
blocked_ids_size,
blocked_ids_suppressed,
counters,
worker_pool_dropped,
dispatch_dropped,
dispatch_dropped_users,
clear_dispatch_dropped,
dispatcher_queue_full,
dispatcher_enqueue_timeout,
dispatcher_job_crashed,
dispatcher_worker_died,
dispatcher_restarts,
dispatcher_restart_discarded,
dispatcher_queue_backlog
]
[blocked_ids_size, blocked_ids_suppressed, counters, worker_pool_dropped]
)
end).
@@ -1386,11 +953,6 @@ 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])
-98
View File
@@ -1,98 +0,0 @@
%% SPDX-License-Identifier: AGPL-3.0-or-later
-module(push_apns).
-typing([eqwalizer]).
-export([send/3]).
-spec send(integer(), map(), map()) -> false | {true, map()}.
send(UserId, Subscription, Payload) ->
case fluxer_gateway_env:get(apns_enabled) of
true ->
send_via_api_rpc(UserId, Subscription, Payload);
_ ->
false
end.
-spec send_via_api_rpc(integer(), map(), map()) -> false | {true, map()}.
send_via_api_rpc(UserId, Subscription, Payload) ->
case extract_subscription(Subscription) of
{ok, DeviceToken, SubscriptionId, AppId, Environment} ->
Request = build_apns_request(
UserId, SubscriptionId, DeviceToken, AppId, Environment, Payload
),
handle_apns_rpc_result(UserId, SubscriptionId, rpc_client:call(Request));
{error, Reason} ->
logger:debug("Push: invalid APNs subscription", #{
user_id => UserId, reason => Reason
}),
false
end.
-spec build_apns_request(integer(), binary(), binary(), binary(), binary(), map()) -> map().
build_apns_request(UserId, SubscriptionId, DeviceToken, AppId, Environment, Payload) ->
#{
<<"type">> => <<"send_apns_push">>,
<<"user_id">> => integer_to_binary(UserId),
<<"subscription_id">> => SubscriptionId,
<<"device_token">> => DeviceToken,
<<"app_id">> => AppId,
<<"provider_environment">> => Environment,
<<"payload">> => Payload
}.
-spec handle_apns_rpc_result(integer(), binary(), {ok, map()} | {error, term()}) ->
false | {true, map()}.
handle_apns_rpc_result(_UserId, _SubscriptionId, {ok, #{<<"success">> := true}}) ->
false;
handle_apns_rpc_result(UserId, SubscriptionId, {ok, #{<<"should_delete">> := true}}) ->
{true, delete_payload(UserId, SubscriptionId)};
handle_apns_rpc_result(_UserId, _SubscriptionId, {ok, Response}) when is_map(Response) ->
false;
handle_apns_rpc_result(UserId, _SubscriptionId, {error, Reason}) ->
logger:debug("Push: APNs RPC failed", #{user_id => UserId, reason => Reason}),
false.
-spec extract_subscription(map()) ->
{ok, binary(), binary(), binary(), binary()} | {error, term()}.
extract_subscription(Subscription) ->
DeviceToken = push_utils:normalize_binary(
maps:get(<<"endpoint">>, Subscription, undefined), undefined
),
SubscriptionId = push_utils:normalize_binary(
maps:get(<<"subscription_id">>, Subscription, undefined), undefined
),
AppId = push_utils:normalize_binary(
maps:get(<<"app_id">>, Subscription, <<"stable">>), <<"stable">>
),
DefaultEnvironment = fluxer_gateway_env:get(apns_default_environment),
Environment = normalize_environment(
maps:get(<<"provider_environment">>, Subscription, DefaultEnvironment)
),
case {DeviceToken, SubscriptionId, AppId, Environment} of
{Token, Id, App, Env} when
is_binary(Token),
byte_size(Token) > 0,
is_binary(Id),
is_binary(App),
is_binary(Env)
->
{ok, Token, Id, App, Env};
_ ->
{error, missing_fields}
end.
-spec normalize_environment(term()) -> binary().
normalize_environment(Value) ->
case push_utils:normalize_binary(Value, <<"production">>) of
<<"development">> -> <<"development">>;
<<"sandbox">> -> <<"development">>;
_ -> <<"production">>
end.
-spec delete_payload(integer(), binary()) -> map().
delete_payload(UserId, SubscriptionId) ->
#{
<<"user_id">> => integer_to_binary(UserId),
<<"subscription_id">> => SubscriptionId
}.
@@ -1,459 +0,0 @@
%% 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(),
relay_consent_accepted := boolean()
}.
-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(<<Salt/binary, ":", UserId/binary>>, ?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(<<Byte:8, Rest/binary>>, 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 => #{},
relay_consent_accepted => false
}.
-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},
{relay_consent_accepted, <<"relay_consent_accepted">>, fun validate_enabled/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(<<Byte:8, Rest/binary>>, 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(<<Byte:8, Rest/binary>>) 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, relay_consent_accepted]
).
-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.
-ifdef(TEST).
-include_lib("eunit/include/eunit.hrl").
an_accepted_relay_notice_is_read_off_the_wire_config_test() ->
{ok, Config} = validate_config(#{<<"relay_consent_accepted">> => true}),
?assertEqual(true, maps:get(relay_consent_accepted, Config)).
a_wire_config_without_a_relay_notice_has_not_been_accepted_test() ->
{ok, Config} = validate_config(#{<<"enabled">> => true}),
?assertEqual(false, maps:get(relay_consent_accepted, Config)).
a_relay_notice_that_is_not_a_boolean_is_refused_test() ->
?assertMatch(
{error, {invalid_field, <<"relay_consent_accepted">>, _}},
validate_config(#{<<"relay_consent_accepted">> => <<"yes">>})
).
-endif.
-874
View File
@@ -1,874 +0,0 @@
%% SPDX-License-Identifier: AGPL-3.0-or-later
-module(push_dispatcher).
-typing([eqwalizer]).
-behaviour(gen_server).
-export([
start_link/0,
enqueue_send_notifications/8,
enqueue_send_notifications/9,
enqueue_clear_notifications/4,
stats/0
]).
-export([init/1, handle_call/3, handle_cast/2, handle_info/2, terminate/2, code_change/3]).
-define(DEFAULT_MAX_INFLIGHT, 256).
-define(DEFAULT_MAX_QUEUE, 10000).
-define(ENQUEUE_TIMEOUT_MS, 1000).
-define(PUSH_COUNTER_TABLE, push_worker_counter).
-define(CNT_QUEUE_FULL, push_loss_queue_full).
-define(CNT_INVALID_JOB, push_loss_invalid_job).
-define(CNT_ENQUEUE_TIMEOUT, push_loss_enqueue_timeout).
-define(CNT_ENQUEUE_FAILED, push_loss_enqueue_failed).
-define(CNT_JOB_CRASHED, push_loss_job_crashed).
-define(CNT_WORKER_DIED, push_loss_worker_died).
-define(CNT_WORKER_POOL, push_loss_worker_pool).
-define(CNT_RESTARTS, push_dispatcher_restarts).
-define(CNT_RESTART_DISCARDED, push_dispatcher_restart_discarded).
-define(CNT_QUEUE_ENQUEUED, push_dispatcher_queue_enqueued).
-define(CNT_QUEUE_DEQUEUED, push_dispatcher_queue_dequeued).
-type push_job() ::
#{
type := message_create,
user_ids := [integer()],
message_data := map(),
markdown_context := map(),
guild_id := integer(),
channel_id := integer(),
message_id := integer(),
guild_name := binary() | undefined,
channel_name := binary() | undefined,
badge_counts_ttl_seconds := non_neg_integer()
}
| #{
type := clear_channel,
user_id := integer(),
channel_id := integer(),
message_id := integer(),
badge_counts_ttl_seconds := non_neg_integer()
}.
-type state() :: #{
queue := queue:queue(push_job()),
queued := non_neg_integer(),
inflight := non_neg_integer(),
workers := #{reference() => true},
max_inflight := pos_integer(),
max_queue := pos_integer(),
started_at => integer()
}.
-type counter_value() :: non_neg_integer() | unavailable.
-type stats() :: #{
queued := non_neg_integer(),
inflight := non_neg_integer(),
counters := live | unavailable,
dispatcher_uptime_seconds := non_neg_integer() | undefined,
dispatcher_restarts := counter_value(),
restart_discarded_jobs := counter_value(),
queue_backlog_lost := counter_value(),
queue_full_dropped := counter_value(),
invalid_job_dropped := counter_value(),
enqueue_timeout := counter_value(),
enqueue_failed := counter_value(),
job_crashed := counter_value(),
worker_died := counter_value(),
worker_pool_dropped := counter_value()
}.
-spec start_link() -> {ok, pid()} | {error, term()} | ignore.
start_link() ->
gen_server:start_link({local, ?MODULE}, ?MODULE, [], []).
-spec enqueue_send_notifications(
[integer()],
map(),
integer(),
integer(),
integer(),
binary() | undefined,
binary() | undefined,
non_neg_integer()
) -> ok | dropped.
enqueue_send_notifications(
UserIds,
MessageData,
GuildId,
ChannelId,
MessageId,
GuildName,
ChannelName,
BadgeCountsTtlSeconds
) ->
enqueue_send_notifications(
UserIds,
MessageData,
#{},
GuildId,
ChannelId,
MessageId,
GuildName,
ChannelName,
BadgeCountsTtlSeconds
).
-spec enqueue_send_notifications(
[integer()],
map(),
map(),
integer(),
integer(),
integer(),
binary() | undefined,
binary() | undefined,
non_neg_integer()
) -> ok | dropped.
enqueue_send_notifications(
UserIds,
MessageData,
MarkdownContext,
GuildId,
ChannelId,
MessageId,
GuildName,
ChannelName,
BadgeCountsTtlSeconds
) ->
Job = send_notifications_job(
UserIds,
MessageData,
MarkdownContext,
GuildId,
ChannelId,
MessageId,
GuildName,
ChannelName,
BadgeCountsTtlSeconds
),
log_enqueue_send_notifications(UserIds, GuildId, ChannelId, MessageId),
safe_enqueue(Job).
-spec enqueue_clear_notifications(integer(), integer(), integer(), non_neg_integer()) ->
ok | dropped.
enqueue_clear_notifications(UserId, ChannelId, MessageId, BadgeCountsTtlSeconds) ->
Job = #{
type => clear_channel,
user_id => UserId,
channel_id => ChannelId,
message_id => MessageId,
badge_counts_ttl_seconds => BadgeCountsTtlSeconds
},
logger:debug(
"Push: enqueuing clear notification job",
#{user_id => UserId, channel_id => ChannelId, message_id => MessageId}
),
safe_enqueue(Job).
-spec stats() -> stats() | #{}.
stats() ->
Enqueued = read_counter(?CNT_QUEUE_ENQUEUED),
try gen_server:call(?MODULE, stats, 1000) of
#{queued := Queued, inflight := Inflight} = Reply ->
StartedAt = maps:get(started_at, Reply, undefined),
stats_with_loss(Queued, Inflight, StartedAt, Enqueued);
_ ->
#{}
catch
exit:_ -> #{};
error:_ -> #{}
end.
-spec stats_with_loss(non_neg_integer(), non_neg_integer(), term(), term()) -> stats().
stats_with_loss(Queued, Inflight, StartedAt, Enqueued) ->
#{
queued => Queued,
inflight => Inflight,
counters => counter_table_status(),
dispatcher_uptime_seconds => uptime_seconds(StartedAt),
dispatcher_restarts => read_counter(?CNT_RESTARTS),
restart_discarded_jobs => read_counter(?CNT_RESTART_DISCARDED),
queue_backlog_lost => queue_backlog_lost(Enqueued, Queued),
queue_full_dropped => read_counter(?CNT_QUEUE_FULL),
invalid_job_dropped => read_counter(?CNT_INVALID_JOB),
enqueue_timeout => read_counter(?CNT_ENQUEUE_TIMEOUT),
enqueue_failed => read_counter(?CNT_ENQUEUE_FAILED),
job_crashed => read_counter(?CNT_JOB_CRASHED),
worker_died => read_counter(?CNT_WORKER_DIED),
worker_pool_dropped => read_counter(?CNT_WORKER_POOL)
}.
-spec uptime_seconds(term()) -> non_neg_integer() | undefined.
uptime_seconds(StartedAt) when is_integer(StartedAt) ->
max(0, erlang:monotonic_time(second) - StartedAt);
uptime_seconds(_StartedAt) ->
undefined.
-spec init([]) -> {ok, state()}.
init([]) ->
erlang:process_flag(fullsweep_after, 10),
bump_counter(?CNT_RESTARTS),
{ok, #{
queue => queue:new(),
queued => 0,
inflight => 0,
workers => #{},
max_inflight => budget_aware_max_inflight(),
max_queue => get_int_or_default(push_dispatcher_max_queue, ?DEFAULT_MAX_QUEUE),
started_at => erlang:monotonic_time(second)
}}.
-spec budget_aware_max_inflight() -> pos_integer().
budget_aware_max_inflight() ->
Configured = get_int_or_default(push_dispatcher_max_inflight, ?DEFAULT_MAX_INFLIGHT),
PushBudget = gateway_http_client:push_max_concurrency(),
DeliveryConcurrency = max(1, push_subscriptions:delivery_concurrency()),
BudgetCap = max(1, PushBudget div DeliveryConcurrency),
min(Configured, BudgetCap).
-spec handle_call(term(), gen_server:from(), state()) ->
{reply, term(), state()}.
handle_call(stats, _From, #{queued := Queued, inflight := Inflight} = State) ->
Reply = #{
queued => Queued,
inflight => Inflight,
started_at => maps:get(started_at, State, undefined)
},
{reply, Reply, State};
handle_call({enqueue, Job}, _From, State) ->
{Result, State1} = handle_enqueue(Job, State),
{reply, Result, State1};
handle_call(_Request, _From, State) ->
{reply, ok, State}.
-spec handle_cast(term(), state()) -> {noreply, state()}.
handle_cast({enqueue, Job}, State) ->
{_Result, State1} = handle_enqueue(Job, State),
{noreply, State1};
handle_cast(_Msg, State) ->
{noreply, State}.
-spec handle_info(term(), state()) -> {noreply, state()}.
handle_info(
{'DOWN', Ref, process, _Pid, Reason},
#{workers := Workers, inflight := Inflight} = State
) ->
case maps:is_key(Ref, Workers) of
true ->
count_worker_down(Reason),
RemainingWorkers = maps:remove(Ref, Workers),
DecrementedInflight = max(0, Inflight - 1),
drain_queue(State#{
workers := RemainingWorkers,
inflight := DecrementedInflight
});
false ->
{noreply, State}
end;
handle_info(_Info, State) ->
{noreply, State}.
-spec terminate(term(), state()) -> ok.
terminate(_Reason, State) ->
bump_counter(?CNT_RESTART_DISCARDED, maps:get(queued, State, 0)).
-spec count_worker_down(term()) -> ok.
count_worker_down(normal) ->
ok;
count_worker_down(_Reason) ->
bump_counter(?CNT_JOB_CRASHED),
bump_counter(?CNT_WORKER_DIED).
-spec code_change(term(), state(), term()) -> {ok, state()}.
code_change(_OldVsn, State, _Extra) ->
erlang:garbage_collect(),
{ok, State}.
-spec maybe_enqueue_or_start(push_job(), state()) -> {ok | dropped, state()}.
maybe_enqueue_or_start(Job, #{inflight := Inflight, max_inflight := MaxInflight} = State) ->
case Inflight < MaxInflight of
true ->
logger:debug(
"Push: starting job immediately",
#{
message_id => maps:get(message_id, Job, undefined),
inflight => Inflight,
max_inflight => MaxInflight
}
),
{ok, start_job(Job, State)};
false ->
logger:debug(
"Push: at capacity, queueing job",
#{
message_id => maps:get(message_id, Job, undefined),
inflight => Inflight,
max_inflight => MaxInflight,
queued => maps:get(queued, State)
}
),
maybe_enqueue(Job, State)
end.
-spec maybe_enqueue(push_job(), state()) -> {ok | dropped, state()}.
maybe_enqueue(Job, #{queued := Queued, max_queue := MaxQueue, queue := Queue0} = State) ->
case Queued < MaxQueue of
true ->
bump_counter(?CNT_QUEUE_ENQUEUED),
Queue1 = queue:in(Job, Queue0),
{ok, State#{queue := Queue1, queued := Queued + 1}};
false ->
bump_counter(?CNT_QUEUE_FULL),
DropCount = bump_drop_count(),
log_queue_full_drop(loss_logging_enabled(), Job, Queued, MaxQueue, DropCount),
{dropped, State}
end.
-spec log_queue_full_drop(
boolean(), push_job(), non_neg_integer(), pos_integer(), non_neg_integer()
) -> ok.
log_queue_full_drop(true, Job, Queued, MaxQueue, _DropCount) ->
logger:error(
"Push: queue full, dropping job",
#{
message_id => maps:get(message_id, Job, undefined),
queued => Queued,
max_queue => MaxQueue,
total_dropped => read_counter(?CNT_QUEUE_FULL)
}
);
log_queue_full_drop(false, Job, Queued, MaxQueue, DropCount) ->
logger:warning(
"Push: queue full, dropping job",
#{
message_id => maps:get(message_id, Job, undefined),
queued => Queued,
max_queue => MaxQueue,
total_dropped => DropCount
}
).
-spec bump_drop_count() -> non_neg_integer().
bump_drop_count() ->
Current =
case erlang:get(push_dispatcher_drop_count) of
N when is_integer(N) -> N;
_ -> 0
end,
Updated = Current + 1,
erlang:put(push_dispatcher_drop_count, Updated),
Updated.
-spec loss_logging_enabled() -> boolean().
loss_logging_enabled() ->
application:get_env(fluxer_gateway, push_loss_logging, false) =:= true.
-spec counter_table_status() -> live | unavailable.
counter_table_status() ->
case ets:info(?PUSH_COUNTER_TABLE, size) of
Size when is_integer(Size) -> live;
_ -> unavailable
end.
-spec queue_backlog_lost(term(), non_neg_integer()) -> counter_value().
queue_backlog_lost(Enqueued, Queued) when is_integer(Enqueued) ->
max(0, trunc(Enqueued - dequeued_total() - Queued));
queue_backlog_lost(_Enqueued, _Queued) ->
unavailable.
-spec dequeued_total() -> non_neg_integer().
dequeued_total() ->
case read_counter(?CNT_QUEUE_DEQUEUED) of
Dequeued when is_integer(Dequeued) -> Dequeued;
unavailable -> 0
end.
-spec read_counter(atom()) -> counter_value().
read_counter(Key) ->
try ets:lookup(?PUSH_COUNTER_TABLE, Key) of
[{Key, Value}] when is_integer(Value), Value >= 0 -> Value;
_ -> unavailable
catch
error:badarg -> unavailable
end.
-spec bump_counter(atom()) -> ok.
bump_counter(Key) ->
bump_counter(Key, 1).
-spec bump_counter(atom(), non_neg_integer()) -> ok.
bump_counter(Key, Increment) ->
try ets:update_counter(?PUSH_COUNTER_TABLE, Key, {2, Increment}) of
_Value -> ok
catch
error:badarg -> insert_missing_counter(Key, Increment)
end.
-spec insert_missing_counter(atom(), non_neg_integer()) -> ok.
insert_missing_counter(Key, Increment) ->
try ets:insert_new(?PUSH_COUNTER_TABLE, {Key, Increment}) of
true -> ok;
false -> retry_bump_counter(Key, Increment)
catch
error:badarg -> ok
end.
-spec retry_bump_counter(atom(), non_neg_integer()) -> ok.
retry_bump_counter(Key, Increment) ->
try ets:update_counter(?PUSH_COUNTER_TABLE, Key, {2, Increment}) of
_Value -> ok
catch
error:badarg -> ok
end.
-spec start_job(push_job(), state()) -> state().
start_job(Job, #{workers := Workers, inflight := Inflight} = State) ->
{_Pid, Ref} =
spawn_monitor(fun() ->
run_job(Job)
end),
State#{
workers := Workers#{Ref => true},
inflight := Inflight + 1
}.
-spec drain_queue(state()) -> {noreply, state()}.
drain_queue(
#{inflight := Inflight, max_inflight := MaxInflight, queue := Queue0, queued := Queued} =
State
) ->
case Inflight < MaxInflight of
true -> drain_available_queue(queue:out(Queue0), Queued, State);
false -> {noreply, State}
end.
-spec send_notifications_job(
[integer()],
map(),
map(),
integer(),
integer(),
integer(),
binary() | undefined,
binary() | undefined,
non_neg_integer()
) -> push_job().
send_notifications_job(
UserIds,
MessageData,
MarkdownContext,
GuildId,
ChannelId,
MessageId,
GuildName,
ChannelName,
BadgeCountsTtlSeconds
) ->
#{
type => message_create,
user_ids => UserIds,
message_data => MessageData,
markdown_context => MarkdownContext,
guild_id => GuildId,
channel_id => ChannelId,
message_id => MessageId,
guild_name => GuildName,
channel_name => ChannelName,
badge_counts_ttl_seconds => BadgeCountsTtlSeconds
}.
-spec log_enqueue_send_notifications([integer()], integer(), integer(), integer()) -> ok.
log_enqueue_send_notifications(UserIds, GuildId, ChannelId, MessageId) ->
logger:debug(
"Push: enqueuing dispatch job",
#{
message_id => MessageId,
channel_id => ChannelId,
guild_id => GuildId,
user_count => length(UserIds)
}
).
-spec drain_available_queue(
{{value, push_job()}, queue:queue(push_job())} | {empty, queue:queue(push_job())},
non_neg_integer(),
state()
) -> {noreply, state()}.
drain_available_queue({{value, Job}, Queue1}, Queued, State) ->
bump_counter(?CNT_QUEUE_DEQUEUED),
State1 = State#{queue := Queue1, queued := max(0, Queued - 1)},
State2 = start_job(Job, State1),
drain_queue(State2);
drain_available_queue({empty, _}, _Queued, State) ->
{noreply, State}.
-spec run_job(push_job()) -> ok.
run_job(#{message_id := MessageId} = Job) ->
try
run_typed_job(maps:get(type, Job, message_create), Job),
logger:debug("Push: worker completed", #{message_id => MessageId}),
ok
catch
Class:Reason:Stacktrace ->
bump_counter(?CNT_JOB_CRASHED),
log_job_crash(loss_logging_enabled(), MessageId, Class, Reason, Stacktrace),
ok
end.
-spec log_job_crash(boolean(), integer(), atom(), term(), list()) -> ok.
log_job_crash(true, MessageId, Class, Reason, Stacktrace) ->
logger:error(
"Push: worker crashed",
#{
message_id => MessageId,
class => Class,
reason => Reason,
stacktrace => Stacktrace
}
);
log_job_crash(false, MessageId, Class, Reason, _Stacktrace) ->
logger:debug(
"Push: worker crashed",
#{message_id => MessageId, class => Class, reason => Reason}
).
-spec run_typed_job(message_create | clear_channel, push_job()) -> ok.
run_typed_job(message_create, #{
user_ids := UserIds,
message_data := MessageData,
markdown_context := MarkdownContext,
guild_id := GuildId,
channel_id := ChannelId,
message_id := MessageId,
guild_name := GuildName,
channel_name := ChannelName,
badge_counts_ttl_seconds := BadgeCountsTtlSeconds
}) ->
logger:debug(
"Push: worker starting send_push_notifications",
#{message_id => MessageId, user_count => length(UserIds)}
),
push_sender:send_push_notifications(#{
user_ids => UserIds,
message_data => MessageData,
markdown_context => MarkdownContext,
guild_id => GuildId,
channel_id => ChannelId,
message_id => MessageId,
guild_name => GuildName,
channel_name => ChannelName,
badge_counts_ttl_seconds => BadgeCountsTtlSeconds
});
run_typed_job(clear_channel, #{
user_id := UserId,
channel_id := ChannelId,
message_id := MessageId,
badge_counts_ttl_seconds := BadgeCountsTtlSeconds
}) ->
logger:debug(
"Push: worker starting clear_channel_notifications",
#{user_id => UserId, channel_id => ChannelId, message_id => MessageId}
),
push_sender:send_clear_channel_notifications(
UserId, ChannelId, MessageId, BadgeCountsTtlSeconds
).
-spec get_int_or_default(atom(), integer()) -> integer().
get_int_or_default(Key, Default) ->
case fluxer_gateway_env:get_optional(Key) of
Value when is_integer(Value), Value > 0 -> Value;
_ -> Default
end.
-spec safe_enqueue(push_job()) -> ok | dropped.
safe_enqueue(Job) ->
try gen_server:call(?MODULE, {enqueue, Job}, ?ENQUEUE_TIMEOUT_MS) of
ok ->
ok;
dropped ->
dropped;
_ ->
bump_counter(?CNT_ENQUEUE_FAILED),
dropped
catch
throw:_Reason ->
bump_counter(?CNT_ENQUEUE_FAILED),
dropped;
error:_Reason ->
bump_counter(?CNT_ENQUEUE_FAILED),
dropped;
exit:Reason ->
count_enqueue_exit(Reason),
dropped
end.
-spec count_enqueue_exit(term()) -> ok.
count_enqueue_exit({timeout, _Call}) ->
bump_counter(?CNT_ENQUEUE_TIMEOUT);
count_enqueue_exit(_Reason) ->
bump_counter(?CNT_ENQUEUE_FAILED).
-spec handle_enqueue(term(), state()) -> {ok | dropped, state()}.
handle_enqueue(Job0, State) ->
case push_job(Job0) of
{ok, Job} ->
maybe_enqueue_or_start(Job, State);
error ->
bump_counter(?CNT_INVALID_JOB),
{dropped, State}
end.
-spec push_job(term()) -> {ok, push_job()} | error.
push_job(#{type := message_create} = Job) ->
push_message_create_job(Job);
push_job(#{type := clear_channel} = Job) ->
push_clear_channel_job(Job);
push_job(_) ->
error.
-spec push_message_create_job(map()) -> {ok, push_job()} | error.
push_message_create_job(
#{
user_ids := UserIds,
message_data := MessageData,
guild_id := GuildId,
channel_id := ChannelId,
message_id := MessageId,
guild_name := GuildName,
channel_name := ChannelName,
badge_counts_ttl_seconds := BadgeCountsTtlSeconds
} = Job
) when
is_map(MessageData),
is_integer(GuildId),
is_integer(ChannelId),
is_integer(MessageId),
is_integer(BadgeCountsTtlSeconds),
BadgeCountsTtlSeconds >= 0
->
push_send_job(UserIds, MessageData, GuildId, ChannelId, MessageId, #{
guild_name => GuildName,
channel_name => ChannelName,
badge_counts_ttl_seconds => BadgeCountsTtlSeconds,
markdown_context => maps:get(markdown_context, Job, #{})
});
push_message_create_job(_) ->
error.
-spec push_clear_channel_job(map()) -> {ok, push_job()} | error.
push_clear_channel_job(#{
user_id := UserId,
channel_id := ChannelId,
message_id := MessageId,
badge_counts_ttl_seconds := BadgeCountsTtlSeconds
}) when
is_integer(UserId),
is_integer(ChannelId),
is_integer(MessageId),
is_integer(BadgeCountsTtlSeconds),
BadgeCountsTtlSeconds >= 0
->
{ok, clear_channel_job(UserId, ChannelId, MessageId, BadgeCountsTtlSeconds)};
push_clear_channel_job(_) ->
error.
-spec clear_channel_job(integer(), integer(), integer(), non_neg_integer()) -> push_job().
clear_channel_job(UserId, ChannelId, MessageId, BadgeCountsTtlSeconds) ->
#{
type => clear_channel,
user_id => UserId,
channel_id => ChannelId,
message_id => MessageId,
badge_counts_ttl_seconds => BadgeCountsTtlSeconds
}.
-spec push_send_job(
term(), map(), integer(), integer(), integer(), #{
guild_name := term(),
channel_name := term(),
badge_counts_ttl_seconds := non_neg_integer(),
markdown_context => map()
}
) -> {ok, push_job()} | error.
push_send_job(UserIds0, MessageData, GuildId, ChannelId, MessageId, Options) ->
case send_job_options(UserIds0, Options) of
{ok, UserIds, GuildName, ChannelName, BadgeCountsTtlSeconds, MarkdownContext} ->
{ok,
send_notifications_job(
UserIds,
MessageData,
MarkdownContext,
GuildId,
ChannelId,
MessageId,
GuildName,
ChannelName,
BadgeCountsTtlSeconds
)};
error ->
error
end.
-spec send_job_options(term(), map()) ->
{ok, [integer()], binary() | undefined, binary() | undefined, non_neg_integer(), map()}
| error.
send_job_options(UserIds0, Options) ->
GuildName0 = maps:get(guild_name, Options),
ChannelName0 = maps:get(channel_name, Options),
BadgeCountsTtlSeconds = maps:get(badge_counts_ttl_seconds, Options),
MarkdownContext = maps:get(markdown_context, Options, #{}),
case
{
push_normalize:integer_list(UserIds0),
optional_binary(GuildName0),
optional_binary(ChannelName0)
}
of
{{ok, UserIds}, {ok, GuildName}, {ok, ChannelName}} when is_map(MarkdownContext) ->
{ok, UserIds, GuildName, ChannelName, BadgeCountsTtlSeconds, MarkdownContext};
_ ->
error
end.
-spec optional_binary(term()) -> {ok, binary() | undefined} | error.
optional_binary(undefined) ->
{ok, undefined};
optional_binary(Value) when is_binary(Value) ->
{ok, Value};
optional_binary(_) ->
error.
-ifdef(TEST).
-include_lib("eunit/include/eunit.hrl").
absent_counter_table_reads_unavailable_test() ->
delete_counter_table(),
?assertEqual(unavailable, counter_table_status()),
?assertEqual(unavailable, read_counter(?CNT_QUEUE_FULL)),
?assertEqual(ok, bump_counter(?CNT_QUEUE_FULL)),
?assertEqual(unavailable, read_counter(?CNT_QUEUE_FULL)).
absent_counter_key_is_distinct_from_zero_test() ->
with_counter_table(fun() ->
?assertEqual(live, counter_table_status()),
?assertEqual(unavailable, read_counter(?CNT_QUEUE_FULL)),
ok = bump_counter(?CNT_RESTART_DISCARDED, 0),
?assertEqual(0, read_counter(?CNT_RESTART_DISCARDED))
end).
bump_counter_creates_then_increments_key_test() ->
with_counter_table(fun() ->
ok = bump_counter(?CNT_QUEUE_FULL),
?assertEqual(1, read_counter(?CNT_QUEUE_FULL)),
ok = bump_counter(?CNT_QUEUE_FULL, 4),
?assertEqual(5, read_counter(?CNT_QUEUE_FULL))
end).
enqueue_exit_separates_timeout_from_failure_test() ->
with_counter_table(fun() ->
ok = count_enqueue_exit({timeout, {gen_server, call, []}}),
ok = count_enqueue_exit({noproc, {gen_server, call, []}}),
?assertEqual(1, read_counter(?CNT_ENQUEUE_TIMEOUT)),
?assertEqual(1, read_counter(?CNT_ENQUEUE_FAILED))
end).
worker_down_counts_only_abnormal_exits_test() ->
with_counter_table(fun() ->
ok = count_worker_down(normal),
?assertEqual(unavailable, read_counter(?CNT_WORKER_DIED)),
ok = count_worker_down(killed),
ok = count_worker_down({shutdown, restarting}),
?assertEqual(2, read_counter(?CNT_WORKER_DIED))
end).
terminate_counts_discarded_queue_depth_test() ->
with_counter_table(fun() ->
?assertEqual(ok, terminate(shutdown, dispatcher_state(7, 10))),
?assertEqual(7, read_counter(?CNT_RESTART_DISCARDED))
end).
invalid_job_enqueue_is_counted_test() ->
with_counter_table(fun() ->
State = dispatcher_state(0, 10),
?assertEqual({dropped, State}, handle_enqueue(#{type => invalid}, State)),
?assertEqual(1, read_counter(?CNT_INVALID_JOB))
end).
uptime_seconds_reports_undefined_without_start_time_test() ->
?assertEqual(undefined, uptime_seconds(undefined)),
?assertEqual(0, uptime_seconds(erlang:monotonic_time(second) + 5)),
?assert(is_integer(uptime_seconds(erlang:monotonic_time(second) - 3))).
abnormal_worker_down_also_counts_a_crashed_job_test() ->
with_counter_table(fun() ->
ok = count_worker_down(normal),
?assertEqual(unavailable, read_counter(?CNT_JOB_CRASHED)),
ok = count_worker_down(killed),
?assertEqual(1, read_counter(?CNT_JOB_CRASHED)),
?assertEqual(1, read_counter(?CNT_WORKER_DIED))
end).
enqueue_records_queue_growth_outside_the_process_test() ->
with_counter_table(fun() ->
{ok, State1} = maybe_enqueue(clear_job(), dispatcher_state(0, 10)),
?assertEqual(1, maps:get(queued, State1)),
?assertEqual(1, read_counter(?CNT_QUEUE_ENQUEUED)),
?assertEqual(unavailable, read_counter(?CNT_QUEUE_DEQUEUED))
end).
queue_full_drop_does_not_record_queue_growth_test() ->
with_counter_table(fun() ->
{dropped, _State1} = maybe_enqueue(clear_job(), dispatcher_state(3, 3)),
?assertEqual(1, read_counter(?CNT_QUEUE_FULL)),
?assertEqual(unavailable, read_counter(?CNT_QUEUE_ENQUEUED))
end).
queue_backlog_lost_survives_an_untrappable_kill_test() ->
with_counter_table(fun() ->
?assertEqual(unavailable, queue_backlog_lost(read_counter(?CNT_QUEUE_ENQUEUED), 0)),
ok = bump_counter(?CNT_QUEUE_ENQUEUED, 7),
ok = bump_counter(?CNT_QUEUE_DEQUEUED, 2),
?assertEqual(0, queue_backlog_lost(read_counter(?CNT_QUEUE_ENQUEUED), 5)),
?assertEqual(5, queue_backlog_lost(read_counter(?CNT_QUEUE_ENQUEUED), 0)),
?assertEqual(0, queue_backlog_lost(read_counter(?CNT_QUEUE_ENQUEUED), 900))
end).
dispatcher_state(Queued, MaxQueue) ->
#{
queue => queue:new(),
queued => Queued,
inflight => 0,
workers => #{},
max_inflight => 1,
max_queue => MaxQueue
}.
clear_job() ->
#{
type => clear_channel,
user_id => 1,
channel_id => 2,
message_id => 3,
badge_counts_ttl_seconds => 0
}.
with_counter_table(Fun) ->
delete_counter_table(),
_ = ets:new(?PUSH_COUNTER_TABLE, [named_table, public, set, {write_concurrency, true}]),
try
Fun()
after
delete_counter_table()
end.
delete_counter_table() ->
try ets:delete(?PUSH_COUNTER_TABLE) of
_ -> ok
catch
error:badarg -> ok
end.
-endif.
+55 -1
View File
@@ -141,7 +141,7 @@ fetch_settings(UserId, GuildId) ->
-spec fetch_settings_rpc(integer(), integer()) -> map().
fetch_settings_rpc(UserId, GuildId) ->
try push_subscriptions:fetch_and_cache_user_guild_settings(UserId, GuildId) of
try fetch_and_cache_user_guild_settings(UserId, GuildId) of
S0 when is_map(S0) -> S0;
_ -> #{}
catch
@@ -150,6 +150,60 @@ fetch_settings_rpc(UserId, GuildId) ->
exit:_ -> #{}
end.
-spec fetch_and_cache_user_guild_settings(integer(), integer()) -> map() | null.
fetch_and_cache_user_guild_settings(UserId, GuildId) ->
Req = #{
<<"type">> => <<"get_user_guild_settings">>,
<<"user_ids">> => [integer_to_binary(UserId)],
<<"guild_id">> => integer_to_binary(GuildId)
},
logger:debug(
"Push: fetching user guild settings via RPC",
#{user_id => UserId, guild_id => GuildId}
),
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, 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(), 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;
_ -> null
end,
case SettingsData of
null ->
logger:debug(
"Push: user guild settings returned null; caching empty sentinel",
#{user_id => UserId, guild_id => GuildId}
),
push_ets_cache:put_user_guild_settings(UserId, GuildId, #{}, Fill),
#{};
Settings ->
logger:debug(
"Push: user guild settings fetched and cached",
#{
user_id => UserId,
guild_id => GuildId,
muted => maps:get(muted, Settings, undefined),
mobile_push => maps:get(mobile_push, Settings, undefined)
}
),
push_ets_cache:put_user_guild_settings(UserId, GuildId, Settings, Fill),
Settings
end.
-spec prefetch_user_guild_settings([integer()], integer(), integer()) -> ok.
prefetch_user_guild_settings(_UserIds, _AuthorId, 0) ->
ok;
@@ -1,301 +0,0 @@
%% SPDX-License-Identifier: AGPL-3.0-or-later
-module(push_endpoint_guard).
-typing([eqwalizer]).
-export([check/1, check/2, check/3]).
-export_type([resolver/0, cache_mode/0, verdict/0]).
-define(RESOLVE_TIMEOUT_MS, 3000).
-define(MAX_HOST_LENGTH, 253).
-define(MAX_LABEL_LENGTH, 63).
-define(ALLOWED_VERDICT_TTL_SECONDS, 300).
-define(REFUSED_VERDICT_TTL_SECONDS, 30).
-type resolver() :: fun((string()) -> {ok, [inet:ip_address()]} | {error, term()}).
-type cache_mode() :: cached | uncached.
-type verdict() :: ok | {error, term()}.
-spec check(binary()) -> verdict().
check(Endpoint) ->
check(Endpoint, fun resolve/1, cached).
-spec check(binary(), resolver()) -> verdict().
check(Endpoint, Resolver) ->
check(Endpoint, Resolver, uncached).
-spec check(binary(), resolver(), cache_mode()) -> verdict().
check(Endpoint, Resolver, CacheMode) ->
case enabled() of
true -> check_endpoint(Endpoint, Resolver, CacheMode);
false -> ok
end.
-spec enabled() -> boolean().
enabled() ->
case fluxer_gateway_env:get(push_endpoint_guard_enabled) of
Enabled when is_boolean(Enabled) -> Enabled;
_ -> true
end.
-spec check_endpoint(binary(), resolver(), cache_mode()) -> verdict().
check_endpoint(Endpoint, Resolver, CacheMode) ->
case parse_endpoint(Endpoint) of
{ok, Host} -> check_host(Host, Resolver, CacheMode);
{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(), cache_mode()) -> verdict().
check_host(Host, Resolver, CacheMode) ->
case inet:parse_address(Host) of
{ok, Address} -> check_addresses([Address]);
{error, _Reason} -> check_hostname(Host, Resolver, CacheMode)
end.
-spec check_hostname(string(), resolver(), cache_mode()) -> verdict().
check_hostname(Host, Resolver, CacheMode) ->
case is_fqdn(Host) of
true -> resolve_and_check(Host, Resolver, CacheMode);
false -> {error, endpoint_rejected}
end.
-spec resolve_and_check(string(), resolver(), cache_mode()) -> verdict().
resolve_and_check(Host, Resolver, uncached) ->
resolve_and_check(Host, Resolver);
resolve_and_check(Host, Resolver, cached) ->
CacheKey = list_to_binary(Host),
case push_ets_cache:get_endpoint_verdict(CacheKey) of
{ok, Verdict} -> Verdict;
undefined -> store_verdict(CacheKey, resolve_and_check(Host, Resolver))
end.
-spec store_verdict(binary(), verdict()) -> verdict().
store_verdict(CacheKey, Verdict) ->
Ttl = verdict_ttl_seconds(Verdict),
ok = push_ets_cache:put_endpoint_verdict(CacheKey, Verdict, Ttl),
Verdict.
-spec verdict_ttl_seconds(verdict()) -> pos_integer().
verdict_ttl_seconds(ok) -> ?ALLOWED_VERDICT_TTL_SECONDS;
verdict_ttl_seconds({error, _Reason}) -> ?REFUSED_VERDICT_TTL_SECONDS.
-spec resolve_and_check(string(), resolver()) -> verdict().
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) -> "".
+4 -175
View File
@@ -10,23 +10,10 @@
put_user_guild_settings/4,
delete_user_guild_settings/2,
reserve_user_guild_settings/2,
get_subscriptions/1,
get_subscriptions_many/1,
put_subscriptions/3,
delete_subscriptions/1,
reserve_subscriptions/1,
get_blocked_ids/1,
put_blocked_ids/2,
put_blocked_ids_fetched/3,
reserve_blocked_ids/1,
get_badge_count/1,
put_badge_count/4,
delete_badge_count/1,
reserve_badge_counts/1,
get_bearer_token/1,
put_bearer_token/3,
get_endpoint_verdict/1,
put_endpoint_verdict/3,
release/1,
rebalance/0,
rebalance_async/0,
@@ -35,20 +22,12 @@
table_size/1
]).
-export_type([fill/0, endpoint_verdict/0]).
-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(ENDPOINT_VERDICTS, push_endpoint_verdicts).
-define(MAX_TABLE_ENTRIES, 500000).
-define(MAX_BEARER_TOKENS, 10000).
-define(MAX_ENDPOINT_VERDICTS, 2048).
-define(MAX_ENDPOINT_HOST_BYTES, 253).
-define(ENDPOINT_VERDICT_EVICT_BATCH, 512).
-define(EVICT_BATCH, 4096).
-define(MAX_EVICT_RESEEKS, 8).
-define(RESERVATION_TTL_MS, 120000).
@@ -61,16 +40,11 @@
]).
-type fill() :: {atom(), pos_integer(), [term()]}.
-type endpoint_verdict() :: ok | {error, 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),
ensure_table(?ENDPOINT_VERDICTS),
ok.
-spec get_user_guild_settings(integer(), integer()) -> map() | undefined.
@@ -98,45 +72,6 @@ delete_user_guild_settings(UserId, GuildId) ->
reserve_user_guild_settings(UserIds, GuildId) ->
reserve(?USER_GUILD_SETTINGS, [{UserId, GuildId} || UserId <- UserIds]).
-spec get_subscriptions(integer()) -> list() | undefined.
get_subscriptions(UserId) ->
try ets:lookup(?SUBSCRIPTIONS, UserId) of
[{UserId, Subs}] when is_list(Subs) -> Subs;
_ -> undefined
catch
error:badarg -> undefined
end.
-spec get_subscriptions_many([integer()]) -> {#{integer() => list()}, [integer()]}.
get_subscriptions_many(UserIds) ->
lists:foldl(
fun add_cached_subscriptions/2,
{#{}, []},
UserIds
).
-spec add_cached_subscriptions(integer(), {#{integer() => list()}, [integer()]}) ->
{#{integer() => list()}, [integer()]}.
add_cached_subscriptions(UserId, {CachedAcc, MissingAcc}) ->
case get_subscriptions(UserId) of
Subscriptions when is_list(Subscriptions) ->
{CachedAcc#{UserId => Subscriptions}, MissingAcc};
undefined ->
{CachedAcc, [UserId | MissingAcc]}
end.
-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) ->
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) ->
try ets:lookup(?BLOCKED_IDS, UserId) of
@@ -184,91 +119,6 @@ app_pos_integer(Key, Default) ->
_ -> Default
end.
-spec get_badge_count(integer()) -> {non_neg_integer(), integer()} | undefined.
get_badge_count(UserId) ->
try ets:lookup(?BADGE_COUNTS, UserId) of
[{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(), 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) ->
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 get_endpoint_verdict(binary()) -> {ok, endpoint_verdict()} | undefined.
get_endpoint_verdict(Host) when is_binary(Host) ->
try ets:lookup(?ENDPOINT_VERDICTS, Host) of
[{_, ok, ExpiresAt}] when is_integer(ExpiresAt) ->
live_endpoint_verdict(ok, ExpiresAt);
[{_, {error, Reason}, ExpiresAt}] when is_integer(ExpiresAt) ->
live_endpoint_verdict({error, Reason}, ExpiresAt);
_ ->
undefined
catch
error:badarg -> undefined
end.
-spec live_endpoint_verdict(endpoint_verdict(), integer()) ->
{ok, endpoint_verdict()} | undefined.
live_endpoint_verdict(Verdict, ExpiresAt) ->
case erlang:system_time(second) < ExpiresAt of
true -> {ok, Verdict};
false -> undefined
end.
-spec put_endpoint_verdict(binary(), endpoint_verdict(), pos_integer()) -> ok.
put_endpoint_verdict(Host, Verdict, TtlSeconds) when
is_binary(Host), is_integer(TtlSeconds), TtlSeconds > 0
->
case byte_size(Host) =< ?MAX_ENDPOINT_HOST_BYTES of
true -> insert_endpoint_verdict(Host, Verdict, TtlSeconds);
false -> ok
end.
-spec insert_endpoint_verdict(binary(), endpoint_verdict(), pos_integer()) -> ok.
insert_endpoint_verdict(Host, Verdict, TtlSeconds) ->
guard_table_size(
?ENDPOINT_VERDICTS, ?MAX_ENDPOINT_VERDICTS, ?ENDPOINT_VERDICT_EVICT_BATCH
),
ExpiresAt = erlang:system_time(second) + TtlSeconds,
try ets:insert(?ENDPOINT_VERDICTS, {Host, Verdict, ExpiresAt}) of
_ -> ok
catch
error:badarg -> ok
end.
-spec write(atom(), tuple()) -> ok.
write(Table, Row) ->
guard_table_size(Table, ?MAX_TABLE_ENTRIES),
@@ -300,8 +150,6 @@ reserve_key(Table, Key, Token, ReservedAt) ->
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) ->
[].
@@ -330,38 +178,23 @@ rebalance_async() ->
rebalance() ->
init(),
_ = rebalance_table(?USER_GUILD_SETTINGS, fun user_id_from_user_guild_key/1),
_ = rebalance_table(?SUBSCRIPTIONS, fun user_id_from_key/1),
_ = rebalance_table(?BLOCKED_IDS, fun user_id_from_key/1),
_ = rebalance_table(?BADGE_COUNTS, fun user_id_from_key/1),
ok.
-spec cache_stats() -> map().
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),
bearer_tokens_size => table_size(?BEARER_TOKENS),
endpoint_verdicts_size => table_size(?ENDPOINT_VERDICTS)
blocked_ids_size => table_size(?BLOCKED_IDS)
}.
-spec evict_tables(map()) -> ok.
evict_tables(MaxEntries) ->
Now = erlang:system_time(second),
select_delete(?BLOCKED_IDS, expired_rows(Now)),
select_delete(?BEARER_TOKENS, expired_rows(Now)),
select_delete(?ENDPOINT_VERDICTS, expired_rows(Now)),
lists:foreach(
fun expire_reservations/1,
[?USER_GUILD_SETTINGS, ?SUBSCRIPTIONS, ?BLOCKED_IDS, ?BADGE_COUNTS]
),
lists:foreach(fun expire_reservations/1, [?USER_GUILD_SETTINGS, ?BLOCKED_IDS]),
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),
evict_table(?ENDPOINT_VERDICTS, ?MAX_ENDPOINT_VERDICTS),
ok.
-spec expired_rows(integer()) -> ets:match_spec().
@@ -383,12 +216,8 @@ select_delete(Table, MatchSpec) ->
-spec guard_table_size(atom(), non_neg_integer()) -> ok.
guard_table_size(Table, MaxEntries) ->
guard_table_size(Table, MaxEntries, ?EVICT_BATCH).
-spec guard_table_size(atom(), non_neg_integer(), pos_integer()) -> ok.
guard_table_size(Table, MaxEntries, EvictBatch) ->
case table_size(Table) >= MaxEntries of
true -> evict_table(Table, max(0, MaxEntries - EvictBatch));
true -> evict_table(Table, max(0, MaxEntries - ?EVICT_BATCH));
false -> ok
end.
-327
View File
@@ -1,327 +0,0 @@
%% SPDX-License-Identifier: AGPL-3.0-or-later
-module(push_fcm).
-typing([eqwalizer]).
-export([send/3]).
-define(FCM_SCOPE, <<"https://www.googleapis.com/auth/firebase.messaging">>).
-define(DEFAULT_TOKEN_URI, <<"https://oauth2.googleapis.com/token">>).
-define(ACCESS_TOKEN_SKEW_SECONDS, 60).
-spec send(integer(), map(), map()) -> false | {true, map()}.
send(UserId, Subscription, Payload) ->
case fluxer_gateway_env:get(fcm_enabled) of
true ->
send_enabled(UserId, Subscription, Payload);
_ ->
false
end.
-spec send_enabled(integer(), map(), map()) -> false | {true, map()}.
send_enabled(UserId, Subscription, Payload) ->
case extract_subscription(Subscription) of
{ok, DeviceToken, SubscriptionId, AppId} ->
send_with_config(UserId, DeviceToken, SubscriptionId, AppId, Payload);
{error, Reason} ->
log_config_error(UserId, <<"invalid_fcm_subscription">>, Reason),
false
end.
-spec send_with_config(integer(), binary(), binary(), binary(), map()) -> false | {true, map()}.
send_with_config(UserId, DeviceToken, SubscriptionId, AppId, Payload) ->
case {resolve_project_id(AppId), resolve_access_token()} of
{{ok, ProjectId}, {ok, AccessToken}} ->
Message = push_fcm_payload:build_message(DeviceToken, Payload),
Body = iolist_to_binary(json:encode(Message)),
Url =
<<"https://fcm.googleapis.com/v1/projects/", ProjectId/binary,
"/messages:send">>,
Headers = [
{<<"Authorization">>, <<"Bearer ", AccessToken/binary>>},
{<<"Content-Type">>, <<"application/json; charset=UTF-8">>}
],
Response = gateway_http_client:request(
push,
post,
Url,
Headers,
Body,
#{content_type => <<"application/json; charset=UTF-8">>}
),
push_fcm_payload:handle_response(UserId, SubscriptionId, Response);
{{error, Reason}, _} ->
log_config_error(UserId, <<"fcm_project_error">>, Reason),
false;
{_, {error, Reason}} ->
log_config_error(UserId, <<"fcm_auth_error">>, Reason),
false
end.
-spec extract_subscription(map()) -> {ok, binary(), binary(), binary()} | {error, term()}.
extract_subscription(Subscription) ->
DeviceToken = push_utils:normalize_binary(
maps:get(<<"endpoint">>, Subscription, undefined), undefined
),
SubscriptionId = push_utils:normalize_binary(
maps:get(<<"subscription_id">>, Subscription, undefined), undefined
),
AppId = push_utils:normalize_binary(
maps:get(<<"app_id">>, Subscription, <<"stable">>), <<"stable">>
),
case {DeviceToken, SubscriptionId, AppId} of
{Token, Id, App} when
is_binary(Token),
byte_size(Token) > 0,
is_binary(Id),
is_binary(App)
->
{ok, Token, Id, App};
_ ->
{error, missing_fields}
end.
-spec resolve_project_id(binary()) -> {ok, binary()} | {error, term()}.
resolve_project_id(AppId) ->
Apps = fluxer_gateway_env:get(fcm_apps),
case find_app(AppId, map_utils:ensure_list(Apps)) of
App when is_map(App) -> resolve_app_project_id(App);
undefined -> resolve_default_project_id()
end.
-spec resolve_app_project_id(map()) -> {ok, binary()} | {error, term()}.
resolve_app_project_id(App) ->
case push_utils:normalize_binary(get_map_value(App, <<"project_id">>), undefined) of
ProjectId when is_binary(ProjectId), byte_size(ProjectId) > 0 -> {ok, ProjectId};
_ -> resolve_default_project_id()
end.
-spec resolve_default_project_id() -> {ok, binary()} | {error, term()}.
resolve_default_project_id() ->
case fluxer_gateway_env:get(fcm_project_id) of
ProjectId when is_binary(ProjectId), byte_size(ProjectId) > 0 -> {ok, ProjectId};
_ -> {error, missing_project_id}
end.
-spec resolve_access_token() -> {ok, binary()} | {error, term()}.
resolve_access_token() ->
maybe
{ok, ServiceAccount} ?= resolve_service_account(),
ClientEmail = maps:get(client_email, ServiceAccount),
TokenUri = maps:get(token_uri, ServiceAccount),
CacheKey = {?MODULE, access_token, ClientEmail},
Now = erlang:system_time(second),
get_or_fetch_token(ServiceAccount, CacheKey, Now, TokenUri)
else
{error, Reason} -> {error, Reason}
end.
-spec get_or_fetch_token(map(), term(), integer(), binary()) ->
{ok, binary()} | {error, term()}.
get_or_fetch_token(ServiceAccount, CacheKey, Now, TokenUri) ->
case push_ets_cache:get_bearer_token(CacheKey) of
{ok, Token, ExpiresAt} when ExpiresAt - ?ACCESS_TOKEN_SKEW_SECONDS > Now ->
{ok, Token};
_ ->
fetch_access_token(ServiceAccount, CacheKey, Now, TokenUri)
end.
-spec fetch_access_token(map(), term(), integer(), binary()) ->
{ok, binary()} | {error, term()}.
fetch_access_token(ServiceAccount, CacheKey, Now, TokenUri) ->
Claims = #{
<<"iss">> => maps:get(client_email, ServiceAccount),
<<"scope">> => ?FCM_SCOPE,
<<"aud">> => TokenUri,
<<"iat">> => Now,
<<"exp">> => Now + 3600
},
Header = #{<<"alg">> => <<"RS256">>, <<"typ">> => <<"JWT">>},
maybe
PrivateKey = maps:get(private_key, ServiceAccount),
{ok, Assertion} ?= push_utils:generate_jwt_from_pem(PrivateKey, Header, Claims),
exchange_assertion_for_token(CacheKey, Now, TokenUri, Assertion)
else
{error, _Reason} -> {error, jwt_signing_failed}
end.
-spec exchange_assertion_for_token(term(), integer(), binary(), binary()) ->
{ok, binary()} | {error, term()}.
exchange_assertion_for_token(CacheKey, Now, TokenUri, Assertion) ->
Body =
<<"grant_type=urn%3Aietf%3Aparams%3Aoauth%3Agrant-type%3Ajwt-bearer&assertion=",
Assertion/binary>>,
Headers = [{<<"Content-Type">>, <<"application/x-www-form-urlencoded">>}],
case
gateway_http_client:request(
push,
post,
TokenUri,
Headers,
Body,
#{content_type => <<"application/x-www-form-urlencoded">>}
)
of
{ok, Status, _Headers, ResponseBody} when Status >= 200, Status < 300 ->
parse_token_response(CacheKey, Now, ResponseBody);
{ok, Status, _Headers, ResponseBody} ->
{error, {token_http_error, Status, ResponseBody}};
{error, Reason} ->
{error, {token_request_failed, Reason}}
end.
-spec parse_token_response(term(), integer(), binary()) -> {ok, binary()} | {error, term()}.
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_ets_cache:put_bearer_token(CacheKey, AccessToken, Now + ExpiresIn),
{ok, AccessToken};
_ ->
{error, invalid_token_response}
end.
-spec resolve_service_account() -> {ok, map()} | {error, term()}.
resolve_service_account() ->
JsonPath = fluxer_gateway_env:get(fcm_service_account_json_path),
maybe
{ok, Json} ?= load_json_config(JsonPath),
build_service_account(Json)
else
{error, Reason} -> {error, Reason}
end.
-spec load_json_config(binary() | term()) -> {ok, map()} | {error, term()}.
load_json_config(Path) when is_binary(Path), byte_size(Path) > 0 ->
read_json_file(Path);
load_json_config(_) ->
{ok, #{}}.
-spec build_service_account(map()) -> {ok, map()} | {error, term()}.
build_service_account(Json) ->
ClientEmail = first_binary([
get_map_value(Json, <<"client_email">>), fluxer_gateway_env:get(fcm_client_email)
]),
TokenUri = first_binary([
get_map_value(Json, <<"token_uri">>),
fluxer_gateway_env:get(fcm_token_uri),
?DEFAULT_TOKEN_URI
]),
PrivateKey = resolve_private_key(Json),
validate_service_account(ClientEmail, TokenUri, PrivateKey).
-spec validate_service_account(
binary() | undefined, binary() | undefined, {ok, binary()} | {error, term()}
) ->
{ok, map()} | {error, term()}.
validate_service_account(Email, Uri, {ok, Key}) when
is_binary(Email), byte_size(Email) > 0, is_binary(Uri), byte_size(Uri) > 0
->
{ok, #{client_email => Email, token_uri => Uri, private_key => Key}};
validate_service_account(undefined, _, _) ->
{error, missing_client_email};
validate_service_account(_, undefined, _) ->
{error, missing_token_uri};
validate_service_account(_, _, {error, Reason}) ->
{error, Reason};
validate_service_account(_, _, _) ->
{error, invalid_service_account}.
-spec resolve_private_key(map()) -> {ok, binary()} | {error, term()}.
resolve_private_key(Json) ->
case
first_binary([
get_map_value(Json, <<"private_key">>), fluxer_gateway_env:get(fcm_private_key)
])
of
Key when is_binary(Key), byte_size(Key) > 0 ->
{ok, normalize_pem(Key)};
_ ->
resolve_private_key_from_file()
end.
-spec resolve_private_key_from_file() -> {ok, binary()} | {error, term()}.
resolve_private_key_from_file() ->
case fluxer_gateway_env:get(fcm_private_key_path) of
Path when is_binary(Path), byte_size(Path) > 0 -> read_pem_file(Path);
_ -> {error, missing_private_key}
end.
-spec find_app(binary(), list()) -> map() | undefined.
find_app(_AppId, []) ->
undefined;
find_app(AppId, [App | Rest]) when is_map(App) ->
case push_utils:normalize_binary(get_map_value(App, <<"app_id">>), <<>>) of
AppId -> App;
_ -> find_app(AppId, Rest)
end;
find_app(AppId, [_ | Rest]) ->
find_app(AppId, Rest).
-spec read_json_file(binary()) -> {ok, map()} | {error, term()}.
read_json_file(Path) ->
case file:read_file(binary_to_list(Path)) of
{ok, Content} -> decode_json_file(Content);
{error, Reason} -> {error, {read_json_failed, Reason}}
end.
-spec decode_json_file(binary()) -> {ok, map()} | {error, invalid_json_file}.
decode_json_file(Content) ->
case decode_json_map(Content) of
Map when is_map(Map) -> {ok, Map};
_ -> {error, invalid_json_file}
end.
-spec read_pem_file(binary()) -> {ok, binary()} | {error, term()}.
read_pem_file(Path) ->
case file:read_file(binary_to_list(Path)) of
{ok, Content} -> {ok, normalize_pem(Content)};
{error, Reason} -> {error, {read_private_key_failed, Reason}}
end.
-spec decode_json_map(binary()) -> map() | undefined.
decode_json_map(Body) when is_binary(Body), byte_size(Body) > 0 ->
try json:decode(Body) of
Map when is_map(Map) -> Map;
_ -> undefined
catch
error:_ -> undefined;
throw:_ -> undefined;
exit:_ -> undefined
end;
decode_json_map(_) ->
undefined.
-spec normalize_expires_in(term()) -> pos_integer().
normalize_expires_in(Value) ->
case guild_data_normalize_schema:int(Value) of
ExpiresIn when is_integer(ExpiresIn), ExpiresIn > 0 -> ExpiresIn;
_ -> 3600
end.
-spec first_binary(list()) -> binary() | undefined.
first_binary([]) ->
undefined;
first_binary([Value | Rest]) ->
case push_utils:normalize_binary(Value, undefined) of
Bin when is_binary(Bin), byte_size(Bin) > 0 -> Bin;
_ -> first_binary(Rest)
end.
-spec get_map_value(map(), binary()) -> term().
get_map_value(Map, Key) when is_map(Map), is_binary(Key) ->
case maps:get(Key, Map, undefined) of
undefined -> maps:get(binary_to_list(Key), Map, undefined);
Value -> Value
end.
-spec normalize_pem(binary()) -> binary().
normalize_pem(Pem) -> binary:replace(Pem, <<"\\n">>, <<"\n">>, [global]).
-spec log_config_error(integer(), binary(), term()) -> ok.
log_config_error(UserId, ReasonCode, Reason) ->
logger:debug(
"Push: FCM delivery unavailable",
#{user_id => UserId, reason => ReasonCode, detail => Reason}
),
ok.
@@ -1,345 +0,0 @@
%% SPDX-License-Identifier: AGPL-3.0-or-later
-module(push_fcm_payload).
-typing([eqwalizer]).
-export([
build_message/2,
handle_response/3,
delete_payload/2
]).
-export_type([fcm_response/0]).
-type fcm_response() :: {ok, integer(), term(), binary()} | {error, term()}.
-spec build_message(binary(), map()) -> map().
build_message(DeviceToken, #{<<"type">> := <<"notification_clear">>} = Payload) ->
build_clear_message(DeviceToken, Payload);
build_message(DeviceToken, Payload) ->
build_notification_message(DeviceToken, Payload).
-spec build_clear_message(binary(), map()) -> map().
build_clear_message(DeviceToken, Payload) ->
Data0 = data_as_strings(maps:get(<<"data">>, Payload, #{})),
Tag = push_utils:normalize_binary(
maps:get(
<<"notification_tag">>, Payload, maps:get(<<"tag">>, Payload, <<"fluxer-message">>)
),
<<"fluxer-message">>
),
Data = maps:merge(
Data0,
data_as_strings(#{
<<"type">> => <<"notification_clear">>,
<<"action">> => <<"clear_channel">>,
<<"notification_tag">> => Tag
})
),
#{
<<"message">> => #{
<<"token">> => DeviceToken,
<<"data">> => Data,
<<"android">> => #{
<<"priority">> => <<"NORMAL">>,
<<"ttl">> => <<"3600s">>,
<<"collapse_key">> => <<"clear:", Tag/binary>>
},
<<"fcm_options">> => #{
<<"analytics_label">> => <<"notification_clear">>
}
}
}.
-spec build_notification_message(binary(), map()) -> map().
build_notification_message(DeviceToken, Payload) ->
Notification = maps:get(<<"notification">>, Payload, #{}),
Title = resolve_title(Notification, Payload),
Body = resolve_body(Notification, Payload),
Tag = push_utils:normalize_binary(
maps:get(<<"tag">>, Payload, <<"fluxer-message">>), <<"fluxer-message">>
),
ImageUrl = resolve_image_url(Notification, Payload),
NotificationBody = maybe_put(
<<"image">>, ImageUrl, #{<<"title">> => Title, <<"body">> => Body}
),
Data = build_notification_data(Payload, Title, Body, Tag, ImageUrl),
AndroidNotification = build_android_notification(Tag, ImageUrl),
wrap_notification_message(DeviceToken, NotificationBody, Data, AndroidNotification).
-spec build_android_notification(binary(), binary() | undefined) -> map().
build_android_notification(Tag, ImageUrl) ->
maybe_put(<<"image">>, ImageUrl, #{
<<"channel_id">> => <<"fluxer_default_push">>,
<<"tag">> => Tag,
<<"click_action">> => <<"FLUXER_MESSAGE">>
}).
-spec wrap_notification_message(binary(), map(), map(), map()) -> map().
wrap_notification_message(DeviceToken, NotificationBody, Data, AndroidNotification) ->
#{
<<"message">> => #{
<<"token">> => DeviceToken,
<<"notification">> => NotificationBody,
<<"data">> => Data,
<<"android">> => #{
<<"priority">> => <<"HIGH">>,
<<"ttl">> => <<"86400s">>,
<<"notification">> => AndroidNotification
},
<<"fcm_options">> => #{
<<"analytics_label">> => <<"message_create">>
}
}
}.
-spec build_notification_data(map(), binary(), binary(), binary(), binary() | undefined) ->
map().
build_notification_data(Payload, Title, Body, Tag, ImageUrl) ->
Data0 = data_as_strings(maps:get(<<"data">>, Payload, #{})),
DataExtra = maybe_put(<<"image_url">>, ImageUrl, #{
<<"title">> => Title,
<<"body">> => Body,
<<"tag">> => Tag
}),
maps:merge(Data0, data_as_strings(DataExtra)).
-spec resolve_title(map(), map()) -> binary().
resolve_title(Notification, Payload) ->
sanitize_text(
push_utils:normalize_binary(
maps:get(<<"title">>, Notification, maps:get(<<"title">>, Payload, <<"Fluxer">>)),
<<"Fluxer">>
)
).
-spec resolve_body(map(), map()) -> binary().
resolve_body(Notification, Payload) ->
sanitize_text(
push_utils:normalize_binary(
maps:get(<<"body">>, Notification, maps:get(<<"body">>, Payload, <<"">>)),
<<"">>
)
).
-spec sanitize_text(binary()) -> binary().
sanitize_text(Bin) when is_binary(Bin) ->
case unicode:characters_to_list(Bin) of
Codepoints when is_list(Codepoints) ->
safe_codepoints_to_binary(Codepoints, Bin);
_ ->
Bin
end.
-spec safe_codepoints_to_binary([integer()], binary()) -> binary().
safe_codepoints_to_binary(Codepoints, Fallback) ->
Filtered = [C || C <- Codepoints, is_safe_codepoint(C)],
case unicode:characters_to_binary(Filtered) of
Result when is_binary(Result) -> Result;
_ -> Fallback
end.
-spec is_safe_codepoint(integer()) -> boolean().
is_safe_codepoint(C) when C =:= $\n; C =:= $\t -> true;
is_safe_codepoint(C) when C >= 0, C =< 16#1F -> false;
is_safe_codepoint(C) when C >= 16#7F, C =< 16#9F -> false;
is_safe_codepoint(C) when C >= 16#200E, C =< 16#200F -> false;
is_safe_codepoint(C) when C >= 16#202A, C =< 16#202E -> false;
is_safe_codepoint(C) when C >= 16#2066, C =< 16#2069 -> false;
is_safe_codepoint(_) -> true.
-spec resolve_image_url(map(), map()) -> binary() | undefined.
resolve_image_url(Notification, Payload) ->
first_binary([
maps:get(<<"image_url">>, Payload, undefined),
maps:get(<<"image">>, Notification, undefined),
maps:get(<<"image_url">>, Notification, undefined)
]).
-spec handle_response(integer(), binary(), fcm_response()) -> false | {true, map()}.
handle_response(_UserId, _SubscriptionId, {ok, Status, _, _}) when
Status >= 200, Status < 300
->
false;
handle_response(UserId, SubscriptionId, {ok, _Status, _, Body}) ->
Reason = fcm_error_code(Body),
case is_permanent_fcm_error(Reason) of
true -> {true, delete_payload(UserId, SubscriptionId)};
false -> false
end;
handle_response(UserId, _SubscriptionId, {error, Reason}) ->
logger:debug("Push: FCM network error", #{user_id => UserId, reason => Reason}),
false.
-spec delete_payload(integer(), binary()) -> map().
delete_payload(UserId, SubscriptionId) ->
#{
<<"user_id">> => integer_to_binary(UserId),
<<"subscription_id">> => SubscriptionId
}.
-spec data_as_strings(term()) -> map().
data_as_strings(Data) when is_map(Data) ->
maps:fold(
fun(Key, Value, Acc) ->
Acc#{push_utils:normalize_binary(Key, <<>>) => stringify_value(Value)}
end,
#{},
Data
);
data_as_strings(_) ->
#{}.
-spec maybe_put(binary(), binary() | undefined, map()) -> map().
maybe_put(_Key, undefined, Map) ->
Map;
maybe_put(Key, Value, Map) when is_binary(Value), byte_size(Value) > 0 ->
Map#{Key => Value};
maybe_put(_Key, _Value, Map) ->
Map.
-spec stringify_value(term()) -> binary().
stringify_value(Value) when is_binary(Value) -> Value;
stringify_value(Value) when is_integer(Value) -> integer_to_binary(Value);
stringify_value(Value) when is_float(Value) -> list_to_binary(io_lib:format("~p", [Value]));
stringify_value(Value) when is_atom(Value) -> atom_to_binary(Value, utf8);
stringify_value(Value) when is_list(Value) -> stringify_list_value(Value);
stringify_value(Value) -> stringify_json_value(Value).
-spec stringify_list_value([term()]) -> binary().
stringify_list_value(Value) ->
case type_conv:to_binary(Value) of
Bin when is_binary(Bin) -> Bin;
undefined -> stringify_json_value(Value)
end.
-spec stringify_json_value(term()) -> binary().
stringify_json_value(Value) ->
iolist_to_binary(json:encode(json_encodable_value(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_encodable_value(Item)}
end,
#{},
Value
);
json_encodable_value(Value) ->
iolist_to_binary(io_lib:format("~p", [Value])).
-spec first_binary(list()) -> binary() | undefined.
first_binary([]) ->
undefined;
first_binary([Value | Rest]) ->
case push_utils:normalize_binary(Value, undefined) of
Bin when is_binary(Bin), byte_size(Bin) > 0 -> Bin;
_ -> first_binary(Rest)
end.
-spec fcm_error_code(binary()) -> binary().
fcm_error_code(Body) ->
case decode_json_map(Body) of
#{<<"error">> := #{<<"details">> := Details}} when is_list(Details) ->
fcm_details_error_code(Details);
#{<<"error">> := #{<<"status">> := Status}} when is_binary(Status) ->
Status;
_ ->
<<"http_error">>
end.
-spec fcm_details_error_code(list()) -> binary().
fcm_details_error_code(Details) ->
case find_fcm_error_code(Details) of
undefined -> <<"http_error">>;
Code -> Code
end.
-spec find_fcm_error_code(list()) -> binary() | undefined.
find_fcm_error_code([]) -> undefined;
find_fcm_error_code([#{<<"errorCode">> := Code} | _]) when is_binary(Code) -> Code;
find_fcm_error_code([_ | Rest]) -> find_fcm_error_code(Rest).
-spec is_permanent_fcm_error(binary()) -> boolean().
is_permanent_fcm_error(<<"UNREGISTERED">>) -> true;
is_permanent_fcm_error(<<"INVALID_ARGUMENT">>) -> true;
is_permanent_fcm_error(_) -> false.
-spec decode_json_map(binary()) -> map() | undefined.
decode_json_map(Body) when is_binary(Body), byte_size(Body) > 0 ->
try json:decode(Body) of
Map when is_map(Map) -> Map;
_ -> undefined
catch
error:_ -> undefined;
throw:_ -> undefined;
exit:_ -> undefined
end;
decode_json_map(_) ->
undefined.
-ifdef(TEST).
-include_lib("eunit/include/eunit.hrl").
build_message_includes_android_chat_notification_fields_test() ->
Payload = #{
<<"title">> => <<"Alice">>,
<<"body">> => <<"Hello">>,
<<"tag">> => <<"channel:123:456">>,
<<"image_url">> => <<"https://cdn.example/image.png">>,
<<"data">> => #{
<<"channel_id">> => <<"123">>,
<<"message_id">> => <<"456">>,
<<"notification_tag">> => <<"channel:123">>,
<<"badge_count">> => 4
},
<<"notification">> => #{
<<"title">> => <<"Alice">>,
<<"body">> => <<"Hello">>
}
},
#{<<"message">> := Message} = build_message(<<"device-token">>, Payload),
?assertEqual(<<"device-token">>, maps:get(<<"token">>, Message)),
?assertEqual(<<"HIGH">>, maps:get(<<"priority">>, maps:get(<<"android">>, Message))),
?assertEqual(<<"86400s">>, maps:get(<<"ttl">>, maps:get(<<"android">>, Message))),
AndroidNotification = maps:get(<<"notification">>, maps:get(<<"android">>, Message)),
?assertEqual(<<"fluxer_default_push">>, maps:get(<<"channel_id">>, AndroidNotification)),
?assertEqual(<<"channel:123:456">>, maps:get(<<"tag">>, AndroidNotification)),
Android = maps:get(<<"android">>, Message),
?assertEqual(false, maps:is_key(<<"collapse_key">>, Android)),
?assertEqual(
<<"https://cdn.example/image.png">>, maps:get(<<"image">>, AndroidNotification)
),
Data = maps:get(<<"data">>, Message),
?assertEqual(<<"4">>, maps:get(<<"badge_count">>, Data)),
?assertEqual(<<"https://cdn.example/image.png">>, maps:get(<<"image_url">>, Data)).
build_clear_message_is_data_only_and_collapsible_test() ->
Payload = #{
<<"type">> => <<"notification_clear">>,
<<"tag">> => <<"channel:123">>,
<<"data">> => #{
<<"channel_id">> => <<"123">>,
<<"message_id">> => <<"456">>,
<<"badge_count">> => 0
}
},
#{<<"message">> := Message} = build_message(<<"device-token">>, Payload),
?assertEqual(false, maps:is_key(<<"notification">>, Message)),
Android = maps:get(<<"android">>, Message),
?assertEqual(<<"NORMAL">>, maps:get(<<"priority">>, Android)),
?assertEqual(<<"3600s">>, maps:get(<<"ttl">>, Android)),
?assertEqual(<<"clear:channel:123">>, maps:get(<<"collapse_key">>, Android)),
Data = maps:get(<<"data">>, Message),
?assertEqual(<<"notification_clear">>, maps:get(<<"type">>, Data)),
?assertEqual(<<"clear_channel">>, maps:get(<<"action">>, Data)),
?assertEqual(<<"channel:123">>, maps:get(<<"notification_tag">>, Data)),
?assertEqual(<<"0">>, maps:get(<<"badge_count">>, Data)).
-endif.
+329 -93
View File
@@ -3,22 +3,22 @@
-module(push_job_publisher).
-typing([eqwalizer]).
-export([publish_message/8, publish_message/10, publish_clear/3, publish_clear/5]).
-export([publish_message/8, publish_clear/3]).
-export([publish_ring/6, request/3]).
-define(SUBJECT_MESSAGE, <<"push.job.message">>).
-define(SUBJECT_CLEAR, <<"push.job.clear">>).
-define(SUBJECT_RING, <<"push.job.ring">>).
-define(JOB_VERSION, 1).
-define(LEGACY_CONFIG_VERSION, 0).
-define(NATS_MAX_PAYLOAD_BYTES, 1048576).
-define(MAX_CALLER_NAME_BYTES, 128).
-type meta() :: #{
kind := message | clear | ring,
kind := push_outbox:kind(),
user_ids := [integer()],
channel_id := integer(),
message_id := integer(),
fallback := push_outbox:fallback()
message_id := integer()
}.
-spec publish_message(
@@ -41,51 +41,12 @@ publish_message(
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,
<<"config_version">> => ?LEGACY_CONFIG_VERSION,
<<"guild_id">> => integer_to_binary(GuildId),
<<"channel_id">> => ChannelIdBin,
<<"message_id">> => MessageIdBin,
<<"channel_id">> => integer_to_binary(ChannelId),
<<"message_id">> => integer_to_binary(MessageId),
<<"notification">> => notification_fields(
MessageData,
MarkdownContext,
@@ -94,35 +55,39 @@ publish_message(
MessageId,
GuildName,
ChannelName
),
<<"user_ids">> => [integer_to_binary(UserId) || UserId <- UserIds]
)
},
publish(?SUBJECT_MESSAGE, Job, #{
publish_recipients(Job, UserIds, #{
kind => message,
user_ids => UserIds,
channel_id => ChannelId,
message_id => MessageId,
fallback => Fallback
message_id => MessageId
}).
-spec publish_recipients(map(), [integer()], meta()) -> ok | {error, term()}.
publish_recipients(Job, UserIds, Meta) ->
Chunk = Job#{<<"user_ids">> => [integer_to_binary(UserId) || UserId <- UserIds]},
case encode(Chunk) of
{ok, Body} when byte_size(Body) > ?NATS_MAX_PAYLOAD_BYTES, length(UserIds) > 1 ->
{Left, Right} = lists:split(length(UserIds) div 2, UserIds),
LeftResult = publish_recipients(Job, Left, Meta),
RightResult = publish_recipients(Job, Right, Meta),
first_error(LeftResult, RightResult);
Encoded ->
publish_encoded(?SUBJECT_MESSAGE, Chunk, Encoded, Meta#{user_ids := UserIds})
end.
-spec first_error(ok | {error, term()}, ok | {error, term()}) -> ok | {error, term()}.
first_error(ok, Second) ->
Second;
first_error({error, Reason}, _Second) ->
{error, Reason}.
-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,
<<"config_version">> => ?LEGACY_CONFIG_VERSION,
<<"user_id">> => integer_to_binary(UserId),
<<"channel_id">> => integer_to_binary(ChannelId),
<<"message_id">> => integer_to_binary(MessageId)
@@ -131,8 +96,7 @@ publish_clear(UserId, ChannelId, MessageId, ConfigVersion, Fallback) ->
kind => clear,
user_ids => [UserId],
channel_id => ChannelId,
message_id => MessageId,
fallback => Fallback
message_id => MessageId
}).
-spec publish_ring(integer(), integer(), integer(), integer(), integer(), map()) ->
@@ -141,7 +105,7 @@ publish_ring(UserId, ChannelId, MessageId, StartedAtMs, ExpiresAtMs, Caller) ->
Job = maps:merge(
#{
<<"v">> => ?JOB_VERSION,
<<"config_version">> => push_delivery_config:config_version(),
<<"config_version">> => ?LEGACY_CONFIG_VERSION,
<<"user_id">> => integer_to_binary(UserId),
<<"channel_id">> => integer_to_binary(ChannelId),
<<"message_id">> => integer_to_binary(MessageId),
@@ -154,8 +118,7 @@ publish_ring(UserId, ChannelId, MessageId, StartedAtMs, ExpiresAtMs, Caller) ->
kind => ring,
user_ids => [UserId],
channel_id => ChannelId,
message_id => MessageId,
fallback => fun ignore_fallback/1
message_id => MessageId
}).
-spec caller_fields(map()) -> map().
@@ -208,10 +171,6 @@ decode_reply(Payload) ->
_:_ -> {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().
@@ -250,13 +209,7 @@ nullable(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.
publish_encoded(Subject, Job, encode(Job), Meta).
-spec encode(map()) -> {ok, binary()} | {error, term()}.
encode(Job) ->
@@ -266,30 +219,42 @@ encode(Job) ->
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 ->
-spec publish_encoded(binary(), map(), {ok, binary()} | {error, term()}, meta()) ->
ok | {error, term()}.
publish_encoded(Subject, _Job, {error, Reason}, #{kind := Kind}) ->
logger:warning("Push job encode failed", #{subject => Subject, reason => Reason}),
push_outbox:record_dropped(Kind, encode_failed),
{error, Reason};
publish_encoded(Subject, _Job, {ok, Body}, #{kind := Kind}) 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
}),
push_outbox:record_dropped(Kind, payload_too_large),
{error, {payload_too_large, byte_size(Body)}};
publish_bounded(Subject, Job, Body, Meta) ->
publish_encoded(Subject, Job, {ok, Body}, #{kind := Kind} = 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}),
record_enqueue_failure(Kind, Reason),
{error, Reason}
end.
-spec record_enqueue_failure(push_outbox:kind(), term()) -> ok.
record_enqueue_failure(_Kind, outbox_unavailable) ->
ok;
record_enqueue_failure(Kind, {outbox_unavailable, timeout}) ->
push_outbox:record_dropped(Kind, enqueue_timeout);
record_enqueue_failure(Kind, _Reason) ->
push_outbox:record_dropped(Kind, outbox_unavailable).
-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, user_ids := UserIds, channel_id := ChannelId, message_id := MessageId} =
Meta,
#{
kind => Kind,
subject => Subject,
@@ -297,8 +262,7 @@ outbox_job(Subject, Job, Body, Meta) ->
body => Body,
user_ids => UserIds,
channel_id => ChannelId,
message_id => MessageId,
fallback => Fallback
message_id => MessageId
}.
-ifdef(TEST).
@@ -373,4 +337,276 @@ caller_fields_caps_the_caller_name_test() ->
?MAX_CALLER_NAME_BYTES, byte_size(maps:get(<<"caller_name">>, Fields))
).
notification_fields_carry_title_body_and_tags_test() ->
Fields = test_notification_fields(
#{<<"content">> => <<"Hello world">>, <<"mentions">> => []},
123,
<<"Server">>,
<<"general">>
),
?assertEqual(<<"Alice (#general, Server)">>, maps:get(<<"title">>, Fields)),
?assertEqual(<<"Hello world">>, maps:get(<<"body">>, Fields)),
?assertEqual(<<"channel:456:789">>, maps:get(<<"tag">>, Fields)),
?assertEqual(<<"channel:456">>, maps:get(<<"notification_tag">>, Fields)),
?assertEqual(<<"/channels/123/456/789">>, maps:get(<<"url">>, Fields)),
?assertEqual(null, maps:get(<<"image_url">>, Fields)).
notification_fields_use_single_sticker_preview_test() ->
MessageData = #{
<<"content">> => <<>>,
<<"mentions">> => [],
<<"stickers">> => [
#{<<"id">> => <<"1">>, <<"name">> => <<"Wave">>, <<"animated">> => false}
]
},
Fields = test_notification_fields(MessageData, 0, undefined, undefined),
?assertEqual(<<"Sticker: Wave">>, maps:get(<<"body">>, Fields)).
notification_fields_use_multiple_sticker_preview_test() ->
MessageData = #{
<<"content">> => <<>>,
<<"mentions">> => [],
<<"stickers">> => [
#{<<"id">> => <<"1">>, <<"name">> => <<"Wave">>, <<"animated">> => false},
#{<<"id">> => <<"2">>, <<"name">> => <<"Dance">>, <<"animated">> => false}
]
},
Fields = test_notification_fields(MessageData, 0, undefined, undefined),
?assertEqual(<<"Stickers: Wave and Dance">>, maps:get(<<"body">>, Fields)).
notification_fields_use_attachment_fallback_test() ->
MessageData = #{
<<"content">> => <<>>,
<<"mentions">> => [],
<<"attachments">> => [#{<<"id">> => <<"1">>, <<"filename">> => <<"report.pdf">>}]
},
Fields = test_notification_fields(MessageData, 0, undefined, undefined),
?assertEqual(<<"Attachment: report.pdf">>, maps:get(<<"body">>, Fields)).
notification_fields_use_embed_fallback_test() ->
MessageData = #{
<<"content">> => <<>>,
<<"mentions">> => [],
<<"embeds">> => [#{<<"title">> => <<"Build">>, <<"description">> => <<"green">>}]
},
Fields = test_notification_fields(MessageData, 0, undefined, undefined),
?assertEqual(<<"Build: green">>, maps:get(<<"body">>, Fields)).
notification_fields_use_markdown_plaintext_context_test() ->
MessageData = #{<<"content">> => <<"**Hi** <@1> <@&2> <#3>">>, <<"mentions">> => []},
Context = #{
<<"preserve_markdown">> => true,
<<"users">> => #{<<"1">> => <<"Ada">>},
<<"roles">> => #{<<"2">> => <<"Ops">>},
<<"channels">> => #{<<"3">> => <<"alerts">>}
},
Fields = test_notification_fields(MessageData, 123, <<"Server">>, <<"general">>, Context),
?assertEqual(<<"**Hi** @Ada @Ops #alerts">>, maps:get(<<"body">>, Fields)).
notification_fields_use_author_nickname_in_guild_title_test() ->
MessageData = #{<<"content">> => <<"Hello">>, <<"mentions">> => []},
Context = #{<<"user_nicknames">> => #{<<"42">> => <<"Guild Alice">>}},
Fields = test_notification_fields(MessageData, 123, <<"Server">>, <<"general">>, Context),
?assertEqual(<<"Guild Alice (#general, Server)">>, maps:get(<<"title">>, Fields)).
notification_fields_use_author_nickname_in_group_dm_title_test() ->
MessageData = #{
<<"content">> => <<"Hello">>,
<<"channel_type">> => 3,
<<"nicks">> => #{<<"42">> => <<"Group Alice">>},
<<"mentions">> => []
},
Fields = test_notification_fields(MessageData, 0, undefined, undefined),
?assertEqual(<<"Group Alice (Group DM)">>, maps:get(<<"title">>, Fields)).
notification_fields_include_safe_attachment_image_test() ->
MessageData = #{
<<"content">> => <<"Photo">>,
<<"mentions">> => [],
<<"attachments">> => [
#{
<<"content_type">> => <<"image/png">>,
<<"proxy_url">> => <<"https://cdn.example/image.png">>
}
]
},
Fields = test_notification_fields(MessageData, 123, <<"Server">>, <<"general">>),
?assertEqual(<<"https://cdn.example/image.png">>, maps:get(<<"image_url">>, Fields)).
notification_fields_omit_sensitive_attachment_image_test() ->
MessageData = #{
<<"content">> => <<"Spoiler">>,
<<"mentions">> => [],
<<"attachments">> => [
#{
<<"content_type">> => <<"image/png">>,
<<"proxy_url">> => <<"https://cdn.example/spoiler.png">>,
<<"flags">> => 8
}
]
},
Fields = test_notification_fields(MessageData, 123, <<"Server">>, <<"general">>),
?assertEqual(null, maps:get(<<"image_url">>, Fields)).
test_notification_fields(MessageData, GuildId, GuildName, ChannelName) ->
test_notification_fields(MessageData, GuildId, GuildName, ChannelName, #{}).
test_notification_fields(MessageData, GuildId, GuildName, ChannelName, MarkdownContext) ->
Author = #{<<"id">> => <<"42">>, <<"username">> => <<"Alice">>, <<"avatar">> => null},
with_endpoint_env(fun() ->
notification_fields(
MessageData#{<<"author">> => Author},
MarkdownContext,
GuildId,
456,
789,
GuildName,
ChannelName
)
end).
with_captured_enqueues(Fun) ->
{Jobs, _Dropped} = with_captured_outbox(fun(_OutboxJob) -> ok end, Fun),
Jobs.
with_captured_outbox(EnqueueResult, Fun) ->
Self = self(),
ok = meck:new(push_outbox, [passthrough, no_link]),
try
ok = meck:expect(push_outbox, enqueue, fun(OutboxJob) ->
Self ! {enqueued, OutboxJob},
EnqueueResult(OutboxJob)
end),
ok = meck:expect(push_outbox, record_dropped, fun(Kind, Reason) ->
Self ! {dropped, Kind, Reason},
ok
end),
Fun(),
{drain_enqueued([]), drain_dropped([])}
after
meck:unload(push_outbox)
end.
drain_enqueued(Acc) ->
receive
{enqueued, OutboxJob} -> drain_enqueued([OutboxJob | Acc])
after 0 ->
lists:reverse(Acc)
end.
drain_dropped(Acc) ->
receive
{dropped, Kind, Reason} -> drain_dropped([{Kind, Reason} | Acc])
after 0 ->
lists:reverse(Acc)
end.
padded_job(UserIds, TargetBytes) ->
Job = #{
<<"v">> => ?JOB_VERSION,
<<"config_version">> => ?LEGACY_CONFIG_VERSION,
<<"pad">> => <<>>
},
{ok, Base} = encode(Job#{<<"user_ids">> => [integer_to_binary(Id) || Id <- UserIds]}),
Padded = Job#{<<"pad">> => binary:copy(<<"a">>, TargetBytes - byte_size(Base))},
{ok, Body} = encode(Padded#{<<"user_ids">> => [integer_to_binary(Id) || Id <- UserIds]}),
?assertEqual(TargetBytes, byte_size(Body)),
Padded.
message_meta(UserIds) ->
#{kind => message, user_ids => UserIds, channel_id => 20, message_id => 30}.
a_body_exactly_at_the_payload_limit_is_sent_as_one_job_test() ->
UserIds = [1, 2],
Job = padded_job(UserIds, ?NATS_MAX_PAYLOAD_BYTES),
Jobs = with_captured_enqueues(fun() ->
?assertEqual(ok, publish_recipients(Job, UserIds, message_meta(UserIds)))
end),
?assertMatch([#{user_ids := [1, 2]}], Jobs),
[#{body := Body}] = Jobs,
?assertEqual(?NATS_MAX_PAYLOAD_BYTES, byte_size(Body)).
a_body_one_byte_over_the_payload_limit_is_split_test() ->
UserIds = [1, 2],
Job = padded_job(UserIds, ?NATS_MAX_PAYLOAD_BYTES + 1),
{Jobs, Dropped} = with_captured_outbox(fun(_OutboxJob) -> ok end, fun() ->
?assertEqual(ok, publish_recipients(Job, UserIds, message_meta(UserIds)))
end),
?assertMatch([#{user_ids := [1]}, #{user_ids := [2]}], Jobs),
?assertEqual([], Dropped).
a_single_recipient_over_the_payload_limit_is_dropped_test() ->
UserIds = [1],
Job = padded_job(UserIds, ?NATS_MAX_PAYLOAD_BYTES + 1),
{Jobs, Dropped} = with_captured_outbox(fun(_OutboxJob) -> ok end, fun() ->
?assertEqual(
{error, {payload_too_large, ?NATS_MAX_PAYLOAD_BYTES + 1}},
publish_recipients(Job, UserIds, message_meta(UserIds))
)
end),
?assertEqual([], Jobs),
?assertEqual([{message, payload_too_large}], Dropped).
a_failed_left_chunk_still_publishes_the_right_chunk_test() ->
UserIds = [1, 2],
Job = padded_job(UserIds, ?NATS_MAX_PAYLOAD_BYTES + 1),
Failure = {outbox_unavailable, shutdown},
FailLeft = fun
(#{user_ids := [1]}) -> {error, Failure};
(_OutboxJob) -> ok
end,
{Jobs, Dropped} = with_captured_outbox(FailLeft, fun() ->
?assertEqual(
{error, Failure}, publish_recipients(Job, UserIds, message_meta(UserIds))
)
end),
?assertMatch([#{user_ids := [1]}, #{user_ids := [2]}], Jobs),
?assertEqual([{message, outbox_unavailable}], Dropped).
an_enqueue_timeout_is_counted_apart_from_an_unavailable_outbox_test() ->
Job = #{<<"v">> => ?JOB_VERSION},
Meta = #{kind => clear, user_ids => [1], channel_id => 20, message_id => 30},
{_TimeoutJobs, TimeoutDropped} = with_captured_outbox(
fun(_OutboxJob) -> {error, {outbox_unavailable, timeout}} end,
fun() -> publish(?SUBJECT_CLEAR, Job, Meta) end
),
?assertEqual([{clear, enqueue_timeout}], TimeoutDropped),
{_NoprocJobs, NoprocDropped} = with_captured_outbox(
fun(_OutboxJob) -> {error, outbox_unavailable} end,
fun() -> publish(?SUBJECT_CLEAR, Job, Meta) end
),
?assertEqual([], NoprocDropped).
oversized_recipient_lists_are_split_under_the_payload_limit_test() ->
UserIds = lists:seq(1000000000000000000, 1000000000000000000 + 59999),
Job = #{<<"v">> => ?JOB_VERSION, <<"config_version">> => ?LEGACY_CONFIG_VERSION},
Meta = #{kind => message, user_ids => UserIds, channel_id => 20, message_id => 30},
Jobs = with_captured_enqueues(fun() ->
?assertEqual(ok, publish_recipients(Job, UserIds, Meta))
end),
?assert(length(Jobs) > 1),
lists:foreach(
fun(#{body := Body, user_ids := ChunkIds, job := ChunkJob}) ->
?assert(byte_size(Body) =< ?NATS_MAX_PAYLOAD_BYTES),
?assertEqual(
[integer_to_binary(UserId) || UserId <- ChunkIds],
maps:get(<<"user_ids">>, ChunkJob)
),
?assertEqual(?LEGACY_CONFIG_VERSION, maps:get(<<"config_version">>, ChunkJob))
end,
Jobs
),
?assertEqual(UserIds, lists:append([ChunkIds || #{user_ids := ChunkIds} <- Jobs])).
recipient_lists_under_the_payload_limit_stay_in_one_job_test() ->
UserIds = [1, 2, 3],
Job = #{<<"v">> => ?JOB_VERSION, <<"config_version">> => ?LEGACY_CONFIG_VERSION},
Meta = #{kind => message, user_ids => UserIds, channel_id => 20, message_id => 30},
Jobs = with_captured_enqueues(fun() ->
?assertEqual(ok, publish_recipients(Job, UserIds, Meta))
end),
?assertMatch(
[#{kind := message, user_ids := [1, 2, 3], subject := ?SUBJECT_MESSAGE}], Jobs
).
-endif.
+1 -634
View File
@@ -3,64 +3,7 @@
-module(push_notification).
-typing([eqwalizer]).
-export([
build_notification_title/5,
build_notification_payload/1,
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(),
guild_id := integer(),
navigate_url := binary(),
badge_value := non_neg_integer(),
target_user_id := integer(),
image_url := binary() | undefined,
image_fields := map(),
tag := binary()
}.
-type notification_input() :: #{
message_data := map(),
guild_id := integer(),
channel_id := integer(),
message_id := integer(),
guild_name := binary() | undefined,
channel_name := binary() | undefined,
author_username := binary(),
author_avatar_url := binary(),
target_user_id := integer(),
badge_count := non_neg_integer(),
content_preview => binary() | undefined,
markdown_context => map()
}.
-export([build_notification_title/5]).
-spec build_notification_title(
binary(), map(), integer(), binary() | undefined, binary() | undefined
@@ -86,355 +29,6 @@ format_guild_title(AuthorUsername, _, undefined) ->
format_guild_title(AuthorUsername, GName, ChanName) ->
iolist_to_binary([AuthorUsername, <<" (#">>, ChanName, <<", ">>, GName, <<")">>]).
-spec build_notification_payload(notification_input()) -> map().
build_notification_payload(
#{
message_data := MessageData,
guild_id := GuildId,
channel_id := ChannelId,
message_id := MessageId,
guild_name := GuildName,
channel_name := ChannelName,
author_username := AuthorUsername,
author_avatar_url := AuthorAvatarUrl,
target_user_id := TargetUserId,
badge_count := BadgeCount
} = Input
) ->
ContentPreview = resolve_content_preview(MessageData, Input),
MarkdownContext = maps:get(markdown_context, Input, #{}),
AuthorName = push_notification_format:resolve_author_name(
MessageData, MarkdownContext, AuthorUsername
),
Title = build_notification_title(
AuthorName, MessageData, GuildId, GuildName, ChannelName
),
Ctx = build_push_ctx(GuildId, ChannelId, MessageId, TargetUserId, BadgeCount, MessageData),
assemble_payload(Ctx, Title, ContentPreview, AuthorAvatarUrl).
-spec resolve_content_preview(map(), notification_input()) -> binary().
resolve_content_preview(MessageData, Input) ->
case maps:get(content_preview, Input, undefined) of
Preview when is_binary(Preview) ->
Preview;
_ ->
MarkdownContext = maps:get(markdown_context, Input, #{}),
push_notification_format:build_content_preview(MessageData, MarkdownContext)
end.
-spec build_channel_tag(integer()) -> binary().
build_channel_tag(ChannelId) ->
<<"channel:", (integer_to_binary(ChannelId))/binary>>.
-spec build_message_tag(integer(), integer()) -> binary().
build_message_tag(ChannelId, MessageId) ->
iolist_to_binary([
<<"channel:">>,
integer_to_binary(ChannelId),
<<":">>,
integer_to_binary(MessageId)
]).
-spec build_push_ctx(integer(), integer(), integer(), integer(), non_neg_integer(), map()) ->
push_ctx().
build_push_ctx(GuildId, ChannelId, MessageId, TargetUserId, BadgeCount, MessageData) ->
ImageUrl = push_notification_format:extract_image_url(MessageData),
#{
channel_id => ChannelId,
message_id => MessageId,
guild_id => GuildId,
navigate_url => push_notification_format:build_url(GuildId, ChannelId, MessageId),
badge_value => max(0, BadgeCount),
target_user_id => TargetUserId,
image_url => ImageUrl,
image_fields => push_notification_format:maybe_image_fields(ImageUrl),
tag => build_message_tag(ChannelId, MessageId)
}.
-spec assemble_payload(push_ctx(), binary(), binary(), binary()) -> map().
assemble_payload(
#{image_fields := ImageFields, tag := Tag} = Ctx,
Title,
ContentPreview,
AuthorAvatarUrl
) ->
Data = build_data(Ctx, AuthorAvatarUrl),
Notification = build_notification_body(Ctx, Title, ContentPreview, AuthorAvatarUrl, Data),
maps:merge(
#{
<<"web_push">> => ?WEB_PUSH_MARKER,
<<"notification">> => Notification,
<<"title">> => Title,
<<"body">> => ContentPreview,
<<"icon">> => AuthorAvatarUrl,
<<"badge">> => push_badge_url(),
<<"tag">> => Tag,
<<"data">> => Data
},
ImageFields
).
-spec build_data(push_ctx(), binary()) -> map().
build_data(
#{
channel_id := ChannelId,
message_id := MessageId,
guild_id := GuildId,
navigate_url := NavigateUrl,
badge_value := BadgeValue,
target_user_id := TargetUserId,
image_url := ImageUrl,
image_fields := ImageFields
},
AuthorAvatarUrl
) ->
BaseData = #{
<<"channel_id">> => integer_to_binary(ChannelId),
<<"author_avatar_url">> => AuthorAvatarUrl,
<<"message_id">> => integer_to_binary(MessageId),
<<"notification_tag">> => build_channel_tag(ChannelId),
<<"guild_id">> =>
case GuildId of
0 -> null;
_ -> integer_to_binary(GuildId)
end,
<<"url">> => NavigateUrl,
<<"badge_count">> => BadgeValue,
<<"target_user_id">> => integer_to_binary(TargetUserId),
<<"has_media">> => ImageUrl =/= undefined
},
maps:merge(BaseData, ImageFields).
-spec build_notification_body(push_ctx(), binary(), binary(), binary(), map()) -> map().
build_notification_body(
#{
navigate_url := NavigateUrl,
badge_value := BadgeValue,
image_fields := ImageFields,
tag := Tag
},
Title,
ContentPreview,
AuthorAvatarUrl,
Data
) ->
BaseNotification = #{
<<"title">> => Title,
<<"body">> => ContentPreview,
<<"icon">> => AuthorAvatarUrl,
<<"badge">> => push_badge_url(),
<<"tag">> => Tag,
<<"navigate">> => NavigateUrl,
<<"app_badge">> => integer_to_binary(BadgeValue),
<<"data">> => Data
},
maps:merge(BaseNotification, ImageFields).
-spec push_badge_url() -> binary().
push_badge_url() ->
push_utils:construct_static_asset_url(<<"marketing/branding/symbol-white.svg">>).
-spec build_clear_notification_payload(integer(), integer(), integer(), non_neg_integer()) ->
map().
build_clear_notification_payload(TargetUserId, ChannelId, MessageId, BadgeCount) ->
BadgeValue = max(0, BadgeCount),
Tag = build_channel_tag(ChannelId),
Data = #{
<<"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),
<<"notification_tag">> => Tag,
<<"tag">> => Tag,
<<"badge_count">> => BadgeValue
},
#{
<<"type">> => ?CLEAR_TYPE,
<<"action">> => ?CLEAR_ACTION,
<<"silent">> => true,
<<"tag">> => Tag,
<<"notification_tag">> => Tag,
<<"data">> => Data,
<<"badge_count">> => BadgeValue,
<<"web_push">> => ?WEB_PUSH_MARKER,
<<"notification">> => #{
<<"tag">> => Tag,
<<"data">> => Data,
<<"silent">> => true,
<<"close">> => true
}
}.
-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").
@@ -465,231 +59,4 @@ build_url_guild_test() ->
<<"/channels/123/456/789">>, push_notification_format:build_url(123, 456, 789)
).
build_notification_payload_test() ->
MessageData = #{<<"content">> => <<"Hello world">>, <<"mentions">> => []},
Result = test_notification_payload(MessageData, 123, <<"Server">>, <<"general">>),
Data = maps:get(<<"data">>, Result),
?assertEqual(<<"Alice (#general, Server)">>, maps:get(<<"title">>, Result)),
?assertEqual(<<"Hello world">>, maps:get(<<"body">>, Result)),
?assertEqual(5, maps:get(<<"badge_count">>, Data)),
?assertEqual(<<"channel:456:789">>, maps:get(<<"tag">>, Result)),
?assertEqual(<<"channel:456">>, maps:get(<<"notification_tag">>, Data)).
build_notification_payload_uses_single_sticker_preview_test() ->
MessageData = #{
<<"content">> => <<>>,
<<"mentions">> => [],
<<"stickers">> => [
#{<<"id">> => <<"1">>, <<"name">> => <<"Wave">>, <<"animated">> => false}
]
},
Result = test_notification_payload(MessageData, 0, undefined, undefined),
?assertEqual(<<"Sticker: Wave">>, maps:get(<<"body">>, Result)),
?assertEqual(
<<"Sticker: Wave">>, maps:get(<<"body">>, maps:get(<<"notification">>, Result))
).
build_notification_payload_uses_multiple_sticker_preview_test() ->
MessageData = #{
<<"content">> => <<>>,
<<"mentions">> => [],
<<"stickers">> => [
#{<<"id">> => <<"1">>, <<"name">> => <<"Wave">>, <<"animated">> => false},
#{<<"id">> => <<"2">>, <<"name">> => <<"Dance">>, <<"animated">> => false}
]
},
Result = test_notification_payload(MessageData, 0, undefined, undefined),
?assertEqual(<<"Stickers: Wave and Dance">>, maps:get(<<"body">>, Result)).
build_notification_payload_uses_attachment_fallback_test() ->
MessageData = #{
<<"content">> => <<>>,
<<"mentions">> => [],
<<"attachments">> => [
#{<<"id">> => <<"1">>, <<"filename">> => <<"report.pdf">>}
]
},
Result = test_notification_payload(MessageData, 0, undefined, undefined),
?assertEqual(<<"Attachment: report.pdf">>, maps:get(<<"body">>, Result)).
build_notification_payload_uses_embed_fallback_test() ->
MessageData = #{
<<"content">> => <<>>,
<<"mentions">> => [],
<<"embeds">> => [
#{<<"title">> => <<"Build">>, <<"description">> => <<"green">>}
]
},
Result = test_notification_payload(MessageData, 0, undefined, undefined),
?assertEqual(<<"Build: green">>, maps:get(<<"body">>, Result)).
build_notification_payload_uses_markdown_plaintext_context_test() ->
MessageData = #{
<<"content">> => <<"**Hi** <@1> <@&2> <#3>">>,
<<"mentions">> => []
},
Context = #{
<<"preserve_markdown">> => true,
<<"users">> => #{<<"1">> => <<"Ada">>},
<<"roles">> => #{<<"2">> => <<"Ops">>},
<<"channels">> => #{<<"3">> => <<"alerts">>}
},
Result = test_notification_payload(
MessageData, 123, <<"Server">>, <<"general">>, Context
),
?assertEqual(<<"**Hi** @Ada @Ops #alerts">>, maps:get(<<"body">>, Result)).
build_notification_payload_uses_author_nickname_in_guild_title_test() ->
MessageData = #{
<<"content">> => <<"Hello">>,
<<"author">> => #{<<"id">> => <<"42">>, <<"username">> => <<"Alice">>},
<<"mentions">> => []
},
Context = #{<<"user_nicknames">> => #{<<"42">> => <<"Guild Alice">>}},
Result = test_notification_payload(
MessageData, 123, <<"Server">>, <<"general">>, Context
),
?assertEqual(<<"Guild Alice (#general, Server)">>, maps:get(<<"title">>, Result)).
build_notification_payload_uses_author_nickname_in_group_dm_title_test() ->
MessageData = #{
<<"content">> => <<"Hello">>,
<<"channel_type">> => 3,
<<"author">> => #{<<"id">> => <<"42">>, <<"username">> => <<"Alice">>},
<<"nicks">> => #{<<"42">> => <<"Group Alice">>},
<<"mentions">> => []
},
Result = test_notification_payload(MessageData, 0, undefined, undefined),
?assertEqual(<<"Group Alice (Group DM)">>, maps:get(<<"title">>, Result)).
build_notification_payload_includes_safe_attachment_image_test() ->
MessageData = #{
<<"content">> => <<"Photo">>,
<<"mentions">> => [],
<<"attachments">> => [
#{
<<"content_type">> => <<"image/png">>,
<<"proxy_url">> => <<"https://cdn.example/image.png">>
}
]
},
Result = test_notification_payload(MessageData, 123, <<"Server">>, <<"general">>),
ImgUrl = <<"https://cdn.example/image.png">>,
DataMap = maps:get(<<"data">>, Result),
?assertEqual(ImgUrl, maps:get(<<"image_url">>, Result)),
?assertEqual(ImgUrl, maps:get(<<"image_url">>, DataMap)),
?assertEqual(true, maps:get(<<"has_media">>, DataMap)).
build_notification_payload_omits_sensitive_attachment_image_test() ->
MessageData = #{
<<"content">> => <<"Spoiler">>,
<<"mentions">> => [],
<<"attachments">> => [
#{
<<"content_type">> => <<"image/png">>,
<<"proxy_url">> => <<"https://cdn.example/spoiler.png">>,
<<"flags">> => 8
}
]
},
Result = test_notification_payload(MessageData, 123, <<"Server">>, <<"general">>),
?assertEqual(false, maps:get(<<"has_media">>, maps:get(<<"data">>, Result))),
?assertEqual(false, maps:is_key(<<"image_url">>, Result)).
build_clear_notification_payload_test() ->
Result = build_clear_notification_payload(999, 456, 789, 2),
Data = maps:get(<<"data">>, Result),
?assertEqual(<<"notification_clear">>, maps:get(<<"type">>, Result)),
?assertEqual(<<"clear_channel">>, maps:get(<<"action">>, Result)),
?assertEqual(<<"channel:456">>, maps:get(<<"tag">>, Result)),
?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, #{}).
test_notification_payload(MessageData, GuildId, GuildName, ChannelName, MarkdownContext) ->
build_notification_payload(#{
message_data => MessageData,
guild_id => GuildId,
channel_id => 456,
message_id => 789,
guild_name => GuildName,
channel_name => ChannelName,
author_username => <<"Alice">>,
author_avatar_url => <<"http://avatar">>,
target_user_id => 999,
badge_count => 5,
markdown_context => MarkdownContext
}).
-endif.
@@ -10,7 +10,6 @@
resolve_author_name/3,
resolve_author_avatar_url/1,
extract_image_url/1,
maybe_image_fields/1,
build_url/3,
truncate_bytes/2
]).
@@ -616,12 +615,6 @@ is_sensitive_media(true, _Flags) ->
is_sensitive_media(_Nsfw, Flags) ->
bitset:any(Flags, 16#18).
-spec maybe_image_fields(binary() | undefined) -> map().
maybe_image_fields(undefined) ->
#{};
maybe_image_fields(ImageUrl) when is_binary(ImageUrl) ->
#{<<"image_url">> => ImageUrl, <<"image">> => ImageUrl}.
-spec build_url(integer(), integer(), integer()) -> binary().
build_url(0, ChannelId, MessageId) ->
build_url_parts([<<"@me">>, integer_to_binary(ChannelId), integer_to_binary(MessageId)]);
+74 -223
View File
@@ -9,26 +9,24 @@
enqueue/1,
truncate_read/3,
note_session_active/1,
delivery_config_changed/0,
record_dropped/2,
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]).
-export_type([job/0, kind/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 | ring.
-type fallback() :: fun(([integer()]) -> term()).
-type job() :: #{
kind := kind(),
subject := binary(),
@@ -36,8 +34,7 @@
body := binary(),
user_ids := [integer()],
channel_id := integer(),
message_id := integer(),
fallback := fallback()
message_id := integer()
}.
-type entry() :: #{
kind := kind(),
@@ -47,35 +44,28 @@
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
attempts := non_neg_integer()
}.
-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()},
dropped := #{{kind(), atom()} => pos_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()
@@ -102,9 +92,9 @@ truncate_read(UserId, ChannelId, MessageId) ->
note_session_active(UserId) ->
broadcast({session_active, UserId}).
-spec delivery_config_changed() -> ok.
delivery_config_changed() ->
gen_server:cast(?MODULE, delivery_config_changed).
-spec record_dropped(kind(), atom()) -> ok.
record_dropped(Kind, Reason) ->
gen_server:cast(?MODULE, {dropped, Kind, Reason}).
-spec stats() -> map().
stats() ->
@@ -126,17 +116,13 @@ init([]) ->
jobs => gb_trees:empty(),
ready => queue:new(),
inflight => #{},
fallback_backlog => queue:new(),
fallback_runners => #{},
next_seq => 0,
reads => #{},
active => #{},
counters => #{},
dropped => #{},
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)
@@ -149,8 +135,7 @@ handle_call({enqueue, Job}, _From, State) when
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)
is_map_key(message_id, Job)
->
{reply, ok, pump(admit(Job, State))};
handle_call(stats, _From, State) ->
@@ -165,18 +150,14 @@ handle_cast({truncate_read, UserId, ChannelId, MessageId}, State) when
{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({dropped, Kind, Reason}, State) when is_atom(Kind), is_atom(Reason) ->
{noreply, count_dropped(Kind, Reason, 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) ->
@@ -200,28 +181,9 @@ code_change(_OldVsn, State, _Extra) ->
-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)
}),
Entry = maps:merge(Job, #{seq => Seq, enqueued_at => now_ms(), attempts => 0}),
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.
shed_to_capacity(insert(Entry, State1)).
-spec insert(entry(), state()) -> state().
insert(#{seq := Seq} = Entry, #{jobs := Jobs, ready := Ready} = State) ->
@@ -272,13 +234,13 @@ take_ready({value, Entry}, Seq, #{jobs := Jobs} = State) ->
dispatch(Entry, State) ->
case prepare(Entry, State) of
{skip, State1} -> State1;
{send, Prepared, State1} -> send_or_fall_back(Prepared, State1)
{send, Prepared, State1} -> send_or_drop(Prepared, State1)
end.
-spec send_or_fall_back(entry(), state()) -> state().
send_or_fall_back(Entry, State) ->
-spec send_or_drop(entry(), state()) -> state().
send_or_drop(Entry, State) ->
case is_expired(Entry, State) of
true -> fall_back(Entry, State);
true -> drop_expired(Entry, State);
false -> start_worker(Entry, State)
end.
@@ -437,23 +399,17 @@ handle_result({error, Reason}, Entry, State) ->
attempts => maps:get(attempts, Entry)
}),
case is_expired(Entry, State) of
true -> fall_back_prepared(Entry, State);
false -> retry_settled(settle_if_stale(Entry, State))
true -> drop_prepared(Entry, State);
false -> schedule_retry(Entry, State)
end.
-spec fall_back_prepared(entry(), state()) -> state().
fall_back_prepared(Entry, State) ->
-spec drop_prepared(entry(), state()) -> state().
drop_prepared(Entry, State) ->
case prepare(Entry, State) of
{skip, State1} -> State1;
{send, Prepared, State1} -> fall_back(Prepared, State1)
{send, Prepared, State1} -> drop_expired(Prepared, State1)
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,
@@ -477,117 +433,21 @@ make_ready(Seq, #{jobs := Jobs, ready := Ready} = State) ->
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),
-spec drop_expired(entry(), state()) -> state().
drop_expired(#{kind := Kind, user_ids := UserIds} = Entry, State) ->
logger:warning("Push outbox job expired undelivered, dropping it", #{
kind => Kind,
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).
count_dropped(Kind, expired, 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(#{kind := ring} = Entry, State) ->
{keep, Entry, State};
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 count_dropped(kind(), atom(), state()) -> state().
count_dropped(Kind, Reason, #{dropped := Dropped} = State) ->
Key = {Kind, Reason},
State#{dropped := Dropped#{Key => maps:get(Key, Dropped, 0) + 1}}.
-spec prune(state()) -> state().
prune(#{reads := Reads, active := Active, max_age_ms := MaxAge} = State) ->
@@ -598,13 +458,7 @@ prune(#{reads := Reads, active := Active, max_age_ms := MaxAge} = State) ->
}.
-spec build_stats(state()) -> map().
build_stats(#{
jobs := Jobs,
inflight := Inflight,
fallback_backlog := Backlog,
fallback_runners := Runners,
counters := Counters
}) ->
build_stats(#{jobs := Jobs, inflight := Inflight, counters := Counters, dropped := Dropped}) ->
maps:merge(
#{
delivered => 0,
@@ -612,22 +466,15 @@ build_stats(#{
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)
dropped => Dropped
}
).
-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;
@@ -672,63 +519,69 @@ app_pos_integer(Key, Default) ->
-ifdef(TEST).
-include_lib("eunit/include/eunit.hrl").
an_expired_entry_is_handed_back_without_the_users_who_read_it_test() ->
Self = self(),
count(Counter, #{counters := Counters}) ->
maps:get(Counter, Counters, 0).
an_expired_entry_is_dropped_without_the_users_who_read_it_test() ->
State = test_state(#{{7, 20} => {30, now_ms()}}),
Entry = test_entry([7, 8], fun(UserIds) -> Self ! {handed_back, UserIds} end),
Result = handle_result({error, timeout}, Entry, State),
?assertEqual(1, count(fallbacks, Result)),
Result = handle_result({error, timeout}, test_entry([7, 8]), State),
?assertEqual(#{{message, expired} => 1}, maps:get(dropped, Result)),
?assertEqual(1, count(truncations, Result)),
receive
{handed_back, UserIds} -> ?assertEqual([8], UserIds)
after 2000 -> erlang:error(no_fallback_ran)
end.
?assertEqual(0, count(retries, Result)).
an_expired_entry_every_recipient_read_is_not_handed_back_test() ->
Self = self(),
an_expired_entry_every_recipient_read_is_not_counted_as_dropped_test() ->
State = test_state(#{{7, 20} => {30, now_ms()}, {8, 20} => {31, now_ms()}}),
Entry = test_entry([7, 8], fun(UserIds) -> Self ! {handed_back, UserIds} end),
Result = handle_result({error, timeout}, Entry, State),
?assertEqual(0, count(fallbacks, Result)),
?assertEqual(2, count(truncations, Result)),
receive
{handed_back, _} -> erlang:error(fallback_ran_for_read_users)
after 200 -> ok
end.
Result = handle_result({error, timeout}, test_entry([7, 8]), State),
?assertEqual(#{}, maps:get(dropped, Result)),
?assertEqual(2, count(truncations, Result)).
an_expired_clear_is_handed_back_unchanged_test() ->
Self = self(),
an_expired_clear_is_dropped_and_counted_by_kind_test() ->
State = test_state(#{{7, 20} => {30, now_ms()}, {8, 20} => {30, now_ms()}}),
Entry = (test_entry([7, 8], fun(UserIds) -> Self ! {handed_back, UserIds} end))#{
kind := clear
},
Entry = (test_entry([7, 8]))#{kind := clear},
Result = handle_result({error, timeout}, Entry, State),
?assertEqual(1, count(fallbacks, Result)),
receive
{handed_back, UserIds} -> ?assertEqual([7, 8], UserIds)
after 2000 -> erlang:error(no_fallback_ran)
end.
?assertEqual(#{{clear, expired} => 1}, maps:get(dropped, Result)),
?assertEqual(0, count(truncations, Result)).
a_failed_entry_within_its_age_is_retried_test() ->
State = test_state(#{}),
Entry = (test_entry([7]))#{enqueued_at := now_ms()},
Result = handle_result({error, timeout}, Entry, State),
?assertEqual(#{}, maps:get(dropped, Result)),
?assertEqual(1, count(retries, Result)),
?assertEqual({value, Entry#{attempts := 4}}, gb_trees:lookup(0, maps:get(jobs, Result))).
an_entry_expired_at_dispatch_is_dropped_test() ->
Result = dispatch(test_entry([7]), test_state(#{})),
?assertEqual(#{{message, expired} => 1}, maps:get(dropped, Result)),
?assertEqual(0, map_size(maps:get(inflight, Result))).
reported_drops_are_counted_by_kind_and_reason_test() ->
{noreply, State1} = handle_cast({dropped, message, payload_too_large}, test_state(#{})),
{noreply, State2} = handle_cast({dropped, message, payload_too_large}, State1),
{noreply, State3} = handle_cast({dropped, ring, outbox_unavailable}, State2),
?assertEqual(
#{{message, payload_too_large} => 2, {ring, outbox_unavailable} => 1},
maps:get(dropped, build_stats(State3))
).
test_state(Reads) ->
#{
jobs => gb_trees:empty(),
ready => queue:new(),
inflight => #{},
fallback_backlog => queue:new(),
fallback_runners => #{},
next_seq => 1,
reads => Reads,
active => #{},
counters => #{},
dropped => #{},
max_queue => ?DEFAULT_MAX_QUEUE,
max_inflight => ?DEFAULT_MAX_INFLIGHT,
max_fallback_runners => ?DEFAULT_MAX_FALLBACK_RUNNERS,
request_timeout_ms => ?DEFAULT_REQUEST_TIMEOUT_MS,
max_age_ms => ?DEFAULT_MAX_AGE_MS,
retry_base_ms => ?DEFAULT_RETRY_BASE_MS
}.
test_entry(UserIds, Fallback) ->
test_entry(UserIds) ->
#{
kind => message,
subject => <<"rpc.push.message">>,
@@ -737,11 +590,9 @@ test_entry(UserIds, Fallback) ->
user_ids => UserIds,
channel_id => 20,
message_id => 30,
fallback => Fallback,
seq => 0,
enqueued_at => now_ms() - ?DEFAULT_MAX_AGE_MS,
attempts => 3,
config_version => undefined
attempts => 3
}.
-endif.
-635
View File
@@ -1,635 +0,0 @@
%% SPDX-License-Identifier: AGPL-3.0-or-later
-module(push_sender).
-typing([eqwalizer]).
-export([
send_to_user_subscriptions/3,
send_clear_to_user_subscriptions/5,
send_push_notifications/1,
send_push_notifications/8,
send_clear_channel_notifications/4
]).
-export_type([send_context/0]).
-define(DEFAULT_BADGE_FETCH_BATCH, 2000).
-define(BADGE_FETCH_MAX_CONSECUTIVE_FAILURES, 3).
-define(DEFAULT_BADGE_FETCH_BUDGET_MS, 120000).
-define(PUSH_COUNTERS, push_worker_counter).
-type send_context() :: #{
message_data := map(),
guild_id := integer(),
channel_id := integer(),
message_id := integer(),
guild_name := binary() | undefined,
channel_name := binary() | undefined,
badge_count := non_neg_integer(),
content_preview := binary(),
markdown_context := map()
}.
-spec send_to_user_subscriptions(integer(), list(), send_context()) -> ok.
send_to_user_subscriptions(UserId, Subscriptions, SendContext) ->
BadgeCount = maps:get(badge_count, SendContext),
NotificationPayload = notification_payload(UserId, SendContext),
logger:debug("Push: sending to user subscriptions", #{
user_id => UserId,
subscription_count => length(Subscriptions),
badge_count => BadgeCount
}),
FailedSubscriptions = send_subscriptions(UserId, NotificationPayload, Subscriptions, []),
handle_failed_subscriptions(UserId, FailedSubscriptions).
-spec notification_payload(integer(), send_context()) -> map().
notification_payload(UserId, SendContext) ->
#{
message_data := MessageData,
guild_id := GuildId,
channel_id := ChannelId,
message_id := MessageId,
guild_name := GuildName,
channel_name := ChannelName,
badge_count := BadgeCount,
content_preview := ContentPreview,
markdown_context := MarkdownContext
} = SendContext,
AuthorData = maps:get(<<"author">>, MessageData, #{}),
AuthorUsername = maps:get(<<"username">>, AuthorData, <<"Unknown">>),
AuthorAvatarUrl = push_notification_format:resolve_author_avatar_url(AuthorData),
push_notification:build_notification_payload(#{
message_data => MessageData,
guild_id => GuildId,
channel_id => ChannelId,
message_id => MessageId,
guild_name => GuildName,
channel_name => ChannelName,
author_username => AuthorUsername,
author_avatar_url => AuthorAvatarUrl,
target_user_id => UserId,
badge_count => BadgeCount,
content_preview => ContentPreview,
markdown_context => MarkdownContext
}).
-spec send_clear_to_user_subscriptions(
integer(), list(), integer(), integer(), non_neg_integer()
) -> ok.
send_clear_to_user_subscriptions(UserId, Subscriptions, ChannelId, MessageId, BadgeCount) ->
Payload = push_notification:build_clear_notification_payload(
UserId, ChannelId, MessageId, BadgeCount
),
FailedSubscriptions = send_subscriptions(UserId, Payload, Subscriptions, []),
handle_failed_subscriptions(UserId, FailedSubscriptions).
-spec send_push_notifications(
[integer()],
map(),
integer(),
integer(),
integer(),
binary() | undefined,
binary() | undefined,
non_neg_integer()
) -> ok.
send_push_notifications(
UserIds,
MessageData,
GuildId,
ChannelId,
MessageId,
GuildName,
ChannelName,
BadgeCountsTtlSeconds
) ->
send_push_notifications(#{
user_ids => UserIds,
message_data => MessageData,
markdown_context => #{},
guild_id => GuildId,
channel_id => ChannelId,
message_id => MessageId,
guild_name => GuildName,
channel_name => ChannelName,
badge_counts_ttl_seconds => BadgeCountsTtlSeconds
}).
-spec send_push_notifications(map()) -> ok.
send_push_notifications(#{
user_ids := UserIds,
message_data := MessageData,
markdown_context := MarkdownContext,
guild_id := GuildId,
channel_id := ChannelId,
message_id := MessageId,
guild_name := GuildName,
channel_name := ChannelName,
badge_counts_ttl_seconds := BadgeCountsTtlSeconds
}) ->
logger:debug("Push: send_push_notifications starting", #{
message_id => MessageId,
channel_id => ChannelId,
guild_id => GuildId,
user_count => length(UserIds)
}),
BadgeCounts = ensure_badge_counts(UserIds, BadgeCountsTtlSeconds),
logger:debug(
"Push: badge counts fetched",
#{message_id => MessageId, badge_count_users => map_size(BadgeCounts)}
),
push_subscriptions:fetch_and_send_subscriptions(
UserIds,
MessageData,
GuildId,
ChannelId,
MessageId,
GuildName,
ChannelName,
MarkdownContext,
BadgeCounts
),
ok.
-spec send_clear_channel_notifications(integer(), integer(), integer(), non_neg_integer()) ->
ok.
send_clear_channel_notifications(UserId, ChannelId, MessageId, BadgeCountsTtlSeconds) ->
BadgeCounts = ensure_badge_counts([UserId], BadgeCountsTtlSeconds),
BadgeCount = maps:get(UserId, BadgeCounts, 0),
push_subscriptions:fetch_and_send_clear_notification(
UserId, ChannelId, MessageId, BadgeCount
),
ok.
-spec handle_failed_subscriptions(integer(), list()) -> ok.
handle_failed_subscriptions(_UserId, []) ->
ok;
handle_failed_subscriptions(UserId, FailedSubscriptions) ->
logger:debug(
"Push: removing failed subscriptions",
#{user_id => UserId, failed_count => length(FailedSubscriptions)}
),
_ = push_subscriptions:delete_failed_subscriptions(FailedSubscriptions),
ok.
-spec send_subscriptions(integer(), map(), list(), [map()]) -> [map()].
send_subscriptions(_UserId, _Payload, [], FailedAcc) ->
lists:reverse(FailedAcc);
send_subscriptions(UserId, Payload, [Subscription | Rest], FailedAcc) ->
case send_notification_to_subscription(UserId, Subscription, Payload) of
{true, FailedSubscription} ->
send_subscriptions(UserId, Payload, Rest, [FailedSubscription | FailedAcc]);
false ->
send_subscriptions(UserId, Payload, Rest, FailedAcc)
end.
-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 => Platform,
web_push_shape => WebPushShape
}),
route_subscription(UserId, Platform, WebPushShape, Subscription, Payload).
-spec route_subscription(integer(), binary(), boolean(), map(), map()) -> false | {true, map()}.
route_subscription(UserId, <<"ios_apns_voip">>, _WebPushShape, _Subscription, _Payload) ->
logger:debug("Push: skipping a VoIP subscription", #{user_id => UserId}),
false;
route_subscription(UserId, _Platform, true, Subscription, Payload) ->
push_sender_delivery:send_webpush_notification(UserId, Subscription, Payload);
route_subscription(UserId, Platform, false, Subscription, Payload) ->
send_platform_notification(UserId, Platform, Subscription, Payload).
-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">>),
push_utils:normalize_binary(Platform, <<"web_push">>).
-spec ensure_badge_counts([integer()], non_neg_integer()) -> map().
ensure_badge_counts(UserIds, TTL) ->
Now = erlang:system_time(second),
{CachedCounts, Missing} = lists:foldl(
fun(UserId, {Acc, MissingAcc}) ->
check_badge_cache(UserId, TTL, Now, Acc, MissingAcc)
end,
{#{}, []},
UserIds
),
case lists:usort(Missing) of
[] -> CachedCounts;
UniqueMissing -> fetch_badge_counts(UniqueMissing, CachedCounts, Now)
end.
-spec check_badge_cache(integer(), non_neg_integer(), integer(), map(), [integer()]) ->
{map(), [integer()]}.
check_badge_cache(UserId, TTL, Now, Acc, MissingAcc) ->
case push_ets_cache:get_badge_count(UserId) of
{Count, Timestamp} when TTL > 0, Now - Timestamp < TTL ->
{Acc#{UserId => Count}, MissingAcc};
_ ->
{Acc, [UserId | MissingAcc]}
end.
-spec fetch_badge_counts([integer()], map(), integer()) -> map().
fetch_badge_counts(UserIds, Counts, CachedAt) ->
Batches = chunk_badge_user_ids(UserIds, badge_fetch_batch_size(), []),
{Counted, FailedBatches, DefaultedUsers, _Consecutive} =
fetch_badge_count_batches(Batches, {Counts, 0, 0, 0}, CachedAt),
report_badge_fetch_failures(FailedBatches, DefaultedUsers),
Counted.
-type badge_batch_acc() :: {map(), non_neg_integer(), non_neg_integer(), non_neg_integer()}.
-spec fetch_badge_count_batches([[integer()]], badge_batch_acc(), integer()) ->
badge_batch_acc().
fetch_badge_count_batches(Batches, Acc, CachedAt) ->
Deadline = erlang:monotonic_time(millisecond) + badge_fetch_budget_ms(),
fetch_badge_count_batches(Batches, Acc, CachedAt, Deadline).
-spec fetch_badge_count_batches([[integer()]], badge_batch_acc(), integer(), integer()) ->
badge_batch_acc().
fetch_badge_count_batches([], Acc, _CachedAt, _Deadline) ->
Acc;
fetch_badge_count_batches(Remaining, Acc, CachedAt, Deadline) ->
case badge_fetch_exhausted(Acc, Deadline) of
true -> abandon_badge_batches(Remaining, Acc);
false -> fetch_next_badge_batch(Remaining, Acc, CachedAt, Deadline)
end.
-spec badge_fetch_exhausted(badge_batch_acc(), integer()) -> boolean().
badge_fetch_exhausted({_Counts, _FailedBatches, _DefaultedUsers, Consecutive}, Deadline) ->
Consecutive >= ?BADGE_FETCH_MAX_CONSECUTIVE_FAILURES orelse
erlang:monotonic_time(millisecond) >= Deadline.
-spec abandon_badge_batches([[integer()]], badge_batch_acc()) -> badge_batch_acc().
abandon_badge_batches(Remaining, {Counts, FailedBatches, DefaultedUsers, Consecutive}) ->
{Counts, FailedBatches + length(Remaining),
DefaultedUsers + badge_batched_user_count(Remaining, 0), Consecutive}.
-spec badge_batched_user_count([[integer()]], non_neg_integer()) -> non_neg_integer().
badge_batched_user_count([], Acc) ->
Acc;
badge_batched_user_count([Batch | Rest], Acc) ->
badge_batched_user_count(Rest, Acc + length(Batch)).
-spec fetch_next_badge_batch([[integer()]], badge_batch_acc(), integer(), integer()) ->
badge_batch_acc().
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]
},
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, Fill),
{Merged, FailedBatches, DefaultedUsers, 0};
{error, _Reason} ->
{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.
report_badge_fetch_failures(0, _DefaultedUsers) ->
ok;
report_badge_fetch_failures(FailedBatches, DefaultedUsers) ->
bump_counter(badge_fetch_calls_failed, FailedBatches),
bump_counter(badge_fetch_users_defaulted, DefaultedUsers),
logger:warning("Push: badge count batches failed; users default to zero badge", #{
failed_batches => FailedBatches, users_defaulted => DefaultedUsers
}).
-spec bump_counter(atom(), integer()) -> ok.
bump_counter(Key, Delta) ->
try
_ = ets:update_counter(?PUSH_COUNTERS, Key, {2, Delta}),
ok
catch
error:badarg -> seed_and_bump_counter(Key, Delta)
end.
-spec seed_and_bump_counter(atom(), integer()) -> ok.
seed_and_bump_counter(Key, Delta) ->
try
_ = ets:insert_new(?PUSH_COUNTERS, {Key, 0}),
_ = ets:update_counter(?PUSH_COUNTERS, Key, {2, Delta}),
ok
catch
error:badarg -> ok
end.
-spec chunk_badge_user_ids([integer()], pos_integer(), [[integer()]]) -> [[integer()]].
chunk_badge_user_ids([], _BatchSize, Acc) ->
lists:reverse(Acc);
chunk_badge_user_ids(UserIds, BatchSize, Acc) ->
{Batch, Rest} = take_badge_user_id_batch(UserIds, BatchSize, []),
chunk_badge_user_ids(Rest, BatchSize, [Batch | Acc]).
-spec take_badge_user_id_batch([integer()], non_neg_integer(), [integer()]) ->
{[integer()], [integer()]}.
take_badge_user_id_batch(Rest, 0, Acc) ->
{lists:reverse(Acc), Rest};
take_badge_user_id_batch([], _Remaining, Acc) ->
{lists:reverse(Acc), []};
take_badge_user_id_batch([UserId | Rest], Remaining, Acc) ->
take_badge_user_id_batch(Rest, Remaining - 1, [UserId | Acc]).
-spec badge_fetch_batch_size() -> pos_integer().
badge_fetch_batch_size() ->
case application:get_env(fluxer_gateway, push_badge_fetch_batch_size, undefined) of
Value when is_integer(Value), Value > 0 -> Value;
_ -> ?DEFAULT_BADGE_FETCH_BATCH
end.
-spec badge_fetch_budget_ms() -> pos_integer().
badge_fetch_budget_ms() ->
case application:get_env(fluxer_gateway, push_badge_fetch_budget_ms, undefined) of
Value when is_integer(Value), Value > 0 -> Value;
_ -> ?DEFAULT_BADGE_FETCH_BUDGET_MS
end.
-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, Fill),
Acc#{UserId => Count}
end,
Counts,
UserIds
).
-spec normalize_badge_count(integer() | term()) -> non_neg_integer().
normalize_badge_count(Value) when is_integer(Value), Value >= 0 -> Value;
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)).
voip_row_is_skipped_test() ->
?assertEqual({0, 0, 0}, routed_target(web_push_row(<<"ios_apns_voip">>))).
voip_row_without_keys_is_skipped_test() ->
?assertEqual({0, 0, 0}, routed_target(legacy_row(<<"ios_apns_voip">>))).
send_subscriptions_skips_voip_rows_test() ->
ok = meck:new(push_sender_delivery, [passthrough, no_link]),
try
ok = meck:expect(push_sender_delivery, send_webpush_notification, fun(_U, _S, _P) ->
false
end),
Rows = [web_push_row(<<"ios_apns_voip">>), web_push_row(<<"ios_apns">>)],
?assertEqual([], send_subscriptions(7, #{}, Rows, [])),
?assertEqual(1, meck:num_calls(push_sender_delivery, send_webpush_notification, '_'))
after
meck:unload(push_sender_delivery)
end.
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, [])).
fetch_badge_counts_in_batches_keeps_successful_batches_test() ->
ok = meck:new(rpc_client, [passthrough, no_link]),
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, _Fill) ->
ok
end),
ok = meck:expect(rpc_client, call, fun(#{<<"user_ids">> := Ids}) ->
case Ids of
[<<"1">>, <<"2">>] ->
{ok, #{<<"badge_counts">> => #{<<"1">> => 3, <<"2">> => 4}}};
_ ->
{error, unavailable}
end
end),
?assertEqual(#{1 => 3, 2 => 4}, fetch_badge_counts([1, 2, 3, 4], #{}, 0))
after
application:unset_env(fluxer_gateway, push_badge_fetch_batch_size),
meck:unload(push_ets_cache),
meck:unload(rpc_client)
end.
fetch_badge_count_batches_stops_after_consecutive_failures_test() ->
ok = meck:new(rpc_client, [passthrough, no_link]),
try
ok = meck:expect(rpc_client, call, fun(_Req) -> {error, unavailable} end),
Batches = [[N] || N <- lists:seq(1, 10)],
?assertEqual(
{#{}, 10, 10, ?BADGE_FETCH_MAX_CONSECUTIVE_FAILURES},
fetch_badge_count_batches(Batches, {#{}, 0, 0, 0}, 0)
),
?assertEqual(?BADGE_FETCH_MAX_CONSECUTIVE_FAILURES, length(meck:history(rpc_client)))
after
meck:unload(rpc_client)
end.
fetch_badge_count_batches_stops_at_the_wall_clock_budget_test() ->
ok = meck:new(rpc_client, [passthrough, no_link]),
application:set_env(fluxer_gateway, push_badge_fetch_budget_ms, 1),
try
ok = meck:expect(rpc_client, call, fun(_Req) ->
{ok, #{<<"badge_counts">> => #{}}}
end),
Batches = [[N] || N <- lists:seq(1, 10)],
Deadline = erlang:monotonic_time(millisecond) - 1,
?assertEqual(
{#{}, 10, 10, 0},
fetch_badge_count_batches(Batches, {#{}, 0, 0, 0}, 0, Deadline)
),
?assertEqual(0, length(meck:history(rpc_client)))
after
application:unset_env(fluxer_gateway, push_badge_fetch_budget_ms),
meck:unload(rpc_client)
end.
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, _Fill) ->
ok
end),
ok = meck:expect(rpc_client, call, fun(#{<<"user_ids">> := Ids}) ->
?assertEqual([<<"1">>, <<"2">>, <<"3">>], Ids),
{ok, #{<<"badge_counts">> => #{<<"1">> => 1}}}
end),
?assertEqual(#{1 => 1, 2 => 0, 3 => 0}, fetch_badge_counts([1, 2, 3], #{}, 0))
after
meck:unload(push_ets_cache),
meck:unload(rpc_client)
end.
badge_fetch_failure_counts_defaulted_users_test() ->
ok = meck:new(rpc_client, [passthrough, no_link]),
ensure_test_counter_table(),
ok = reset_test_counter(badge_fetch_calls_failed),
ok = reset_test_counter(badge_fetch_users_defaulted),
try
ok = meck:expect(rpc_client, call, fun(_Req) -> {error, unavailable} end),
?assertEqual(#{}, fetch_badge_counts([1, 2, 3], #{}, 0)),
?assertEqual(1, counter_value(badge_fetch_calls_failed)),
?assertEqual(3, counter_value(badge_fetch_users_defaulted))
after
meck:unload(rpc_client)
end.
counter_value(Key) ->
try ets:lookup(?PUSH_COUNTERS, Key) of
[{Key, Value}] when is_integer(Value) -> Value;
_ -> 0
catch
error:badarg -> 0
end.
ensure_test_counter_table() ->
case ets:info(?PUSH_COUNTERS, name) of
undefined ->
_ = ets:new(?PUSH_COUNTERS, [named_table, public, set, {write_concurrency, true}]),
ok;
_ ->
ok
end.
reset_test_counter(Key) ->
try ets:insert(?PUSH_COUNTERS, {Key, 0}) of
_ -> ok
catch
error:badarg -> ok
end.
-endif.
@@ -1,962 +0,0 @@
%% SPDX-License-Identifier: AGPL-3.0-or-later
-module(push_sender_delivery).
-typing([eqwalizer]).
-export([
send_webpush_notification/3,
is_transient_error/1
]).
-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).
-define(MAX_RETRY_DELAY_MS, 2000).
-define(OVERLOAD_BASE_DELAY_MS, 250).
-define(OVERLOAD_MAX_DELAY_MS, 4000).
-define(VAPID_TOKEN_TTL_SECONDS, 43200).
-define(VAPID_TOKEN_SKEW_SECONDS, 60).
-define(DEFAULT_MANAGED_RELAY_HOSTS, ["push.fluxer.com"]).
-define(MANAGED_RELAY_PATH_PREFIXES, [
"/relay/v1/apns/",
"/relay/v1/apns-voip/",
"/relay/v1/fcm/"
]).
-type push_response() :: {ok, integer(), term(), binary()} | {error, term()}.
-spec send_webpush_notification(integer(), map(), map()) -> false | {true, map()}.
send_webpush_notification(UserId, Subscription, Payload) ->
case extract_subscription_fields(Subscription) of
{ok, Endpoint, P256dhKey, AuthKey, SubscriptionId} ->
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_to_consented_endpoint(
UserId, Endpoint, P256dhKey, AuthKey, SubscriptionId, Payload
);
{error, Reason} ->
log_endpoint_rejected(UserId, SubscriptionId, Reason),
false
end.
-spec send_to_consented_endpoint(integer(), binary(), binary(), binary(), binary(), map()) ->
false | {true, map()}.
send_to_consented_endpoint(UserId, Endpoint, P256dhKey, AuthKey, SubscriptionId, Payload) ->
case relay_consent_missing(Endpoint) of
false ->
send_with_vapid(UserId, Endpoint, P256dhKey, AuthKey, SubscriptionId, Payload);
true ->
log_endpoint_rejected(UserId, SubscriptionId, relay_consent_required),
false
end.
-spec relay_consent_missing(binary()) -> boolean().
relay_consent_missing(Endpoint) ->
not relay_consent_accepted() andalso is_managed_relay_endpoint(Endpoint).
-spec relay_consent_accepted() -> boolean().
relay_consent_accepted() ->
env_relay_consent_accepted() orelse instance_relay_consent_accepted().
-spec env_relay_consent_accepted() -> boolean().
env_relay_consent_accepted() ->
case fluxer_gateway_env:get(push_relay_consent_accepted) of
Accepted when is_boolean(Accepted) -> Accepted;
_ -> false
end.
-spec instance_relay_consent_accepted() -> boolean().
instance_relay_consent_accepted() ->
case maps:get(relay_consent_accepted, push_delivery_config:config(), false) of
Accepted when is_boolean(Accepted) -> Accepted;
_ -> false
end.
-spec is_managed_relay_endpoint(binary()) -> boolean().
is_managed_relay_endpoint(Endpoint) ->
case safe_parse_endpoint(Endpoint) of
{ok, Parsed} ->
Scheme = lower_string(to_string(maps:get(scheme, Parsed, ""))),
Host = lower_string(to_string(maps:get(host, Parsed, ""))),
Path = to_string(maps:get(path, Parsed, "")),
Scheme =:= "https" andalso
lists:member(Host, managed_relay_hosts()) andalso
is_managed_relay_path(Path);
error ->
false
end.
-spec safe_parse_endpoint(binary()) -> {ok, map()} | error.
safe_parse_endpoint(Endpoint) ->
try uri_string:parse(binary_to_list(Endpoint)) of
Parsed when is_map(Parsed) -> {ok, Parsed};
_ -> error
catch
_:_ -> error
end.
-spec is_managed_relay_path(string()) -> boolean().
is_managed_relay_path(Path) ->
lists:any(fun(Prefix) -> lists:prefix(Prefix, Path) end, ?MANAGED_RELAY_PATH_PREFIXES).
-spec managed_relay_hosts() -> [string()].
managed_relay_hosts() ->
case fluxer_gateway_env:get(push_managed_relay_hosts) of
Hosts when is_list(Hosts) -> [lower_string(to_string(Host)) || Host <- Hosts];
_ -> ?DEFAULT_MANAGED_RELAY_HOSTS
end.
-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) -> "".
-spec lower_string(string()) -> string().
lower_string(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 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) ->
maybe
{ok, VapidEmail, VapidPublicKey, VapidPrivateKey} ?= ensure_vapid_credentials(),
Aud = push_utils:extract_origin(Endpoint),
{ok, VapidToken} ?=
cached_vapid_token(Aud, VapidEmail, VapidPublicKey, VapidPrivateKey),
Headers = build_push_headers(VapidToken, VapidPublicKey, push_urgency(Payload)),
send_encrypted_push(
UserId,
SubscriptionId,
Endpoint,
Headers,
Payload,
P256dhKey,
AuthKey,
push_sender_retry:initial_record_size(),
0
)
else
{error, {vapid_credentials, Reason}} ->
log_vapid_unavailable(UserId, Reason),
false;
{error, _} ->
false
end.
-spec cached_vapid_token(binary(), binary(), binary(), binary()) ->
{ok, binary()} | {error, term()}.
cached_vapid_token(Aud, VapidEmail, VapidPublicKey, VapidPrivateKey) ->
cached_vapid_token(
Aud,
VapidEmail,
VapidPublicKey,
VapidPrivateKey,
erlang:system_time(second)
).
-spec cached_vapid_token(binary(), binary(), binary(), binary(), integer()) ->
{ok, binary()} | {error, term()}.
cached_vapid_token(Aud, VapidEmail, VapidPublicKey, VapidPrivateKey, Now) ->
CacheKey = vapid_cache_key(Aud, VapidEmail, VapidPublicKey),
case push_ets_cache:get_bearer_token(CacheKey) of
{ok, Token, ExpiresAt} when
is_binary(Token), ExpiresAt - ?VAPID_TOKEN_SKEW_SECONDS > Now
->
{ok, Token};
_ ->
generate_cached_vapid_token(
CacheKey, Aud, VapidEmail, VapidPublicKey, VapidPrivateKey, Now
)
end.
-spec generate_cached_vapid_token(
term(), binary(), binary(), binary(), binary(), integer()
) -> {ok, binary()} | {error, term()}.
generate_cached_vapid_token(CacheKey, Aud, VapidEmail, VapidPublicKey, VapidPrivateKey, Now) ->
ExpiresAt = Now + ?VAPID_TOKEN_TTL_SECONDS,
VapidClaims = #{
<<"sub">> => <<"mailto:", VapidEmail/binary>>,
<<"aud">> => Aud,
<<"exp">> => ExpiresAt
},
case safe_generate_vapid_token(VapidClaims, VapidPublicKey, VapidPrivateKey) of
{ok, Token} ->
push_ets_cache:put_bearer_token(CacheKey, Token, ExpiresAt),
{ok, Token};
{error, Reason} ->
{error, Reason}
end.
-spec vapid_cache_key(binary(), binary(), binary()) -> term().
vapid_cache_key(Aud, VapidEmail, VapidPublicKey) ->
{?MODULE, vapid_token, Aud, VapidEmail, VapidPublicKey}.
-spec safe_generate_vapid_token(map(), binary(), binary()) -> {ok, binary()} | {error, term()}.
safe_generate_vapid_token(VapidClaims, VapidPublicKey, VapidPrivateKey) ->
try
{ok, push_utils:generate_vapid_token(VapidClaims, VapidPublicKey, VapidPrivateKey)}
catch
error:Err -> {error, {error, Err}};
throw:Thr -> {error, {throw, Thr}};
exit:Ex -> {error, {exit, Ex}}
end.
-spec send_encrypted_push(
integer(),
binary(),
binary(),
[{binary(), binary()}],
map(),
binary(),
binary(),
pos_integer(),
non_neg_integer()
) -> false | {true, map()}.
send_encrypted_push(
UserId,
SubscriptionId,
Endpoint,
Headers,
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),
handle_encrypted_response(
UserId,
SubscriptionId,
Endpoint,
Headers,
Payload,
P256dhKey,
AuthKey,
RecordSize,
Attempt,
Response
);
{error, _EncryptError} ->
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()}],
map(),
binary(),
binary(),
pos_integer(),
non_neg_integer(),
push_response()
) -> false | {true, map()}.
handle_encrypted_response(
UserId,
SubscriptionId,
Endpoint,
Headers,
Payload,
P256dhKey,
AuthKey,
RecordSize,
Attempt,
Response
) ->
Ctx = #{
user_id => UserId,
subscription_id => SubscriptionId,
endpoint => Endpoint,
headers => Headers,
payload => Payload,
p256dh_key => P256dhKey,
auth_key => AuthKey
},
RetryResult = push_sender_retry:maybe_retry_with_smaller_record_size(
Response, RecordSize, Attempt
),
case RetryResult of
{retry, NextRecordSize} ->
retry_encrypted_push(Ctx, NextRecordSize, Attempt);
no_retry ->
maybe_retry_transient(Ctx, RecordSize, Attempt, Response)
end.
-spec maybe_retry_transient(map(), pos_integer(), non_neg_integer(), push_response()) ->
false | {true, map()}.
maybe_retry_transient(
#{user_id := UserId, subscription_id := SubscriptionId} = Ctx,
RecordSize,
Attempt,
Response
) ->
case is_local_backpressure(Response) of
true ->
maybe_retry_overload(Ctx, RecordSize, Response);
false ->
maybe_retry_transient_error(
Ctx, RecordSize, Attempt, Response, UserId, SubscriptionId
)
end.
-spec maybe_retry_transient_error(
map(), pos_integer(), non_neg_integer(), push_response(), integer(), binary()
) -> false | {true, map()}.
maybe_retry_transient_error(Ctx, RecordSize, Attempt, Response, UserId, SubscriptionId) ->
case is_transient_error(Response) andalso Attempt < ?MAX_TRANSIENT_RETRIES of
true ->
Delay = retry_delay(Attempt),
ok = gateway_retry_timer:wait(Delay),
retry_encrypted_push(Ctx, RecordSize, Attempt);
false ->
handle_push_response(UserId, SubscriptionId, Response)
end.
-spec is_local_backpressure(push_response()) -> boolean().
is_local_backpressure({error, overloaded}) -> true;
is_local_backpressure({error, circuit_open}) -> true;
is_local_backpressure(_) -> false.
-spec maybe_retry_overload(map(), pos_integer(), push_response()) -> false | {true, map()}.
maybe_retry_overload(Ctx, RecordSize, Response) ->
OverloadAttempt = maps:get(overload_attempt, Ctx, 0),
case OverloadAttempt < ?MAX_OVERLOAD_RETRIES of
true ->
Delay = overload_retry_delay(OverloadAttempt),
ok = gateway_retry_timer:wait(Delay),
retry_overload(Ctx, RecordSize, OverloadAttempt + 1);
false ->
log_overload_dropped(Ctx, Response),
false
end.
-spec retry_overload(map(), pos_integer(), non_neg_integer()) -> false | {true, map()}.
retry_overload(
#{
endpoint := Endpoint,
headers := Headers,
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),
handle_overload_retry_response(Ctx, RecordSize, OverloadAttempt, Response);
{error, _EncryptError} ->
false
end.
-spec handle_overload_retry_response(map(), pos_integer(), non_neg_integer(), push_response()) ->
false | {true, map()}.
handle_overload_retry_response(Ctx, RecordSize, OverloadAttempt, Response) ->
case is_local_backpressure(Response) of
true ->
maybe_retry_overload(
Ctx#{overload_attempt => OverloadAttempt}, RecordSize, Response
);
false ->
handle_push_response(
maps:get(user_id, Ctx), maps:get(subscription_id, Ctx), Response
)
end.
-spec overload_retry_delay(non_neg_integer()) -> pos_integer().
overload_retry_delay(Attempt) ->
Base = min(?OVERLOAD_MAX_DELAY_MS, ?OVERLOAD_BASE_DELAY_MS * (1 bsl Attempt)),
Jitter = rand:uniform(max(1, Base div 4)),
min(?OVERLOAD_MAX_DELAY_MS, Base + Jitter - 1).
-spec log_overload_dropped(map(), push_response()) -> ok.
log_overload_dropped(#{user_id := UserId, subscription_id := SubscriptionId}, Response) ->
Reason = backpressure_reason(Response),
logger:warning(
"Push: notification dropped after overload retries exhausted",
#{user_id => UserId, subscription_id => SubscriptionId, reason => Reason}
),
ok.
-spec backpressure_reason(push_response()) -> binary().
backpressure_reason({error, overloaded}) -> <<"client_overloaded">>;
backpressure_reason({error, circuit_open}) -> <<"circuit_open">>;
backpressure_reason(_) -> <<"backpressure">>.
-spec retry_encrypted_push(map(), pos_integer(), non_neg_integer()) -> false | {true, map()}.
retry_encrypted_push(
#{
user_id := UserId,
subscription_id := SubscriptionId,
endpoint := Endpoint,
headers := Headers,
payload := Payload,
p256dh_key := P256dhKey,
auth_key := AuthKey
},
RecordSize,
Attempt
) ->
send_encrypted_push(
UserId,
SubscriptionId,
Endpoint,
Headers,
Payload,
P256dhKey,
AuthKey,
RecordSize,
Attempt + 1
).
-spec is_transient_error(push_response()) -> boolean().
is_transient_error({ok, Status, _, _}) when Status >= 500 -> true;
is_transient_error({ok, 429, _, _}) -> true;
is_transient_error({error, overloaded}) -> false;
is_transient_error({error, circuit_open}) -> false;
is_transient_error({error, timeout}) -> true;
is_transient_error({error, {timeout, _}}) -> true;
is_transient_error({error, closed}) -> true;
is_transient_error({error, econnrefused}) -> true;
is_transient_error({error, econnreset}) -> true;
is_transient_error({error, ehostunreach}) -> true;
is_transient_error({error, enetunreach}) -> true;
is_transient_error({error, etimedout}) -> true;
is_transient_error({error, {failed_connect, _}}) -> true;
is_transient_error({error, nxdomain}) -> true;
is_transient_error({error, _}) -> false;
is_transient_error(_) -> false.
-spec retry_delay(non_neg_integer()) -> pos_integer().
retry_delay(Attempt) ->
Base = min(?MAX_RETRY_DELAY_MS, ?BASE_RETRY_DELAY_MS * (1 bsl Attempt)),
Jitter = rand:uniform(max(1, Base div 4)),
min(?MAX_RETRY_DELAY_MS, Base + Jitter - 1).
-spec handle_push_response(integer(), binary(), push_response()) -> false | {true, map()}.
handle_push_response(UserId, SubscriptionId, {ok, Status, _, _}) when
Status >= 200, Status < 300
->
logger:debug(
"Push: delivery succeeded",
#{user_id => UserId, subscription_id => SubscriptionId, status => Status}
),
false;
handle_push_response(UserId, SubscriptionId, {ok, 410, _, _}) ->
log_and_delete(UserId, SubscriptionId, <<"expired">>);
handle_push_response(UserId, SubscriptionId, {ok, 404, _, _}) ->
log_and_delete(UserId, SubscriptionId, <<"not_found">>);
handle_push_response(UserId, SubscriptionId, {ok, Status, _, Body}) ->
log_http_error(UserId, SubscriptionId, Status, Body);
handle_push_response(UserId, SubscriptionId, {error, overloaded}) ->
log_push_error(
UserId, SubscriptionId, <<"client_overloaded">>, "Push: HTTP client overloaded"
);
handle_push_response(UserId, SubscriptionId, {error, circuit_open}) ->
log_push_error(UserId, SubscriptionId, <<"circuit_open">>, "Push: circuit breaker open");
handle_push_response(UserId, SubscriptionId, {error, Reason}) ->
logger:debug(
"Push: network error",
#{user_id => UserId, subscription_id => SubscriptionId, reason => Reason}
),
false.
-spec log_and_delete(integer(), binary(), binary()) -> {true, map()}.
log_and_delete(UserId, SubscriptionId, Reason) ->
logger:debug(
"Push: subscription gone, will delete",
#{user_id => UserId, subscription_id => SubscriptionId, reason => Reason}
),
{true, delete_payload(UserId, SubscriptionId)}.
-spec log_push_error(integer(), binary(), binary(), string()) -> false.
log_push_error(UserId, SubscriptionId, _Reason, Message) ->
logger:debug(
Message,
#{user_id => UserId, subscription_id => SubscriptionId}
),
false.
-spec log_http_error(integer(), binary(), integer(), binary()) -> false.
log_http_error(UserId, SubscriptionId, Status, Body) ->
logger:debug(
"Push: delivery failed with HTTP error",
#{user_id => UserId, subscription_id => SubscriptionId, status => Status, body => Body}
),
false.
-spec ensure_vapid_credentials() ->
{ok, binary(), binary(), binary()} | {error, {vapid_credentials, string()}}.
ensure_vapid_credentials() ->
Email = fluxer_gateway_env:get(vapid_email),
Public = fluxer_gateway_env:get(vapid_public_key),
Private = fluxer_gateway_env:get(vapid_private_key),
case {Email, Public, Private} of
{Email0, Public0, Private0} when
is_binary(Email0),
is_binary(Public0),
is_binary(Private0),
byte_size(Public0) > 0,
byte_size(Private0) > 0
->
{ok, Email0, Public0, Private0};
_ ->
{error, {vapid_credentials, "Missing VAPID credentials"}}
end.
-spec log_vapid_unavailable(integer(), term()) -> ok.
log_vapid_unavailable(UserId, Reason) ->
Now = erlang:monotonic_time(second),
Last =
case persistent_term:get({?MODULE, vapid_warn_at}, undefined) of
undefined -> 0;
V -> V
end,
case Now - Last >= 60 of
true ->
persistent_term:put({?MODULE, vapid_warn_at}, Now),
logger:error(
"Push: VAPID credentials unavailable; pushes are being silently dropped",
#{user_id => UserId, reason => Reason}
);
false ->
ok
end,
ok.
-spec extract_subscription_fields(map()) ->
{ok, binary(), binary(), binary(), binary()} | {error, string()}.
extract_subscription_fields(Subscription) ->
Endpoint = maps:get(<<"endpoint">>, Subscription, undefined),
P256dhKey = maps:get(<<"p256dh_key">>, Subscription, undefined),
AuthKey = maps:get(<<"auth_key">>, Subscription, undefined),
SubscriptionId = maps:get(<<"subscription_id">>, Subscription, undefined),
case {Endpoint, P256dhKey, AuthKey, SubscriptionId} of
{E, P, A, S} when is_binary(E), is_binary(P), is_binary(A), is_binary(S) ->
{ok, E, P, A, S};
_ ->
{error, "missing keys"}
end.
-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>>}
].
-spec request_push_endpoint(binary(), [{binary(), binary()}], binary()) ->
{ok, non_neg_integer(), [{binary(), binary()}], binary()} | {error, term()}.
request_push_endpoint(Endpoint, Headers, Body) ->
gateway_http_client:request(
push,
post,
Endpoint,
Headers,
Body,
#{content_type => <<"application/octet-stream">>}
).
-spec delete_payload(integer(), binary()) -> map().
delete_payload(UserId, SubscriptionId) ->
#{
<<"user_id">> => integer_to_binary(UserId),
<<"subscription_id">> => SubscriptionId
}.
-ifdef(TEST).
-include_lib("eunit/include/eunit.hrl").
is_transient_error_5xx_test() ->
?assertEqual(true, is_transient_error({ok, 500, [], <<>>})),
?assertEqual(true, is_transient_error({ok, 502, [], <<>>})),
?assertEqual(true, is_transient_error({ok, 503, [], <<>>})).
is_transient_error_429_test() ->
?assertEqual(true, is_transient_error({ok, 429, [], <<>>})).
is_transient_error_network_test() ->
?assertEqual(true, is_transient_error({error, timeout})),
?assertEqual(true, is_transient_error({error, econnrefused})).
is_transient_error_local_backpressure_test() ->
?assertEqual(false, is_transient_error({error, overloaded})),
?assertEqual(false, is_transient_error({error, circuit_open})).
is_transient_error_permanent_test() ->
?assertEqual(false, is_transient_error({ok, 200, [], <<>>})),
?assertEqual(false, is_transient_error({ok, 201, [], <<>>})),
?assertEqual(false, is_transient_error({ok, 400, [], <<>>})),
?assertEqual(false, is_transient_error({ok, 404, [], <<>>})),
?assertEqual(false, is_transient_error({ok, 410, [], <<>>})).
retry_delay_exponential_backoff_test() ->
D0 = retry_delay(0),
D1 = retry_delay(1),
D2 = retry_delay(2),
?assert(D0 >= ?BASE_RETRY_DELAY_MS),
?assert(D0 < ?BASE_RETRY_DELAY_MS * 2),
?assert(D1 >= ?BASE_RETRY_DELAY_MS * 2),
?assert(D1 < ?BASE_RETRY_DELAY_MS * 3),
?assert(D2 =< ?MAX_RETRY_DELAY_MS).
retry_delay_capped_test() ->
D10 = retry_delay(10),
?assert(D10 =< ?MAX_RETRY_DELAY_MS),
?assert(D10 >= ?MAX_RETRY_DELAY_MS - (?MAX_RETRY_DELAY_MS div 4)).
maybe_retry_transient_waits_before_retry_test() ->
Self = self(),
Endpoint = <<"https://push.example/sub-1">>,
RecordSize = push_sender_retry:initial_record_size(),
Ctx = #{
user_id => 42,
subscription_id => <<"sub-1">>,
endpoint => Endpoint,
headers => [],
payload => #{},
p256dh_key => <<"p256dh">>,
auth_key => <<"auth">>
},
ok = meck:new(gateway_retry_timer, [passthrough, no_link]),
ok = meck:new(push_utils, [passthrough, no_link]),
ok = meck:new(gateway_http_client, [passthrough, no_link]),
try
ok = meck:expect(gateway_retry_timer, wait, fun(DelayMs) ->
Self ! {retry_wait, DelayMs},
ok
end),
ok = meck:expect(push_utils, encrypt_payload, fun(
<<"{}">>, <<"p256dh">>, <<"auth">>, ActualRecordSize
) when ActualRecordSize =:= RecordSize ->
{ok, <<"encrypted">>}
end),
ok = meck:expect(
gateway_http_client,
request,
fun retry_request_meck/6
),
?assertEqual(false, maybe_retry_transient(Ctx, RecordSize, 0, {error, timeout})),
%% retry_delay/1 adds rand:uniform(Base div 4) of jitter, so the first
%% retry is anywhere in [Base, Base + Base div 4 - 1].
{retry_wait, DelayMs} = receive_retry_wait(),
?assert(DelayMs >= ?BASE_RETRY_DELAY_MS),
?assert(DelayMs =< ?BASE_RETRY_DELAY_MS + (?BASE_RETRY_DELAY_MS div 4) - 1),
?assertEqual(ok, receive_retried_push_request(Endpoint)),
?assert(meck:validate(gateway_retry_timer)),
?assert(meck:validate(push_utils)),
?assert(meck:validate(gateway_http_client))
after
meck:unload(gateway_http_client),
meck:unload(push_utils),
meck:unload(gateway_retry_timer)
end.
-spec retry_request_meck(atom(), atom(), binary(), list(), binary(), term()) ->
{ok, non_neg_integer(), list(), binary()}.
retry_request_meck(push, post, Endpoint, [], <<"encrypted">>, _Opts) ->
self() ! {retried_push_request, Endpoint},
{ok, 201, [], <<>>}.
extract_subscription_fields_ok_test() ->
Sub = #{
<<"endpoint">> => <<"https://push.example.com/sub1">>,
<<"p256dh_key">> => <<"key1">>,
<<"auth_key">> => <<"auth1">>,
<<"subscription_id">> => <<"sub1">>
},
?assertMatch({ok, _, _, _, _}, extract_subscription_fields(Sub)).
extract_subscription_fields_missing_test() ->
?assertMatch({error, _}, extract_subscription_fields(#{})).
cached_vapid_token_reuses_unexpired_token_test() ->
Aud = <<"https://push.example">>,
VapidEmail = <<"[email protected]">>,
VapidPublicKey = <<"public">>,
VapidPrivateKey = <<"private">>,
CacheKey = vapid_cache_key(Aud, VapidEmail, VapidPublicKey),
erase_token_cache(CacheKey),
Self = self(),
ok = meck:new(push_utils, [passthrough, no_link]),
try
ok = meck:expect(push_utils, generate_vapid_token, fun(
_Claims, _PublicKey, _PrivateKey
) ->
Self ! vapid_generated,
<<"cached-token">>
end),
?assertEqual(
{ok, <<"cached-token">>},
cached_vapid_token(Aud, VapidEmail, VapidPublicKey, VapidPrivateKey, 1000)
),
?assertEqual(
{ok, <<"cached-token">>},
cached_vapid_token(Aud, VapidEmail, VapidPublicKey, VapidPrivateKey, 1001)
),
?assertEqual(1, drain_vapid_generated(0)),
?assert(meck:validate(push_utils))
after
meck:unload(push_utils),
erase_token_cache(CacheKey)
end.
drain_vapid_generated(Count) ->
receive
vapid_generated -> drain_vapid_generated(Count + 1)
after 0 ->
Count
end.
receive_retry_wait() ->
receive
{retry_wait, DelayMs} -> {retry_wait, DelayMs}
after 100 ->
timeout
end.
receive_retried_push_request(ExpectedEndpoint) ->
receive
{retried_push_request, ExpectedEndpoint} -> ok
after 100 ->
timeout
end.
erase_token_cache(Key) ->
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.com/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) -> <<"[email protected]">>;
vapid_env_meck(vapid_public_key) -> <<"public-key">>;
vapid_env_meck(vapid_private_key) -> <<"private-key">>;
vapid_env_meck(push_relay_consent_accepted) -> true;
vapid_env_meck(push_managed_relay_hosts) -> [<<"push.fluxer.com">>];
vapid_env_meck(Key) -> meck:passthrough([Key]).
capture_request_meck(push, post, _Endpoint, Headers, Body, _Opts) ->
self() ! {captured_push, Headers, Body},
{ok, 201, [], <<>>}.
a_managed_relay_endpoint_is_refused_without_operator_consent_test() ->
?assertEqual(no_push_request, attempt_push(false, managed_relay_endpoint(<<"apns">>))).
every_managed_relay_leg_is_refused_without_operator_consent_test() ->
lists:foreach(
fun(Leg) ->
?assertEqual(no_push_request, attempt_push(false, managed_relay_endpoint(Leg)))
end,
[<<"apns">>, <<"apns-voip">>, <<"fcm">>]
).
a_managed_relay_endpoint_is_delivered_once_the_operator_consents_test() ->
Endpoint = managed_relay_endpoint(<<"apns">>),
?assertEqual(Endpoint, attempt_push(true, Endpoint)).
a_notice_accepted_in_the_instance_config_lets_the_managed_relay_send_through_test() ->
Endpoint = managed_relay_endpoint(<<"apns">>),
?assertEqual(Endpoint, attempt_push(false, true, Endpoint)).
a_unified_push_endpoint_is_delivered_whatever_the_operator_accepted_test() ->
Endpoint = <<"https://ntfy.sh/upZzH87cT9jJCc?up=1">>,
?assertEqual(Endpoint, attempt_push(false, Endpoint)),
?assertEqual(Endpoint, attempt_push(true, Endpoint)).
a_relay_we_do_not_operate_is_delivered_without_consent_test() ->
Endpoint = <<"https://push.example.org/relay/v1/apns/stable/production/token">>,
?assertEqual(Endpoint, attempt_push(false, Endpoint)).
managed_relay_endpoint(Leg) ->
<<"https://push.fluxer.com/relay/v1/", Leg/binary, "/stable/production/",
(binary:copy(<<"a">>, 64))/binary>>.
attempt_push(EnvConsent, Endpoint) ->
attempt_push(EnvConsent, false, Endpoint).
attempt_push(EnvConsent, InstanceConsent, Endpoint) ->
{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]),
ok = meck:new(push_delivery_config, [passthrough, no_link]),
try
ok = meck:expect(push_endpoint_guard, check, fun(_Endpoint) -> ok end),
ok = meck:expect(push_delivery_config, config, fun() ->
#{relay_consent_accepted => InstanceConsent}
end),
ok = meck:expect(fluxer_gateway_env, get, consent_env_meck(EnvConsent)),
ok = meck:expect(push_utils, generate_vapid_token, fun(_Claims, _Public, _Private) ->
<<"vapid-token">>
end),
ok = meck:expect(gateway_http_client, request, fun requested_endpoint_meck/6),
?assertEqual(
false, send_webpush_notification(42, Subscription, alert_payload(<<"Hello">>))
),
receive
{push_requested, Requested} -> Requested
after 100 ->
no_push_request
end
after
meck:unload(push_delivery_config),
meck:unload(push_endpoint_guard),
meck:unload(gateway_http_client),
meck:unload(push_utils),
meck:unload(fluxer_gateway_env)
end.
consent_env_meck(EnvConsent) ->
fun
(push_relay_consent_accepted) -> EnvConsent;
(push_managed_relay_hosts) -> [<<"push.fluxer.com">>];
(Key) -> vapid_env_meck(Key)
end.
-spec requested_endpoint_meck(atom(), atom(), binary(), list(), binary(), term()) ->
{ok, non_neg_integer(), list(), binary()}.
requested_endpoint_meck(push, post, Endpoint, _Headers, _Body, _Opts) ->
self() ! {push_requested, Endpoint},
{ok, 201, [], <<>>}.
-endif.
@@ -1,166 +0,0 @@
%% SPDX-License-Identifier: AGPL-3.0-or-later
-module(push_sender_retry).
-typing([eqwalizer]).
-export([
maybe_retry_with_smaller_record_size/3,
initial_record_size/0
]).
-export_type([push_response/0]).
-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() -> pos_integer().
initial_record_size() ->
?PUSH_RECORD_SIZE.
-spec maybe_retry_with_smaller_record_size(
push_response(), pos_integer(), non_neg_integer()
) ->
no_retry | {retry, pos_integer()}.
maybe_retry_with_smaller_record_size(_Response, _CurrentRecordSize, Attempt) when
Attempt >= ?MAX_PAYLOAD_RETRY_ATTEMPTS
->
no_retry;
maybe_retry_with_smaller_record_size(
{ok, 413, _ResponseHeaders, ResponseBody}, CurrentRecordSize, _Attempt
) ->
case next_record_size_for_payload_too_large(CurrentRecordSize, ResponseBody) of
undefined -> no_retry;
NextRecordSize -> {retry, NextRecordSize}
end;
maybe_retry_with_smaller_record_size(_Response, _CurrentRecordSize, _Attempt) ->
no_retry.
-spec next_record_size_for_payload_too_large(pos_integer(), binary()) ->
pos_integer() | undefined.
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);
_ ->
sanitize_next_record_size(?CONSTRAINED_PUSH_RECORD_SIZE, CurrentRecordSize)
end.
-spec sanitize_next_record_size(integer(), pos_integer()) ->
pos_integer() | undefined.
sanitize_next_record_size(CandidateRecordSize, CurrentRecordSize) when
is_integer(CandidateRecordSize)
->
ClampedRecordSize = erlang:max(?MIN_PUSH_RECORD_SIZE, CandidateRecordSize),
case ClampedRecordSize < CurrentRecordSize of
true -> ClampedRecordSize;
false -> undefined
end.
-spec parse_constrained_overage_bytes(binary()) -> non_neg_integer() | undefined.
parse_constrained_overage_bytes(ResponseBody) ->
case decode_push_error_body(ResponseBody) of
#{<<"message">> := Message} -> parse_constrained_overage_from_message(Message);
_ -> undefined
end.
-spec parse_constrained_overage_from_message(binary() | list()) ->
non_neg_integer() | undefined.
parse_constrained_overage_from_message(Message) when is_list(Message) ->
parse_constrained_overage_from_message(list_to_binary(Message));
parse_constrained_overage_from_message(Message) when is_binary(Message) ->
case
re:run(Message, <<"too long by ([0-9]+) bytes">>, [caseless, {capture, [1], binary}])
of
{match, [OverageBytesBin]} -> parse_non_neg_integer(OverageBytesBin);
_ -> undefined
end.
-spec decode_push_error_body(binary()) -> map() | undefined.
decode_push_error_body(ResponseBody) when
is_binary(ResponseBody), byte_size(ResponseBody) > 0
->
try json:decode(ResponseBody) of
ParsedBody when is_map(ParsedBody) -> ParsedBody;
_ -> undefined
catch
error:_ -> undefined;
throw:_ -> undefined;
exit:_ -> undefined
end;
decode_push_error_body(_ResponseBody) ->
undefined.
-spec parse_non_neg_integer(binary()) -> non_neg_integer() | undefined.
parse_non_neg_integer(Value) ->
case guild_data_normalize_schema:int(Value) of
ParsedValue when ParsedValue >= 0 -> ParsedValue;
_ -> undefined
end.
-ifdef(TEST).
-include_lib("eunit/include/eunit.hrl").
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 441 bytes\"}"
>>,
?assertEqual(
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\"}">>,
?assertEqual(
?CONSTRAINED_PUSH_RECORD_SIZE,
next_record_size_for_payload_too_large(?PUSH_RECORD_SIZE, ResponseBody)
),
?assertEqual(
undefined,
next_record_size_for_payload_too_large(
?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.
File diff suppressed because it is too large Load Diff
-269
View File
@@ -7,27 +7,12 @@
construct_avatar_url/2,
construct_static_asset_url/1,
get_default_avatar_url/1,
extract_origin/1,
generate_vapid_token/3,
assert_vapid_pair/2,
generate_jwt_from_pem/3,
base64url_encode/1,
base64url_decode/1,
encrypt_payload/4,
plaintext_budget/1,
decode_subscription_key/1,
hkdf_expand/4,
hkdf_expand_loop/6,
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(),
@@ -111,260 +96,6 @@ strip_trailing_slashes(Value) ->
_ -> Value
end.
-spec extract_origin(binary()) -> binary().
extract_origin(Url) ->
case binary:split(Url, <<"://">>) of
[Protocol, Rest] ->
extract_origin_host(Protocol, Rest, Url);
_ ->
Url
end.
-spec extract_origin_host(binary(), binary(), binary()) -> binary().
extract_origin_host(Protocol, Rest, Url) ->
case binary:split(Rest, <<"/">>) of
[Host | _] -> <<Protocol/binary, "://", Host/binary>>;
_ -> Url
end.
-spec generate_vapid_token(map(), binary(), binary()) -> binary().
generate_vapid_token(Claims, PublicKeyB64Url, PrivateKeyB64Url) ->
try
ensure_crypto_started(),
PrivRaw = decode_or_error(PrivateKeyB64Url, invalid_private_key),
PubRaw = decode_or_error(PublicKeyB64Url, invalid_public_key),
JWK = build_ec_jwk(PrivRaw, PubRaw),
Header = #{<<"alg">> => <<"ES256">>, <<"typ">> => <<"JWT">>},
sign_and_compact(JWK, Header, Claims)
catch
C:R:_Stack ->
erlang:error({vapid_token_generation_failed, C, R})
end.
-spec ensure_crypto_started() -> ok.
ensure_crypto_started() ->
{ok, _} = application:ensure_all_started(crypto),
{ok, _} = application:ensure_all_started(public_key),
{ok, _} = application:ensure_all_started(jose),
ok.
-spec decode_or_error(binary(), atom()) -> binary().
decode_or_error(B64Url, ErrorAtom) ->
case base64url_decode(B64Url) of
error -> erlang:error(ErrorAtom);
Decoded -> Decoded
end.
-spec build_ec_jwk(binary(), binary()) -> term().
build_ec_jwk(PrivRaw, PubRaw) ->
<<4, X:32/binary, Y:32/binary>> = PubRaw,
JWKMap = #{
<<"kty">> => <<"EC">>,
<<"crv">> => <<"P-256">>,
<<"d">> => base64url_encode(PrivRaw),
<<"x">> => base64url_encode(X),
<<"y">> => base64url_encode(Y)
},
unwrap_jwk(jose_jwk:from_map(JWKMap)).
-spec unwrap_jwk(term()) -> term().
unwrap_jwk({JW, _Fields}) -> JW;
unwrap_jwk(JW) -> JW.
-spec assert_vapid_pair(binary(), binary()) -> ok.
assert_vapid_pair(PublicKeyB64Url, PrivateKeyB64Url) ->
ensure_crypto_started(),
PubRaw = decode_or_error(PublicKeyB64Url, invalid_public_key),
PrivRaw = decode_or_error(PrivateKeyB64Url, invalid_private_key),
case {PubRaw, PrivRaw} of
{<<4, _:64/binary>>, <<_:32/binary>>} ->
assert_vapid_scalar_derives_point(PublicKeyB64Url, PubRaw, PrivRaw);
_ ->
erlang:error({vapid_keys_malformed, byte_size(PubRaw), byte_size(PrivRaw)})
end.
-spec assert_vapid_scalar_derives_point(binary(), binary(), binary()) -> ok.
assert_vapid_scalar_derives_point(PublicKeyB64Url, PubRaw, PrivRaw) ->
Derived =
try crypto:generate_key(ecdh, prime256v1, PrivRaw) of
{Point, _} -> Point
catch
_:_ -> undefined
end,
case Derived of
PubRaw -> ok;
_ -> erlang:error({vapid_keys_mismatched, PublicKeyB64Url})
end.
-spec sign_and_compact(term(), map(), map()) -> binary().
sign_and_compact(JWK, Header, Claims) ->
JWS = jose_jwt:sign(JWK, Header, Claims),
case jose_jws:compact(JWS) of
{_Meta, Bin} when is_binary(Bin) -> Bin;
Other -> erlang:error({unexpected_compact_return, Other})
end.
-spec generate_jwt_from_pem(binary(), map(), map()) -> {ok, binary()} | {error, term()}.
generate_jwt_from_pem(Pem, Header, Claims) when
is_binary(Pem), is_map(Header), is_map(Claims)
->
try
{ok, _} = application:ensure_all_started(crypto),
{ok, _} = application:ensure_all_started(public_key),
{ok, _} = application:ensure_all_started(jose),
JWK0 = jose_jwk:from_pem(Pem),
JWK =
case JWK0 of
{JW, _Fields} -> JW;
JW -> JW
end,
JWS = jose_jwt:sign(JWK, Header, Claims),
Compact0 = jose_jws:compact(JWS),
CompactBin =
case Compact0 of
{_Meta, Bin} when is_binary(Bin) -> Bin;
Other -> erlang:error({unexpected_compact_return, Other})
end,
{ok, CompactBin}
catch
C:R:Stack ->
logger:error(
"Push: JWT signing from PEM failed",
#{class => C, reason => R, stack => Stack}
),
{error, jwt_signing_failed}
end.
-spec base64url_encode(binary()) -> binary().
base64url_encode(Data) ->
jose_base64url:encode(Data).
-spec base64url_decode(binary()) -> binary() | error.
base64url_decode(Data) ->
case jose_base64url:decode(Data) of
{ok, Decoded} -> Decoded;
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) ->
try
PeerPub = decode_subscription_key(PeerPubB64),
AuthSecret = decode_subscription_key(AuthSecretB64),
RecordSize =
case RecordSize0 of
0 -> push_sender_retry:initial_record_size();
_ -> RecordSize0
end,
Salt = crypto:strong_rand_bytes(16),
{LocalPub, LocalPriv} = generate_local_ecdh_key(),
<<4, _/binary>> = PeerPub,
Keys = derive_encryption_keys(PeerPub, LocalPub, LocalPriv, AuthSecret, Salt),
encrypt_and_build_body(Message, Salt, RecordSize, LocalPub, Keys)
catch
C:R:Stack ->
logger:error(
"Push: encrypt_payload failed",
#{class => C, reason => R, stack => Stack}
),
{error, encryption_failed}
end.
-spec generate_local_ecdh_key() -> {binary(), binary()}.
generate_local_ecdh_key() ->
case crypto:generate_key(ecdh, prime256v1) of
{LocalPub, LocalPriv} when is_binary(LocalPub), is_binary(LocalPriv) ->
{LocalPub, LocalPriv};
Other ->
erlang:error({invalid_ecdh_key, Other})
end.
-spec derive_encryption_keys(binary(), binary(), binary(), binary(), binary()) ->
{binary(), binary()}.
derive_encryption_keys(PeerPub, LocalPub, LocalPriv, AuthSecret, Salt) ->
Secret = crypto:compute_key(ecdh, PeerPub, LocalPriv, prime256v1),
PRKInfo = <<"WebPush: info", 0, PeerPub/binary, LocalPub/binary>>,
IKM = hkdf_expand(Secret, AuthSecret, PRKInfo, 32),
CEK = hkdf_expand(IKM, Salt, <<"Content-Encoding: aes128gcm", 0>>, 16),
Nonce = hkdf_expand(IKM, Salt, <<"Content-Encoding: nonce", 0>>, 12),
{CEK, Nonce}.
-spec encrypt_and_build_body(
binary(), binary(), non_neg_integer(), binary(), {binary(), binary()}
) -> {ok, binary()} | {error, term()}.
encrypt_and_build_body(Message, Salt, RecordSize, LocalPub, {CEK, Nonce}) ->
HeaderLen = 16 + 4 + 1 + byte_size(LocalPub),
RecordLen = RecordSize - 16,
Data0 = <<Message/binary, 16#02>>,
Required = RecordLen - HeaderLen,
case byte_size(Data0) > Required of
true ->
{error, max_pad_exceeded};
false ->
Data = pad_data(Data0, Required),
{Cipher, Tag} = crypto:crypto_one_time_aead(
aes_gcm, CEK, Nonce, Data, <<>>, 16, true
),
Ciphertext = <<Cipher/binary, Tag/binary>>,
Body =
<<Salt/binary, RecordSize:32/big-unsigned-integer, (byte_size(LocalPub)):8,
LocalPub/binary, Ciphertext/binary>>,
{ok, Body}
end.
-spec pad_data(binary(), non_neg_integer()) -> binary().
pad_data(Data0, Required) ->
PadLen = Required - byte_size(Data0),
Padding =
case PadLen of
0 -> <<>>;
_ -> binary:copy(<<0>>, PadLen)
end,
<<Data0/binary, Padding/binary>>.
-spec decode_subscription_key(binary()) -> binary().
decode_subscription_key(B64) when is_binary(B64) ->
Padded =
case byte_size(B64) rem 4 of
0 -> B64;
Rem -> <<B64/binary, (binary:copy(<<"=">>, 4 - Rem))/binary>>
end,
case jose_base64url:decode(Padded) of
{ok, Decoded} ->
Decoded;
_ ->
decode_base64_subscription_key(Padded)
end.
-spec decode_base64_subscription_key(binary()) -> binary().
decode_base64_subscription_key(Padded) ->
try base64:decode(Padded) of
Decoded when is_binary(Decoded) -> Decoded
catch
_:_ -> erlang:error(decode_key_error)
end.
-spec hkdf_expand(binary(), binary(), binary(), pos_integer()) -> binary().
hkdf_expand(IKM, Salt, Info, Length) ->
PRK = crypto:mac(hmac, sha256, Salt, IKM),
hkdf_expand_loop(PRK, Info, Length, 1, <<>>, <<>>).
-spec hkdf_expand_loop(binary(), binary(), pos_integer(), pos_integer(), binary(), binary()) ->
binary().
hkdf_expand_loop(_PRK, _Info, Length, _I, _Tprev, Acc) when byte_size(Acc) >= Length ->
binary:part(Acc, 0, Length);
hkdf_expand_loop(PRK, Info, Length, I, Tprev, Acc) ->
T = crypto:mac(hmac, sha256, PRK, <<Tprev/binary, Info/binary, I:8/integer>>),
hkdf_expand_loop(PRK, Info, Length, I + 1, T, <<Acc/binary, T/binary>>).
-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);
+18 -15
View File
@@ -66,24 +66,10 @@ handle_call_or_ignore(Ref, Reason, State, Calls) ->
-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);
{stop, normal, State#{socket_pid => undefined, socket_mref => undefined}};
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(),
@@ -353,6 +339,23 @@ handle_socket_down_delays_session_offline_test() ->
end,
SocketPid ! stop.
handle_socket_down_client_closed_stops_the_session_test() ->
SocketRef = make_ref(),
SocketPid = spawn_test_proc(),
State0 = (build_test_session_state(50000, #{}))#{
presence_pid => self(),
socket_pid => SocketPid,
socket_mref => SocketRef
},
{stop, normal, State1} = handle_process_down(
SocketRef, {shutdown, client_closed}, State0
),
?assertEqual(undefined, maps:get(socket_pid, State1)),
?assertEqual(undefined, maps:get(socket_mref, State1)),
?assertEqual(maps:get(resume_timer, State0), maps:get(resume_timer, State1)),
?assertEqual(maps:get(offline_timer, State0), maps:get(offline_timer, State1)),
SocketPid ! stop.
spawn_test_proc() ->
spawn(fun test_proc_loop/0).
@@ -118,7 +118,6 @@ singleton_pids() ->
presence_manager,
guild_manager,
call_manager,
push_dispatcher,
push,
gateway_nats_rpc,
gateway_nats_pool,
@@ -80,16 +80,31 @@ cluster_static_peers_accepts_valid_node_names_test() ->
maps:get(cluster_static_peers, Config)
).
push_endpoint_guard_enabled_defaults_on_test() ->
Config = fluxer_gateway_config:load(),
?assertEqual(true, maps:get(push_endpoint_guard_enabled, Config)).
push_endpoint_guard_enabled_can_be_turned_off_test() ->
with_env("FLUXER_GATEWAY_PUSH_ENDPOINT_GUARD_ENABLED", "false", fun() ->
push_clear_switch_reads_the_enrolled_env_name_test() ->
?assertEqual(
true, maps:get(push_enrolled_clear_notifications_enabled, fluxer_gateway_config:load())
),
with_env("FLUXER_GATEWAY_PUSH_ENROLLED_CLEAR_NOTIFICATIONS_ENABLED", "false", fun() ->
Config = fluxer_gateway_config:load(),
?assertEqual(false, maps:get(push_endpoint_guard_enabled, Config))
?assertEqual(false, maps:get(push_enrolled_clear_notifications_enabled, Config))
end).
gateway_config_holds_no_direct_push_delivery_keys_test() ->
Config = fluxer_gateway_config:load(),
lists:foreach(
fun(Key) -> ?assertNot(maps:is_key(Key, Config)) end,
[
push_clear_notifications_enabled,
push_endpoint_guard_enabled,
push_managed_relay_hosts,
push_relay_consent_accepted,
vapid_public_key,
apns_enabled,
fcm_enabled,
gateway_http_push_max_concurrency
]
).
presence_push_buffer_env_defaults_test() ->
with_env("FLUXER_GATEWAY_PRESENCE_PUSH_BUFFER_MAX_ENTRIES", "7", fun() ->
with_env("FLUXER_GATEWAY_PRESENCE_PUSH_BUFFER_MAX_BYTES", "4096", fun() ->
@@ -99,55 +114,20 @@ presence_push_buffer_env_defaults_test() ->
end)
end).
env_only_push_and_http_runtime_config_test() ->
ApnsAppsJson =
"[{\"app_id\":\"ios-stable\",\"topic\":\"app.fluxer\",\"environment\":\"production\"}]",
FcmAppsJson = "[{\"app_id\":\"android-stable\",\"project_id\":\"fluxer-fcm\"}]",
env_only_http_runtime_config_test() ->
with_envs(
[
{"FLUXER_GATEWAY_SHUTDOWN_DRAIN_WAIT_MS", "1234"},
{"FLUXER_GATEWAY_HTTP_RPC_MAX_CONCURRENCY", "42"},
{"FLUXER_GATEWAY_HTTP_FAILURE_THRESHOLD", "9"},
{"FLUXER_GATEWAY_HTTP_RECOVERY_TIMEOUT_MS", "6000"},
{"FLUXER_PUSH_APNS_ENABLED", "true"},
{"FLUXER_PUSH_APNS_TEAM_ID", "TEAMID"},
{"FLUXER_PUSH_APNS_KEY_ID", "KEYID"},
{"FLUXER_PUSH_APNS_PRIVATE_KEY_PATH", "/etc/fluxer/apns.p8"},
{"FLUXER_PUSH_APNS_DEFAULT_ENVIRONMENT", "development"},
{"FLUXER_PUSH_APNS_APPS", ApnsAppsJson},
{"FLUXER_PUSH_FCM_ENABLED", "true"},
{"FLUXER_PUSH_FCM_PROJECT_ID", "fluxer-fcm"},
{"FLUXER_PUSH_FCM_SERVICE_ACCOUNT_JSON_PATH", "/etc/fluxer/fcm.json"},
{"FLUXER_PUSH_FCM_TOKEN_URI", "https://oauth2.example/token"},
{"FLUXER_PUSH_FCM_APPS", FcmAppsJson}
{"FLUXER_GATEWAY_HTTP_RECOVERY_TIMEOUT_MS", "6000"}
],
fun() ->
Config = fluxer_gateway_config:load(),
?assertEqual(1234, maps:get(shutdown_drain_wait_ms, Config)),
?assertEqual(42, maps:get(gateway_http_rpc_max_concurrency, Config)),
?assertEqual(9, maps:get(gateway_http_failure_threshold, Config)),
?assertEqual(6000, maps:get(gateway_http_recovery_timeout_ms, Config)),
?assertEqual(true, maps:get(apns_enabled, Config)),
?assertEqual(<<"TEAMID">>, maps:get(apns_team_id, Config)),
?assertEqual(<<"KEYID">>, maps:get(apns_key_id, Config)),
?assertEqual(<<"development">>, maps:get(apns_default_environment, Config)),
?assertEqual(
[
#{
<<"app_id">> => <<"ios-stable">>,
<<"topic">> => <<"app.fluxer">>,
<<"environment">> => <<"production">>
}
],
maps:get(apns_apps, Config)
),
?assertEqual(true, maps:get(fcm_enabled, Config)),
?assertEqual(<<"fluxer-fcm">>, maps:get(fcm_project_id, Config)),
?assertEqual(<<"https://oauth2.example/token">>, maps:get(fcm_token_uri, Config)),
?assertEqual(
[#{<<"app_id">> => <<"android-stable">>, <<"project_id">> => <<"fluxer-fcm">>}],
maps:get(fcm_apps, Config)
)
?assertEqual(6000, maps:get(gateway_http_recovery_timeout_ms, Config))
end
).
@@ -176,7 +176,7 @@ init_push_role_includes_cluster_handoff_test() ->
),
{ok, {_SupFlags, Children}} = fluxer_gateway_sup:init([]),
Ids = child_ids(Children),
?assert(lists:member(push_dispatcher, Ids)),
?assert(lists:member(push_outbox, Ids)),
?assert(lists:member(push, Ids)),
?assert(lists:member(gateway_cluster_handoff, Ids)),
?assertNot(lists:member(session_manager, Ids)),
@@ -25,7 +25,7 @@ allow_circuit_request_uses_recovery_timeout_test() ->
update_circuit_state_uses_failure_threshold_test() ->
cleanup_circuit_tables(),
ensure_circuit_tables(),
Key = {push, <<"push.example.test">>},
Key = {rpc, <<"threshold.example.test">>},
record(Key, failure(), 2),
?assertEqual(closed, circuit_state(Key)),
record(Key, failure(), 1),
@@ -62,7 +62,7 @@ open_circuit_half_opens_then_recloses_on_success_test() ->
circuit_window_is_kept_out_of_the_state_record_test() ->
cleanup_circuit_tables(),
ensure_circuit_tables(),
Key = {push, <<"window.example.test">>},
Key = {rpc, <<"window.example.test">>},
record(Key, failure(), 1),
?assertMatch([{Key, closed, undefined, _}], ets:lookup(?CIRCUIT_TABLE, Key)),
?assertMatch([{Key, [{true, _}]}], ets:lookup(?CIRCUIT_WINDOW_TABLE, Key)),
@@ -1,178 +0,0 @@
%% SPDX-License-Identifier: AGPL-3.0-or-later
-module(push_dispatcher_tests).
-typing([eqwalizer]).
-include_lib("eunit/include/eunit.hrl").
enqueue_saturation_bounds_inflight_and_queue_test_() ->
{timeout, 30, fun enqueue_saturation_bounds_inflight_and_queue/0}.
worker_down_drains_queue_without_exceeding_inflight_limit_test_() ->
{timeout, 30, fun worker_down_drains_queue_without_exceeding_inflight_limit/0}.
invalid_jobs_do_not_change_dispatcher_state_test() ->
State0 = dispatcher_state(2, 4),
{noreply, State1} = push_dispatcher:handle_cast({enqueue, #{type => invalid}}, State0),
?assertEqual(State0, State1).
enqueue_call_reports_drop_when_queue_full_test_() ->
{timeout, 30, fun enqueue_call_reports_drop_when_queue_full/0}.
enqueue_saturation_bounds_inflight_and_queue() ->
with_push_sender_blocked(fun() ->
State0 = dispatcher_state(2, 3),
State1 = enqueue_jobs(lists:seq(1, 10), State0),
?assertEqual(2, maps:get(inflight, State1)),
?assertEqual(3, maps:get(queued, State1)),
?assertEqual(3, queue:len(maps:get(queue, State1))),
Started = collect_started(2),
assert_no_push_started(),
release_started(Started),
drain_worker_downs(length(Started))
end).
worker_down_drains_queue_without_exceeding_inflight_limit() ->
with_push_sender_blocked(fun() ->
State0 = dispatcher_state(2, 10),
State1 = enqueue_jobs(lists:seq(1, 5), State0),
Started0 = collect_started(2),
?assertEqual(2, maps:get(inflight, State1)),
?assertEqual(3, maps:get(queued, State1)),
[First | Rest] = Started0,
release_started([First]),
Down = wait_worker_down(First),
{noreply, State2} = push_dispatcher:handle_info(Down, State1),
Started1 = collect_started(1),
?assertEqual(2, maps:get(inflight, State2)),
?assertEqual(2, maps:get(queued, State2)),
assert_no_push_started(),
release_started(Rest ++ Started1),
drain_worker_downs(length(Rest ++ Started1))
end).
enqueue_call_reports_drop_when_queue_full() ->
with_push_sender_blocked(fun() ->
State0 = dispatcher_state(1, 1),
{reply, ok, State1} =
push_dispatcher:handle_call({enqueue, send_job(1)}, from(), State0),
Started = collect_started(1),
{reply, ok, State2} =
push_dispatcher:handle_call({enqueue, send_job(2)}, from(), State1),
{reply, dropped, State3} =
push_dispatcher:handle_call({enqueue, send_job(3)}, from(), State2),
?assertEqual(1, maps:get(inflight, State3)),
?assertEqual(1, maps:get(queued, State3)),
?assertEqual(1, queue:len(maps:get(queue, State3))),
release_started(Started),
drain_worker_downs(length(Started))
end).
dispatcher_state(MaxInflight, MaxQueue) ->
#{
queue => queue:new(),
queued => 0,
inflight => 0,
workers => #{},
max_inflight => MaxInflight,
max_queue => MaxQueue
}.
from() ->
{self(), make_ref()}.
enqueue_jobs(MessageIds, State0) ->
lists:foldl(fun enqueue_job/2, State0, MessageIds).
enqueue_job(MessageId, State) ->
{noreply, NewState} = push_dispatcher:handle_cast(
{enqueue, send_job(MessageId)}, State
),
NewState.
send_job(MessageId) ->
#{
type => message_create,
user_ids => [1000 + MessageId],
message_data => #{<<"id">> => integer_to_binary(MessageId)},
guild_id => 10,
channel_id => 20,
message_id => MessageId,
guild_name => <<"guild">>,
channel_name => <<"channel">>,
badge_counts_ttl_seconds => 60
}.
with_push_sender_blocked(Fun) ->
meck:new(push_sender, [passthrough, no_link]),
Parent = self(),
persistent_term:put({?MODULE, push_parent}, Parent),
meck:expect(
push_sender,
send_push_notifications,
fun push_send_notifications_mock/1
),
try
Fun()
after
_ = persistent_term:erase({?MODULE, push_parent}),
meck:unload(push_sender)
end.
push_send_notifications_mock(#{message_id := MessageId}) ->
Parent = persistent_term:get({?MODULE, push_parent}),
Parent ! {push_started, self(), MessageId},
receive
{release_push_worker, Parent} -> ok
after 30000 ->
ok
end.
collect_started(Count) ->
collect_started(Count, []).
collect_started(0, Acc) ->
lists:reverse(Acc);
collect_started(Count, Acc) ->
receive
{push_started, Pid, MessageId} ->
collect_started(Count - 1, [{Pid, MessageId} | Acc])
after 5000 ->
?assert(false, {push_started_timeout, Count})
end.
assert_no_push_started() ->
receive
{push_started, Pid, MessageId} ->
?assert(false, {unexpected_push_started, Pid, MessageId})
after 100 ->
ok
end.
release_started(Started) ->
lists:foreach(
fun({Pid, _MessageId}) ->
Pid ! {release_push_worker, self()}
end,
Started
).
wait_worker_down({Pid, _MessageId}) ->
receive
{'DOWN', Ref, process, Pid, Reason} ->
{'DOWN', Ref, process, Pid, Reason}
after 5000 ->
?assert(false, {worker_down_timeout, Pid})
end.
drain_worker_downs(Count) ->
lists:foreach(
fun(_) ->
receive
{'DOWN', _Ref, process, _Pid, _Reason} -> ok
after 5000 ->
?assert(false, worker_down_drain_timeout)
end
end,
lists:seq(1, Count)
).
@@ -1,390 +0,0 @@
%% 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">>).
-define(VERDICT_TABLE, push_endpoint_verdicts).
-define(MAX_VERDICTS, 2048).
resolves_to(Addresses) ->
fun(_Host) -> {ok, Addresses} end.
fails_with(Reason) ->
fun(_Host) -> {error, Reason} end.
never_resolves() ->
fun(_Host) -> erlang:error(resolver_called) 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://[email protected]/sub">>, resolves_to([{1, 1, 1, 1}])
)
),
?assertEqual(
{error, endpoint_rejected},
push_endpoint_guard:check(
<<"https://user:[email protected]/sub">>, resolves_to([{1, 1, 1, 1}])
)
).
a_repeat_lookup_for_the_same_host_does_not_resolve_again_test() ->
with_verdict_cache(fun() ->
Counter = counters:new(1, []),
Endpoint = endpoint("cache-repeat.example.com"),
Resolver = counting_resolver(Counter, [{93, 184, 216, 34}]),
?assertEqual(ok, push_endpoint_guard:check(Endpoint, Resolver, cached)),
?assertEqual(ok, push_endpoint_guard:check(Endpoint, Resolver, cached)),
?assertEqual(ok, push_endpoint_guard:check(Endpoint, Resolver, cached)),
?assertEqual(1, counters:get(Counter, 1))
end).
a_cache_hit_returns_the_verdict_the_uncached_path_returns_test() ->
with_verdict_cache(fun() ->
lists:foreach(
fun assert_cached_matches_uncached/1,
[
{"cache-allow.example.com", resolves_to([{93, 184, 216, 34}]), ok},
{"cache-block.example.com", resolves_to([{10, 0, 0, 1}]),
{error, endpoint_blocked}},
{"cache-empty.example.com", resolves_to([]), {error, nxdomain}},
{"cache-timeout.example.com", fails_with(timeout), {error, timeout}}
]
)
end).
assert_cached_matches_uncached({Host, Resolver, Expected}) ->
Endpoint = endpoint(Host),
?assertEqual(Expected, push_endpoint_guard:check(Endpoint, Resolver)),
?assertEqual(Expected, push_endpoint_guard:check(Endpoint, Resolver, cached)),
?assertEqual(Expected, push_endpoint_guard:check(Endpoint, never_resolves(), cached)).
an_uncached_check_never_writes_the_cache_test() ->
with_verdict_cache(fun() ->
Endpoint = endpoint("cache-bypass.example.com"),
?assertEqual(ok, push_endpoint_guard:check(Endpoint, resolves_to([{1, 1, 1, 1}]))),
?assertEqual(0, push_ets_cache:table_size(?VERDICT_TABLE))
end).
an_ip_literal_is_never_cached_test() ->
with_verdict_cache(fun() ->
?assertEqual(
ok, push_endpoint_guard:check(<<"https://93.184.216.34/sub">>, never_resolves())
),
?assertEqual(0, push_ets_cache:table_size(?VERDICT_TABLE))
end).
a_cached_verdict_expires_test() ->
with_verdict_cache(fun() ->
Counter = counters:new(1, []),
Endpoint = endpoint("cache-expiry.example.com"),
Resolver = counting_resolver(Counter, [{93, 184, 216, 34}]),
?assertEqual(ok, push_endpoint_guard:check(Endpoint, Resolver, cached)),
?assertEqual(ok, push_endpoint_guard:check(Endpoint, Resolver, cached)),
?assertEqual(1, counters:get(Counter, 1)),
expire_verdict(<<"cache-expiry.example.com">>),
?assertEqual(ok, push_endpoint_guard:check(Endpoint, Resolver, cached)),
?assertEqual(2, counters:get(Counter, 1))
end).
a_refused_verdict_expires_sooner_than_an_allowed_one_test() ->
with_verdict_cache(fun() ->
Allowed = endpoint("cache-ttl-allowed.example.com"),
Refused = endpoint("cache-ttl-refused.example.com"),
?assertEqual(
ok, push_endpoint_guard:check(Allowed, resolves_to([{93, 184, 216, 34}]), cached)
),
?assertEqual(
{error, timeout}, push_endpoint_guard:check(Refused, fails_with(timeout), cached)
),
AllowedExpiry = expires_at(<<"cache-ttl-allowed.example.com">>),
RefusedExpiry = expires_at(<<"cache-ttl-refused.example.com">>),
?assert(RefusedExpiry < AllowedExpiry)
end).
many_distinct_hosts_cannot_grow_the_cache_without_bound_test() ->
with_verdict_cache(fun() ->
Resolver = resolves_to([{93, 184, 216, 34}]),
lists:foreach(
fun(N) -> flood_one_host(N, Resolver) end,
lists:seq(1, 20000)
),
Size = push_ets_cache:table_size(?VERDICT_TABLE),
?assert(Size >= 1500),
?assert(Size =< ?MAX_VERDICTS),
?assert(verdict_table_bytes() =< 4 * 1024 * 1024)
end).
flood_one_host(N, Resolver) ->
Endpoint = endpoint("flood-" ++ integer_to_list(N) ++ ".example.com"),
?assertEqual(ok, push_endpoint_guard:check(Endpoint, Resolver, cached)),
?assert(push_ets_cache:table_size(?VERDICT_TABLE) =< ?MAX_VERDICTS).
the_guard_is_enabled_by_default_test() ->
?assertEqual(
{error, endpoint_blocked},
push_endpoint_guard:check(?ENDPOINT, resolves_to([{10, 0, 0, 1}]))
).
a_disabled_guard_passes_everything_through_without_resolving_test() ->
with_verdict_cache(fun() ->
with_guard_disabled(fun() ->
Never = never_resolves(),
?assertEqual(ok, push_endpoint_guard:check(?ENDPOINT, Never)),
?assertEqual(ok, push_endpoint_guard:check(?ENDPOINT, Never, cached)),
?assertEqual(
ok, push_endpoint_guard:check(<<"https://169.254.169.254/latest">>, Never)
),
?assertEqual(
ok, push_endpoint_guard:check(<<"http://push.example.com/sub">>, Never)
),
?assertEqual(ok, push_endpoint_guard:check(<<"not-a-url">>, Never)),
?assertEqual(0, push_ets_cache:table_size(?VERDICT_TABLE))
end)
end).
counting_resolver(Counter, Addresses) ->
fun(_Host) ->
counters:add(Counter, 1, 1),
{ok, Addresses}
end.
endpoint(Host) ->
list_to_binary("https://" ++ Host ++ "/sub").
expires_at(Host) ->
[{Host, _Verdict, ExpiresAt}] = ets:lookup(?VERDICT_TABLE, Host),
ExpiresAt.
expire_verdict(Host) ->
[{Host, Verdict, _ExpiresAt}] = ets:lookup(?VERDICT_TABLE, Host),
true = ets:insert(?VERDICT_TABLE, {Host, Verdict, erlang:system_time(second) - 1}),
ok.
verdict_table_bytes() ->
ets:info(?VERDICT_TABLE, memory) * erlang:system_info(wordsize).
with_verdict_cache(Fun) ->
ok = push_ets_cache:init(),
true = ets:delete_all_objects(?VERDICT_TABLE),
try
Fun()
after
ets:delete_all_objects(?VERDICT_TABLE)
end.
with_guard_disabled(Fun) ->
Original = fluxer_gateway_env:get(push_endpoint_guard_enabled),
_ = fluxer_gateway_env:patch(#{push_endpoint_guard_enabled => false}),
try
Fun()
after
_ = fluxer_gateway_env:patch(#{push_endpoint_guard_enabled => Original})
end.
ip_literals_skip_dns_and_are_screened_directly_test() ->
Never = never_resolves(),
?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 = never_resolves(),
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}])
)
).
@@ -1,129 +0,0 @@
%% SPDX-License-Identifier: AGPL-3.0-or-later
-module(push_ets_cache_stress_tests).
-typing([eqwalizer]).
-include_lib("eunit/include/eunit.hrl").
-define(USER_COUNT, 10000).
-define(CONCURRENT_WRITERS, 8).
-define(WRITES_PER_WORKER, 2000).
get_subscriptions_many_large_mixed_set_test_() ->
{timeout, 15, fun get_subscriptions_many_large_mixed_set/0}.
concurrent_subscription_reads_and_writes_keep_cache_consistent_test_() ->
{timeout, 15, fun concurrent_subscription_reads_and_writes_keep_cache_consistent/0}.
get_subscriptions_many_large_mixed_set() ->
cleanup_tables(),
ok = push_ets_cache:init(),
try
lists:foreach(fun put_subscription_for_even_user/1, lists:seq(1, ?USER_COUNT)),
{Cached, Missing} = push_ets_cache:get_subscriptions_many(lists:seq(1, ?USER_COUNT)),
?assertEqual(?USER_COUNT div 2, map_size(Cached)),
?assertEqual(?USER_COUNT div 2, length(Missing)),
?assertEqual([#{endpoint => <<"endpoint-2">>}], maps:get(2, Cached)),
?assertEqual(true, lists:member(1, Missing)),
?assertEqual(false, lists:member(2, Missing))
after
cleanup_tables()
end.
concurrent_subscription_reads_and_writes_keep_cache_consistent() ->
cleanup_tables(),
ok = push_ets_cache:init(),
Ref = make_ref(),
Parent = self(),
try
_Writers = spawn_writers(Ref, Parent),
_Readers = spawn_readers(Ref, Parent),
collect_done(Ref, writer_done, ?CONCURRENT_WRITERS),
collect_done(Ref, reader_done, ?CONCURRENT_WRITERS),
ExpectedSize = ?CONCURRENT_WRITERS * ?WRITES_PER_WORKER,
?assertEqual(ExpectedSize, push_ets_cache:table_size(push_subscriptions)),
{Cached, Missing} = push_ets_cache:get_subscriptions_many(lists:seq(1, ExpectedSize)),
?assertEqual(ExpectedSize, map_size(Cached)),
?assertEqual([], Missing)
after
cleanup_tables()
end.
put_subscription_for_even_user(UserId) when UserId rem 2 =:= 0 ->
ok = put_subscriptions(UserId, [#{endpoint => endpoint(UserId)}]);
put_subscription_for_even_user(_UserId) ->
ok.
spawn_writers(Ref, Parent) ->
[
spawn(fun() ->
writer_loop(WriterIndex, ?WRITES_PER_WORKER),
Parent ! {Ref, writer_done}
end)
|| WriterIndex <- lists:seq(0, ?CONCURRENT_WRITERS - 1)
].
writer_loop(WriterIndex, Count) ->
Start = WriterIndex * Count + 1,
End = Start + Count - 1,
lists:foreach(
fun(UserId) ->
ok = put_subscriptions(UserId, [#{endpoint => endpoint(UserId)}])
end,
lists:seq(Start, End)
).
spawn_readers(Ref, Parent) ->
[
spawn(fun() ->
reader_loop(ReaderIndex),
Parent ! {Ref, reader_done}
end)
|| ReaderIndex <- lists:seq(0, ?CONCURRENT_WRITERS - 1)
].
reader_loop(ReaderIndex) ->
Offset = ReaderIndex * 100,
lists:foreach(
fun(Iteration) ->
Start = 1 + ((Offset + Iteration * 97) rem ?USER_COUNT),
_ = push_ets_cache:get_subscriptions_many(lists:seq(Start, Start + 99)),
ok
end,
lists:seq(1, 100)
).
collect_done(_Ref, _Message, 0) ->
ok;
collect_done(Ref, Message, Remaining) ->
receive
{Ref, Message} ->
collect_done(Ref, Message, Remaining - 1)
after 5000 ->
?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>>.
cleanup_tables() ->
delete_table(push_user_guild_settings),
delete_table(push_subscriptions),
delete_table(push_blocked_ids),
delete_table(push_badge_counts),
ok.
delete_table(Table) ->
try ets:delete(Table) of
_ -> ok
catch
throw:_ -> ok;
error:_ -> ok;
exit:_ -> ok
end.
+17 -125
View File
@@ -8,10 +8,7 @@ init_creates_tables_test() ->
cleanup_tables(),
ok = push_ets_cache:init(),
?assertNotEqual(undefined, ets:whereis(push_user_guild_settings)),
?assertNotEqual(undefined, ets:whereis(push_subscriptions)),
?assertNotEqual(undefined, ets:whereis(push_blocked_ids)),
?assertNotEqual(undefined, ets:whereis(push_badge_counts)),
?assertNotEqual(undefined, ets:whereis(push_endpoint_verdicts)),
cleanup_tables().
init_idempotent_test() ->
@@ -31,16 +28,6 @@ user_guild_settings_test() ->
?assertEqual(undefined, push_ets_cache:get_user_guild_settings(1, 2)),
cleanup_tables().
subscriptions_test() ->
cleanup_tables(),
ok = push_ets_cache:init(),
?assertEqual(undefined, push_ets_cache:get_subscriptions(1)),
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)),
cleanup_tables().
blocked_ids_test() ->
cleanup_tables(),
ok = push_ets_cache:init(),
@@ -49,46 +36,28 @@ blocked_ids_test() ->
?assertEqual([2, 3], push_ets_cache:get_blocked_ids(1)),
cleanup_tables().
badge_count_test() ->
cleanup_tables(),
ok = push_ets_cache:init(),
?assertEqual(undefined, push_ets_cache:get_badge_count(1)),
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)),
cleanup_tables().
badge_count_keeps_fresher_timestamp_test() ->
cleanup_tables(),
ok = push_ets_cache:init(),
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 = seed_badge_count(1, 7, 3000),
?assertEqual({7, 3000}, push_ets_cache:get_badge_count(1)),
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 = seed_subscriptions(1, []),
ok = seed_subscriptions(2, []),
ok = push_ets_cache:put_blocked_ids(1, []),
ok = push_ets_cache:put_blocked_ids(2, []),
Stats = push_ets_cache:cache_stats(),
?assertEqual(2, maps:get(push_subscriptions_size, Stats)),
?assertEqual(0, maps:get(user_guild_settings_size, Stats)),
?assertEqual(
#{blocked_ids_size => 2, user_guild_settings_size => 0},
Stats
),
cleanup_tables().
evict_tables_test() ->
cleanup_tables(),
ok = push_ets_cache:init(),
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)),
lists:foreach(
fun(I) -> ok = push_ets_cache:put_user_guild_settings(I, 1, #{}) end,
lists:seq(1, 10)
),
?assertEqual(10, push_ets_cache:table_size(push_user_guild_settings)),
ok = push_ets_cache:evict_tables(#{user_guild_settings => 5}),
?assertEqual(5, push_ets_cache:table_size(push_user_guild_settings)),
cleanup_tables().
rebalance_evicts_remote_owned_entries_test() ->
@@ -100,92 +69,19 @@ 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 = seed_subscriptions(LocalUserId, [local]),
ok = seed_subscriptions(RemoteUserId, [remote]),
ok = push_ets_cache:put_blocked_ids(LocalUserId, [1]),
ok = push_ets_cache:put_blocked_ids(RemoteUserId, [2]),
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(),
?assertEqual([local], push_ets_cache:get_subscriptions(LocalUserId)),
?assertEqual(undefined, push_ets_cache:get_subscriptions(RemoteUserId)),
?assertEqual([1], push_ets_cache:get_blocked_ids(LocalUserId)),
?assertEqual(undefined, push_ets_cache:get_blocked_ids(RemoteUserId)),
?assertEqual(#{local => true}, push_ets_cache:get_user_guild_settings(LocalUserId, 10)),
?assertEqual(undefined, push_ets_cache:get_user_guild_settings(RemoteUserId, 10)),
persistent_term:erase({gateway_cluster_membership, members}),
persistent_term:erase({gateway_cluster_membership, members_by_role}),
cleanup_tables().
endpoint_verdict_round_trip_test() ->
cleanup_tables(),
ok = push_ets_cache:init(),
?assertEqual(undefined, push_ets_cache:get_endpoint_verdict(<<"push.example.com">>)),
ok = push_ets_cache:put_endpoint_verdict(<<"push.example.com">>, ok, 300),
?assertEqual({ok, ok}, push_ets_cache:get_endpoint_verdict(<<"push.example.com">>)),
ok = push_ets_cache:put_endpoint_verdict(
<<"bad.example.com">>, {error, endpoint_blocked}, 30
),
?assertEqual(
{ok, {error, endpoint_blocked}},
push_ets_cache:get_endpoint_verdict(<<"bad.example.com">>)
),
cleanup_tables().
endpoint_verdicts_expire_and_are_reclaimed_test() ->
cleanup_tables(),
ok = push_ets_cache:init(),
ok = push_ets_cache:put_endpoint_verdict(<<"stale.example.com">>, ok, 300),
Stale = erlang:system_time(second) - 1,
true = ets:insert(push_endpoint_verdicts, {<<"stale.example.com">>, ok, Stale}),
?assertEqual(undefined, push_ets_cache:get_endpoint_verdict(<<"stale.example.com">>)),
ok = push_ets_cache:evict_tables(#{}),
?assertEqual([], ets:lookup(push_endpoint_verdicts, <<"stale.example.com">>)),
cleanup_tables().
an_oversized_host_is_never_cached_test() ->
cleanup_tables(),
ok = push_ets_cache:init(),
Oversized = binary:copy(<<"a">>, 254),
ok = push_ets_cache:put_endpoint_verdict(Oversized, ok, 300),
?assertEqual(undefined, push_ets_cache:get_endpoint_verdict(Oversized)),
?assertEqual(0, push_ets_cache:table_size(push_endpoint_verdicts)),
AtLimit = binary:copy(<<"a">>, 253),
ok = push_ets_cache:put_endpoint_verdict(AtLimit, ok, 300),
?assertEqual({ok, ok}, push_ets_cache:get_endpoint_verdict(AtLimit)),
cleanup_tables().
endpoint_verdicts_stay_bounded_under_max_length_hosts_test() ->
cleanup_tables(),
ok = push_ets_cache:init(),
lists:foreach(fun seed_max_length_verdict/1, lists:seq(1, 20000)),
Size = push_ets_cache:table_size(push_endpoint_verdicts),
?assert(Size >= 1500),
?assert(Size =< 2048),
Bytes = ets:info(push_endpoint_verdicts, memory) * erlang:system_info(wordsize),
?assert(Bytes =< 4 * 1024 * 1024),
cleanup_tables().
seed_max_length_verdict(N) ->
Suffix = integer_to_binary(N),
Host = <<(binary:copy(<<"a">>, 253 - byte_size(Suffix)))/binary, Suffix/binary>>,
ok = push_ets_cache:put_endpoint_verdict(Host, ok, 300),
?assert(push_ets_cache:table_size(push_endpoint_verdicts) =< 2048).
endpoint_verdicts_are_reported_in_cache_stats_test() ->
cleanup_tables(),
ok = push_ets_cache:init(),
ok = push_ets_cache:put_endpoint_verdict(<<"push.example.com">>, ok, 300),
Stats = push_ets_cache:cache_stats(),
?assertEqual(1, maps:get(endpoint_verdicts_size, Stats)),
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([
@@ -203,11 +99,7 @@ find_split_user_ids(Members, RemoteNode) ->
cleanup_tables() ->
delete_table(push_user_guild_settings),
delete_table(push_subscriptions),
delete_table(push_blocked_ids),
delete_table(push_badge_counts),
delete_table(push_bearer_tokens),
delete_table(push_endpoint_verdicts),
ok.
delete_table(Table) ->
+127 -138
View File
@@ -4,68 +4,124 @@
-typing([eqwalizer]).
-include_lib("eunit/include/eunit.hrl").
clear_channel_notifications_disabled_by_default_test() ->
erase_persistent_term(push_noop),
erase_persistent_term(push_clear_notifications_enabled),
?assertEqual(ok, push:clear_channel_notifications(1, 2, 3)).
clear_channel_notifications_are_enabled_by_default_test() ->
{Truncates, Casts} = capture_clear(undefined, #{}),
?assertEqual([{1, 2, 3}], Truncates),
?assertEqual([{clear_channel_notifications, 1, 2, 3}], Casts).
a_saturated_dispatcher_retries_the_clear_before_dropping_it_test() ->
ok = meck:new(push_dispatcher, [passthrough, no_link]),
try
ok = meck:expect(
push_dispatcher, enqueue_clear_notifications, fun(_U, _C, _M, _T) -> dropped end
),
State = #{badge_counts_ttl_seconds => 0},
?assertEqual(
{noreply, State},
push:handle_info({retry_clear_notifications, 1, 2, 3, 0}, State)
),
receive
{retry_clear_notifications, 1, 2, 3, 1} -> ok
after 2000 -> erlang:error(no_retry_scheduled)
end
after
meck:unload(push_dispatcher)
end.
the_clear_env_switch_turns_off_the_clear_cast_test() ->
{Truncates, Casts} = capture_clear(undefined, #{
push_enrolled_clear_notifications_enabled => false
}),
?assertEqual([{1, 2, 3}], Truncates),
?assertEqual([], Casts).
a_clear_is_dropped_only_after_the_retry_budget_is_spent_test() ->
ok = meck:new(push_dispatcher, [passthrough, no_link]),
try
ok = meck:expect(
push_dispatcher, enqueue_clear_notifications, fun(_U, _C, _M, _T) -> dropped end
),
State = #{badge_counts_ttl_seconds => 0},
?assertEqual(
{noreply, State},
push:handle_info({retry_clear_notifications, 1, 2, 3, 3}, State)
),
receive
{retry_clear_notifications, _, _, _, _} -> erlang:error(retried_past_budget)
after 700 -> ok
end
after
meck:unload(push_dispatcher)
end.
the_clear_operator_switch_turns_off_the_clear_cast_test() ->
{Truncates, Casts} = capture_clear(false, #{
push_enrolled_clear_notifications_enabled => true
}),
?assertEqual([{1, 2, 3}], Truncates),
?assertEqual([], Casts).
an_enrolled_user_gets_clears_while_the_fleet_switch_is_off_test() ->
the_clear_operator_switch_turns_on_the_clear_cast_test() ->
{Truncates, Casts} = capture_clear(true, #{
push_enrolled_clear_notifications_enabled => false
}),
?assertEqual([{1, 2, 3}], Truncates),
?assertEqual([{clear_channel_notifications, 1, 2, 3}], Casts).
clears_do_nothing_while_push_is_disabled_test() ->
{Truncates, Casts} = capture_clear(undefined, #{push_enabled => false}),
?assertEqual([], Truncates),
?assertEqual([], Casts).
capture_clear(OperatorChoice, Env) ->
Self = self(),
erase_persistent_term(push_noop),
erase_persistent_term(push_enrolled_clear_notifications_enabled),
persistent_term:put(push_clear_notifications_enabled, false),
put_operator_choice(OperatorChoice),
Stub = spawn(fun() -> push_stub(Self) end),
true = register(push, Stub),
Modules = [fluxer_gateway_env, push_outbox, gateway_node_router],
lists:foreach(fun(Module) -> ok = meck:new(Module, [passthrough, no_link]) end, Modules),
try
?assertEqual(ok, push:clear_channel_notifications(1, 2, 3))
after
erase_persistent_term(push_clear_notifications_enabled)
end.
the_cohort_switch_can_be_turned_off_on_its_own_test() ->
erase_persistent_term(push_noop),
persistent_term:put(push_enrolled_clear_notifications_enabled, false),
try
?assertEqual(ok, push:clear_channel_notifications(1, 2, 3))
ok = meck:expect(fluxer_gateway_env, get, fun(Key) ->
maps:get(Key, maps:merge(#{push_enabled => true}, Env), undefined)
end),
ok = meck:expect(push_outbox, truncate_read, fun(UserId, ChannelId, MessageId) ->
Self ! {truncated, {UserId, ChannelId, MessageId}},
ok
end),
ok = meck:expect(gateway_node_router, owner_node_result, fun(_Key, push) ->
{ok, node()}
end),
?assertEqual(ok, push:clear_channel_notifications(1, 2, 3)),
Stub ! {flush, Self},
receive
flushed -> ok
after 5000 -> error(push_stub_timeout)
end,
{drain_tagged(truncated), drain_tagged(cast)}
after
lists:foreach(fun meck:unload/1, Modules),
unregister(push),
exit(Stub, kill),
erase_persistent_term(push_enrolled_clear_notifications_enabled)
end.
put_operator_choice(undefined) ->
ok;
put_operator_choice(Choice) ->
persistent_term:put(push_enrolled_clear_notifications_enabled, Choice).
push_stub(Parent) ->
receive
{'$gen_cast', Msg} ->
Parent ! {cast, Msg},
push_stub(Parent);
{flush, From} ->
From ! flushed,
push_stub(Parent)
end.
drain_tagged(Tag) ->
receive
{Tag, Value} -> [Value | drain_tagged(Tag)]
after 0 ->
[]
end.
a_clear_publishes_through_the_outbox_test() ->
State = #{max_entries => 10},
[Job] = with_captured_enqueues(fun() ->
?assertEqual(
{noreply, State},
push:handle_cast({clear_channel_notifications, 1, 2, 3}, State)
)
end),
?assertMatch(
#{
kind := clear,
subject := <<"push.job.clear">>,
user_ids := [1],
channel_id := 2,
message_id := 3,
job := #{
<<"v">> := 1,
<<"config_version">> := 0,
<<"user_id">> := <<"1">>,
<<"channel_id">> := <<"2">>,
<<"message_id">> := <<"3">>
}
},
Job
).
legacy_cache_invalidation_casts_are_ignored_test() ->
State = #{max_entries => 10},
?assertEqual({noreply, State}, push:handle_cast({invalidate_user_subscriptions, 1}, State)),
?assertEqual({noreply, State}, push:handle_cast({invalidate_user_badge_count, 1}, State)).
push_owner_key_prefers_first_recipient_test() ->
?assertEqual(
42,
@@ -217,29 +273,26 @@ prefetch_user_guild_settings_skips_direct_messages_test() ->
end),
?assertEqual([], settings_requests()).
init_logs_a_mismatched_vapid_pair_test() ->
with_push_env(
fun() ->
{Pub, _} = generate_vapid_pair(),
{_, OtherPriv} = generate_vapid_pair(),
patch_vapid(true, Pub, OtherPriv),
{ok, Pid} = with_captured_logs(fun() -> push:start_link() end),
?assert(is_process_alive(Pid)),
?assert(any_error_log_mentions("FLUXER_VAPID_PUBLIC_KEY")),
gen_server:stop(Pid)
end
).
with_captured_enqueues(Fun) ->
Self = self(),
ok = meck:new(push_outbox, [passthrough, no_link]),
try
ok = meck:expect(push_outbox, enqueue, fun(OutboxJob) ->
Self ! {enqueued, OutboxJob},
ok
end),
Fun(),
drain_enqueued([])
after
meck:unload(push_outbox)
end.
init_accepts_a_matching_vapid_pair_test() ->
with_push_env(
fun() ->
{Pub, Priv} = generate_vapid_pair(),
patch_vapid(true, Pub, Priv),
{ok, Pid} = push:start_link(),
?assert(is_process_alive(Pid)),
gen_server:stop(Pid)
end
).
drain_enqueued(Acc) ->
receive
{enqueued, OutboxJob} -> drain_enqueued([OutboxJob | Acc])
after 0 ->
lists:reverse(Acc)
end.
with_rpc_client_stub(Result, Fun) ->
Self = self(),
@@ -265,70 +318,6 @@ settings_requests() ->
[]
end.
with_push_env(Fun) ->
push_ets_cache:init(),
push_worker_pool:init_counter(),
OldConfig = fluxer_gateway_env:get_map(),
OldTrap = erlang:process_flag(trap_exit, true),
try
Fun()
after
_ = erlang:process_flag(trap_exit, OldTrap),
flush_exit_signals(),
_ = fluxer_gateway_env:update(fun(_) -> OldConfig end)
end.
with_captured_logs(Fun) ->
Self = self(),
ok = logger:add_primary_filter(
capture_logs, {
fun(Event, Pid) ->
Pid ! {captured_log, Event},
stop
end,
Self
}
),
try
Fun()
after
_ = logger:remove_primary_filter(capture_logs)
end.
any_error_log_mentions(Needle) ->
receive
{captured_log, #{level := error, msg := {string, Message}}} ->
case string:find(Message, Needle) of
nomatch -> any_error_log_mentions(Needle);
_ -> true
end;
{captured_log, _} ->
any_error_log_mentions(Needle)
after 0 ->
false
end.
patch_vapid(Enabled, Pub, Priv) ->
_ = fluxer_gateway_env:patch(#{
push_enabled => Enabled,
vapid_public_key => push_utils:base64url_encode(Pub),
vapid_private_key => push_utils:base64url_encode(Priv)
}),
ok.
generate_vapid_pair() ->
case crypto:generate_key(ecdh, prime256v1) of
{<<4, _:64/binary>> = Pub, <<_:32/binary>> = Priv} -> {Pub, Priv};
_ -> generate_vapid_pair()
end.
flush_exit_signals() ->
receive
{'EXIT', _, _} -> flush_exit_signals()
after 0 ->
ok
end.
erase_persistent_term(Key) ->
try persistent_term:erase(Key) of
_ -> ok
-94
View File
@@ -4,16 +4,6 @@
-typing([eqwalizer]).
-include_lib("eunit/include/eunit.hrl").
extract_origin_test() ->
?assertEqual(
<<"https://example.com">>,
push_utils:extract_origin(<<"https://example.com/path/to/resource">>)
),
?assertEqual(
<<"http://localhost:8080">>, push_utils:extract_origin(<<"http://localhost:8080/api">>)
),
?assertEqual(<<"invalid">>, push_utils:extract_origin(<<"invalid">>)).
get_default_avatar_url_test() ->
Url = push_utils:get_default_avatar_url(<<"123">>),
?assert(is_binary(Url)),
@@ -32,87 +22,3 @@ wrap_avatar_index_test() ->
?assertEqual(1, push_utils:wrap_avatar_index(1)),
?assertEqual(0, push_utils:wrap_avatar_index(6)),
?assertEqual(1, push_utils:wrap_avatar_index(7)).
base64url_encode_test() ->
Encoded = push_utils:base64url_encode(<<"test">>),
?assert(is_binary(Encoded)).
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),
Info = <<"test info">>,
Result = push_utils:hkdf_expand(IKM, Salt, Info, 32),
?assertEqual(32, byte_size(Result)).
assert_vapid_pair_accepts_generated_pair_test() ->
{Pub, Priv} = generate_vapid_pair(),
?assertEqual(
ok,
push_utils:assert_vapid_pair(
push_utils:base64url_encode(Pub),
push_utils:base64url_encode(Priv)
)
).
assert_vapid_pair_rejects_mismatched_scalar_test() ->
{Pub, _} = generate_vapid_pair(),
{_, OtherPriv} = generate_vapid_pair(),
?assertError(
{vapid_keys_mismatched, _},
push_utils:assert_vapid_pair(
push_utils:base64url_encode(Pub),
push_utils:base64url_encode(OtherPriv)
)
).
assert_vapid_pair_rejects_malformed_public_point_test() ->
{Pub, Priv} = generate_vapid_pair(),
?assertError(
{vapid_keys_malformed, 64, 32},
push_utils:assert_vapid_pair(
push_utils:base64url_encode(binary:part(Pub, 0, 64)),
push_utils:base64url_encode(Priv)
)
).
assert_vapid_pair_rejects_short_scalar_test() ->
{Pub, Priv} = generate_vapid_pair(),
?assertError(
{vapid_keys_malformed, 65, 31},
push_utils:assert_vapid_pair(
push_utils:base64url_encode(Pub),
push_utils:base64url_encode(binary:part(Priv, 0, 31))
)
).
generate_vapid_pair() ->
case crypto:generate_key(ecdh, prime256v1) of
{<<4, _:64/binary>> = Pub, <<_:32/binary>> = Priv} -> {Pub, Priv};
_ -> generate_vapid_pair()
end.
-4
View File
@@ -387,7 +387,6 @@ async fn run_message_job(
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,
@@ -432,7 +431,6 @@ async fn run_clear_job(state: &AppState, sends: &Semaphore, job: ClearJob) -> an
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,
@@ -455,7 +453,6 @@ async fn run_ring_job(state: &AppState, sends: &Semaphore, job: RingJob) -> anyh
user_id = %job.user_id,
channel_id = %job.channel_id,
message_id = %job.message_id,
config_version = job.config_version,
expires_at_ms = job.expires_at_ms,
"push ring dropped past its ring window"
);
@@ -487,7 +484,6 @@ async fn run_ring_job(state: &AppState, sends: &Semaphore, job: RingJob) -> anyh
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,
+35 -3
View File
@@ -14,7 +14,6 @@ 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,
@@ -38,7 +37,6 @@ pub struct NotificationFields {
#[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,
@@ -47,7 +45,6 @@ pub struct ClearJob {
#[derive(Clone, Debug, Deserialize, Eq, PartialEq)]
pub struct RingJob {
pub v: u8,
pub config_version: u64,
pub user_id: String,
pub channel_id: String,
pub message_id: String,
@@ -99,3 +96,38 @@ fn supported(version: u8) -> Result<(), JobError> {
}
Err(JobError::UnsupportedVersion(version))
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::{Value, json};
fn clear_job(config_version: Option<u64>) -> Vec<u8> {
let mut job = json!({
"v": 1,
"user_id": "1",
"channel_id": "2",
"message_id": "3",
});
if let Some(version) = config_version {
job["config_version"] = Value::from(version);
}
job.to_string().into_bytes()
}
#[test]
fn a_job_from_a_gateway_that_sends_a_config_version_decodes() {
assert_eq!(
decode_clear(&clear_job(Some(7))).expect("decodes").user_id,
"1"
);
}
#[test]
fn a_job_without_a_config_version_decodes() {
assert_eq!(
decode_clear(&clear_job(None)).expect("decodes").user_id,
"1"
);
}
}
+1 -1
View File
@@ -11,9 +11,9 @@ mod metrics;
mod payload;
mod providers;
mod relay;
mod relay_consent;
mod resolver;
mod retry;
mod rollout;
mod rpc;
mod secret;
pub mod server;
+16 -30
View File
@@ -1,7 +1,7 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
use crate::relay::reject::{REASON_COUNT, Reason};
use crate::rollout::{RolloutOutcome, RolloutSnapshot};
use crate::relay_consent::ConsentUpdate;
use fluxer_svc::metrics::now_ms;
use std::fmt::{self, Write as _};
use std::sync::atomic::{AtomicU64, Ordering};
@@ -343,7 +343,7 @@ 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();
const CONSENT_UPDATE_COUNT: usize = ConsentUpdate::ALL.len();
pub struct Metrics {
jobs_received: [AtomicU64; JOB_KIND_COUNT],
@@ -357,7 +357,7 @@ pub struct Metrics {
own_relay_shortcuts: AtomicU64,
auth_tokens_minted: [AtomicU64; AUTH_PROVIDER_COUNT],
rpc_requests: [[AtomicU64; RPC_OUTCOME_COUNT]; RPC_METHOD_COUNT],
rollout_updates: [AtomicU64; ROLLOUT_OUTCOME_COUNT],
relay_consent_updates: [AtomicU64; CONSENT_UPDATE_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],
@@ -368,9 +368,7 @@ pub struct Metrics {
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,
relay_consent_accepted: AtomicU64,
start_ms: i64,
}
@@ -389,7 +387,7 @@ impl Metrics {
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],
relay_consent_updates: [const { AtomicU64::new(0) }; CONSENT_UPDATE_COUNT],
delivery_routes: [const { [const { AtomicU64::new(0) }; SEND_RESULT_COUNT] };
DELIVERY_ROUTE_COUNT],
relay_served: [const { [const { AtomicU64::new(0) }; RELAY_RESULT_COUNT] };
@@ -403,9 +401,7 @@ impl Metrics {
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),
relay_consent_accepted: AtomicU64::new(0),
start_ms: now_ms(),
}
}
@@ -457,17 +453,13 @@ impl Metrics {
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_relay_consent_update(&self, outcome: ConsentUpdate) {
self.relay_consent_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_relay_consent_accepted(&self, accepted: bool) {
self.relay_consent_accepted
.store(u64::from(accepted), ORDERING);
}
pub fn record_delivery_route(&self, route: DeliveryRoute, result: SendResult) {
@@ -575,10 +567,10 @@ impl Metrics {
render_labelled_counter(
out,
"fluxer_push_rollout_updates_total",
"fluxer_push_relay_consent_updates_total",
"result",
RolloutOutcome::ALL.map(RolloutOutcome::label),
&self.rollout_updates,
ConsentUpdate::ALL.map(ConsentUpdate::label),
&self.relay_consent_updates,
)?;
render_labelled_histogram(
out,
@@ -642,16 +634,10 @@ impl Metrics {
&self.rings_suppressed,
)?;
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,
"fluxer_push_relay_consent_accepted",
&self.relay_consent_accepted,
)?;
writeln!(out, "# TYPE fluxer_push_uptime_seconds gauge")?;
-2
View File
@@ -605,7 +605,6 @@ mod tests {
fn message_job(image_url: Option<&str>) -> MessageJob {
MessageJob {
v: 1,
config_version: 7,
guild_id: "0".to_owned(),
channel_id: CHANNEL_ID.to_owned(),
message_id: MESSAGE_ID.to_owned(),
@@ -626,7 +625,6 @@ mod tests {
fn clear_job() -> ClearJob {
ClearJob {
v: 1,
config_version: 7,
user_id: USER_ID.to_owned(),
channel_id: CHANNEL_ID.to_owned(),
message_id: MESSAGE_ID.to_owned(),
+3 -13
View File
@@ -118,18 +118,10 @@ pub async fn send(state: &AppState, sub: &Subscription, envelope: &Value) -> Sen
}
fn relay_consent_missing(state: &AppState, endpoint: &str) -> bool {
!relay_consent_accepted(state)
!state.relay_consent.accepted()
&& own_relay::is_managed(endpoint, &state.cfg.managed_relay_hosts)
}
fn relay_consent_accepted(state: &AppState) -> bool {
state.cfg.relay_consent_accepted
|| state
.rollout
.snapshot()
.is_some_and(|held| held.relay_consent_accepted)
}
fn in_process_hop(endpoint: &str, hosts: &[String]) -> Option<own_relay::Hop> {
own_relay::parse(endpoint, hosts).filter(|hop| matches!(hop.leg, own_relay::Leg::Apns))
}
@@ -276,8 +268,7 @@ mod consent_tests {
#[tokio::test]
async fn a_notice_accepted_in_the_instance_config_lets_the_send_through() {
let state = state(false);
state.rollout.update(&serde_json::json!({
"enabled": true,
state.relay_consent.update(&serde_json::json!({
"config_version": 1,
"relay_consent_accepted": true,
}));
@@ -290,8 +281,7 @@ mod consent_tests {
#[tokio::test]
async fn an_instance_config_that_has_not_accepted_still_refuses_the_send() {
let state = state(false);
state.rollout.update(&serde_json::json!({
"enabled": true,
state.relay_consent.update(&serde_json::json!({
"config_version": 1,
"relay_consent_accepted": false,
}));
+314
View File
@@ -0,0 +1,314 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
use crate::metrics::Metrics;
use crate::rpc::RpcClient;
use fluxer_svc::transport::{Transport, TransportMessage, TransportSubscriber};
use serde_json::Value;
use std::sync::RwLock;
use std::time::Duration;
use tokio::time::{Instant, MissedTickBehavior};
use tracing::{info, warn};
pub const RECONCILE_INTERVAL: Duration = Duration::from_secs(30);
const SUBJECT: &str = "config.push.delivery";
const MESSAGE_TYPE: &str = "push_service_delivery_config";
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
#[repr(usize)]
pub enum ConsentUpdate {
Updated,
Unchanged,
Stale,
Rejected,
}
impl ConsentUpdate {
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",
}
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
struct HeldConsent {
accepted: bool,
config_version: u64,
}
impl HeldConsent {
fn parse(config: &Value) -> Option<Self> {
let accepted = match config.get("relay_consent_accepted") {
None | Some(Value::Null) => false,
Some(Value::Bool(flag)) => *flag,
Some(_) => return None,
};
let config_version = match config.get("config_version") {
None | Some(Value::Null) => 0,
Some(value) => value.as_u64()?,
};
Some(Self {
accepted,
config_version,
})
}
}
pub struct RelayConsentStore {
env_accepted: bool,
held: RwLock<Option<HeldConsent>>,
}
impl RelayConsentStore {
pub fn new(env_accepted: bool) -> Self {
Self {
env_accepted,
held: RwLock::new(None),
}
}
pub fn accepted(&self) -> bool {
self.env_accepted || self.held().is_some_and(|held| held.accepted)
}
fn held(&self) -> Option<HeldConsent> {
*self.held.read().expect("relay consent lock poisoned")
}
fn apply(&self, payload: &[u8]) -> ConsentUpdate {
let Ok(value) = serde_json::from_slice::<Value>(payload) else {
warn!("relay consent payload is not JSON");
return ConsentUpdate::Rejected;
};
let Some(config) = config_object(&value) else {
warn!("relay consent payload has no config object of its type");
return ConsentUpdate::Rejected;
};
self.update(config)
}
pub(crate) fn update(&self, config: &Value) -> ConsentUpdate {
let Some(offered) = HeldConsent::parse(config) else {
warn!("relay consent config rejected as invalid");
return ConsentUpdate::Rejected;
};
let mut held = self.held.write().expect("relay consent lock poisoned");
match *held {
Some(current) if offered.config_version < current.config_version => {
warn!(
highest = current.config_version,
offered = offered.config_version,
"relay consent ignored a lower config_version"
);
return ConsentUpdate::Stale;
}
Some(current) if current == offered => return ConsentUpdate::Unchanged,
_ => {}
}
info!(
accepted = offered.accepted,
config_version = offered.config_version,
"relay consent updated"
);
*held = Some(offered);
ConsentUpdate::Updated
}
}
pub async fn run_subscriber<T: Transport>(
transport: T,
rpc: &RpcClient,
store: &RelayConsentStore,
metrics: &Metrics,
reconcile_every: Duration,
) {
loop {
let mut subscriber = match transport.subscribe(SUBJECT).await {
Ok(subscriber) => subscriber,
Err(error) => {
warn!(error = %error, subject = SUBJECT, "relay consent subscribe failed");
transport.wait_for_reconnect().await;
continue;
}
};
info!(subject = SUBJECT, "listening for relay consent updates");
let outcome = fetch(rpc, store, metrics).await;
info!(
outcome = outcome.label(),
"relay consent 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(metrics, store, outcome);
}
_ = reconcile.tick() => {
fetch(rpc, store, metrics).await;
}
}
}
warn!(
subject = SUBJECT,
"relay consent subscription ended, will re-subscribe"
);
}
}
async fn fetch(rpc: &RpcClient, store: &RelayConsentStore, metrics: &Metrics) -> ConsentUpdate {
let outcome = match rpc.push_service_delivery_config().await {
Ok(config) => store.update(&config),
Err(error) => {
warn!(error = %error, "relay consent read failed");
ConsentUpdate::Rejected
}
};
record(metrics, store, outcome);
outcome
}
fn record(metrics: &Metrics, store: &RelayConsentStore, outcome: ConsentUpdate) {
metrics.record_relay_consent_update(outcome);
metrics.record_relay_consent_accepted(store.accepted());
}
fn config_object(value: &Value) -> Option<&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)
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn legacy_document(config_version: u64, accepted: bool) -> Value {
json!({
"enabled": true,
"rollout_basis_points": 10000,
"rollout_salt": "push-service-delivery-v1",
"included_user_ids": [],
"excluded_user_ids": [],
"config_version": config_version,
"relay_consent_accepted": accepted,
"relay_consent_accepted_at": null,
"relay_consent_accepted_by": null,
})
}
#[test]
fn the_operator_relay_consent_is_read_off_the_legacy_document() {
let store = RelayConsentStore::new(false);
assert_eq!(
store.update(&legacy_document(3, true)),
ConsentUpdate::Updated
);
assert!(store.accepted());
}
#[test]
fn a_document_without_the_consent_field_has_not_consented() {
let store = RelayConsentStore::new(false);
assert_eq!(
store.update(&json!({"config_version": 3})),
ConsentUpdate::Updated
);
assert!(!store.accepted());
}
#[test]
fn a_consent_field_that_is_not_a_boolean_is_refused() {
let store = RelayConsentStore::new(false);
assert_eq!(
store.update(&json!({"relay_consent_accepted": "yes"})),
ConsentUpdate::Rejected
);
assert!(store.held().is_none());
}
#[test]
fn a_config_version_that_is_not_an_integer_is_refused() {
let store = RelayConsentStore::new(false);
assert_eq!(
store.update(&json!({"config_version": -1, "relay_consent_accepted": true})),
ConsentUpdate::Rejected
);
assert!(!store.accepted());
}
#[test]
fn an_older_config_version_never_overwrites_a_newer_one() {
let store = RelayConsentStore::new(false);
store.update(&legacy_document(5, true));
assert_eq!(
store.update(&legacy_document(4, false)),
ConsentUpdate::Stale
);
assert!(store.accepted());
assert_eq!(
store.update(&legacy_document(6, false)),
ConsentUpdate::Updated
);
assert!(!store.accepted());
}
#[test]
fn the_same_document_again_is_unchanged() {
let store = RelayConsentStore::new(false);
store.update(&legacy_document(5, true));
assert_eq!(
store.update(&legacy_document(5, true)),
ConsentUpdate::Unchanged
);
}
#[test]
fn the_env_override_accepts_whatever_the_document_says() {
let store = RelayConsentStore::new(true);
assert!(store.accepted());
store.update(&legacy_document(1, false));
assert!(store.accepted());
}
#[test]
fn a_nats_message_is_unwrapped_from_its_envelope() {
let store = RelayConsentStore::new(false);
let message = json!({"type": MESSAGE_TYPE, "config": legacy_document(2, true)});
assert_eq!(
store.apply(message.to_string().as_bytes()),
ConsentUpdate::Updated
);
assert!(store.accepted());
}
#[test]
fn a_nats_message_of_another_type_is_refused() {
let store = RelayConsentStore::new(false);
let message = json!({"type": "something_else", "config": legacy_document(2, true)});
assert_eq!(
store.apply(message.to_string().as_bytes()),
ConsentUpdate::Rejected
);
assert!(!store.accepted());
}
}
-326
View File
@@ -1,326 +0,0 @@
// 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<Self>;
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,
pub relay_consent_accepted: bool,
}
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<Self> {
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,
relay_consent_accepted: parse_flag(config, "relay_consent_accepted")?,
})
}
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<bool> {
parse_flag(config, "enabled")
}
fn parse_flag(config: &Value, key: &str) -> Option<bool> {
match config.get(key) {
None | Some(Value::Null) => Some(false),
Some(Value::Bool(flag)) => Some(*flag),
Some(_) => None,
}
}
pub fn parse_config_version(config: &Value) -> Option<u64> {
match config.get("config_version") {
None | Some(Value::Null) => Some(0),
Some(value) => value.as_u64(),
}
}
struct Current<C> {
held: Option<Arc<C>>,
highest_version: u64,
}
pub struct RolloutStore<C> {
current: RwLock<Current<C>>,
}
impl<C> Default for RolloutStore<C> {
fn default() -> Self {
Self {
current: RwLock::new(Current {
held: None,
highest_version: 0,
}),
}
}
}
impl<C: RolloutConfig> RolloutStore<C> {
pub fn new() -> Self {
Self::default()
}
pub fn snapshot(&self) -> Option<Arc<C>> {
self.current
.read()
.expect("rollout config lock poisoned")
.held
.clone()
}
fn apply(&self, payload: &[u8]) -> RolloutOutcome {
let Ok(value) = serde_json::from_slice::<Value>(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)
}
pub(crate) 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<T: Transport, C: RolloutConfig>(
transport: T,
rpc: &RpcClient,
store: &RolloutStore<C>,
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<C: RolloutConfig>(
rpc: &RpcClient,
store: &RolloutStore<C>,
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<C: RolloutConfig>(
metrics: &Metrics,
store: &RolloutStore<C>,
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)
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn the_operator_relay_consent_is_read_off_the_instance_config() {
let config = json!({
"enabled": true,
"config_version": 3,
"relay_consent_accepted": true,
});
let snapshot = RolloutSnapshot::parse(&config).expect("the config parses");
assert!(snapshot.relay_consent_accepted);
}
#[test]
fn an_instance_config_without_the_consent_field_has_not_consented() {
let config = json!({"enabled": true, "config_version": 3});
let snapshot = RolloutSnapshot::parse(&config).expect("the config parses");
assert!(!snapshot.relay_consent_accepted);
}
#[test]
fn a_consent_field_that_is_not_a_boolean_is_refused() {
let config = json!({"enabled": true, "relay_consent_accepted": "yes"});
assert!(RolloutSnapshot::parse(&config).is_none());
}
}
+8 -3
View File
@@ -62,8 +62,13 @@ impl RpcClient {
}
}
pub async fn rollout_config(&self, method: RpcMethod) -> Result<Value, RpcError> {
let data: RolloutConfigData = self.call(method, &json!({"type": method.label()})).await?;
pub async fn push_service_delivery_config(&self) -> Result<Value, RpcError> {
let data: PushServiceDeliveryConfigData = self
.call(
RpcMethod::GetPushServiceDeliveryConfig,
&json!({"type": "get_push_service_delivery_config"}),
)
.await?;
Ok(data.config)
}
@@ -190,7 +195,7 @@ impl RpcClient {
}
#[derive(Deserialize)]
struct RolloutConfigData {
struct PushServiceDeliveryConfigData {
config: Value,
}
+8 -6
View File
@@ -4,7 +4,7 @@ use crate::config::{Config, DeliveryConfig};
use crate::delivery;
use crate::metrics::Metrics;
use crate::relay;
use crate::rollout::{self, RolloutSnapshot, RolloutStore};
use crate::relay_consent::{self, RelayConsentStore};
use crate::rpc::RpcClient;
use crate::secret::SecretString;
use crate::tokens::TokenCache;
@@ -91,7 +91,7 @@ pub struct AppState {
pub(crate) cfg: DeliveryConfig,
pub(crate) metrics: Arc<Metrics>,
pub(crate) sidecar: Arc<Sidecar>,
pub(crate) rollout: RolloutStore<RolloutSnapshot>,
pub(crate) relay_consent: RelayConsentStore,
pub(crate) rpc: RpcClient,
pub(crate) http: reqwest::Client,
pub(crate) web_push_http: reqwest::Client,
@@ -103,10 +103,12 @@ pub struct AppState {
impl AppState {
pub(crate) fn try_new(cfg: DeliveryConfig) -> anyhow::Result<Self> {
let metrics = Arc::new(Metrics::new());
let relay_consent = RelayConsentStore::new(cfg.relay_consent_accepted);
metrics.record_relay_consent_accepted(relay_consent.accepted());
let http = vendor::http_client()?;
Ok(Self {
rpc: RpcClient::new(&cfg.rpc, http.clone(), Arc::clone(&metrics)),
rollout: RolloutStore::new(),
relay_consent,
sidecar: Arc::new(Sidecar::new(Arc::clone(&metrics))),
web_push_http: vendor::web_push_http_client()?,
apns_http: vendor::apns_http_client()?,
@@ -144,12 +146,12 @@ async fn run_delivery(cfg: DeliveryConfig) -> anyhow::Result<()> {
let transport = transport.clone();
let state = Arc::clone(&state);
async move {
rollout::run_rollout_subscriber(
relay_consent::run_subscriber(
transport,
&state.rpc,
&state.rollout,
&state.relay_consent,
&state.metrics,
rollout::RECONCILE_INTERVAL,
relay_consent::RECONCILE_INTERVAL,
)
.await
}
-1
View File
@@ -27,7 +27,6 @@
"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/VoiceNoiseSuppressionSchemas.ts": ["exports"],
"packages/schema/src/domains/download/DownloadSchemas.ts": ["exports"],
"packages/schema/src/domains/geolocation/GeolocationSchemas.ts": ["exports"],
@@ -23,10 +23,7 @@ import {
GatewayRolloutConfigResponse,
GatewayRolloutConfigUpdateRequest,
} from '@fluxer/schema/src/domains/admin/GatewayRolloutSchemas';
import {
PushServiceDeliveryConfigResponse,
PushServiceDeliveryConfigUpdateRequest,
} from '@fluxer/schema/src/domains/admin/PushServiceDeliverySchemas';
import {PushRelayConfigResponse, PushRelayConfigUpdateRequest} from '@fluxer/schema/src/domains/admin/PushRelaySchemas';
import {
VoiceNoiseSuppressionConfigResponse,
VoiceNoiseSuppressionConfigUpdateRequest,
@@ -656,7 +653,7 @@ export const InstanceConfigResponse = z.object({
sso: SsoConfigResponse,
gateway_rollout: GatewayRolloutConfigResponse,
voice_noise_suppression: VoiceNoiseSuppressionConfigResponse,
push_service_delivery: PushServiceDeliveryConfigResponse,
push_relay: PushRelayConfigResponse,
domain_migration: DomainMigrationConfigResponse,
altcha_captcha: AltchaCaptchaConfigResponse,
experiment_delivery: ExperimentDeliveryConfigResponse,
@@ -695,7 +692,7 @@ const InstancePolicyUpdateSchema = z.object({
export const InstanceConfigUpdateRequest = z.object({
gateway_rollout: GatewayRolloutConfigUpdateRequest.nullish(),
voice_noise_suppression: VoiceNoiseSuppressionConfigUpdateRequest.nullish(),
push_service_delivery: PushServiceDeliveryConfigUpdateRequest.nullish(),
push_relay: PushRelayConfigUpdateRequest.nullish(),
domain_migration: DomainMigrationConfigUpdateRequest.nullish(),
altcha_captcha: AltchaCaptchaConfigUpdateRequest.nullish(),
experiment_delivery: ExperimentDeliveryConfigUpdateRequest.nullish(),
@@ -0,0 +1,123 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import {
LegacyPushServiceDeliveryWire,
PushRelayConfigSchema,
PushRelayConfigUpdateRequest,
toLegacyPushServiceDeliveryWire,
} from '@fluxer/schema/src/domains/admin/PushRelaySchemas';
import {describe, expect, test} from 'vitest';
const ADMIN_USER_ID = '1500000000000000001';
const ACCEPTED_AT = '2026-09-27T10:11:12.000Z';
describe('push relay consent', () => {
test('reads back as unaccepted by default', () => {
expect(PushRelayConfigSchema.parse({})).toEqual({
relay_consent_accepted: false,
relay_consent_accepted_at: null,
relay_consent_accepted_by: null,
});
});
test('a stored push service delivery row keeps its consent', () => {
expect(
PushRelayConfigSchema.parse({
enabled: true,
config_version: 12,
rollout_basis_points: 10000,
rollout_salt: 'push-service-delivery-v1',
included_user_ids: [],
excluded_user_ids: ['1500000000000000003'],
relay_consent_accepted: true,
relay_consent_accepted_at: ACCEPTED_AT,
relay_consent_accepted_by: ADMIN_USER_ID,
}),
).toEqual({
relay_consent_accepted: true,
relay_consent_accepted_at: ACCEPTED_AT,
relay_consent_accepted_by: ADMIN_USER_ID,
});
});
test('a stored row that predates relay consent reads back as not accepted', () => {
expect(
PushRelayConfigSchema.parse({
enabled: true,
config_version: 4,
rollout_basis_points: 10000,
rollout_salt: 'push-service-delivery-v1',
included_user_ids: [],
excluded_user_ids: [],
}),
).toEqual({
relay_consent_accepted: false,
relay_consent_accepted_at: null,
relay_consent_accepted_by: null,
});
});
test('an accepted consent round-trips', () => {
const accepted = {
relay_consent_accepted: true,
relay_consent_accepted_at: ACCEPTED_AT,
relay_consent_accepted_by: ADMIN_USER_ID,
};
expect(PushRelayConfigSchema.parse(accepted)).toEqual(accepted);
});
test('the update request takes the consent flag on its own', () => {
expect(PushRelayConfigUpdateRequest.parse({relay_consent_accepted: true})).toEqual({
relay_consent_accepted: true,
});
});
test('the update request drops a client-supplied acceptance stamp', () => {
expect(
PushRelayConfigUpdateRequest.parse({
relay_consent_accepted: true,
relay_consent_accepted_at: '2020-01-01T00:00:00.000Z',
relay_consent_accepted_by: ADMIN_USER_ID,
}),
).toEqual({relay_consent_accepted: true});
});
test.each([
{relay_consent_accepted_at: 'yesterday'},
{relay_consent_accepted_at: '2026-09-27'},
{relay_consent_accepted_by: 'not-an-id'},
{relay_consent_accepted: 'yes'},
])('rejects a malformed stored consent: %j', (value) => {
expect(PushRelayConfigSchema.safeParse(value).success).toBe(false);
});
});
describe('legacy push service delivery wire', () => {
test('pins every rollout field to full enrolment and keeps the consent and version', () => {
const wire = toLegacyPushServiceDeliveryWire(
{relay_consent_accepted: true, relay_consent_accepted_at: ACCEPTED_AT, relay_consent_accepted_by: ADMIN_USER_ID},
7,
);
expect(wire).toEqual({
enabled: true,
config_version: 7,
rollout_basis_points: 10000,
rollout_salt: 'push-service-delivery-v1',
included_user_ids: [],
excluded_user_ids: [],
relay_consent_accepted: true,
relay_consent_accepted_at: ACCEPTED_AT,
relay_consent_accepted_by: ADMIN_USER_ID,
});
expect(LegacyPushServiceDeliveryWire.parse(wire)).toEqual(wire);
});
test('rejects a document that would take anyone off the push service', () => {
const wire = toLegacyPushServiceDeliveryWire(PushRelayConfigSchema.parse({}), 0);
expect(LegacyPushServiceDeliveryWire.safeParse({...wire, enabled: false}).success).toBe(false);
expect(LegacyPushServiceDeliveryWire.safeParse({...wire, rollout_basis_points: 5000}).success).toBe(false);
expect(LegacyPushServiceDeliveryWire.safeParse({...wire, excluded_user_ids: ['1500000000000000003']}).success).toBe(
false,
);
});
});
@@ -0,0 +1,56 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import {z} from 'zod';
const LEGACY_PUSH_SERVICE_DELIVERY_ROLLOUT_BASIS_POINTS = 10000;
const LEGACY_PUSH_SERVICE_DELIVERY_SALT = 'push-service-delivery-v1';
export const PushRelayConfigSchema = z.object({
relay_consent_accepted: z.boolean().default(false),
relay_consent_accepted_at: z.iso.datetime().nullable().default(null),
relay_consent_accepted_by: z
.string()
.regex(/^\d{1,20}$/u)
.nullable()
.default(null),
});
export type PushRelayConfig = z.infer<typeof PushRelayConfigSchema>;
export const PushRelayConfigResponse = PushRelayConfigSchema;
export type PushRelayConfigResponse = z.infer<typeof PushRelayConfigResponse>;
export const PushRelayConfigUpdateRequest = z.object({
relay_consent_accepted: z.boolean().optional(),
});
export type PushRelayConfigUpdateRequest = z.infer<typeof PushRelayConfigUpdateRequest>;
export const LegacyPushServiceDeliveryWire = PushRelayConfigSchema.extend({
enabled: z.literal(true),
config_version: z.number().int().min(0),
rollout_basis_points: z.literal(LEGACY_PUSH_SERVICE_DELIVERY_ROLLOUT_BASIS_POINTS),
rollout_salt: z.literal(LEGACY_PUSH_SERVICE_DELIVERY_SALT),
included_user_ids: z.tuple([]),
excluded_user_ids: z.tuple([]),
});
export type LegacyPushServiceDeliveryWire = z.infer<typeof LegacyPushServiceDeliveryWire>;
export function toLegacyPushServiceDeliveryWire(
config: PushRelayConfig,
configVersion: number,
): LegacyPushServiceDeliveryWire {
return {
enabled: true,
config_version: configVersion,
rollout_basis_points: LEGACY_PUSH_SERVICE_DELIVERY_ROLLOUT_BASIS_POINTS,
rollout_salt: LEGACY_PUSH_SERVICE_DELIVERY_SALT,
included_user_ids: [],
excluded_user_ids: [],
relay_consent_accepted: config.relay_consent_accepted,
relay_consent_accepted_at: config.relay_consent_accepted_at,
relay_consent_accepted_by: config.relay_consent_accepted_by,
};
}
@@ -1,89 +0,0 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import {
DEFAULT_PUSH_SERVICE_DELIVERY_CONFIG,
type PushServiceDeliveryConfig,
PushServiceDeliveryConfigSchema,
PushServiceDeliveryConfigUpdateRequest,
pushServiceDeliveryEnrols,
} from '@fluxer/schema/src/domains/admin/PushServiceDeliverySchemas';
import {describe, expect, test} from 'vitest';
const ADMIN_USER_ID = '1500000000000000001';
const TARGETED_USER_ID = '1500000000000000002';
function createConfig(overrides: Partial<PushServiceDeliveryConfig> = {}): PushServiceDeliveryConfig {
return {
...DEFAULT_PUSH_SERVICE_DELIVERY_CONFIG,
included_user_ids: [],
excluded_user_ids: [],
...overrides,
};
}
describe('push service delivery relay consent', () => {
test('a stored configuration that predates relay consent reads back as not accepted', () => {
expect(
PushServiceDeliveryConfigSchema.parse({
enabled: true,
config_version: 4,
rollout_basis_points: 10000,
rollout_salt: 'push-service-delivery-v1',
included_user_ids: [],
excluded_user_ids: [],
}),
).toMatchObject({
relay_consent_accepted: false,
relay_consent_accepted_at: null,
relay_consent_accepted_by: null,
});
});
test('the defaults export carries the unaccepted consent', () => {
expect(DEFAULT_PUSH_SERVICE_DELIVERY_CONFIG.relay_consent_accepted).toBe(false);
expect(DEFAULT_PUSH_SERVICE_DELIVERY_CONFIG.relay_consent_accepted_at).toBeNull();
expect(DEFAULT_PUSH_SERVICE_DELIVERY_CONFIG.relay_consent_accepted_by).toBeNull();
});
test('an accepted consent round-trips through the stored schema', () => {
const accepted = {
...DEFAULT_PUSH_SERVICE_DELIVERY_CONFIG,
relay_consent_accepted: true,
relay_consent_accepted_at: '2026-09-27T10:11:12.000Z',
relay_consent_accepted_by: ADMIN_USER_ID,
};
expect(PushServiceDeliveryConfigSchema.parse(accepted)).toEqual(accepted);
});
test('the update request takes the consent flag on its own', () => {
expect(PushServiceDeliveryConfigUpdateRequest.parse({relay_consent_accepted: true})).toEqual({
relay_consent_accepted: true,
});
});
test('the update request refuses a client-supplied acceptance stamp', () => {
expect(
PushServiceDeliveryConfigUpdateRequest.parse({
relay_consent_accepted: true,
relay_consent_accepted_at: '2020-01-01T00:00:00.000Z',
relay_consent_accepted_by: ADMIN_USER_ID,
}),
).toEqual({relay_consent_accepted: true});
});
test.each([
{relay_consent_accepted_at: 'yesterday'},
{relay_consent_accepted_at: '2026-09-27'},
{relay_consent_accepted_by: 'not-an-id'},
{relay_consent_accepted: 'yes'},
])('rejects a malformed stored consent: %j', (value) => {
expect(PushServiceDeliveryConfigSchema.safeParse(value).success).toBe(false);
});
test('consent alone enrols nobody and refusing it excludes nobody', () => {
const withConsent = createConfig({enabled: false, relay_consent_accepted: true});
const withoutConsent = createConfig({enabled: true, rollout_basis_points: 10000});
expect(pushServiceDeliveryEnrols(withConsent, TARGETED_USER_ID)).toBe(false);
expect(pushServiceDeliveryEnrols(withoutConsent, TARGETED_USER_ID)).toBe(true);
});
});
@@ -1,63 +0,0 @@
// 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,
relay_consent_accepted: z.boolean(),
relay_consent_accepted_at: z.iso.datetime().nullable(),
relay_consent_accepted_by: PushServiceDeliveryTargetIdSchema.nullable(),
};
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([]),
relay_consent_accepted: pushServiceDeliveryConfigFields.relay_consent_accepted.default(false),
relay_consent_accepted_at: pushServiceDeliveryConfigFields.relay_consent_accepted_at.default(null),
relay_consent_accepted_by: pushServiceDeliveryConfigFields.relay_consent_accepted_by.default(null),
});
export type PushServiceDeliveryConfig = z.infer<typeof PushServiceDeliveryConfigSchema>;
export const DEFAULT_PUSH_SERVICE_DELIVERY_CONFIG: PushServiceDeliveryConfig = PushServiceDeliveryConfigSchema.parse(
{},
);
export const PushServiceDeliveryConfigUpdateRequest = z
.object(pushServiceDeliveryConfigFields)
.omit({config_version: true, relay_consent_accepted_at: true, relay_consent_accepted_by: true})
.partial();
export type PushServiceDeliveryConfigUpdateRequest = z.infer<typeof PushServiceDeliveryConfigUpdateRequest>;
export const PushServiceDeliveryConfigResponse = PushServiceDeliveryConfigSchema;
export type PushServiceDeliveryConfigResponse = z.infer<typeof PushServiceDeliveryConfigResponse>;
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;
}
@@ -2,7 +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 {LegacyPushServiceDeliveryWire} from '@fluxer/schema/src/domains/admin/PushRelaySchemas';
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';
@@ -529,7 +529,7 @@ export const RpcResponse = z.discriminatedUnion('type', [
.describe('Response type for push service delivery configuration'),
data: z
.object({
config: PushServiceDeliveryConfigResponse.describe('Push service delivery configuration'),
config: LegacyPushServiceDeliveryWire.describe('Push service delivery configuration'),
})
.describe('Push service delivery config result'),
}),
+1 -1
View File
@@ -227,7 +227,7 @@ pub fn resolve_cloudflare_public_url(public_url_arg: Option<&str>) -> Result<Str
}
}
bail!(
"Missing Cloudflare tunnel public URL. Run `pnpm dev:tunnel:configure -- --public-url https://...` or pass `pnpm dev -- --cloudflare-tunnel --public-url https://...`."
"Missing Cloudflare tunnel public URL. Run `pnpm dev:tunnel:configure --public-url https://...` or pass `pnpm dev --cloudflare-tunnel --public-url https://...`."
);
}