diff --git a/Cargo.lock b/Cargo.lock index bb2d529..89aef42 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -529,6 +529,7 @@ dependencies = [ "tracing", "tracing-subscriber", "tracing-test", + "uuid 1.19.0", "validator", "wiremock", ] @@ -2561,6 +2562,7 @@ version = "1.19.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e2e054861b4bd027cd373e18e8d8d8e6548085000e41290d95ce0c373a654b4a" dependencies = [ + "getrandom 0.3.4", "js-sys", "wasm-bindgen", ] diff --git a/Cargo.toml b/Cargo.toml index a6dde39..7cc2a07 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -10,6 +10,7 @@ axum = { version = "0.8", features = ["macros"] } tokio = { version = "1", features = ["full"] } tower = "0.5" parking_lot = "0.12" +uuid = { version = "1", features = ["v4"] } tower-http = { version = "0.6", features = ["compression-gzip", "cors", "trace", "normalize-path"] } reqwest = { version = "0.13", features = ["json", "gzip"] } serde = { version = "1", features = ["derive"] } diff --git a/src/config/settings.rs b/src/config/settings.rs index 70f1b30..ba2ef9e 100644 --- a/src/config/settings.rs +++ b/src/config/settings.rs @@ -108,6 +108,10 @@ pub struct AppSettings { pub api_poll_timeout_seconds: u64, #[serde(default = "default_allow_origins")] pub allow_origins: Vec, + // A zero interval would panic the flush task silently. + #[serde(default = "default_usage_flush_interval")] + #[validate(range(min = 1))] + pub usage_flush_interval_seconds: u64, #[serde(default)] pub server: ServerSettings, #[serde(default)] @@ -132,6 +136,10 @@ fn default_allow_origins() -> Vec { vec!["*".to_string()] } +fn default_usage_flush_interval() -> u64 { + 60 +} + impl Default for AppSettings { fn default() -> Self { Self { @@ -141,6 +149,7 @@ impl Default for AppSettings { api_poll_frequency_seconds: default_api_poll_frequency(), api_poll_timeout_seconds: default_api_poll_timeout(), allow_origins: default_allow_origins(), + usage_flush_interval_seconds: default_usage_flush_interval(), server: ServerSettings::default(), logging: LoggingSettings::default(), health_check: HealthCheckSettings::default(), @@ -194,6 +203,16 @@ mod tests { assert!(settings.validate().is_err()); } + #[test] + fn test_config_with_zero_usage_flush_interval_is_invalid() { + // Given an interval that would panic the flush task + let settings: AppSettings = + serde_json::from_str(r#"{"usage_flush_interval_seconds": 0}"#).unwrap(); + + // Then + assert!(settings.validate().is_err()); + } + #[test] fn test_config_with_only_a_proxy_key_is_valid() { // Given a config file that relies entirely on the proxy config diff --git a/src/lib.rs b/src/lib.rs index 2bfc72a..ce2f691 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -2,7 +2,9 @@ pub mod cache; pub mod config; pub mod environments; pub mod error; +pub mod middleware; pub mod models; pub mod routes; pub mod services; pub mod state; +pub mod usage; diff --git a/src/main.rs b/src/main.rs index 405c5fd..894d18c 100644 --- a/src/main.rs +++ b/src/main.rs @@ -30,20 +30,25 @@ async fn main() -> anyhow::Result<()> { ); } - let (app, environment_service) = create_router(settings.clone()); + let (app, state) = create_router(settings.clone()); // Refreshes must never overlap: a delayed older poll finishing after a // newer one could restore removed environments or rotated keys. The // poll loop is serial, so it just has to start after the initial // refresh completes. info!("Loading initial environment data..."); - environment_service.refresh_environment_caches().await; + state.environments.refresh_environment_caches().await; - let polling_service = environment_service.clone(); + let polling_service = state.environments.clone(); tokio::spawn(async move { polling_service.poll_environments().await; }); + let usage = state.usage.clone(); + tokio::spawn(async move { + usage.flush_periodically().await; + }); + let addr = SocketAddr::from(( settings .server diff --git a/src/routes/cors.rs b/src/middleware/cors.rs similarity index 100% rename from src/routes/cors.rs rename to src/middleware/cors.rs diff --git a/src/middleware/mod.rs b/src/middleware/mod.rs new file mode 100644 index 0000000..d8d781e --- /dev/null +++ b/src/middleware/mod.rs @@ -0,0 +1,2 @@ +pub mod cors; +pub mod usage; diff --git a/src/middleware/usage.rs b/src/middleware/usage.rs new file mode 100644 index 0000000..21417f5 --- /dev/null +++ b/src/middleware/usage.rs @@ -0,0 +1,40 @@ +use crate::routes::{ENVIRONMENT_DOCUMENT_PATH, FLAGS_PATH, IDENTITIES_PATH}; +use crate::state::AppState; +use crate::usage::Resource; +use axum::extract::{MatchedPath, Request, State}; +use axum::middleware::Next; +use axum::response::Response; + +/// Count a request once the handler has answered: only a 2xx is usage. +pub async fn track_usage(State(state): State, request: Request, next: Next) -> Response { + let resource = request + .extensions() + .get::() + .and_then(|path| resource_for(path.as_str())); + let environment_key = request + .headers() + .get("X-Environment-Key") + .and_then(|value| value.to_str().ok()) + .map(str::to_owned); + + let response = next.run(request).await; + + if !response.status().is_success() { + return response; + } + if let (Some(resource), Some(environment_key)) = (resource, environment_key) { + if let Some(keys) = state.environments.resolve(&environment_key) { + state.usage.track(&keys.client_key, resource); + } + } + response +} + +fn resource_for(route: &str) -> Option { + match route { + FLAGS_PATH => Some(Resource::Flags), + IDENTITIES_PATH => Some(Resource::Identities), + ENVIRONMENT_DOCUMENT_PATH => Some(Resource::EnvironmentDocument), + _ => None, + } +} diff --git a/src/routes/environment_document.rs b/src/routes/environment_document.rs index 7e68e3a..da361f0 100644 --- a/src/routes/environment_document.rs +++ b/src/routes/environment_document.rs @@ -1,14 +1,15 @@ use crate::error::{EdgeProxyError, Result}; use crate::routes::extractors::extract_environment_key; -use crate::state::AppState; +use crate::services::EnvironmentService; use axum::{ extract::State, http::{HeaderMap, header}, response::IntoResponse, }; +use std::sync::Arc; pub async fn get_environment_document( - State(service): State, + State(service): State>, headers: HeaderMap, ) -> Result { let environment_key = extract_environment_key(&headers)?; diff --git a/src/routes/flags.rs b/src/routes/flags.rs index 136f868..b51ed0c 100644 --- a/src/routes/flags.rs +++ b/src/routes/flags.rs @@ -1,12 +1,13 @@ use crate::error::Result; use crate::routes::extractors::extract_environment_key; -use crate::state::AppState; +use crate::services::EnvironmentService; use axum::{ Json, extract::{Query, State}, http::HeaderMap, }; use serde::Deserialize; +use std::sync::Arc; #[derive(Deserialize)] pub struct FlagsQuery { @@ -14,7 +15,7 @@ pub struct FlagsQuery { } pub async fn get_flags( - State(service): State, + State(service): State>, headers: HeaderMap, Query(query): Query, ) -> Result> { diff --git a/src/routes/health.rs b/src/routes/health.rs index 14738b8..2720652 100644 --- a/src/routes/health.rs +++ b/src/routes/health.rs @@ -1,7 +1,8 @@ -use crate::state::AppState; +use crate::services::EnvironmentService; use axum::{Json, extract::State, http::StatusCode, response::IntoResponse}; use chrono::Utc; use serde::{Deserialize, Serialize}; +use std::sync::Arc; #[derive(Debug, Clone, Serialize, Deserialize)] pub struct HealthCheckResponse { @@ -28,7 +29,7 @@ impl HealthCheckResponse { } } -pub async fn health_check(State(service): State) -> impl IntoResponse { +pub async fn health_check(State(service): State>) -> impl IntoResponse { let last_updated = service.last_updated_at.read().await; match *last_updated { diff --git a/src/routes/identities.rs b/src/routes/identities.rs index 2fabaaf..83904f5 100644 --- a/src/routes/identities.rs +++ b/src/routes/identities.rs @@ -1,13 +1,14 @@ use crate::error::Result; use crate::models::{IdentityResponse, IdentityWithTraits}; use crate::routes::extractors::extract_environment_key; -use crate::state::AppState; +use crate::services::EnvironmentService; use axum::{ Json, extract::{Query, State}, http::HeaderMap, }; use serde::Deserialize; +use std::sync::Arc; #[derive(Deserialize)] pub struct IdentitiesQuery { @@ -15,7 +16,7 @@ pub struct IdentitiesQuery { } pub async fn get_identities( - State(service): State, + State(service): State>, headers: HeaderMap, Query(query): Query, ) -> Result> { @@ -30,7 +31,7 @@ pub async fn get_identities( } pub async fn post_identities( - State(service): State, + State(service): State>, headers: HeaderMap, Json(identity): Json, ) -> Result> { diff --git a/src/routes/mod.rs b/src/routes/mod.rs index fc419b0..8bcda0d 100644 --- a/src/routes/mod.rs +++ b/src/routes/mod.rs @@ -1,4 +1,3 @@ -pub mod cors; pub mod environment_document; pub mod extractors; pub mod flags; @@ -6,18 +5,27 @@ pub mod health; pub mod identities; use crate::config::AppSettings; -use crate::services::EnvironmentService; +use crate::middleware::{cors, usage}; +use crate::services::{EnvironmentService, UsageProcessor}; +use crate::state::AppState; use axum::{ Router, - middleware::map_response, + middleware::{from_fn_with_state, map_response}, routing::{get, post}, }; use std::sync::Arc; use tower_http::{compression::CompressionLayer, normalize_path::NormalizePath, trace::TraceLayer}; -pub fn create_router(settings: AppSettings) -> (Router, Arc) { +pub(crate) const FLAGS_PATH: &str = "/api/v1/flags"; +pub(crate) const IDENTITIES_PATH: &str = "/api/v1/identities"; +pub(crate) const ENVIRONMENT_DOCUMENT_PATH: &str = "/api/v1/environment-document"; + +pub fn create_router(settings: AppSettings) -> (Router, AppState) { let cors = cors::layer(&settings.allow_origins); - let environment_service = Arc::new(EnvironmentService::new(settings)); + let state = AppState { + usage: Arc::new(UsageProcessor::new(&settings)), + environments: Arc::new(EnvironmentService::new(settings)), + }; let router = Router::new() // Health check routes @@ -26,25 +34,26 @@ pub fn create_router(settings: AppSettings) -> (Router, Arc) .route("/proxy/health/readiness", get(health::health_check)) .route("/proxy/health/liveness", get(health::liveness_check)) // Flags routes (with and without trailing slash) - .route("/api/v1/flags", get(flags::get_flags)) + .route(FLAGS_PATH, get(flags::get_flags)) // Identities routes (with and without trailing slash) - .route("/api/v1/identities", get(identities::get_identities)) - .route("/api/v1/identities", post(identities::post_identities)) + .route(IDENTITIES_PATH, get(identities::get_identities)) + .route(IDENTITIES_PATH, post(identities::post_identities)) // Environment document route .route( - "/api/v1/environment-document", + ENVIRONMENT_DOCUMENT_PATH, get(environment_document::get_environment_document), ) // Middleware layers + .layer(from_fn_with_state(state.clone(), usage::track_usage)) .layer(CompressionLayer::new()) .layer(cors) .layer(map_response(cors::merge_vary)) .layer(TraceLayer::new_for_http()) - .with_state(environment_service.clone()); + .with_state(state.clone()); // Trailing-slash normalization must wrap the router itself: axum matches // routes before `Router::layer` middleware runs let app = Router::new().fallback_service(NormalizePath::trim_trailing_slash(router)); - (app, environment_service) + (app, state) } diff --git a/src/services/environment.rs b/src/services/environment.rs index 423a89b..a742fb4 100644 --- a/src/services/environment.rs +++ b/src/services/environment.rs @@ -30,6 +30,9 @@ impl EnvironmentService { let client = Client::builder() .timeout(Duration::from_secs(settings.api_poll_timeout_seconds)) .gzip(true) + // reqwest keeps custom headers across redirects; never let one + // carry the environment or proxy key to another host. + .redirect(reqwest::redirect::Policy::none()) .build() .expect("Failed to create HTTP client"); @@ -153,6 +156,10 @@ impl EnvironmentService { Ok(response.json().await?) } + pub fn resolve(&self, environment_key: &str) -> Option> { + self.environments.resolve(environment_key) + } + fn resolve_key(&self, environment_key: &str) -> Result> { self.environments .resolve(environment_key) @@ -210,6 +217,11 @@ impl EnvironmentService { .client .get(&next_url) .header("X-Environment-Key", server_side_key); + // Core excludes marked fetches from API usage; the proxy reports + // served requests instead. + if let Some(proxy_key) = &self.settings.proxy_key { + request = request.header("X-Proxy-Key", proxy_key); + } // If-Modified-Since is meaningful only on the first request; the // upstream pagination cursor (page_id) drives subsequent fetches. if document.is_none() { @@ -226,7 +238,7 @@ impl EnvironmentService { response.error_for_status_ref()?; - let next_link = parse_next_link(response.headers(), &self.settings.api_url); + let next_link = parse_next_link(response.headers(), &self.settings.api_url)?; let body: serde_json::Value = response.json().await?; match document.as_mut() { @@ -439,12 +451,15 @@ fn merge_paginated_overrides(base: &mut serde_json::Value, page: serde_json::Val /// Parse the `Link` response header for the next-page URL (RFC 5988). /// -/// Returns an absolute URL, resolving relative targets against `api_url`. -fn parse_next_link(headers: &HeaderMap, api_url: &str) -> Option { - let base = Url::parse(api_url).ok()?; +fn parse_next_link(headers: &HeaderMap, api_url: &str) -> Result> { + let Ok(base) = Url::parse(api_url) else { + return Ok(None); + }; for header_value in headers.get_all(reqwest::header::LINK).iter() { - let raw = header_value.to_str().ok()?; + let Ok(raw) = header_value.to_str() else { + return Ok(None); + }; for segment in raw.split(',') { let segment = segment.trim(); let target = match (segment.find('<'), segment.find('>')) { @@ -455,12 +470,18 @@ fn parse_next_link(headers: &HeaderMap, api_url: &str) -> Option { if !is_next_rel(params) { continue; } - if let Ok(absolute) = base.join(target) { - return Some(absolute.into()); + let Ok(absolute) = base.join(target) else { + continue; + }; + if absolute.origin() != base.origin() { + return Err(EdgeProxyError::ServiceUnavailable(format!( + "refusing cross-origin pagination link {absolute}" + ))); } + return Ok(Some(absolute.into())); } } - None + Ok(None) } fn is_next_rel(params: &str) -> bool { @@ -498,40 +519,58 @@ mod tests { ); let next = parse_next_link(&headers, "https://edge.api.flagsmith.com/api/v1").unwrap(); assert_eq!( - next, + next.unwrap(), "https://edge.api.flagsmith.com/api/v1/environment-document/?page_id=identity_override%3A1%3Aabc" ); } #[test] - fn parse_next_link_absolute() { - let headers = - link("; rel=\"next\""); + fn parse_next_link_absolute_same_origin() { + let headers = link( + "; rel=\"next\"", + ); let next = parse_next_link(&headers, "https://edge.api.flagsmith.com/api/v1").unwrap(); assert_eq!( - next, - "https://example.test/api/v1/environment-document/?page_id=x" + next.unwrap(), + "https://edge.api.flagsmith.com/api/v1/environment-document/?page_id=x" ); } + #[test] + fn parse_next_link_rejects_cross_origin() { + let headers = + link("; rel=\"next\""); + assert!(parse_next_link(&headers, "https://edge.api.flagsmith.com/api/v1").is_err()); + } + #[test] fn parse_next_link_picks_next_among_multiple_rels() { let headers = link("; rel=\"prev\", ; rel=\"next\""); let next = parse_next_link(&headers, "https://edge.api.flagsmith.com/api/v1").unwrap(); - assert_eq!(next, "https://edge.api.flagsmith.com/api/v1/page/next"); + assert_eq!( + next.unwrap(), + "https://edge.api.flagsmith.com/api/v1/page/next" + ); } #[test] fn parse_next_link_returns_none_when_only_other_rels() { let headers = link("; rel=\"prev\", ; rel=\"self\""); - assert!(parse_next_link(&headers, "https://edge.api.flagsmith.com/api/v1").is_none()); + assert!( + parse_next_link(&headers, "https://edge.api.flagsmith.com/api/v1") + .unwrap() + .is_none() + ); } #[test] fn parse_next_link_handles_unquoted_rel() { let headers = link("; rel=next"); let next = parse_next_link(&headers, "https://edge.api.flagsmith.com/api/v1").unwrap(); - assert_eq!(next, "https://edge.api.flagsmith.com/api/v1/page/next"); + assert_eq!( + next.unwrap(), + "https://edge.api.flagsmith.com/api/v1/page/next" + ); } #[test] diff --git a/src/services/mod.rs b/src/services/mod.rs index efaf4e3..f9517e1 100644 --- a/src/services/mod.rs +++ b/src/services/mod.rs @@ -1,4 +1,6 @@ pub mod environment; pub mod feature_utils; +pub mod usage; pub use environment::EnvironmentService; +pub use usage::UsageProcessor; diff --git a/src/services/usage.rs b/src/services/usage.rs new file mode 100644 index 0000000..cab67e5 --- /dev/null +++ b/src/services/usage.rs @@ -0,0 +1,131 @@ +use crate::config::settings::AppSettings; +use crate::usage::{Resource, UsageBatch, UsageCounts}; +use parking_lot::Mutex; +use reqwest::Client; +use std::sync::Arc; +use std::time::Duration; +use tracing::error; + +enum Outcome { + Accepted, + Rejected, + Failed, +} + +pub struct UsageProcessor { + counts: UsageCounts, + pending_batches: Mutex>, + client: Client, + api_url: String, + proxy_key: Option, + flush_interval: Duration, +} + +impl UsageProcessor { + /// Mirrors MAX_USAGE_ROWS on the usage endpoint. + const MAX_ROWS_PER_FLUSH: usize = 1000; + + pub fn new(settings: &AppSettings) -> Self { + let client = Client::builder() + .timeout(Duration::from_secs(settings.api_poll_timeout_seconds)) + // reqwest keeps custom headers across redirects; never let one + // carry the proxy key to another host. + .redirect(reqwest::redirect::Policy::none()) + .build() + .expect("Failed to create HTTP client"); + Self { + counts: UsageCounts::default(), + pending_batches: Mutex::default(), + client, + api_url: settings.api_url.clone(), + proxy_key: settings.proxy_key.clone(), + flush_interval: Duration::from_secs(settings.usage_flush_interval_seconds), + } + } + + pub fn track(&self, client_key: &str, resource: Resource) { + if self.proxy_key.is_none() { + return; + } + self.counts.increment(client_key, resource); + } + + /// Returns false when any batch was not accepted. + pub async fn flush(&self) -> bool { + let Some(proxy_key) = &self.proxy_key else { + return true; + }; + + let mut batches = std::mem::take(&mut *self.pending_batches.lock()); + if batches.is_empty() { + let mut rows = self.counts.drain(); + while !rows.is_empty() { + let chunk = rows.drain(..rows.len().min(Self::MAX_ROWS_PER_FLUSH)); + batches.push(UsageBatch::new(chunk.collect())); + } + } + + let mut all_success = true; + for batch in batches { + match self.post(proxy_key, &batch).await { + Outcome::Accepted => {} + Outcome::Rejected => all_success = false, + Outcome::Failed => { + self.pending_batches.lock().push(batch); + all_success = false; + } + } + } + all_success + } + + async fn post(&self, proxy_key: &str, batch: &UsageBatch) -> Outcome { + let url = format!("{}/proxy/usage/", self.api_url); + let result = self + .client + .post(&url) + .header("X-Proxy-Key", proxy_key) + .header("Idempotency-Key", &batch.id) + .json(&batch.rows) + .send() + .await; + match result { + Ok(response) if response.status().is_success() => Outcome::Accepted, + // Retrying cannot heal a rejection: drop the batch rather than + // resend it forever. + Ok(response) if response.status().is_client_error() => { + error!( + "Usage report rejected with {}: dropping {} rows", + response.status(), + batch.rows.len() + ); + Outcome::Rejected + } + Ok(response) => { + error!("Failed to report usage: {}", response.status()); + Outcome::Failed + } + Err(e) => { + error!("Failed to report usage: {}", e); + Outcome::Failed + } + } + } + + pub async fn flush_periodically(self: Arc) { + if self.proxy_key.is_none() { + return; + } + let mut interval = tokio::time::interval(self.flush_interval); + // A flush that overruns the interval (core slow or down) must not be + // followed by a burst of catch-up flushes hammering it further. + interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay); + // The first tick completes immediately, before anything is counted. + interval.tick().await; + + loop { + interval.tick().await; + self.flush().await; + } + } +} diff --git a/src/state.rs b/src/state.rs index 43ae606..c946f9c 100644 --- a/src/state.rs +++ b/src/state.rs @@ -1,4 +1,9 @@ -use crate::services::EnvironmentService; +use crate::services::{EnvironmentService, UsageProcessor}; +use axum::extract::FromRef; use std::sync::Arc; -pub type AppState = Arc; +#[derive(Clone, FromRef)] +pub struct AppState { + pub environments: Arc, + pub usage: Arc, +} diff --git a/src/usage.rs b/src/usage.rs new file mode 100644 index 0000000..3ca7642 --- /dev/null +++ b/src/usage.rs @@ -0,0 +1,136 @@ +use std::collections::HashMap; + +use parking_lot::Mutex; +use serde::Serialize; + +#[derive(Serialize, Clone, Copy, PartialEq, Eq, Hash, Debug)] +#[serde(rename_all = "kebab-case")] +pub enum Resource { + Flags, + Identities, + EnvironmentDocument, +} + +/// One row of the `POST /proxy/usage/` body. +#[derive(Serialize, Debug, PartialEq)] +pub struct UsageRow { + pub client_side_key: String, + pub resource: Resource, + pub count: u64, +} + +#[derive(Default)] +pub struct UsageCounts { + count_by_environment_and_resource: Mutex>, +} + +impl UsageCounts { + pub fn increment(&self, client_key: &str, resource: Resource) { + let mut count_by_environment_and_resource = self.count_by_environment_and_resource.lock(); + *count_by_environment_and_resource + .entry((client_key.to_string(), resource)) + .or_default() += 1; + } + + /// Take everything counted so far, leaving the map empty. + pub fn drain(&self) -> Vec { + let count_by_environment_and_resource = + std::mem::take(&mut *self.count_by_environment_and_resource.lock()); + count_by_environment_and_resource + .into_iter() + .map(|((client_side_key, resource), count)| UsageRow { + client_side_key, + resource, + count, + }) + .collect() + } +} + +pub struct UsageBatch { + pub id: String, + pub rows: Vec, +} + +impl UsageBatch { + pub fn new(rows: Vec) -> Self { + Self { + id: uuid::Uuid::new_v4().to_string(), + rows, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn increment_aggregates_by_key_and_resource() { + // Given + let counts = UsageCounts::default(); + + // When + counts.increment("client_a", Resource::Flags); + counts.increment("client_a", Resource::Flags); + counts.increment("client_a", Resource::Identities); + counts.increment("client_b", Resource::Flags); + + // Then + let mut rows = counts.drain(); + rows.sort_by_key(|row| (row.client_side_key.clone(), format!("{:?}", row.resource))); + assert_eq!( + rows, + vec![ + UsageRow { + client_side_key: "client_a".to_string(), + resource: Resource::Flags, + count: 2, + }, + UsageRow { + client_side_key: "client_a".to_string(), + resource: Resource::Identities, + count: 1, + }, + UsageRow { + client_side_key: "client_b".to_string(), + resource: Resource::Flags, + count: 1, + }, + ] + ); + } + + #[test] + fn drain_leaves_nothing_behind() { + // Given + let counts = UsageCounts::default(); + counts.increment("client", Resource::Flags); + + // When + counts.drain(); + + // Then + assert!(counts.drain().is_empty()); + } + + #[test] + fn usage_row_serializes_to_the_contract_shape() { + // Given + let row = UsageRow { + client_side_key: "client".to_string(), + resource: Resource::EnvironmentDocument, + count: 3, + }; + + // Then + assert_eq!( + serde_json::to_value(&row).unwrap(), + serde_json::json!({ + "client_side_key": "client", + "resource": "environment-document", + "count": 3, + }) + ); + } +} diff --git a/tests/test_cors.rs b/tests/test_cors.rs index c34414c..0474c4f 100644 --- a/tests/test_cors.rs +++ b/tests/test_cors.rs @@ -17,8 +17,9 @@ async fn server_allowing(allow_origins: &[&str]) -> TestServer { allow_origins: allow_origins.iter().map(|o| o.to_string()).collect(), ..AppSettings::new() }; - let (app, service) = create_router(settings); - service + let (app, state) = create_router(settings); + state + .environments .cache .put_environment(&environment_1_api_key(), environment_1()) .await; diff --git a/tests/test_health.rs b/tests/test_health.rs index 8016c49..ef99212 100644 --- a/tests/test_health.rs +++ b/tests/test_health.rs @@ -21,9 +21,9 @@ async fn test_liveness_check() { async fn test_health_check_returns_200_if_cache_was_updated_recently() { // Given let settings = AppSettings::new(); - let (app, service) = create_router(settings); + let (app, state) = create_router(settings); { - let mut last_updated = service.last_updated_at.write().await; + let mut last_updated = state.environments.last_updated_at.write().await; *last_updated = Some(Utc::now()); } let server = TestServer::new(app).unwrap(); @@ -62,9 +62,9 @@ async fn test_health_check_returns_503_if_cache_was_not_updated() { async fn test_health_check_returns_503_if_cache_is_stale() { // Given let settings = AppSettings::new(); - let (app, service) = create_router(settings); + let (app, state) = create_router(settings); { - let mut last_updated = service.last_updated_at.write().await; + let mut last_updated = state.environments.last_updated_at.write().await; *last_updated = Some(Utc::now() - chrono::Duration::days(10)); } let server = TestServer::new(app).unwrap(); @@ -88,10 +88,10 @@ async fn test_health_check_returns_200_if_cache_is_never_stale() { settings.health_check = HealthCheckSettings { environment_update_grace_period_seconds: None, }; - let (app, service) = create_router(settings); + let (app, state) = create_router(settings); let last_update_time = Utc::now() - chrono::Duration::days(10); { - let mut last_updated = service.last_updated_at.write().await; + let mut last_updated = state.environments.last_updated_at.write().await; *last_updated = Some(last_update_time); } let server = TestServer::new(app).unwrap(); diff --git a/tests/test_server.rs b/tests/test_server.rs index 2a599af..2ddf01a 100644 --- a/tests/test_server.rs +++ b/tests/test_server.rs @@ -16,9 +16,10 @@ async fn setup_test_server() -> TestServer { ..AppSettings::new() }; - let (app, service) = create_router(settings); + let (app, state) = create_router(settings); - service + state + .environments .cache .put_environment(&environment_1_api_key(), environment_1()) .await; @@ -208,8 +209,9 @@ async fn test_get_environment_document() { }], ..AppSettings::new() }; - let (app, service) = create_router(settings); - service + let (app, state) = create_router(settings); + state + .environments .cache .put_environment("client_key", environment_1()) .await; diff --git a/tests/test_usage_tracking.rs b/tests/test_usage_tracking.rs new file mode 100644 index 0000000..c43ec85 --- /dev/null +++ b/tests/test_usage_tracking.rs @@ -0,0 +1,529 @@ +use axum::http::StatusCode; +use axum_test::TestServer; +use edge_proxy::config::settings::{AppSettings, EnvironmentKeyPair}; +use edge_proxy::routes::create_router; +use edge_proxy::services::{EnvironmentService, UsageProcessor}; +use edge_proxy::usage::Resource; +use serde_json::{Value, json}; +use wiremock::matchers::{header, method, path}; +use wiremock::{Mock, MockServer, Request, ResponseTemplate}; + +const PROXY_KEY: &str = "pk.test_proxy_key"; +const CLIENT_KEY: &str = "config_client_key"; +const SERVER_KEY: &str = "ser.config_key"; + +fn settings(api_url: &str, proxy_key: Option<&str>, pairs: Vec) -> AppSettings { + AppSettings { + environment_key_pairs: pairs, + proxy_key: proxy_key.map(str::to_string), + api_url: api_url.to_string(), + ..AppSettings::default() + } +} + +fn config_body() -> Value { + json!([{ + "id": 30, + "name": "Test Environment", + "client_side_key": CLIENT_KEY, + "server_side_keys": [ + {"key": SERVER_KEY, "active": true, "expires_at": null} + ], + "updated_at": "2026-08-15T08:57:43.311081Z", + "project_id": 35, + "organisation_id": 82, + }]) +} + +fn document_body() -> Value { + json!({ + "id": 1, + "api_key": CLIENT_KEY, + "name": "Test", + "updated_at": "2026-08-22T00:00:00Z", + "allow_client_traits": true, + "hide_sensitive_data": false, + "hide_disabled_flags": null, + "use_identity_composite_key_for_hashing": true, + "use_identity_overrides_in_local_eval": true, + "project": { + "id": 1, + "name": "project-1", + "hide_disabled_flags": false, + "segments": [], + "server_key_only_feature_ids": [], + "organisation": { + "id": 1, + "name": "org-1", + "feature_analytics": false, + "persist_trait_data": true, + "stop_serving_flags": false, + }, + }, + "feature_states": [ + { + "multivariate_feature_state_values": [], + "feature_state_value": "config_value", + "feature": {"id": 1, "name": "config_flag", "type": "STANDARD"}, + "enabled": true, + "featurestate_uuid": "fs-uuid-1", + } + ], + "identity_overrides": [], + }) +} + +async fn mount_config(mock_server: &MockServer) { + Mock::given(method("GET")) + .and(path("/proxy/config/")) + .and(header("X-Proxy-Key", PROXY_KEY)) + .respond_with(ResponseTemplate::new(200).set_body_json(config_body())) + .mount(mock_server) + .await; +} + +async fn mount_document(mock_server: &MockServer) { + Mock::given(method("GET")) + .and(path("/environment-document/")) + .respond_with(ResponseTemplate::new(200).set_body_json(document_body())) + .mount(mock_server) + .await; +} + +async fn mount_usage(mock_server: &MockServer, status: u16, up_to: Option) { + let mut mock = Mock::given(method("POST")) + .and(path("/proxy/usage/")) + .and(header("X-Proxy-Key", PROXY_KEY)) + .respond_with(ResponseTemplate::new(status)); + if let Some(n) = up_to { + mock = mock.up_to_n_times(n); + } + mock.mount(mock_server).await; +} + +async fn requests_to(mock_server: &MockServer, url_path: &str) -> Vec { + mock_server + .received_requests() + .await + .unwrap() + .into_iter() + .filter(|request| request.url.path() == url_path) + .collect() +} + +/// The usage rows of the request body, sorted by resource for +/// order-independent assertions. +fn usage_rows(request: &Request) -> Vec { + let mut rows: Vec = serde_json::from_slice(&request.body).unwrap(); + rows.sort_by_key(|row| row["resource"].as_str().unwrap().to_string()); + rows +} + +#[tokio::test] +async fn test_served_requests_flush_aggregated_usage() { + // Given a discovered environment served through the full router + let mock_server = MockServer::start().await; + mount_config(&mock_server).await; + mount_document(&mock_server).await; + mount_usage(&mock_server, 204, None).await; + let (app, state) = create_router(settings(&mock_server.uri(), Some(PROXY_KEY), vec![])); + state.environments.refresh_environment_caches().await; + let server = TestServer::new(app).unwrap(); + + // When SDK traffic arrives under both keys, then usage is flushed + server + .get("/api/v1/flags") + .add_header("X-Environment-Key", CLIENT_KEY) + .await + .assert_status_ok(); + server + .get("/api/v1/flags") + .add_header("X-Environment-Key", CLIENT_KEY) + .await + .assert_status_ok(); + server + .get("/api/v1/identities") + .add_query_param("identifier", "user_1") + .add_header("X-Environment-Key", CLIENT_KEY) + .await + .assert_status_ok(); + server + .post("/api/v1/identities") + .json(&json!({"identifier": "user_2"})) + .add_header("X-Environment-Key", CLIENT_KEY) + .await + .assert_status_ok(); + server + .get("/api/v1/environment-document") + .add_header("X-Environment-Key", SERVER_KEY) + .await + .assert_status_ok(); + let flushed = state.usage.flush().await; + + // Then one POST reports everything, keyed by the client key even for + // requests that presented the server key, under an idempotency key + assert!(flushed); + let posts = requests_to(&mock_server, "/proxy/usage/").await; + assert_eq!(posts.len(), 1); + assert!(!posts[0].headers["Idempotency-Key"].is_empty()); + assert_eq!( + usage_rows(&posts[0]), + vec![ + json!({"client_side_key": CLIENT_KEY, "resource": "environment-document", "count": 1}), + json!({"client_side_key": CLIENT_KEY, "resource": "flags", "count": 2}), + json!({"client_side_key": CLIENT_KEY, "resource": "identities", "count": 2}), + ] + ); +} + +#[tokio::test] +async fn test_unresolved_keys_are_never_counted() { + // Given a proxy serving one environment + let mock_server = MockServer::start().await; + mount_config(&mock_server).await; + mount_document(&mock_server).await; + let (app, state) = create_router(settings(&mock_server.uri(), Some(PROXY_KEY), vec![])); + state.environments.refresh_environment_caches().await; + let server = TestServer::new(app).unwrap(); + + // When requests present an unknown key or none at all + server + .get("/api/v1/flags") + .add_header("X-Environment-Key", "unknown_key") + .await + .assert_status_unauthorized(); + server + .get("/api/v1/flags") + .await + .assert_status_unauthorized(); + + // Then there is nothing to flush and no request is made + assert!(state.usage.flush().await); + assert!(requests_to(&mock_server, "/proxy/usage/").await.is_empty()); +} + +#[tokio::test] +async fn test_failed_requests_are_not_counted() { + // Given one environment that serves and one whose document never loaded + let mock_server = MockServer::start().await; + let mut environments = config_body(); + environments.as_array_mut().unwrap().push(json!({ + "id": 31, + "name": "Broken Environment", + "client_side_key": "broken_client", + "server_side_keys": [ + {"key": "ser.broken_key", "active": true, "expires_at": null} + ], + "updated_at": "2026-08-15T08:57:43.311081Z", + "project_id": 35, + "organisation_id": 82, + })); + Mock::given(method("GET")) + .and(path("/proxy/config/")) + .respond_with(ResponseTemplate::new(200).set_body_json(environments)) + .mount(&mock_server) + .await; + Mock::given(method("GET")) + .and(path("/environment-document/")) + .and(header("X-Environment-Key", "ser.broken_key")) + .respond_with(ResponseTemplate::new(500)) + .mount(&mock_server) + .await; + mount_document(&mock_server).await; + let (app, state) = create_router(settings(&mock_server.uri(), Some(PROXY_KEY), vec![])); + state.environments.refresh_environment_caches().await; + let server = TestServer::new(app).unwrap(); + + // When requests fail after their key resolved + server + .get("/api/v1/flags") + .add_query_param("feature", "missing") + .add_header("X-Environment-Key", CLIENT_KEY) + .await + .assert_status(StatusCode::NOT_FOUND); + server + .get("/api/v1/flags") + .add_header("X-Environment-Key", "broken_client") + .await + .assert_status(StatusCode::SERVICE_UNAVAILABLE); + + // Then nothing is reported + assert!(state.usage.flush().await); + assert!(requests_to(&mock_server, "/proxy/usage/").await.is_empty()); +} + +#[tokio::test] +async fn test_failed_flush_retries_the_same_batch_before_new_counts() { + // Given a served request and a usage endpoint that fails once + let mock_server = MockServer::start().await; + mount_config(&mock_server).await; + mount_document(&mock_server).await; + mount_usage(&mock_server, 500, Some(1)).await; + mount_usage(&mock_server, 204, None).await; + let config = settings(&mock_server.uri(), Some(PROXY_KEY), vec![]); + let service = EnvironmentService::new(config.clone()); + let usage = UsageProcessor::new(&config); + service.refresh_environment_caches().await; + usage.track(CLIENT_KEY, Resource::Flags); + + // When the first flush fails, another request is served, and two + // more flushes run + assert!(!usage.flush().await); + usage.track(CLIENT_KEY, Resource::Flags); + assert!(usage.flush().await); + assert!(usage.flush().await); + + // Then the failed batch is resent unchanged under its original key, + // and the count served meanwhile follows as its own batch + let posts = requests_to(&mock_server, "/proxy/usage/").await; + assert_eq!(posts.len(), 3); + let keys: Vec<&str> = posts + .iter() + .map(|post| post.headers["Idempotency-Key"].to_str().unwrap()) + .collect(); + assert_eq!(keys[0], keys[1]); + assert_ne!(keys[1], keys[2]); + let one_flag = vec![json!({"client_side_key": CLIENT_KEY, "resource": "flags", "count": 1})]; + assert_eq!(usage_rows(&posts[1]), one_flag); + assert_eq!(usage_rows(&posts[2]), one_flag); +} + +#[tokio::test] +async fn test_flush_without_proxy_key_is_inert() { + // Given a statically configured proxy with no proxy key + let mock_server = MockServer::start().await; + mount_document(&mock_server).await; + let config = settings( + &mock_server.uri(), + None, + vec![EnvironmentKeyPair { + client_side_key: CLIENT_KEY.to_string(), + server_side_key: SERVER_KEY.to_string(), + }], + ); + let service = EnvironmentService::new(config.clone()); + let usage = UsageProcessor::new(&config); + service.refresh_environment_caches().await; + usage.track(CLIENT_KEY, Resource::Flags); + + // When / Then: flushing succeeds without reporting anything + assert!(usage.flush().await); + assert!(requests_to(&mock_server, "/proxy/usage/").await.is_empty()); +} + +#[tokio::test] +async fn test_rejected_flush_drops_the_batch_instead_of_retrying_it() { + // Given a served request and a usage endpoint that rejects the batch + let mock_server = MockServer::start().await; + mount_config(&mock_server).await; + mount_document(&mock_server).await; + mount_usage(&mock_server, 400, None).await; + let config = settings(&mock_server.uri(), Some(PROXY_KEY), vec![]); + let service = EnvironmentService::new(config.clone()); + let usage = UsageProcessor::new(&config); + service.refresh_environment_caches().await; + usage.track(CLIENT_KEY, Resource::Flags); + + // When the flush is rejected + assert!(!usage.flush().await); + + // Then the rows are dropped, not resent forever: the next flush has + // nothing to send + assert!(usage.flush().await); + assert_eq!(requests_to(&mock_server, "/proxy/usage/").await.len(), 1); +} + +#[tokio::test] +async fn test_flush_chunks_batches_to_the_server_cap() { + // Given served requests for more environments than one batch may hold + let environments: Vec = (0..1001) + .map(|n| { + json!({ + "id": n, + "name": format!("env {n}"), + "client_side_key": format!("client_{n}"), + "server_side_keys": [ + {"key": format!("ser.key_{n}"), "active": true, "expires_at": null} + ], + "updated_at": "2026-08-15T08:57:43.311081Z", + "project_id": 35, + "organisation_id": 82, + }) + }) + .collect(); + let mock_server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/proxy/config/")) + .respond_with(ResponseTemplate::new(200).set_body_json(Value::Array(environments))) + .mount(&mock_server) + .await; + mount_document(&mock_server).await; + mount_usage(&mock_server, 204, None).await; + let config = settings(&mock_server.uri(), Some(PROXY_KEY), vec![]); + let service = EnvironmentService::new(config.clone()); + let usage = UsageProcessor::new(&config); + service.refresh_environment_caches().await; + for n in 0..1001 { + usage.track(&format!("client_{n}"), Resource::Flags); + } + + // When + assert!(usage.flush().await); + + // Then the rows arrive split across two accepted requests, each its + // own batch + let posts = requests_to(&mock_server, "/proxy/usage/").await; + let row_counts: Vec = posts.iter().map(|post| usage_rows(post).len()).collect(); + assert_eq!(row_counts, vec![1000, 1]); + assert_ne!( + posts[0].headers["Idempotency-Key"], + posts[1].headers["Idempotency-Key"] + ); +} + +#[tokio::test] +async fn test_static_environment_with_proxy_key_is_billed_like_a_discovered_one() { + // Given a proxy key, and a static environment alongside a discovered one + let mock_server = MockServer::start().await; + mount_config(&mock_server).await; + mount_document(&mock_server).await; + mount_usage(&mock_server, 204, None).await; + let config = settings( + &mock_server.uri(), + Some(PROXY_KEY), + vec![EnvironmentKeyPair { + client_side_key: "static_client".to_string(), + server_side_key: "ser.static_key".to_string(), + }], + ); + let service = EnvironmentService::new(config.clone()); + let usage = UsageProcessor::new(&config); + service.refresh_environment_caches().await; + + // When both environments serve a request, and usage is flushed + usage.track("static_client", Resource::Flags); + usage.track(CLIENT_KEY, Resource::Flags); + assert!(usage.flush().await); + + // Then every document fetch is marked and both environments are reported + let fetches = requests_to(&mock_server, "/environment-document/").await; + assert!(!fetches.is_empty()); + assert!( + fetches + .iter() + .all(|request| request.headers["X-Proxy-Key"] == PROXY_KEY) + ); + let posts = requests_to(&mock_server, "/proxy/usage/").await; + assert_eq!(posts.len(), 1); + let mut rows = usage_rows(&posts[0]); + rows.sort_by_key(|row| row["client_side_key"].as_str().unwrap().to_string()); + assert_eq!( + rows, + vec![ + json!({"client_side_key": CLIENT_KEY, "resource": "flags", "count": 1}), + json!({"client_side_key": "static_client", "resource": "flags", "count": 1}), + ] + ); +} + +#[tokio::test] +async fn test_document_fetch_carries_the_proxy_key() { + // Given a document endpoint that only answers requests marked as the + // proxy's own + let mock_server = MockServer::start().await; + mount_config(&mock_server).await; + Mock::given(method("GET")) + .and(path("/environment-document/")) + .and(header("X-Proxy-Key", PROXY_KEY)) + .respond_with(ResponseTemplate::new(200).set_body_json(document_body())) + .mount(&mock_server) + .await; + let service = EnvironmentService::new(settings(&mock_server.uri(), Some(PROXY_KEY), vec![])); + + // When / Then + assert!(service.refresh_environment_caches().await); + assert!(service.get_environment(CLIENT_KEY).await.is_ok()); +} + +#[tokio::test] +async fn test_document_fetch_without_proxy_key_is_unmarked() { + // Given a statically configured proxy with no proxy key + let mock_server = MockServer::start().await; + mount_document(&mock_server).await; + let service = EnvironmentService::new(settings( + &mock_server.uri(), + None, + vec![EnvironmentKeyPair { + client_side_key: CLIENT_KEY.to_string(), + server_side_key: SERVER_KEY.to_string(), + }], + )); + + // When + assert!(service.refresh_environment_caches().await); + + // Then the document fetch is byte-identical to today's + let fetches = requests_to(&mock_server, "/environment-document/").await; + assert!(!fetches.is_empty()); + assert!( + fetches + .iter() + .all(|request| !request.headers.contains_key("X-Proxy-Key")) + ); +} + +#[tokio::test] +async fn test_document_fetch_does_not_follow_redirects() { + // Given a document endpoint that redirects to another host + let mock_server = MockServer::start().await; + let other_host = MockServer::start().await; + mount_config(&mock_server).await; + Mock::given(method("GET")) + .and(path("/environment-document/")) + .respond_with(ResponseTemplate::new(302).insert_header( + "Location", + format!("{}/environment-document/", other_host.uri()), + )) + .mount(&mock_server) + .await; + Mock::given(method("GET")) + .respond_with(ResponseTemplate::new(200).set_body_json(document_body())) + .mount(&other_host) + .await; + let service = EnvironmentService::new(settings(&mock_server.uri(), Some(PROXY_KEY), vec![])); + + // When + let refreshed = service.refresh_environment_caches().await; + + // Then the fetch fails and the keys never reach the other host + assert!(!refreshed); + assert!(other_host.received_requests().await.unwrap().is_empty()); +} + +#[tokio::test] +async fn test_usage_flush_does_not_follow_redirects() { + // Given a usage endpoint that redirects to another host + let mock_server = MockServer::start().await; + let other_host = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/proxy/usage/")) + .respond_with( + ResponseTemplate::new(307) + .insert_header("Location", format!("{}/proxy/usage/", other_host.uri())), + ) + .mount(&mock_server) + .await; + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(204)) + .mount(&other_host) + .await; + let usage = UsageProcessor::new(&settings(&mock_server.uri(), Some(PROXY_KEY), vec![])); + usage.track(CLIENT_KEY, Resource::Flags); + + // When + let flushed = usage.flush().await; + + // Then the batch is not accepted and the key never reaches the other host + assert!(!flushed); + assert!(other_host.received_requests().await.unwrap().is_empty()); +}