feat(push): ring Android calls and harden the relay (#2915)

This commit is contained in:
Hampus
2026-09-24 00:07:15 +02:00
committed by GitHub
parent b16989d567
commit 5fde6eb484
25 changed files with 1050 additions and 120 deletions
+9
View File
@@ -17,6 +17,7 @@ const DEFAULT_QUEUE_CAPACITY: usize = 10_000;
const DEFAULT_SEND_CONCURRENCY: usize = 256;
const DEFAULT_RELAY_MAX_CONCURRENT: usize = 1_024;
const DEFAULT_RELAY_MAX_BODY_BYTES: usize = 2_816;
const DEFAULT_TRUSTED_PROXY_HOPS: usize = 1;
const DEFAULT_DEVICE_TOKEN_BUCKET_ENTRIES: usize = 1_000_000;
const DEFAULT_DEVICE_TOKEN_BUCKET_PER_MINUTE: u32 = 60;
const DEFAULT_DEVICE_TOKEN_BUCKET_BURST: u32 = 20;
@@ -198,6 +199,7 @@ pub struct RelayConfig {
pub max_body_bytes: usize,
pub trust_client_ip_header: bool,
pub client_ip_header_name: String,
pub trusted_proxy_hops: usize,
pub device_token_bucket: BucketConfig,
pub source_bucket: Option<BucketConfig>,
pub apns: Option<ApnsConfig>,
@@ -297,6 +299,13 @@ impl RelayConfig {
.get("FLUXER_CLIENT_IP_HEADER_NAME")
.unwrap_or(DEFAULT_CLIENT_IP_HEADER_NAME)
.to_ascii_lowercase(),
trusted_proxy_hops: parse_number(
"FLUXER_PUSH_RELAY_TRUSTED_PROXY_HOPS",
env.get("FLUXER_PUSH_RELAY_TRUSTED_PROXY_HOPS"),
DEFAULT_TRUSTED_PROXY_HOPS,
0,
8,
)?,
device_token_bucket: bucket_config(
&env,
"FLUXER_PUSH_RELAY_TOKEN_BUCKET",
+10 -3
View File
@@ -58,10 +58,17 @@ enum Audience {
impl Audience {
fn admits(self, subscription: &Subscription) -> bool {
let voip = subscription.platform() == Some(Platform::IosApnsVoip);
match self {
Self::Standard => !voip,
Self::Ring => voip,
Self::Standard => subscription.platform() != Some(Platform::IosApnsVoip),
Self::Ring => Self::rings(subscription),
}
}
fn rings(subscription: &Subscription) -> bool {
match subscription.platform() {
Some(Platform::IosApnsVoip) => true,
Some(Platform::AndroidFcm) => subscription.is_web_push_registration(),
_ => false,
}
}
}
+6
View File
@@ -53,6 +53,12 @@ pub struct RingJob {
pub message_id: String,
pub started_at_ms: i64,
pub expires_at_ms: i64,
#[serde(default)]
pub caller_id: Option<String>,
#[serde(default)]
pub caller_name: Option<String>,
#[serde(default)]
pub caller_avatar_url: Option<String>,
}
#[derive(Debug, Error)]
+38 -1
View File
@@ -20,6 +20,11 @@ const SHRUNK_BODY_MAX_BYTES: usize = 40;
const MINIMAL_TITLE_MAX_BYTES: usize = 120;
const MEDIA_KEYS: [&str; 2] = ["image_url", "image"];
const ICON_KEYS: [&str; 3] = ["icon", "badge", "author_avatar_url"];
const RING_CALLER_ID_KEY: &str = "caller_id";
const RING_CALLER_NAME_KEY: &str = "caller_name";
const RING_CALLER_AVATAR_KEY: &str = "caller_avatar_url";
const RING_AVATAR_KEYS: [&str; 1] = [RING_CALLER_AVATAR_KEY];
const RING_IDENTITY_KEYS: [&str; 2] = [RING_CALLER_ID_KEY, RING_CALLER_NAME_KEY];
const MINIMAL_DATA_KEYS: [&str; 7] = [
"channel_id",
"message_id",
@@ -106,13 +111,15 @@ pub fn web_push_clear(job: &ClearJob, badge_count: u32) -> Value {
}
pub fn web_push_call_ring(job: &RingJob) -> Value {
let data = json!({
let mut data = json!({
"type": RING_TYPE,
"channel_id": job.channel_id,
"message_id": job.message_id,
"target_user_id": job.user_id,
"started_at_ms": job.started_at_ms,
"expires_at_ms": job.expires_at_ms,
});
put_caller(&mut data, job);
json!({
"web_push": WEB_PUSH_MARKER,
"type": RING_TYPE,
@@ -120,6 +127,22 @@ pub fn web_push_call_ring(job: &RingJob) -> Value {
})
}
fn put_caller(data: &mut Value, job: &RingJob) {
let Some(object) = data.as_object_mut() else {
return;
};
let caller = [
(RING_CALLER_ID_KEY, job.caller_id.as_deref()),
(RING_CALLER_NAME_KEY, job.caller_name.as_deref()),
(RING_CALLER_AVATAR_KEY, job.caller_avatar_url.as_deref()),
];
for (key, value) in caller {
if let Some(value) = value.filter(|value| !value.is_empty()) {
object.insert(key.to_owned(), value.into());
}
}
}
pub fn fcm_message(device_token: &str, envelope: &Value) -> Value {
if is_clear(envelope) {
return fcm_clear_message(device_token, envelope);
@@ -397,6 +420,9 @@ pub fn fit(envelope: &Value, budget: usize) -> (Vec<u8>, Option<PayloadShrink>)
if serialized.len() <= budget {
return (serialized, None);
}
if matches!(record_kind(envelope), RecordKind::Ring) {
return fit_ring(envelope, budget);
}
let mut working = envelope.clone();
for step in PayloadShrink::ALL {
working = shrink(&working, step, budget);
@@ -408,6 +434,17 @@ pub fn fit(envelope: &Value, budget: usize) -> (Vec<u8>, Option<PayloadShrink>)
(serialize(&working), Some(PayloadShrink::Minimal))
}
fn fit_ring(envelope: &Value, budget: usize) -> (Vec<u8>, Option<PayloadShrink>) {
let mut working = envelope.clone();
drop_keys(&mut working, &RING_AVATAR_KEYS);
let serialized = serialize(&working);
if serialized.len() <= budget {
return (serialized, Some(PayloadShrink::Icons));
}
drop_keys(&mut working, &RING_IDENTITY_KEYS);
(serialize(&working), Some(PayloadShrink::Minimal))
}
fn shrink(envelope: &Value, step: PayloadShrink, budget: usize) -> Value {
let mut working = envelope.clone();
match step {
+4 -1
View File
@@ -49,7 +49,10 @@ impl From<VendorOutcome> for SendOutcome {
fn from(outcome: VendorOutcome) -> Self {
match outcome {
VendorOutcome::Accepted => Self::Accepted,
VendorOutcome::Unreachable => Self::transient("transport"),
VendorOutcome::Unreachable(unreachable) if unreachable.is_permanent() => {
Self::permanent(unreachable.label())
}
VendorOutcome::Unreachable(unreachable) => Self::transient(unreachable.label()),
VendorOutcome::Refused(refusal) => match refusal.dead_token {
Some(dead_token) => Self::TokenInvalid {
reason: dead_token.label(),
+16 -5
View File
@@ -6,12 +6,13 @@ use crate::providers::SendOutcome;
use crate::resolver;
use crate::server::AppState;
use crate::subscription::Subscription;
use crate::vendor::is_transient_status;
use crate::vendor::{Unreachable, is_transient_status};
use rand::RngExt as _;
use reqwest::header::{AUTHORIZATION, CONTENT_ENCODING, CONTENT_TYPE};
use serde_json::Value;
use std::net::IpAddr;
use std::time::Duration;
use tracing::warn;
use url::{Host, Url};
pub const RECORD_SIZE: usize = 2816;
@@ -90,10 +91,20 @@ pub async fn send(state: &AppState, sub: &Subscription, envelope: &Value) -> Sen
let status = match response {
Ok(response) => response.status().as_u16(),
Err(_) if attempt >= MAX_TRANSIENT_RETRIES => {
return SendOutcome::transient("transport");
}
Err(_) => {
Err(error) => {
let unreachable = Unreachable::of(&error);
warn!(
error = %error,
kind = unreachable.label(),
endpoint = %origin_of(&sub.endpoint),
"web push request did not complete"
);
if unreachable.is_permanent() {
return SendOutcome::permanent(unreachable.label());
}
if attempt >= MAX_TRANSIENT_RETRIES {
return SendOutcome::transient(unreachable.label());
}
tokio::time::sleep(retry_delay(attempt)).await;
attempt += 1;
continue;
+32 -4
View File
@@ -11,7 +11,7 @@ fn resolve(cfg: &RelayConfig, peer: SocketAddr, headers: &HeaderMap) -> Option<I
headers
.get(&cfg.client_ip_header_name)
.and_then(|value| value.to_str().ok())
.and_then(nearest_entry)
.and_then(|value| entry_from_right(value, cfg.trusted_proxy_hops))
.map(|ip| ip.to_canonical())
}
@@ -19,9 +19,37 @@ pub fn for_rate_limit(cfg: &RelayConfig, peer: SocketAddr, headers: &HeaderMap)
resolve(cfg, peer, headers).unwrap_or_else(|| peer.ip().to_canonical())
}
fn nearest_entry(value: &str) -> Option<IpAddr> {
fn entry_from_right(value: &str, skip: usize) -> Option<IpAddr> {
value
.rsplit(',')
.filter_map(|entry| entry.trim().trim_matches(['[', ']']).parse().ok())
.next()
.map(|entry| entry.trim().trim_matches(['[', ']']))
.nth(skip)
.and_then(|entry| entry.parse().ok())
}
#[cfg(test)]
mod tests {
use super::*;
const EDGE: &str = "203.0.113.9";
const INSTANCE: &str = "198.51.100.7";
#[test]
fn the_rightmost_entry_is_the_edge_not_the_sending_instance() {
let chain = format!("{INSTANCE}, {EDGE}");
assert_eq!(entry_from_right(&chain, 0).unwrap().to_string(), EDGE);
assert_eq!(entry_from_right(&chain, 1).unwrap().to_string(), INSTANCE);
}
#[test]
fn a_longer_chain_still_names_the_sending_instance() {
let chain = format!("1.2.3.4, 5.6.7.8, {INSTANCE}, {EDGE}");
assert_eq!(entry_from_right(&chain, 1).unwrap().to_string(), INSTANCE);
}
#[test]
fn a_short_chain_yields_nothing_rather_than_a_wrong_answer() {
assert!(entry_from_right(EDGE, 1).is_none());
assert!(entry_from_right("", 1).is_none());
}
}
+44 -3
View File
@@ -37,7 +37,8 @@ const TTL_HEADER: &str = "ttl";
const URGENCY_HEADER: &str = "urgency";
const JSON_CONTENT_TYPE: &str = "application/json";
const DIGEST_BYTES: usize = 8;
const APNS_DEVICE_TOKEN_LEN: usize = 64;
const MIN_APNS_DEVICE_TOKEN_LEN: usize = 64;
const MAX_APNS_DEVICE_TOKEN_LEN: usize = 256;
const MAX_FCM_DEVICE_TOKEN_LEN: usize = 512;
const MAX_TTL_SECONDS: i64 = 86_400;
const BODY_READ_TIMEOUT: Duration = Duration::from_secs(15);
@@ -396,7 +397,7 @@ fn finish(
})?;
let (result, verdict) = match outcome {
VendorOutcome::Accepted => (RelayResult::Accepted, Ok(())),
VendorOutcome::Unreachable => (
VendorOutcome::Unreachable(_) => (
RelayResult::Failed,
Err(Rejection::new(Reason::ProviderUnavailable)),
),
@@ -425,7 +426,8 @@ fn refusal_reason(refusal: &Refusal) -> Reason {
fn device_token_is_shaped(leg: RelayLeg, device_token: &str) -> bool {
match leg {
RelayLeg::Apns | RelayLeg::ApnsVoip => {
device_token.len() == APNS_DEVICE_TOKEN_LEN
(MIN_APNS_DEVICE_TOKEN_LEN..=MAX_APNS_DEVICE_TOKEN_LEN).contains(&device_token.len())
&& device_token.len().is_multiple_of(2)
&& device_token.bytes().all(|byte| byte.is_ascii_hexdigit())
}
RelayLeg::Fcm => {
@@ -485,3 +487,42 @@ fn digest(value: &str) -> String {
}
BASE64_URL_SAFE_NO_PAD.encode(&Sha256::digest(value.as_bytes())[..DIGEST_BYTES])
}
#[cfg(test)]
mod tests {
use super::*;
fn hex(len: usize) -> String {
"a".repeat(len)
}
#[test]
fn apns_accepts_every_token_length_apple_hands_out() {
for len in [64, 128, 160, 200, 256] {
assert!(
device_token_is_shaped(RelayLeg::Apns, &hex(len)),
"{len} hex characters must be accepted"
);
}
assert!(device_token_is_shaped(RelayLeg::Apns, &"A".repeat(64)));
assert!(device_token_is_shaped(RelayLeg::ApnsVoip, &hex(160)));
}
#[test]
fn apns_rejects_tokens_that_are_not_even_length_hex() {
for token in [hex(62), hex(63), hex(161), hex(258), "z".repeat(64)] {
assert!(
!device_token_is_shaped(RelayLeg::Apns, &token),
"{token} must be rejected"
);
}
}
#[test]
fn payload_too_large_answers_413() {
assert_eq!(
Reason::PayloadTooLarge.status(),
StatusCode::PAYLOAD_TOO_LARGE
);
}
}
+4 -4
View File
@@ -45,10 +45,10 @@ impl Reason {
pub fn status(self) -> StatusCode {
match self {
Self::BadRequest
| Self::PayloadTooLarge
| Self::DeviceTokenInvalid
| Self::AppUnknown => StatusCode::BAD_REQUEST,
Self::BadRequest | Self::DeviceTokenInvalid | Self::AppUnknown => {
StatusCode::BAD_REQUEST
}
Self::PayloadTooLarge => StatusCode::PAYLOAD_TOO_LARGE,
Self::DeviceTokenGone => StatusCode::GONE,
Self::RateLimited => StatusCode::TOO_MANY_REQUESTS,
Self::ProviderUnavailable => StatusCode::BAD_GATEWAY,
+83 -3
View File
@@ -8,6 +8,7 @@ use reqwest::header::{AUTHORIZATION, CONTENT_TYPE};
use reqwest::redirect::Policy;
use serde_json::Value;
use std::time::Duration;
use tracing::warn;
const HTTP_TIMEOUT: Duration = Duration::from_secs(10);
const APNS_TIMEOUT: Duration = Duration::from_secs(5);
@@ -15,6 +16,7 @@ const CONNECT_TIMEOUT: Duration = Duration::from_secs(5);
const APNS_TOPIC_HEADER: &str = "apns-topic";
pub const FCM_CONTENT_TYPE: &str = "application/json; charset=UTF-8";
const TOO_MANY_REQUESTS: u16 = 429;
const DNS_ERROR_MARKER: &str = "dns error";
const MAX_ERROR_BODY_BYTES: usize = 8_192;
const HTTP_ERROR: &str = "http_error";
const UNREGISTERED: &str = "UNREGISTERED";
@@ -58,7 +60,45 @@ pub struct ApnsRequest<'a> {
pub enum VendorOutcome {
Accepted,
Refused(Refusal),
Unreachable,
Unreachable(Unreachable),
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum Unreachable {
Dns,
Transport,
}
impl Unreachable {
pub fn label(self) -> &'static str {
match self {
Self::Dns => "dns",
Self::Transport => "transport",
}
}
pub fn is_permanent(self) -> bool {
matches!(self, Self::Dns)
}
pub fn of(error: &reqwest::Error) -> Self {
if names_no_host(error) {
Self::Dns
} else {
Self::Transport
}
}
}
fn names_no_host(error: &reqwest::Error) -> bool {
let mut current = std::error::Error::source(error);
while let Some(error) = current {
if error.to_string().contains(DNS_ERROR_MARKER) {
return true;
}
current = error.source();
}
false
}
#[derive(Debug, Eq, PartialEq)]
@@ -139,8 +179,17 @@ async fn outcome(
response: reqwest::Result<reqwest::Response>,
refusal: fn(u16, &[u8]) -> Refusal,
) -> VendorOutcome {
let Ok(response) = response else {
return VendorOutcome::Unreachable;
let response = match response {
Ok(response) => response,
Err(error) => {
let unreachable = Unreachable::of(&error);
warn!(
error = %error,
kind = unreachable.label(),
"vendor request did not complete"
);
return VendorOutcome::Unreachable(unreachable);
}
};
if response.status().is_success() {
return VendorOutcome::Accepted;
@@ -224,3 +273,34 @@ pub async fn read_error_body(response: reqwest::Response) -> Vec<u8> {
body.truncate(MAX_ERROR_BODY_BYTES);
body
}
#[cfg(test)]
mod tests {
use super::*;
async fn error_for(url: &str) -> reqwest::Error {
reqwest::Client::builder()
.no_proxy()
.connect_timeout(CONNECT_TIMEOUT)
.build()
.expect("the http client builds")
.post(url)
.send()
.await
.expect_err("the request cannot complete")
}
#[tokio::test]
async fn a_host_that_does_not_resolve_is_permanent() {
let error = error_for("https://push.invalid/relay/v1/apns/stable/production/token").await;
assert_eq!(Unreachable::of(&error), Unreachable::Dns);
assert!(Unreachable::of(&error).is_permanent());
}
#[tokio::test]
async fn a_refused_connection_stays_retryable() {
let error = error_for("http://127.0.0.1:1/").await;
assert_eq!(Unreachable::of(&error), Unreachable::Transport);
assert!(!Unreachable::of(&error).is_permanent());
}
}