Skip to main content

autoendpoint/
server.rs

1//! Main application server
2#![forbid(unsafe_code)]
3use std::sync::Arc;
4use std::sync::atomic::{AtomicUsize, Ordering};
5use std::time::Duration;
6
7use actix_cors::Cors;
8use actix_web::{
9    App, HttpServer, dev, http::StatusCode, middleware::ErrorHandlers, web, web::Data,
10};
11use cadence::StatsdClient;
12use fernet::MultiFernet;
13use serde_json::json;
14
15#[cfg(feature = "bigtable")]
16use autopush_common::db::bigtable::BigTableClientImpl;
17#[cfg(feature = "postgres")]
18use autopush_common::db::postgres::PgClientImpl;
19#[cfg(feature = "redis")]
20use autopush_common::db::redis::RedisClientImpl;
21#[cfg(feature = "reliable_report")]
22use autopush_common::reliability::PushReliability;
23use autopush_common::{
24    db::{DbSettings, StorageType, client::DbClient, spawn_pool_periodic_reporter},
25    metric_name::MetricName,
26    middleware::sentry::SentryWrapper,
27};
28
29use crate::error::{ApiError, ApiErrorKind, ApiResult};
30use crate::metrics;
31#[cfg(feature = "stub")]
32use crate::routers::stub::router::StubRouter;
33use crate::routers::{apns::router::ApnsRouter, fcm::router::FcmRouter};
34use crate::routes::{
35    health::{health_route, lb_heartbeat_route, log_check, status_route, version_route},
36    registration::{
37        get_channels_route, new_channel_route, register_uaid_route, unregister_channel_route,
38        unregister_user_route, update_token_route,
39    },
40    webpush::{delete_notification_route, webpush_route},
41};
42use crate::settings::Settings;
43use crate::settings::VapidTracker;
44
45#[derive(Clone)]
46pub struct AppState {
47    /// Server Data
48    pub metrics: Arc<StatsdClient>,
49    pub settings: Settings,
50    pub fernet: MultiFernet,
51    pub db: Box<dyn DbClient>,
52    pub http: reqwest::Client,
53    // requests that are currently being exchanged between nodes.
54    pub in_flight_requests: Arc<AtomicUsize>,
55    // subscription updates that are currently being processed.
56    pub in_process_subscription_updates: Arc<AtomicUsize>,
57    pub fcm_router: Arc<FcmRouter>,
58    pub apns_router: Arc<ApnsRouter>,
59    #[cfg(feature = "stub")]
60    pub stub_router: Arc<StubRouter>,
61    #[cfg(feature = "reliable_report")]
62    pub reliability: Arc<PushReliability>,
63    pub reliability_filter: VapidTracker,
64}
65
66#[cfg(test)]
67impl AppState {
68    /// build a fake AppState with a MockDbClient and default settings, for use in tests.
69    pub(crate) async fn test_default(mock_db: autopush_common::db::mock::MockDbClient) -> Self {
70        let settings = Settings {
71            auth_keys: "HJVPy4ZwF4Yz_JdvXTL8hRcwIhv742vC60Tg5Ycrvw8=".to_owned(),
72            ..Default::default()
73        };
74        let metrics = Arc::new(crate::metrics::Metrics::sink());
75        let db = mock_db.into_boxed_arc();
76        #[cfg(feature = "reliable_report")]
77        let reliability =
78            Arc::new(PushReliability::new(&None, db.box_clone(), &metrics.clone(), 0).unwrap());
79        let fernet = settings.make_fernet();
80        Self {
81            metrics: metrics.clone(),
82            settings,
83            fernet,
84            db: db.clone(),
85            http: reqwest::Client::new(),
86            in_flight_requests: Arc::new(AtomicUsize::new(0)),
87            in_process_subscription_updates: Arc::new(AtomicUsize::new(0)),
88            fcm_router: Arc::new(
89                FcmRouter::new(
90                    crate::routers::fcm::settings::FcmSettings::default(),
91                    url::Url::parse("https://example.com").unwrap(),
92                    reqwest::Client::new(),
93                    metrics.clone(),
94                    db.clone(),
95                    #[cfg(feature = "reliable_report")]
96                    reliability.clone(),
97                )
98                .await
99                .unwrap(),
100            ),
101            apns_router: Arc::new(
102                ApnsRouter::new(
103                    crate::routers::apns::settings::ApnsSettings::default(),
104                    url::Url::parse("https://example.com").unwrap(),
105                    metrics.clone(),
106                    db.clone(),
107                    #[cfg(feature = "reliable_report")]
108                    reliability.clone(),
109                )
110                .await
111                .unwrap(),
112            ),
113            #[cfg(feature = "stub")]
114            stub_router: Arc::new(
115                StubRouter::new(crate::routers::stub::settings::StubSettings::default()).unwrap(),
116            ),
117            #[cfg(feature = "reliable_report")]
118            reliability,
119            reliability_filter: VapidTracker(Vec::new()),
120        }
121    }
122}
123
124pub struct Server;
125
126impl Server {
127    pub async fn with_settings(settings: Settings) -> ApiResult<dev::Server> {
128        let metrics = Arc::new(metrics::metrics_from_settings(&settings)?);
129        let bind_address = format!("{}:{}", settings.host, settings.port);
130        let fernet = settings.make_fernet();
131        let endpoint_url = settings.endpoint_url();
132        let reliability_filter = VapidTracker(
133            settings
134                .tracking_keys()
135                .map_err(|e| ApiErrorKind::General(format!("Configuration Error: {e}")))?,
136        );
137        let db_settings = DbSettings {
138            dsn: settings.db_dsn.clone(),
139            db_settings: if settings.db_settings.is_empty() {
140                warn!("❗ Using obsolete message_table and router_table args");
141                // backfill from the older arguments.
142                json!({"message_table": settings.message_table_name, "router_table":settings.router_table_name}).to_string()
143            } else {
144                settings.db_settings.clone()
145            },
146        };
147        let db: Box<dyn DbClient> = match StorageType::from_dsn(&db_settings.dsn) {
148            #[cfg(feature = "bigtable")]
149            StorageType::BigTable => {
150                debug!("Using BigTable");
151                let client = BigTableClientImpl::new(metrics.clone(), &db_settings)?;
152                Box::new(client)
153            }
154            #[cfg(feature = "postgres")]
155            StorageType::Postgres => {
156                debug!("Using Postgres");
157                let client = PgClientImpl::new(metrics.clone(), &db_settings)?;
158                // client.spawn_sweeper(Duration::from_secs(30));
159                Box::new(client)
160            }
161            #[cfg(feature = "redis")]
162            StorageType::Redis => Box::new(RedisClientImpl::new(metrics.clone(), &db_settings)?),
163            _ => {
164                debug!("No idea what {:?} is", &db_settings.dsn);
165                return Err(ApiErrorKind::General(
166                    "Invalid or Unsupported DSN specified".to_owned(),
167                )
168                .into());
169            }
170        };
171        #[cfg(feature = "reliable_report")]
172        let reliability = Arc::new(
173            PushReliability::new(
174                &settings.reliability_dsn,
175                db.clone(),
176                &metrics,
177                settings.reliability_retry_count,
178            )
179            .map_err(|e| {
180                ApiErrorKind::General(format!("Could not initialize Reliability Report: {e:?}"))
181            })?,
182        );
183        let http = reqwest::ClientBuilder::new()
184            .connect_timeout(Duration::from_millis(settings.connection_timeout_millis))
185            .timeout(Duration::from_millis(settings.request_timeout_millis))
186            .pool_max_idle_per_host(settings.pool_max_idle_per_host)
187            .pool_idle_timeout(Duration::from_secs(settings.pool_idle_timeout_secs))
188            .build()
189            .expect("Could not generate request client");
190        let fcm_router = Arc::new(
191            FcmRouter::new(
192                settings.fcm.clone(),
193                endpoint_url.clone(),
194                http.clone(),
195                metrics.clone(),
196                db.clone(),
197                #[cfg(feature = "reliable_report")]
198                reliability.clone(),
199            )
200            .await?,
201        );
202        let apns_router = Arc::new(
203            ApnsRouter::new(
204                settings.apns.clone(),
205                endpoint_url.clone(),
206                metrics.clone(),
207                db.clone(),
208                #[cfg(feature = "reliable_report")]
209                reliability.clone(),
210            )
211            .await?,
212        );
213        #[cfg(feature = "stub")]
214        let stub_router = Arc::new(StubRouter::new(settings.stub.clone())?);
215        let in_flight_requests = Arc::new(AtomicUsize::new(0));
216        let in_process_subscription_updates = Arc::new(AtomicUsize::new(0));
217        let app_state = AppState {
218            metrics: metrics.clone(),
219            settings,
220            fernet,
221            db,
222            http,
223            in_flight_requests: in_flight_requests.clone(),
224            in_process_subscription_updates: in_process_subscription_updates.clone(),
225            fcm_router,
226            apns_router,
227            #[cfg(feature = "stub")]
228            stub_router,
229            #[cfg(feature = "reliable_report")]
230            reliability,
231            reliability_filter,
232        };
233
234        spawn_pool_periodic_reporter(
235            Duration::from_secs(10),
236            app_state.db.clone(),
237            app_state.metrics.clone(),
238        );
239
240        // Periodically report in-flight request gauge
241        {
242            let metrics = app_state.metrics.clone();
243            let in_flight = in_flight_requests;
244            tokio::spawn(async move {
245                let mut interval = tokio::time::interval(Duration::from_secs(10));
246                loop {
247                    interval.tick().await;
248                    let count = in_flight.load(Ordering::Relaxed);
249                    if let Err(e) = cadence::Gauged::gauge(
250                        metrics.as_ref(),
251                        MetricName::InFlightNodeRequests.as_ref(),
252                        count as u64,
253                    ) {
254                        debug!("Failed to report in-flight metric: {}", e);
255                    }
256                }
257            });
258        }
259
260        let server = HttpServer::new(move || {
261            // These have a bad habit of being reset. Specify them explicitly.
262            let cors = Cors::default()
263                .allow_any_origin()
264                .allow_any_header()
265                .allowed_methods(vec![
266                    actix_web::http::Method::DELETE,
267                    actix_web::http::Method::GET,
268                    actix_web::http::Method::POST,
269                    actix_web::http::Method::PUT,
270                ])
271                .max_age(3600);
272            let app = App::new()
273                // Actix 4 recommends wrapping structures wtih web::Data (internally an Arc)
274                .app_data(Data::new(app_state.clone()))
275                // Extractor configuration
276                .app_data(web::PayloadConfig::new(app_state.settings.max_data_bytes))
277                .app_data(web::JsonConfig::default().limit(app_state.settings.max_data_bytes))
278                // Middleware
279                .wrap(ErrorHandlers::new().handler(StatusCode::NOT_FOUND, ApiError::render_404))
280                // Our modified Sentry wrapper which does some blocking of non-reportable errors.
281                .wrap(SentryWrapper::<ApiError>::new(
282                    metrics.clone(),
283                    "api_error".to_owned(),
284                    app_state.settings.disable_sentry,
285                ))
286                .wrap(cors)
287                // Endpoints
288                .service(
289                    web::resource(["/wpush/{api_version}/{token}", "/wpush/{token}"])
290                        .route(web::post().to(webpush_route)),
291                )
292                .service(
293                    web::resource("/m/{message_id}")
294                        .route(web::delete().to(delete_notification_route)),
295                )
296                .service(
297                    web::resource("/v1/{router_type}/{app_id}/registration")
298                        .route(web::post().to(register_uaid_route)),
299                )
300                .service(
301                    web::resource("/v1/{router_type}/{app_id}/registration/{uaid}")
302                        .route(web::put().to(update_token_route))
303                        .route(web::get().to(get_channels_route))
304                        .route(web::delete().to(unregister_user_route)),
305                )
306                .service(
307                    web::resource("/v1/{router_type}/{app_id}/registration/{uaid}/subscription")
308                        .route(web::post().to(new_channel_route)),
309                )
310                .service(
311                    web::resource(
312                        "/v1/{router_type}/{app_id}/registration/{uaid}/subscription/{chid}",
313                    )
314                    .route(web::delete().to(unregister_channel_route)),
315                )
316                // Health checks
317                .service(web::resource("/status").route(web::get().to(status_route)))
318                .service(web::resource("/health").route(web::get().to(health_route)))
319                // legacy
320                .service(web::resource("/v1/err").route(web::get().to(log_check)))
321                // standardized
322                .service(web::resource("/__error__").route(web::get().to(log_check)))
323                // Dockerflow
324                .service(web::resource("/__heartbeat__").route(web::get().to(health_route)))
325                .service(web::resource("/__lbheartbeat__").route(web::get().to(lb_heartbeat_route)))
326                .service(web::resource("/__version__").route(web::get().to(version_route)));
327            #[cfg(feature = "reliable_report")]
328            // Note: Only the endpoint returns the Prometheus "/metrics" collection report. This report contains all metrics for both
329            // connection and endpoint, inclusive. It is served here mostly for simplicity's sake (since the endpoint handles more
330            // HTTP requests than the connection server, and this will simplify metric collection and reporting.)
331            let app = app.service(
332                web::resource("/metrics")
333                    .route(web::get().to(crate::routes::reliability::report_handler)),
334            );
335            app
336        })
337        .bind(bind_address)?
338        .run();
339
340        Ok(server)
341    }
342}