diff --git a/src/api/config_edit.rs b/src/api/config_edit.rs index 7b087c6..a431072 100644 --- a/src/api/config_edit.rs +++ b/src/api/config_edit.rs @@ -9,8 +9,11 @@ use super::ApiShared; use super::config_store::{ EDITABLE_SECTIONS, EDITABLE_SERVER_FIELDS, compute_snapshot_revision, is_editable_section, load_candidate_snapshot, load_config_snapshot, render_server_listeners, - render_top_level_section, resolve_single_source_owner, upsert_toml_table, write_atomic, + render_top_level_section, resolve_single_source_owner, upsert_toml_table, + write_atomic_if_unchanged, }; +#[cfg(test)] +use super::config_store::write_atomic; use super::model::ApiFailure; use crate::config::ProxyConfig; use crate::config::hot_reload::classify_config_changes; @@ -43,7 +46,10 @@ pub(super) struct PatchConfigResponse { } struct PreparedConfigPatch { + config_path: PathBuf, + expected_revision: String, owner_path: PathBuf, + expected_owner_contents: String, owner_contents: String, desired_config: Arc, response: PatchConfigResponse, @@ -77,7 +83,14 @@ pub(super) async fn patch_config( } else { None }; - write_atomic(prepared.owner_path, prepared.owner_contents).await?; + write_atomic_if_unchanged( + prepared.config_path, + prepared.expected_revision, + prepared.owner_path, + prepared.expected_owner_contents, + prepared.owner_contents, + ) + .await?; if let Some(reservation) = reservation { prepared.response.reload = Some(reservation.enqueue(prepared.desired_config)); } @@ -111,7 +124,14 @@ pub(super) async fn apply_patch_to_path( expected_revision: Option, ) -> Result { let prepared = prepare_patch_to_path(config_path, patch_json, expected_revision).await?; - write_atomic(prepared.owner_path, prepared.owner_contents).await?; + write_atomic_if_unchanged( + prepared.config_path, + prepared.expected_revision, + prepared.owner_path, + prepared.expected_owner_contents, + prepared.owner_contents, + ) + .await?; Ok(prepared.response) } @@ -197,6 +217,7 @@ async fn prepare_patch_to_path( .get(&owner_path) .cloned() .ok_or_else(|| ApiFailure::internal("config source owner is missing from snapshot"))?; + let expected_owner_contents = owner_contents.clone(); for section in &touched { if *section == "server" { let rendered = render_server_listeners(&requested_cfg)?; @@ -233,7 +254,10 @@ async fn prepare_patch_to_path( deferred_process_fields(&old_cfg, &new_cfg).map_err(ApiFailure::bad_request)?; Ok(PreparedConfigPatch { + config_path: config_path.to_path_buf(), + expected_revision: current, owner_path, + expected_owner_contents, owner_contents, desired_config: Arc::new(new_cfg), response: PatchConfigResponse { diff --git a/src/api/config_edit/tests.rs b/src/api/config_edit/tests.rs index 6979f47..fb2f9a4 100644 --- a/src/api/config_edit/tests.rs +++ b/src/api/config_edit/tests.rs @@ -323,6 +323,30 @@ async fn patch_writes_the_included_section_owner_only() { ); } +#[tokio::test] +async fn prepared_patch_rejects_external_edit_before_commit() { + let (path, _directory) = temp_config("[censorship]\ntls_domain = \"old.example\"\n"); + let patch: Json = serde_json::json!({ + "censorship": {"tls_domain": "api.example"} + }); + let prepared = prepare_patch_to_path(&path, &patch, None).await.unwrap(); + let external = "[censorship]\ntls_domain = \"external.example\"\n"; + tokio::fs::write(&path, external).await.unwrap(); + + let error = write_atomic_if_unchanged( + prepared.config_path, + prepared.expected_revision, + prepared.owner_path, + prepared.expected_owner_contents, + prepared.owner_contents, + ) + .await + .unwrap_err(); + + assert_eq!(error.code, "revision_conflict"); + assert_eq!(tokio::fs::read_to_string(&path).await.unwrap(), external); +} + #[tokio::test] async fn patch_rejects_multiple_source_owners_without_writing() { let dir = tempfile::tempdir().unwrap(); diff --git a/src/api/config_store.rs b/src/api/config_store.rs index 7e97ec7..6d1a90d 100644 --- a/src/api/config_store.rs +++ b/src/api/config_store.rs @@ -10,13 +10,16 @@ use super::model::ApiFailure; // Source-preserving TOML rendering and atomic persistence helpers. mod persistence; +// Compare-and-replace file persistence and metadata preservation. +mod atomic; #[cfg(test)] use persistence::{find_toml_table_bounds, render_access_section, save_sections_to_disk}; pub(in crate::api) use persistence::{ render_server_listeners, render_top_level_section, save_access_sections_to_disk, - upsert_toml_table, write_atomic, + save_access_sections_to_disk_if_revision, upsert_toml_table, }; +pub(in crate::api) use atomic::{write_atomic, write_atomic_if_unchanged}; #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub(super) enum AccessSection { @@ -54,22 +57,21 @@ pub(super) fn parse_if_match(headers: &hyper::HeaderMap) -> Option { .map(|value| value.trim_matches('"').to_string()) } -pub(super) async fn ensure_expected_revision( +/// Loads one mutation base and validates its revision from the same source snapshot. +pub(super) async fn load_config_for_mutation( config_path: &Path, expected_revision: Option<&str>, -) -> Result<(), ApiFailure> { - let Some(expected) = expected_revision else { - return Ok(()); - }; - let current = current_revision(config_path).await?; - if current != expected { +) -> Result<(ProxyConfig, String), ApiFailure> { + let loaded = load_config_snapshot(config_path, false).await?; + let revision = compute_snapshot_revision(&loaded); + if expected_revision.is_some_and(|expected| expected != revision) { return Err(ApiFailure::new( hyper::StatusCode::CONFLICT, "revision_conflict", "Config revision mismatch", )); } - Ok(()) + Ok((loaded.config, revision)) } pub(super) async fn current_revision(config_path: &Path) -> Result { diff --git a/src/api/config_store/atomic.rs b/src/api/config_store/atomic.rs new file mode 100644 index 0000000..d13aa23 --- /dev/null +++ b/src/api/config_store/atomic.rs @@ -0,0 +1,185 @@ +use std::io::{Read, Write}; +use std::path::{Path, PathBuf}; + +#[cfg(unix)] +use std::os::unix::fs::{MetadataExt, OpenOptionsExt, PermissionsExt}; + +use super::compute_source_revision; +use crate::api::model::ApiFailure; +use crate::config::ProxyConfig; + +enum AtomicWriteError { + Conflict, + ReadGraph(String), + Io(std::io::Error), +} + +struct ExistingTarget { + contents: String, + metadata: std::fs::Metadata, +} + +/// Replaces one config source through a durable same-directory rename. +pub(in crate::api) async fn write_atomic( + path: PathBuf, + contents: String, +) -> Result<(), ApiFailure> { + tokio::task::spawn_blocking(move || write_atomic_sync(&path, None, &contents)) + .await + .map_err(|error| ApiFailure::internal(format!("failed to join writer: {error}")))? + .map_err(|error| ApiFailure::internal(format!("failed to write config: {error}"))) +} + +/// Replaces one source only if both its graph revision and owner contents are unchanged. +pub(in crate::api) async fn write_atomic_if_unchanged( + config_path: PathBuf, + expected_revision: String, + path: PathBuf, + expected_contents: String, + contents: String, +) -> Result<(), ApiFailure> { + tokio::task::spawn_blocking(move || { + let graph = ProxyConfig::read_source_graph(&config_path) + .map_err(|error| AtomicWriteError::ReadGraph(error.to_string()))?; + if compute_source_revision(&graph) != expected_revision { + return Err(AtomicWriteError::Conflict); + } + write_atomic_sync(&path, Some(&expected_contents), &contents).map_err(|error| { + if error.kind() == std::io::ErrorKind::AlreadyExists { + AtomicWriteError::Conflict + } else { + AtomicWriteError::Io(error) + } + }) + }) + .await + .map_err(|error| ApiFailure::internal(format!("failed to join writer: {error}")))? + .map_err(|error| match error { + AtomicWriteError::Conflict => revision_conflict(), + AtomicWriteError::ReadGraph(error) => { + ApiFailure::internal(format!("failed to verify config graph: {error}")) + } + AtomicWriteError::Io(error) => { + ApiFailure::internal(format!("failed to write config: {error}")) + } + }) +} + +fn revision_conflict() -> ApiFailure { + ApiFailure::new( + hyper::StatusCode::CONFLICT, + "revision_conflict", + "Config revision changed before persistence", + ) +} + +fn open_existing_target(path: &Path) -> std::io::Result> { + let mut options = std::fs::OpenOptions::new(); + options.read(true); + #[cfg(unix)] + options.custom_flags(libc::O_CLOEXEC | libc::O_NOFOLLOW); + let mut file = match options.open(path) { + Ok(file) => file, + Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None), + Err(error) => return Err(error), + }; + let metadata = file.metadata()?; + if !metadata.is_file() { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidInput, + "config target must be a regular file", + )); + } + let mut contents = String::new(); + file.read_to_string(&mut contents)?; + Ok(Some(ExistingTarget { contents, metadata })) +} + +fn same_target(left: &std::fs::Metadata, right: &std::fs::Metadata) -> bool { + #[cfg(unix)] + { + left.dev() == right.dev() && left.ino() == right.ino() + } + #[cfg(not(unix))] + { + left.len() == right.len() && left.modified().ok() == right.modified().ok() + } +} + +fn write_atomic_sync( + path: &Path, + expected_contents: Option<&str>, + contents: &str, +) -> std::io::Result<()> { + let parent = path.parent().unwrap_or_else(|| Path::new(".")); + std::fs::create_dir_all(parent)?; + let existing = open_existing_target(path)?; + if expected_contents.is_some_and(|expected| { + existing + .as_ref() + .is_none_or(|target| target.contents != expected) + }) { + return Err(std::io::Error::new( + std::io::ErrorKind::AlreadyExists, + "config source changed before persistence", + )); + } + + let tmp_name = format!( + ".{}.tmp-{}", + path.file_name() + .and_then(|name| name.to_str()) + .unwrap_or("config.toml"), + rand::random::() + ); + let tmp_path = parent.join(tmp_name); + + let write_result = (|| { + let mut options = std::fs::OpenOptions::new(); + options.create_new(true).write(true); + #[cfg(unix)] + options.mode(0o600); + let mut file = options.open(&tmp_path)?; + #[cfg(unix)] + if let Some(existing) = existing.as_ref() { + use nix::unistd::{Gid, Uid, fchown}; + + fchown( + &file, + Some(Uid::from_raw(existing.metadata.uid())), + Some(Gid::from_raw(existing.metadata.gid())), + ) + .map_err(|error| std::io::Error::from_raw_os_error(error as i32))?; + file.set_permissions(std::fs::Permissions::from_mode( + existing.metadata.mode() & 0o7777, + ))?; + } + file.write_all(contents.as_bytes())?; + file.sync_all()?; + let current = open_existing_target(path)?; + let target_unchanged = match (&existing, ¤t) { + (Some(expected), Some(current)) => { + same_target(&expected.metadata, ¤t.metadata) + && expected.contents == current.contents + } + (None, None) => true, + _ => false, + }; + if !target_unchanged { + return Err(std::io::Error::new( + std::io::ErrorKind::AlreadyExists, + "config target changed during persistence", + )); + } + std::fs::rename(&tmp_path, path)?; + if let Ok(dir) = std::fs::File::open(parent) { + let _ = dir.sync_all(); + } + Ok(()) + })(); + + if write_result.is_err() { + let _ = std::fs::remove_file(&tmp_path); + } + write_result +} diff --git a/src/api/config_store/persistence.rs b/src/api/config_store/persistence.rs index 1b6ee39..d19e72e 100644 --- a/src/api/config_store/persistence.rs +++ b/src/api/config_store/persistence.rs @@ -1,9 +1,5 @@ use std::collections::BTreeMap; -use std::io::{Read, Write}; -use std::path::{Path, PathBuf}; - -#[cfg(unix)] -use std::os::unix::fs::{MetadataExt, OpenOptionsExt, PermissionsExt}; +use std::path::Path; use chrono::{DateTime, Utc}; use serde::Serialize; @@ -13,9 +9,12 @@ use crate::config::{ProxyConfig, RateLimitBps}; #[cfg(test)] use super::compute_revision; use super::{ - AccessSection, compute_snapshot_revision, compute_source_revision, load_candidate_snapshot, - load_config_snapshot, resolve_single_source_owner, toml_path_exists, + AccessSection, compute_snapshot_revision, load_candidate_snapshot, load_config_snapshot, + resolve_single_source_owner, toml_path_exists, }; +use super::atomic::write_atomic_if_unchanged; +#[cfg(test)] +use super::atomic::write_atomic; use crate::api::model::ApiFailure; /// Re-render the given top-level tables from `cfg` and upsert each into the @@ -398,58 +397,6 @@ fn find_all_table_blocks(source: &str, table_name: &str) -> Vec<(usize, usize)> blocks } -/// Replaces one config source through a durable same-directory rename. -pub(in crate::api) async fn write_atomic( - path: PathBuf, - contents: String, -) -> Result<(), ApiFailure> { - tokio::task::spawn_blocking(move || write_atomic_sync(&path, None, &contents)) - .await - .map_err(|e| ApiFailure::internal(format!("failed to join writer: {}", e)))? - .map_err(|e| ApiFailure::internal(format!("failed to write config: {}", e))) -} - -/// Replaces one source only if both its graph revision and owner contents are unchanged. -pub(in crate::api) async fn write_atomic_if_unchanged( - config_path: PathBuf, - expected_revision: String, - path: PathBuf, - expected_contents: String, - contents: String, -) -> Result<(), ApiFailure> { - tokio::task::spawn_blocking(move || { - let graph = ProxyConfig::read_source_graph(&config_path) - .map_err(|error| AtomicWriteError::ReadGraph(error.to_string()))?; - if compute_source_revision(&graph) != expected_revision { - return Err(AtomicWriteError::Conflict); - } - write_atomic_sync(&path, Some(&expected_contents), &contents) - .map_err(AtomicWriteError::Io) - }) - .await - .map_err(|error| ApiFailure::internal(format!("failed to join writer: {error}")))? - .map_err(|error| match error { - AtomicWriteError::Conflict => revision_conflict(), - AtomicWriteError::ReadGraph(error) => { - ApiFailure::internal(format!("failed to verify config graph: {error}")) - } - AtomicWriteError::Io(error) => { - ApiFailure::internal(format!("failed to write config: {error}")) - } - }) -} - -enum AtomicWriteError { - Conflict, - ReadGraph(String), - Io(std::io::Error), -} - -struct ExistingTarget { - contents: String, - metadata: std::fs::Metadata, -} - fn revision_conflict() -> ApiFailure { ApiFailure::new( hyper::StatusCode::CONFLICT, @@ -457,115 +404,3 @@ fn revision_conflict() -> ApiFailure { "Config revision changed before persistence", ) } - -fn open_existing_target(path: &Path) -> std::io::Result> { - let mut options = std::fs::OpenOptions::new(); - options.read(true); - #[cfg(unix)] - options.custom_flags(libc::O_CLOEXEC | libc::O_NOFOLLOW); - let mut file = match options.open(path) { - Ok(file) => file, - Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None), - Err(error) => return Err(error), - }; - let metadata = file.metadata()?; - if !metadata.is_file() { - return Err(std::io::Error::new( - std::io::ErrorKind::InvalidInput, - "config target must be a regular file", - )); - } - let mut contents = String::new(); - file.read_to_string(&mut contents)?; - Ok(Some(ExistingTarget { contents, metadata })) -} - -fn same_target(left: &std::fs::Metadata, right: &std::fs::Metadata) -> bool { - #[cfg(unix)] - { - left.dev() == right.dev() && left.ino() == right.ino() - } - #[cfg(not(unix))] - { - left.len() == right.len() && left.modified().ok() == right.modified().ok() - } -} - -fn write_atomic_sync( - path: &Path, - expected_contents: Option<&str>, - contents: &str, -) -> std::io::Result<()> { - let parent = path.parent().unwrap_or_else(|| Path::new(".")); - std::fs::create_dir_all(parent)?; - let existing = open_existing_target(path)?; - if expected_contents.is_some_and(|expected| { - existing - .as_ref() - .is_none_or(|target| target.contents != expected) - }) { - return Err(std::io::Error::new( - std::io::ErrorKind::AlreadyExists, - "config source changed before persistence", - )); - } - - let tmp_name = format!( - ".{}.tmp-{}", - path.file_name() - .and_then(|s| s.to_str()) - .unwrap_or("config.toml"), - rand::random::() - ); - let tmp_path = parent.join(tmp_name); - - let write_result = (|| { - let mut file = std::fs::OpenOptions::new() - .create_new(true) - .write(true) - #[cfg(unix)] - .mode(0o600) - .open(&tmp_path)?; - #[cfg(unix)] - if let Some(existing) = existing.as_ref() { - use nix::unistd::{Gid, Uid, fchown}; - - fchown( - &file, - Some(Uid::from_raw(existing.metadata.uid())), - Some(Gid::from_raw(existing.metadata.gid())), - ) - .map_err(|error| std::io::Error::from_raw_os_error(error as i32))?; - file.set_permissions(std::fs::Permissions::from_mode( - existing.metadata.mode() & 0o7777, - ))?; - } - file.write_all(contents.as_bytes())?; - file.sync_all()?; - let current = open_existing_target(path)?; - let target_unchanged = match (&existing, ¤t) { - (Some(expected), Some(current)) => { - same_target(&expected.metadata, ¤t.metadata) - && expected.contents == current.contents - } - (None, None) => true, - _ => false, - }; - if !target_unchanged { - return Err(std::io::Error::new( - std::io::ErrorKind::AlreadyExists, - "config target changed during persistence", - )); - } - std::fs::rename(&tmp_path, path)?; - if let Ok(dir) = std::fs::File::open(parent) { - let _ = dir.sync_all(); - } - Ok(()) - })(); - - if write_result.is_err() { - let _ = std::fs::remove_file(&tmp_path); - } - write_result -} diff --git a/src/api/config_store/tests.rs b/src/api/config_store/tests.rs index ef35909..ccbcc12 100644 --- a/src/api/config_store/tests.rs +++ b/src/api/config_store/tests.rs @@ -260,6 +260,60 @@ async fn access_mutation_writes_only_the_single_included_owner() { assert_eq!(revision, current_revision(&root).await.unwrap()); } +#[tokio::test] +async fn access_mutation_rejects_source_graph_change_after_snapshot() { + let dir = tempfile::tempdir().unwrap(); + let root = dir.path().join("config.toml"); + let included = dir.path().join("users.toml"); + let root_body = "include = \"users.toml\"\n[censorship]\ntls_domain = \"one.example\"\n"; + let external_root = + "include = \"users.toml\"\n[censorship]\ntls_domain = \"two.example\"\n"; + let included_body = "[access.users]\nalice = \"00000000000000000000000000000000\"\n"; + tokio::fs::write(&root, root_body).await.unwrap(); + tokio::fs::write(&included, included_body).await.unwrap(); + let (mut cfg, revision) = load_config_for_mutation(&root, None).await.unwrap(); + cfg.access.users.insert( + "bob".to_string(), + "11111111111111111111111111111111".to_string(), + ); + tokio::fs::write(&root, external_root).await.unwrap(); + + let error = save_access_sections_to_disk_if_revision( + &root, + &cfg, + &[AccessSection::Users], + Some(&revision), + ) + .await + .unwrap_err(); + + assert_eq!(error.code, "revision_conflict"); + assert_eq!(tokio::fs::read_to_string(&root).await.unwrap(), external_root); + assert_eq!( + tokio::fs::read_to_string(&included).await.unwrap(), + included_body + ); +} + +#[cfg(unix)] +#[tokio::test] +async fn atomic_write_preserves_existing_file_mode() { + use std::os::unix::fs::{MetadataExt, PermissionsExt}; + + let dir = tempfile::tempdir().unwrap(); + let path = dir.path().join("config.toml"); + tokio::fs::write(&path, "old").await.unwrap(); + std::fs::set_permissions(&path, std::fs::Permissions::from_mode(0o640)).unwrap(); + let before = std::fs::metadata(&path).unwrap(); + + write_atomic(path.clone(), "new".to_string()).await.unwrap(); + + let after = std::fs::metadata(&path).unwrap(); + assert_eq!(after.mode() & 0o7777, 0o640); + assert_eq!(after.uid(), before.uid()); + assert_eq!(after.gid(), before.gid()); +} + #[tokio::test] async fn access_mutation_rejects_sections_with_different_source_owners() { let dir = tempfile::tempdir().unwrap(); diff --git a/src/api/handler/user_routes.rs b/src/api/handler/user_routes.rs index ca50e0d..f265213 100644 --- a/src/api/handler/user_routes.rs +++ b/src/api/handler/user_routes.rs @@ -134,8 +134,8 @@ pub(super) async fn handle( } let expected_revision = parse_if_match(req.headers()); let _mutation_guard = shared.mutation_lock.lock().await; - let disk_cfg = load_config_from_disk(&shared.config_path).await?; - ensure_expected_revision(&shared.config_path, expected_revision.as_deref()).await?; + let (disk_cfg, _) = + load_config_for_mutation(&shared.config_path, expected_revision.as_deref()).await?; if !disk_cfg.access.users.contains_key(user) { return Ok(error_response( request_id, diff --git a/src/api/mod.rs b/src/api/mod.rs index edf288a..9d9ce68 100644 --- a/src/api/mod.rs +++ b/src/api/mod.rs @@ -59,7 +59,7 @@ mod web_runtime; mod web_status; use config_store::{ - current_revision, ensure_expected_revision, load_config_for_reload, load_config_from_disk, + current_revision, load_config_for_mutation, load_config_for_reload, load_config_from_disk, parse_if_match, }; use events::ApiEventStore; diff --git a/src/api/users.rs b/src/api/users.rs index 5228277..fc6286f 100644 --- a/src/api/users.rs +++ b/src/api/users.rs @@ -9,8 +9,8 @@ use crate::stats::Stats; use super::ApiShared; use super::config_store::{ - AccessSection, current_revision, ensure_expected_revision, load_config_from_disk, - save_access_sections_to_disk, + AccessSection, current_revision, load_config_for_mutation, + save_access_sections_to_disk_if_revision, }; use super::model::{ ApiFailure, CreateUserRequest, CreateUserResponse, PatchUserRequest, RotateSecretRequest, diff --git a/src/api/users/create.rs b/src/api/users/create.rs index 3ee3d4a..40a48df 100644 --- a/src/api/users/create.rs +++ b/src/api/users/create.rs @@ -42,8 +42,8 @@ pub(in crate::api) async fn create_user( let expiration = parse_optional_expiration(body.expiration_rfc3339.as_deref())?; let _guard = shared.mutation_lock.lock().await; - let mut cfg = load_config_from_disk(&shared.config_path).await?; - ensure_expected_revision(&shared.config_path, expected_revision.as_deref()).await?; + let (mut cfg, base_revision) = + load_config_for_mutation(&shared.config_path, expected_revision.as_deref()).await?; if cfg.access.users.contains_key(&body.username) { return Err(ApiFailure::new( @@ -122,8 +122,13 @@ pub(in crate::api) async fn create_user( touched_sections.push(AccessSection::UserEnabled); } - let revision = - save_access_sections_to_disk(&shared.config_path, &cfg, &touched_sections).await?; + let revision = save_access_sections_to_disk_if_revision( + &shared.config_path, + &cfg, + &touched_sections, + Some(&base_revision), + ) + .await?; shared .proxy_shared .stage_user( diff --git a/src/api/users/lifecycle.rs b/src/api/users/lifecycle.rs index 546b53e..e962fb6 100644 --- a/src/api/users/lifecycle.rs +++ b/src/api/users/lifecycle.rs @@ -15,8 +15,8 @@ pub(in crate::api) async fn rotate_secret( } let _guard = shared.mutation_lock.lock().await; - let mut cfg = load_config_from_disk(&shared.config_path).await?; - ensure_expected_revision(&shared.config_path, expected_revision.as_deref()).await?; + let (mut cfg, base_revision) = + load_config_for_mutation(&shared.config_path, expected_revision.as_deref()).await?; if !cfg.access.users.contains_key(user) { return Err(ApiFailure::new( @@ -29,8 +29,13 @@ pub(in crate::api) async fn rotate_secret( cfg.access.users.insert(user.to_string(), secret.clone()); cfg.validate() .map_err(|e| ApiFailure::bad_request(format!("config validation failed: {}", e)))?; - let revision = - save_access_sections_to_disk(&shared.config_path, &cfg, &[AccessSection::Users]).await?; + let revision = save_access_sections_to_disk_if_revision( + &shared.config_path, + &cfg, + &[AccessSection::Users], + Some(&base_revision), + ) + .await?; shared .proxy_shared .stage_user(user, &secret, cfg.access.is_user_enabled(user)) @@ -67,8 +72,8 @@ pub(in crate::api) async fn delete_user( shared: &ApiShared, ) -> Result<(String, String), ApiFailure> { let _guard = shared.mutation_lock.lock().await; - let mut cfg = load_config_from_disk(&shared.config_path).await?; - ensure_expected_revision(&shared.config_path, expected_revision.as_deref()).await?; + let (mut cfg, base_revision) = + load_config_for_mutation(&shared.config_path, expected_revision.as_deref()).await?; if !cfg.access.users.contains_key(user) { return Err(ApiFailure::new( @@ -111,8 +116,13 @@ pub(in crate::api) async fn delete_user( cfg.validate() .map_err(|e| ApiFailure::bad_request(format!("config validation failed: {}", e)))?; - let revision = - save_access_sections_to_disk(&shared.config_path, &cfg, &touched_sections).await?; + let revision = save_access_sections_to_disk_if_revision( + &shared.config_path, + &cfg, + &touched_sections, + Some(&base_revision), + ) + .await?; let deleted_incarnation = shared.proxy_shared.delete_user(user).incarnation; let configured_users = cfg.access.users.keys().cloned().collect(); if let Err(error) = shared diff --git a/src/api/users/update.rs b/src/api/users/update.rs index 4ff090c..1619e10 100644 --- a/src/api/users/update.rs +++ b/src/api/users/update.rs @@ -32,8 +32,8 @@ pub(in crate::api) async fn patch_user( } let expiration = parse_patch_expiration(&body.expiration_rfc3339)?; let _guard = shared.mutation_lock.lock().await; - let mut cfg = load_config_from_disk(&shared.config_path).await?; - ensure_expected_revision(&shared.config_path, expected_revision.as_deref()).await?; + let (mut cfg, base_revision) = + load_config_for_mutation(&shared.config_path, expected_revision.as_deref()).await?; if !cfg.access.users.contains_key(user) { return Err(ApiFailure::new( @@ -168,7 +168,13 @@ pub(in crate::api) async fn patch_user( let revision = if touched_sections.is_empty() { current_revision(&shared.config_path).await? } else { - save_access_sections_to_disk(&shared.config_path, &cfg, &touched_sections).await? + save_access_sections_to_disk_if_revision( + &shared.config_path, + &cfg, + &touched_sections, + Some(&base_revision), + ) + .await? }; if touches_users || touches_user_enabled { let secret = cfg @@ -212,8 +218,8 @@ pub(in crate::api) async fn set_user_enabled( shared: &ApiShared, ) -> Result<(UserInfo, String), ApiFailure> { let _guard = shared.mutation_lock.lock().await; - let mut cfg = load_config_from_disk(&shared.config_path).await?; - ensure_expected_revision(&shared.config_path, expected_revision.as_deref()).await?; + let (mut cfg, base_revision) = + load_config_for_mutation(&shared.config_path, expected_revision.as_deref()).await?; if !cfg.access.users.contains_key(user) { return Err(ApiFailure::new( @@ -231,9 +237,13 @@ pub(in crate::api) async fn set_user_enabled( cfg.validate() .map_err(|e| ApiFailure::bad_request(format!("config validation failed: {}", e)))?; - let revision = - save_access_sections_to_disk(&shared.config_path, &cfg, &[AccessSection::UserEnabled]) - .await?; + let revision = save_access_sections_to_disk_if_revision( + &shared.config_path, + &cfg, + &[AccessSection::UserEnabled], + Some(&base_revision), + ) + .await?; let secret = cfg .access .users diff --git a/src/config/load.rs b/src/config/load.rs index 4b19197..c497a4b 100644 --- a/src/config/load.rs +++ b/src/config/load.rs @@ -36,7 +36,9 @@ mod validate_server; mod validate_web; mod validation; -use self::includes::{hash_rendered_snapshot, normalize_config_path, preprocess_includes}; +use self::includes::{ + hash_rendered_snapshot, normalize_config_path, preprocess_includes, read_config_source, +}; use self::normalize::{ is_valid_ad_tag, is_valid_tls_domain_name, normalize_domain_to_ascii, normalize_exclusive_mask_target, normalize_mask_host_to_ascii, parse_exclusive_mask_target, @@ -175,13 +177,12 @@ impl ProxyConfig { source_overrides: &BTreeMap, ) -> Result { let path = path.as_ref(); - let normalized_path = normalize_config_path(path); - let content = source_overrides - .get(&normalized_path) - .cloned() - .map(Ok) - .unwrap_or_else(|| std::fs::read_to_string(path)) - .map_err(|e| ProxyError::Config(e.to_string()))?; + let initial_path = normalize_config_path(path); + let (normalized_path, content) = if let Some(content) = source_overrides.get(&initial_path) { + (initial_path, content.clone()) + } else { + read_config_source(path)? + }; let base_dir = path.parent().unwrap_or(Path::new(".")); let mut source_files = BTreeSet::new(); source_files.insert(normalized_path.clone()); diff --git a/src/config/load/includes.rs b/src/config/load/includes.rs index 89780b2..aef22ef 100644 --- a/src/config/load/includes.rs +++ b/src/config/load/includes.rs @@ -1,7 +1,11 @@ use std::collections::{BTreeMap, BTreeSet}; use std::hash::{DefaultHasher, Hash, Hasher}; +use std::io::Read; use std::path::{Path, PathBuf}; +#[cfg(unix)] +use std::os::unix::fs::{MetadataExt, OpenOptionsExt}; + use crate::error::{ProxyError, Result}; pub(super) fn normalize_config_path(path: &Path) -> PathBuf { @@ -22,6 +26,74 @@ pub(super) fn hash_rendered_snapshot(rendered: &str) -> u64 { hasher.finish() } +pub(super) fn read_config_source(path: &Path) -> Result<(PathBuf, String)> { + let mut options = std::fs::OpenOptions::new(); + options.read(true); + #[cfg(unix)] + options.custom_flags(libc::O_CLOEXEC); + let mut file = options + .open(path) + .map_err(|error| ProxyError::Config(error.to_string()))?; + let opened_metadata = file + .metadata() + .map_err(|error| ProxyError::Config(error.to_string()))?; + if !opened_metadata.is_file() { + return Err(ProxyError::Config(format!( + "config source `{}` must be a regular file", + path.display() + ))); + } + let normalized = normalize_config_path(path); + let current_metadata = std::fs::metadata(&normalized) + .map_err(|error| ProxyError::Config(error.to_string()))?; + if !same_file_identity(&opened_metadata, ¤t_metadata) { + return Err(ProxyError::Config(format!( + "config source `{}` changed while it was opened", + path.display() + ))); + } + let mut contents = String::new(); + file.read_to_string(&mut contents) + .map_err(|error| ProxyError::Config(error.to_string()))?; + let completed_metadata = file + .metadata() + .map_err(|error| ProxyError::Config(error.to_string()))?; + if !same_file_version(&opened_metadata, &completed_metadata) { + return Err(ProxyError::Config(format!( + "config source `{}` changed while it was read", + path.display() + ))); + } + Ok((normalized, contents)) +} + +fn same_file_identity(left: &std::fs::Metadata, right: &std::fs::Metadata) -> bool { + #[cfg(unix)] + { + left.dev() == right.dev() && left.ino() == right.ino() + } + #[cfg(not(unix))] + { + left.len() == right.len() && left.modified().ok() == right.modified().ok() + } +} + +fn same_file_version(left: &std::fs::Metadata, right: &std::fs::Metadata) -> bool { + #[cfg(unix)] + { + same_file_identity(left, right) + && left.len() == right.len() + && left.mtime() == right.mtime() + && left.mtime_nsec() == right.mtime_nsec() + && left.ctime() == right.ctime() + && left.ctime_nsec() == right.ctime_nsec() + } + #[cfg(not(unix))] + { + same_file_identity(left, right) + } +} + pub(super) fn preprocess_includes( content: &str, base_dir: &Path, @@ -42,14 +114,18 @@ pub(super) fn preprocess_includes( let path_str = rest.trim().trim_matches('"'); let resolved = base_dir.join(path_str); let normalized = normalize_config_path(&resolved); + let cached = source_contents.get(&normalized).cloned(); + let (normalized, included) = if let Some(included) = + source_overrides.get(&normalized).cloned().or(cached) + { + (normalized, included) + } else { + read_config_source(&resolved)? + }; source_files.insert(normalized.clone()); - let included = source_overrides - .get(&normalized) - .cloned() - .map(Ok) - .unwrap_or_else(|| std::fs::read_to_string(&resolved)) - .map_err(|e| ProxyError::Config(e.to_string()))?; - source_contents.insert(normalized, included.clone()); + source_contents + .entry(normalized) + .or_insert_with(|| included.clone()); let included_dir = resolved.parent().unwrap_or(base_dir); output.push_str(&preprocess_includes( &included, diff --git a/src/config/load/runtime_web.rs b/src/config/load/runtime_web.rs index d99ea57..3b701c1 100644 --- a/src/config/load/runtime_web.rs +++ b/src/config/load/runtime_web.rs @@ -5,7 +5,18 @@ use std::path::Path; use std::sync::Arc; #[cfg(unix)] -use std::os::unix::fs::OpenOptionsExt; +use std::ffi::OsString; +#[cfg(unix)] +use std::os::unix::ffi::OsStringExt; +#[cfg(unix)] +use std::os::unix::fs::MetadataExt; + +#[cfg(unix)] +use nix::dir::Dir; +#[cfg(unix)] +use nix::fcntl::{openat, OFlag}; +#[cfg(unix)] +use nix::sys::stat::Mode; use bytes::Bytes; use hmac::{Hmac, Mac}; @@ -13,6 +24,10 @@ use sha2::{Digest, Sha256}; use super::*; +// Path-based static snapshot fallback for platforms without directory descriptors. +#[cfg(not(unix))] +mod static_site_fallback; + const WEB_CAPABILITY_CONTEXT: &[u8] = b"tdesktop-web-proxy-bridge-v1\n"; const WEB_DEBUG_FINGERPRINT_CONTEXT: &[u8] = b"telemt-web-debug-key-fingerprint-v1\0"; const MAX_WEB_STATIC_DEPTH: usize = 64; @@ -193,34 +208,31 @@ fn load_static_site( total_files: &mut usize, total_bytes: &mut usize, ) -> Result { - let root_metadata = fs::symlink_metadata(root).map_err(|error| { - ProxyError::Config(format!( - "failed to inspect WEB static directory `{}`: {error}", - root.display() - )) - })?; - if root_metadata.file_type().is_symlink() || !root_metadata.is_dir() { - return Err(ProxyError::Config(format!( - "WEB static directory `{}` must be a real directory, not a symlink", - root.display() - ))); - } - let canonical_root = fs::canonicalize(root).map_err(|error| { - ProxyError::Config(format!( - "failed to canonicalize WEB static directory `{}`: {error}", - root.display() - )) - })?; let mut assets = BTreeMap::new(); - load_static_directory( - &canonical_root, - &canonical_root, - &mut assets, - total_files, - total_bytes, - limits, - 0, - )?; + #[cfg(unix)] + { + let directory = open_static_root(root)?; + load_static_directory( + directory, + Path::new(""), + root, + &mut assets, + total_files, + total_bytes, + limits, + 0, + )?; + } + #[cfg(not(unix))] + { + static_site_fallback::load_static_site_by_path( + root, + limits, + &mut assets, + total_files, + total_bytes, + )?; + } if !assets.contains_key(&format!("/{index}")) { return Err(ProxyError::Config(format!( "WEB static directory `{}` does not contain index `{index}`", @@ -233,54 +245,94 @@ fn load_static_site( }) } +#[cfg(unix)] +fn open_static_root(root: &Path) -> Result { + Dir::open( + root, + OFlag::O_RDONLY | OFlag::O_DIRECTORY | OFlag::O_NOFOLLOW | OFlag::O_CLOEXEC, + Mode::empty(), + ) + .map_err(|error| { + ProxyError::Config(format!( + "WEB static directory `{}` must be a real directory, not a symlink: {error}", + root.display() + )) + }) +} + +#[cfg(unix)] fn load_static_directory( + mut directory: Dir, + relative: &Path, root: &Path, - directory: &Path, assets: &mut BTreeMap, total_files: &mut usize, total_bytes: &mut usize, limits: &WebLimitsConfig, depth: usize, ) -> Result<()> { - let entries = fs::read_dir(directory).map_err(|error| { - ProxyError::Config(format!( - "failed to read WEB static directory `{}`: {error}", - directory.display() - )) - })?; - for entry in entries { + let mut entries = Vec::new(); + for entry in directory.iter() { let entry = entry.map_err(|error| { - ProxyError::Config(format!("failed to read WEB static entry: {error}")) + ProxyError::Config(format!( + "failed to read WEB static directory `{}`: {error}", + root.join(relative).display() + )) })?; + let name = entry.file_name().to_bytes(); + if name == b"." || name == b".." { + continue; + } if *total_files >= limits.max_static_files { return Err(ProxyError::Config( "WEB static entries exceed process-wide web.limits.max_static_files".to_string(), )); } *total_files += 1; - let path = entry.path(); - let file_type = entry.file_type().map_err(|error| { + entries.push(OsString::from_vec(name.to_vec())); + } + entries.sort_unstable(); + + for name in entries { + let relative_path = relative.join(&name); + let display_path = root.join(&relative_path); + let descriptor = openat( + &directory, + name.as_os_str(), + OFlag::O_RDONLY | OFlag::O_NOFOLLOW | OFlag::O_CLOEXEC, + Mode::empty(), + ) + .map_err(|error| { ProxyError::Config(format!( - "failed to inspect WEB static entry `{}`: {error}", - path.display() + "failed to open WEB static entry `{}` without following symlinks: {error}", + display_path.display() )) })?; - if file_type.is_symlink() { - return Err(ProxyError::Config(format!( - "WEB static entry `{}` must not be a symlink", - path.display() - ))); - } - if file_type.is_dir() { + let file = fs::File::from(descriptor); + let metadata = file.metadata().map_err(|error| { + ProxyError::Config(format!( + "failed to inspect WEB static entry `{}`: {error}", + display_path.display() + )) + })?; + if metadata.is_dir() { if depth >= MAX_WEB_STATIC_DEPTH { return Err(ProxyError::Config(format!( "WEB static directory `{}` exceeds the maximum nesting depth", - path.display() + display_path.display() ))); } + let descriptor = file.into(); + let child = Dir::from_fd(descriptor).map_err(|error| { + ProxyError::Config(format!( + "failed to open WEB static directory `{}`: {error}", + display_path.display() + )) + })?; load_static_directory( + child, + &relative_path, root, - &path, assets, total_files, total_bytes, @@ -289,83 +341,107 @@ fn load_static_directory( )?; continue; } - if !file_type.is_file() { - return Err(ProxyError::Config(format!( - "WEB static entry `{}` must be a regular file", - path.display() - ))); - } - let mut options = fs::OpenOptions::new(); - options.read(true); - #[cfg(unix)] - options.custom_flags(libc::O_CLOEXEC | libc::O_NOFOLLOW); - let file = options.open(&path).map_err(|error| { - ProxyError::Config(format!( - "failed to open WEB static file `{}`: {error}", - path.display() - )) - })?; - let metadata = file.metadata().map_err(|error| { - ProxyError::Config(format!( - "failed to inspect WEB static file `{}`: {error}", - path.display() - )) - })?; if !metadata.is_file() { return Err(ProxyError::Config(format!( - "WEB static entry `{}` changed before it was opened", - path.display() + "WEB static entry `{}` must be a regular file", + display_path.display() ))); } - let file_len = usize::try_from(metadata.len()).map_err(|_| { - ProxyError::Config(format!("WEB static file `{}` is too large", path.display())) - })?; - if file_len > limits.max_static_file_bytes { - return Err(ProxyError::Config(format!( - "WEB static file `{}` exceeds web.limits.max_static_file_bytes", - path.display() - ))); - } - *total_bytes = total_bytes.checked_add(file_len).ok_or_else(|| { - ProxyError::Config("WEB static snapshot byte count overflowed usize".to_string()) - })?; - if *total_bytes > limits.max_static_bytes { - return Err(ProxyError::Config( - "WEB static snapshots exceed process-wide web.limits.max_static_bytes".to_string(), - )); - } - let relative = path.strip_prefix(root).map_err(|_| { - ProxyError::Config("WEB static path escaped its configured root".to_string()) - })?; - let route = static_route(relative)?; - let mut body = Vec::with_capacity(file_len); - file.take(limits.max_static_file_bytes as u64 + 1) - .read_to_end(&mut body) - .map_err(|error| { - ProxyError::Config(format!( - "failed to read WEB static file `{}`: {error}", - path.display() - )) - })?; - if body.len() != file_len { - return Err(ProxyError::Config(format!( - "WEB static file `{}` changed while its snapshot was built", - path.display() - ))); - } - let etag = format!("\"{}\"", hex::encode(Sha256::digest(&body))); - assets.insert( - route, - WebStaticAsset { - body: Bytes::from(body), - content_type: static_content_type(&path), - etag, - }, - ); + load_static_file( + file, + &metadata, + &relative_path, + &display_path, + assets, + total_bytes, + limits, + )?; } Ok(()) } +fn load_static_file( + mut file: fs::File, + metadata: &fs::Metadata, + relative: &Path, + display_path: &Path, + assets: &mut BTreeMap, + total_bytes: &mut usize, + limits: &WebLimitsConfig, +) -> Result<()> { + let file_len = usize::try_from(metadata.len()).map_err(|_| { + ProxyError::Config(format!( + "WEB static file `{}` is too large", + display_path.display() + )) + })?; + if file_len > limits.max_static_file_bytes { + return Err(ProxyError::Config(format!( + "WEB static file `{}` exceeds web.limits.max_static_file_bytes", + display_path.display() + ))); + } + *total_bytes = total_bytes.checked_add(file_len).ok_or_else(|| { + ProxyError::Config("WEB static snapshot byte count overflowed usize".to_string()) + })?; + if *total_bytes > limits.max_static_bytes { + return Err(ProxyError::Config( + "WEB static snapshots exceed process-wide web.limits.max_static_bytes".to_string(), + )); + } + let route = static_route(relative)?; + let mut body = Vec::with_capacity(file_len); + file.by_ref() + .take(limits.max_static_file_bytes as u64 + 1) + .read_to_end(&mut body) + .map_err(|error| { + ProxyError::Config(format!( + "failed to read WEB static file `{}`: {error}", + display_path.display() + )) + })?; + let final_metadata = file.metadata().map_err(|error| { + ProxyError::Config(format!( + "failed to recheck WEB static file `{}`: {error}", + display_path.display() + )) + })?; + if body.len() != file_len || !static_file_version_matches(metadata, &final_metadata) { + return Err(ProxyError::Config(format!( + "WEB static file `{}` changed while its snapshot was built", + display_path.display() + ))); + } + let etag = format!("\"{}\"", hex::encode(Sha256::digest(&body))); + assets.insert( + route, + WebStaticAsset { + body: Bytes::from(body), + content_type: static_content_type(relative), + etag, + }, + ); + Ok(()) +} + +#[cfg(unix)] +fn static_file_version_matches(before: &fs::Metadata, after: &fs::Metadata) -> bool { + before.dev() == after.dev() + && before.ino() == after.ino() + && before.len() == after.len() + && before.mtime() == after.mtime() + && before.mtime_nsec() == after.mtime_nsec() + && before.ctime() == after.ctime() + && before.ctime_nsec() == after.ctime_nsec() +} + +#[cfg(not(unix))] +fn static_file_version_matches(before: &fs::Metadata, after: &fs::Metadata) -> bool { + before.len() == after.len() + && before.modified().ok() == after.modified().ok() + && before.created().ok() == after.created().ok() +} + fn static_route(relative: &Path) -> Result { let mut route = String::new(); for component in relative.components() { @@ -425,4 +501,40 @@ mod tests { "IpJrt3e7sKtzPyoXy6w-Zj6GGEvsvclN66JzQEfPYLA" ); } + + #[cfg(unix)] + #[test] + fn static_snapshot_remains_anchored_after_root_path_replacement() { + use std::os::unix::fs::symlink; + + let temp = tempfile::tempdir().unwrap(); + let root = temp.path().join("site"); + let detached = temp.path().join("detached"); + let replacement = temp.path().join("replacement"); + fs::create_dir(&root).unwrap(); + fs::write(root.join("index.html"), b"original").unwrap(); + fs::create_dir(&replacement).unwrap(); + fs::write(replacement.join("index.html"), b"replacement").unwrap(); + + let directory = open_static_root(&root).unwrap(); + fs::rename(&root, &detached).unwrap(); + symlink(&replacement, &root).unwrap(); + + let mut assets = BTreeMap::new(); + let mut total_files = 0; + let mut total_bytes = 0; + load_static_directory( + directory, + Path::new(""), + &root, + &mut assets, + &mut total_files, + &mut total_bytes, + &WebLimitsConfig::default(), + 0, + ) + .unwrap(); + + assert_eq!(assets["/index.html"].body.as_ref(), b"original"); + } } diff --git a/src/config/load/runtime_web/static_site_fallback.rs b/src/config/load/runtime_web/static_site_fallback.rs new file mode 100644 index 0000000..6a1d667 --- /dev/null +++ b/src/config/load/runtime_web/static_site_fallback.rs @@ -0,0 +1,129 @@ +use std::collections::BTreeMap; +use std::fs; +use std::path::Path; + +use super::*; + +pub(super) fn load_static_site_by_path( + root: &Path, + limits: &WebLimitsConfig, + assets: &mut BTreeMap, + total_files: &mut usize, + total_bytes: &mut usize, +) -> Result<()> { + let root_metadata = fs::symlink_metadata(root).map_err(|error| { + ProxyError::Config(format!( + "failed to inspect WEB static directory `{}`: {error}", + root.display() + )) + })?; + if root_metadata.file_type().is_symlink() || !root_metadata.is_dir() { + return Err(ProxyError::Config(format!( + "WEB static directory `{}` must be a real directory, not a symlink", + root.display() + ))); + } + let canonical_root = fs::canonicalize(root).map_err(|error| { + ProxyError::Config(format!( + "failed to canonicalize WEB static directory `{}`: {error}", + root.display() + )) + })?; + load_static_directory( + &canonical_root, + &canonical_root, + assets, + total_files, + total_bytes, + limits, + 0, + ) +} + +fn load_static_directory( + root: &Path, + directory: &Path, + assets: &mut BTreeMap, + total_files: &mut usize, + total_bytes: &mut usize, + limits: &WebLimitsConfig, + depth: usize, +) -> Result<()> { + let entries = fs::read_dir(directory).map_err(|error| { + ProxyError::Config(format!( + "failed to read WEB static directory `{}`: {error}", + directory.display() + )) + })?; + for entry in entries { + let entry = entry.map_err(|error| { + ProxyError::Config(format!("failed to read WEB static entry: {error}")) + })?; + if *total_files >= limits.max_static_files { + return Err(ProxyError::Config( + "WEB static entries exceed process-wide web.limits.max_static_files".to_string(), + )); + } + *total_files += 1; + let path = entry.path(); + let file_type = entry.file_type().map_err(|error| { + ProxyError::Config(format!( + "failed to inspect WEB static entry `{}`: {error}", + path.display() + )) + })?; + if file_type.is_symlink() { + return Err(ProxyError::Config(format!( + "WEB static entry `{}` must not be a symlink", + path.display() + ))); + } + if file_type.is_dir() { + if depth >= MAX_WEB_STATIC_DEPTH { + return Err(ProxyError::Config(format!( + "WEB static directory `{}` exceeds the maximum nesting depth", + path.display() + ))); + } + load_static_directory( + root, + &path, + assets, + total_files, + total_bytes, + limits, + depth + 1, + )?; + continue; + } + if !file_type.is_file() { + return Err(ProxyError::Config(format!( + "WEB static entry `{}` must be a regular file", + path.display() + ))); + } + let file = fs::File::open(&path).map_err(|error| { + ProxyError::Config(format!( + "failed to open WEB static file `{}`: {error}", + path.display() + )) + })?; + let metadata = file.metadata().map_err(|error| { + ProxyError::Config(format!( + "failed to inspect WEB static file `{}`: {error}", + path.display() + )) + })?; + if !metadata.is_file() { + return Err(ProxyError::Config(format!( + "WEB static entry `{}` changed before it was opened", + path.display() + ))); + } + let relative = path.strip_prefix(root).map_err(|_| { + ProxyError::Config("WEB static path escaped its configured root".to_string()) + })?; + load_static_file(file, &metadata, relative, &path, assets, total_bytes, limits)?; + } + Ok(()) +} diff --git a/src/conntrack_control.rs b/src/conntrack_control.rs index 1cb5b7d..8024cb2 100644 --- a/src/conntrack_control.rs +++ b/src/conntrack_control.rs @@ -1,6 +1,5 @@ use std::collections::BTreeSet; use std::net::IpAddr; -use std::path::PathBuf; use std::sync::Arc; use std::time::Duration; @@ -14,6 +13,8 @@ use crate::config::{ConntrackBackend, ConntrackMode, ProxyConfig}; use crate::proxy::middle_relay::note_global_relay_pressure; use crate::proxy::shared_state::{ConntrackCloseEvent, ConntrackCloseReason, ProxySharedState}; use crate::stats::Stats; +#[cfg(unix)] +use crate::util::trusted_command::resolve_trusted_helper; const CONNTRACK_EVENT_QUEUE_CAPACITY: usize = 32_768; const PRESSURE_RELEASE_TICKS: u8 = 3; @@ -381,13 +382,15 @@ fn pick_backend(configured: ConntrackBackend) -> Option { } fn command_exists(binary: &str) -> bool { - let Some(path_var) = std::env::var_os("PATH") else { - return false; - }; - std::env::split_paths(&path_var).any(|dir| { - let candidate: PathBuf = dir.join(binary); - candidate.exists() && candidate.is_file() - }) + #[cfg(unix)] + { + resolve_trusted_helper(binary).is_some() + } + #[cfg(not(unix))] + { + let _ = binary; + false + } } fn listener_port_set(cfg: &ProxyConfig) -> Vec { @@ -651,10 +654,14 @@ async fn delete_conntrack_entry(event: ConntrackCloseEvent) -> DeleteOutcome { } async fn run_command(binary: &str, args: &[&str], stdin: Option) -> Result<(), String> { - if !command_exists(binary) { + #[cfg(unix)] + let Some(command_path) = resolve_trusted_helper(binary) else { return Err(format!("{binary} is not available")); - } - let mut command = Command::new(binary); + }; + #[cfg(not(unix))] + return Err(format!("{binary} is not available")); + #[cfg(unix)] + let mut command = Command::new(command_path); command.args(args); if stdin.is_some() { command.stdin(std::process::Stdio::piped()); diff --git a/src/main.rs b/src/main.rs index c0453d9..8e2f307 100644 --- a/src/main.rs +++ b/src/main.rs @@ -27,6 +27,7 @@ mod protocol; mod proxy; mod quota_state; mod service; +mod slot_budget; mod startup; mod stats; mod stream; diff --git a/src/proxy/handshake.rs b/src/proxy/handshake.rs index 96cefd3..e0d76c6 100644 --- a/src/proxy/handshake.rs +++ b/src/proxy/handshake.rs @@ -75,7 +75,7 @@ pub(crate) use self::auth_probe::{ auth_probe_saturation_is_throttled_for_testing_in_shared, auth_probe_saturation_state_for_testing_in_shared, auth_probe_saturation_state_lock_for_testing_in_shared, auth_probe_state_for_testing_in_shared, - clear_auth_probe_state_for_testing_in_shared, + auth_probe_slots_for_testing_in_shared, clear_auth_probe_state_for_testing_in_shared, clear_unknown_sni_warn_state_for_testing_in_shared, clear_warned_secrets_for_testing_in_shared, should_emit_unknown_sni_warn_for_testing_in_shared, warned_secrets_for_testing_in_shared, }; @@ -89,13 +89,16 @@ const WARNED_SECRET_MAX_ENTRIES: usize = 1_024; const AUTH_PROBE_TRACK_RETENTION_SECS: u64 = 10 * 60; #[cfg(test)] -const AUTH_PROBE_TRACK_MAX_ENTRIES: usize = 256; +pub(super) const AUTH_PROBE_TRACK_MAX_ENTRIES: usize = 256; #[cfg(not(test))] -const AUTH_PROBE_TRACK_MAX_ENTRIES: usize = 65_536; +pub(super) const AUTH_PROBE_TRACK_MAX_ENTRIES: usize = 65_536; const AUTH_PROBE_PRUNE_SCAN_LIMIT: usize = 1_024; const AUTH_PROBE_BACKOFF_START_FAILS: u32 = 4; const AUTH_PROBE_SATURATION_GRACE_FAILS: u32 = 2; -const STICKY_HINT_MAX_ENTRIES: usize = 65_536; +#[cfg(test)] +pub(super) const STICKY_HINT_MAX_ENTRIES: usize = 256; +#[cfg(not(test))] +pub(super) const STICKY_HINT_MAX_ENTRIES: usize = 65_536; const CANDIDATE_HINT_TRACK_CAP: usize = 64; const OVERLOAD_CANDIDATE_BUDGET_HINTED: usize = 16; const OVERLOAD_CANDIDATE_BUDGET_UNHINTED: usize = 8; diff --git a/src/proxy/handshake/auth_candidates.rs b/src/proxy/handshake/auth_candidates.rs index 48954ee..fb7fe1f 100644 --- a/src/proxy/handshake/auth_candidates.rs +++ b/src/proxy/handshake/auth_candidates.rs @@ -76,27 +76,48 @@ pub(super) fn sticky_hint_record_success_in( user_id: u32, sni: Option<&str>, ) { - if shared.handshake.sticky_user_by_ip.len() > STICKY_HINT_MAX_ENTRIES { - shared.handshake.sticky_user_by_ip.clear(); - } - shared.handshake.sticky_user_by_ip.insert(peer_ip, user_id); - - if shared.handshake.sticky_user_by_ip_prefix.len() > STICKY_HINT_MAX_ENTRIES { - shared.handshake.sticky_user_by_ip_prefix.clear(); - } - shared - .handshake - .sticky_user_by_ip_prefix - .insert(ip_prefix_hint_key(peer_ip), user_id); + bounded_sticky_hint_upsert( + &shared.handshake.sticky_user_by_ip, + &shared.handshake.sticky_user_by_ip_slots, + peer_ip, + user_id, + ); + bounded_sticky_hint_upsert( + &shared.handshake.sticky_user_by_ip_prefix, + &shared.handshake.sticky_user_by_ip_prefix_slots, + ip_prefix_hint_key(peer_ip), + user_id, + ); if let Some(sni) = sni { - if shared.handshake.sticky_user_by_sni_hash.len() > STICKY_HINT_MAX_ENTRIES { - shared.handshake.sticky_user_by_sni_hash.clear(); + bounded_sticky_hint_upsert( + &shared.handshake.sticky_user_by_sni_hash, + &shared.handshake.sticky_user_by_sni_hash_slots, + sni_hint_hash(sni), + user_id, + ); + } +} + +fn bounded_sticky_hint_upsert( + entries: &DashMap, + slots: &crate::slot_budget::SlotBudget, + key: K, + user_id: u32, +) where + K: Eq + Hash, +{ + match entries.entry(key) { + Entry::Occupied(mut entry) => { + entry.insert(user_id); + } + Entry::Vacant(entry) => { + let Some(slot) = slots.try_acquire() else { + return; + }; + entry.insert(user_id); + slot.commit(); } - shared - .handshake - .sticky_user_by_sni_hash - .insert(sni_hint_hash(sni), user_id); } } @@ -343,6 +364,61 @@ mod web_mode_tests { } } +#[cfg(test)] +mod bounded_registry_tests { + use std::sync::Arc; + + use super::*; + + #[test] + fn parallel_sticky_hints_never_exceed_their_hard_caps() { + const ATTEMPTS: usize = 10_000; + + let shared = ProxySharedState::new(); + std::thread::scope(|scope| { + for worker in 0..16 { + let shared = Arc::clone(&shared); + scope.spawn(move || { + for index in (worker..ATTEMPTS).step_by(16) { + let octets = (index as u32).to_be_bytes(); + let peer_ip = IpAddr::V4(std::net::Ipv4Addr::new( + octets[1], octets[2], octets[3], worker as u8, + )); + sticky_hint_record_success_in( + shared.as_ref(), + peer_ip, + index as u32, + Some(&format!("host-{index}.example")), + ); + } + }); + } + }); + + assert_eq!(shared.handshake.sticky_user_by_ip.len(), STICKY_HINT_MAX_ENTRIES); + assert_eq!( + shared.handshake.sticky_user_by_ip_prefix.len(), + STICKY_HINT_MAX_ENTRIES + ); + assert_eq!( + shared.handshake.sticky_user_by_sni_hash.len(), + STICKY_HINT_MAX_ENTRIES + ); + assert_eq!( + shared.handshake.sticky_user_by_ip_slots.used(), + shared.handshake.sticky_user_by_ip.len() + ); + assert_eq!( + shared.handshake.sticky_user_by_ip_prefix_slots.used(), + shared.handshake.sticky_user_by_ip_prefix.len() + ); + assert_eq!( + shared.handshake.sticky_user_by_sni_hash_slots.used(), + shared.handshake.sticky_user_by_sni_hash.len() + ); + } +} + pub(super) fn decode_user_secrets_in( shared: &ProxySharedState, config: &ProxyConfig, diff --git a/src/proxy/handshake/auth_probe.rs b/src/proxy/handshake/auth_probe.rs index 13e0428..daec4cc 100644 --- a/src/proxy/handshake/auth_probe.rs +++ b/src/proxy/handshake/auth_probe.rs @@ -98,9 +98,14 @@ pub(super) fn auth_probe_is_throttled_in( }; if auth_probe_state_expired(&entry, now) { drop(entry); - state.remove_if(&peer_ip, |_, current| { - auth_probe_state_expired(current, now) - }); + if state + .remove_if(&peer_ip, |_, current| { + auth_probe_state_expired(current, now) + }) + .is_some() + { + shared.handshake.auth_probe_slots.release(); + } return false; } now < entry.blocked_until @@ -118,9 +123,14 @@ pub(super) fn auth_probe_saturation_grace_exhausted_in( }; if auth_probe_state_expired(&entry, now) { drop(entry); - state.remove_if(&peer_ip, |_, current| { - auth_probe_state_expired(current, now) - }); + if state + .remove_if(&peer_ip, |_, current| { + auth_probe_state_expired(current, now) + }) + .is_some() + { + shared.handshake.auth_probe_slots.release(); + } return false; } @@ -216,7 +226,13 @@ pub(super) fn auth_probe_record_failure_in( ) { let peer_ip = normalize_auth_probe_ip(peer_ip); let state = &shared.handshake.auth_probe; - auth_probe_record_failure_with_state_in(shared, state, peer_ip, now); + auth_probe_record_failure_with_state_and_budget_in( + shared, + state, + Some(&shared.handshake.auth_probe_slots), + peer_ip, + now, + ); } pub(super) fn auth_probe_record_failure_with_state_in( @@ -224,6 +240,16 @@ pub(super) fn auth_probe_record_failure_with_state_in( state: &DashMap, peer_ip: IpAddr, now: Instant, +) { + auth_probe_record_failure_with_state_and_budget_in(shared, state, None, peer_ip, now); +} + +fn auth_probe_record_failure_with_state_and_budget_in( + shared: &ProxySharedState, + state: &DashMap, + slots: Option<&crate::slot_budget::SlotBudget>, + peer_ip: IpAddr, + now: Instant, ) { let make_new_state = || AuthProbeState { fail_streak: 1, @@ -279,6 +305,9 @@ pub(super) fn auth_probe_record_failure_with_state_in( }) .is_some() { + if let Some(slots) = slots { + slots.release(); + } break; } continue; @@ -347,9 +376,15 @@ pub(super) fn auth_probe_record_failure_with_state_in( } for stale_key in stale_keys { - state.remove_if(&stale_key, |_, current| { - auth_probe_state_expired(current, now) - }); + if state + .remove_if(&stale_key, |_, current| { + auth_probe_state_expired(current, now) + }) + .is_some() + && let Some(slots) = slots + { + slots.release(); + } } if state.len() < AUTH_PROBE_TRACK_MAX_ENTRIES { @@ -360,19 +395,38 @@ pub(super) fn auth_probe_record_failure_with_state_in( auth_probe_note_saturation_in(shared, now); return; }; - state.remove_if(&evict_key, |_, current| { - current.fail_streak == evict_fail_streak && current.last_seen == evict_last_seen - }); + if state + .remove_if(&evict_key, |_, current| { + current.fail_streak == evict_fail_streak + && current.last_seen == evict_last_seen + }) + .is_some() + && let Some(slots) = slots + { + slots.release(); + } auth_probe_note_saturation_in(shared, now); } } + let slot = if let Some(slots) = slots { + let Some(slot) = slots.try_acquire() else { + auth_probe_note_saturation_in(shared, now); + return; + }; + Some(slot) + } else { + None + }; match state.entry(peer_ip) { Entry::Occupied(mut entry) => { update_existing(entry.get_mut()); } Entry::Vacant(entry) => { entry.insert(make_new_state()); + if let Some(slot) = slot { + slot.commit(); + } } } } @@ -380,125 +434,15 @@ pub(super) fn auth_probe_record_failure_with_state_in( pub(super) fn auth_probe_record_success_in(shared: &ProxySharedState, peer_ip: IpAddr) { let peer_ip = normalize_auth_probe_ip(peer_ip); let state = &shared.handshake.auth_probe; - state.remove(&peer_ip); -} - -#[cfg(test)] -pub(crate) fn auth_probe_record_failure_for_testing( - shared: &ProxySharedState, - peer_ip: IpAddr, - now: Instant, -) { - auth_probe_record_failure_in(shared, peer_ip, now); -} - -#[cfg(test)] -pub(crate) fn auth_probe_fail_streak_for_testing_in_shared( - shared: &ProxySharedState, - peer_ip: IpAddr, -) -> Option { - let peer_ip = normalize_auth_probe_ip(peer_ip); - shared - .handshake - .auth_probe - .get(&peer_ip) - .map(|entry| entry.fail_streak) -} - -#[cfg(test)] -pub(crate) fn clear_auth_probe_state_for_testing_in_shared(shared: &ProxySharedState) { - shared.handshake.auth_probe.clear(); - match shared.handshake.auth_probe_saturation.lock() { - Ok(mut saturation) => { - *saturation = None; - } - Err(poisoned) => { - let mut saturation = poisoned.into_inner(); - *saturation = None; - shared.handshake.auth_probe_saturation.clear_poison(); - } + if state.remove(&peer_ip).is_some() { + shared.handshake.auth_probe_slots.release(); } } #[cfg(test)] -pub(crate) fn auth_probe_state_for_testing_in_shared( - shared: &ProxySharedState, -) -> &DashMap { - &shared.handshake.auth_probe -} - +mod testing; #[cfg(test)] -pub(crate) fn auth_probe_saturation_state_for_testing_in_shared( - shared: &ProxySharedState, -) -> &Mutex> { - &shared.handshake.auth_probe_saturation -} - -#[cfg(test)] -pub(crate) fn auth_probe_saturation_state_lock_for_testing_in_shared( - shared: &ProxySharedState, -) -> std::sync::MutexGuard<'_, Option> { - shared - .handshake - .auth_probe_saturation - .lock() - .unwrap_or_else(|poisoned| poisoned.into_inner()) -} - -#[cfg(test)] -pub(crate) fn clear_unknown_sni_warn_state_for_testing_in_shared(shared: &ProxySharedState) { - let mut guard = shared - .handshake - .unknown_sni_warn_next_allowed - .lock() - .unwrap_or_else(|poisoned| poisoned.into_inner()); - *guard = None; -} - -#[cfg(test)] -pub(crate) fn should_emit_unknown_sni_warn_for_testing_in_shared( - shared: &ProxySharedState, - now: Instant, -) -> bool { - should_emit_unknown_sni_warn_in(shared, now) -} - -#[cfg(test)] -pub(crate) fn clear_warned_secrets_for_testing_in_shared(shared: &ProxySharedState) { - if let Ok(mut guard) = shared.handshake.invalid_secret_warned.lock() { - guard.clear(); - } -} - -#[cfg(test)] -pub(crate) fn warned_secrets_for_testing_in_shared( - shared: &ProxySharedState, -) -> &Mutex> { - &shared.handshake.invalid_secret_warned -} - -#[cfg(test)] -pub(crate) fn auth_probe_is_throttled_for_testing_in_shared( - shared: &ProxySharedState, - peer_ip: IpAddr, -) -> bool { - auth_probe_is_throttled_in(shared, peer_ip, Instant::now()) -} - -#[cfg(test)] -pub(crate) fn auth_probe_saturation_is_throttled_for_testing_in_shared( - shared: &ProxySharedState, -) -> bool { - auth_probe_saturation_is_throttled_in(shared, Instant::now()) -} - -#[cfg(test)] -pub(crate) fn auth_probe_saturation_is_throttled_at_for_testing_in_shared( - shared: &ProxySharedState, - now: Instant, -) -> bool { - auth_probe_saturation_is_throttled_in(shared, now) -} +pub(crate) use testing::*; #[inline] pub(super) fn find_matching_tls_domain<'a>(config: &'a ProxyConfig, sni: &str) -> Option<&'a str> { diff --git a/src/proxy/handshake/auth_probe/testing.rs b/src/proxy/handshake/auth_probe/testing.rs new file mode 100644 index 0000000..9d5c4e8 --- /dev/null +++ b/src/proxy/handshake/auth_probe/testing.rs @@ -0,0 +1,137 @@ +use super::*; + +pub(crate) fn auth_probe_record_failure_for_testing( + shared: &ProxySharedState, + peer_ip: IpAddr, + now: Instant, +) { + auth_probe_record_failure_in(shared, peer_ip, now); +} + +pub(crate) fn auth_probe_fail_streak_for_testing_in_shared( + shared: &ProxySharedState, + peer_ip: IpAddr, +) -> Option { + let peer_ip = normalize_auth_probe_ip(peer_ip); + shared + .handshake + .auth_probe + .get(&peer_ip) + .map(|entry| entry.fail_streak) +} + +pub(crate) fn clear_auth_probe_state_for_testing_in_shared(shared: &ProxySharedState) { + shared.handshake.auth_probe.clear(); + shared.handshake.auth_probe_slots.reset_for_testing(); + match shared.handshake.auth_probe_saturation.lock() { + Ok(mut saturation) => { + *saturation = None; + } + Err(poisoned) => { + let mut saturation = poisoned.into_inner(); + *saturation = None; + shared.handshake.auth_probe_saturation.clear_poison(); + } + } +} + +pub(crate) fn auth_probe_state_for_testing_in_shared( + shared: &ProxySharedState, +) -> &DashMap { + &shared.handshake.auth_probe +} + +pub(crate) fn auth_probe_slots_for_testing_in_shared(shared: &ProxySharedState) -> usize { + shared.handshake.auth_probe_slots.used() +} + +pub(crate) fn auth_probe_saturation_state_for_testing_in_shared( + shared: &ProxySharedState, +) -> &Mutex> { + &shared.handshake.auth_probe_saturation +} + +pub(crate) fn auth_probe_saturation_state_lock_for_testing_in_shared( + shared: &ProxySharedState, +) -> std::sync::MutexGuard<'_, Option> { + shared + .handshake + .auth_probe_saturation + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) +} + +pub(crate) fn clear_unknown_sni_warn_state_for_testing_in_shared(shared: &ProxySharedState) { + let mut guard = shared + .handshake + .unknown_sni_warn_next_allowed + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + *guard = None; +} + +pub(crate) fn should_emit_unknown_sni_warn_for_testing_in_shared( + shared: &ProxySharedState, + now: Instant, +) -> bool { + should_emit_unknown_sni_warn_in(shared, now) +} + +pub(crate) fn clear_warned_secrets_for_testing_in_shared(shared: &ProxySharedState) { + if let Ok(mut guard) = shared.handshake.invalid_secret_warned.lock() { + guard.clear(); + } +} + +pub(crate) fn warned_secrets_for_testing_in_shared( + shared: &ProxySharedState, +) -> &Mutex> { + &shared.handshake.invalid_secret_warned +} + +pub(crate) fn auth_probe_is_throttled_for_testing_in_shared( + shared: &ProxySharedState, + peer_ip: IpAddr, +) -> bool { + auth_probe_is_throttled_in(shared, peer_ip, Instant::now()) +} + +pub(crate) fn auth_probe_saturation_is_throttled_for_testing_in_shared( + shared: &ProxySharedState, +) -> bool { + auth_probe_saturation_is_throttled_in(shared, Instant::now()) +} + +pub(crate) fn auth_probe_saturation_is_throttled_at_for_testing_in_shared( + shared: &ProxySharedState, + now: Instant, +) -> bool { + auth_probe_saturation_is_throttled_in(shared, now) +} + +#[test] +fn parallel_distinct_failures_respect_exact_auth_probe_capacity() { + const ATTEMPTS: usize = 10_000; + + let shared = ProxySharedState::new(); + std::thread::scope(|scope| { + for worker in 0..16 { + let shared = Arc::clone(&shared); + scope.spawn(move || { + for index in (worker..ATTEMPTS).step_by(16) { + let octets = (index as u32).to_be_bytes(); + let peer_ip = IpAddr::V4(std::net::Ipv4Addr::new( + octets[1], octets[2], octets[3], worker as u8, + )); + auth_probe_record_failure_in(shared.as_ref(), peer_ip, Instant::now()); + } + }); + } + }); + + assert_eq!(shared.handshake.auth_probe.len(), AUTH_PROBE_TRACK_MAX_ENTRIES); + assert_eq!( + auth_probe_slots_for_testing_in_shared(shared.as_ref()), + AUTH_PROBE_TRACK_MAX_ENTRIES + ); +} diff --git a/src/proxy/shared_state.rs b/src/proxy/shared_state.rs index 706388d..a3c7e67 100644 --- a/src/proxy/shared_state.rs +++ b/src/proxy/shared_state.rs @@ -16,6 +16,7 @@ use crate::proxy::user_admission::{ UserAdmissionAuthority, UserAdmissionPublication, UserCredentialId, UserIncarnation, UserMutationResult, UserSessionRegistration, }; +use crate::slot_budget::SlotBudget; const HANDSHAKE_RECENT_USER_RING_LEN: usize = 64; const MASKING_FALLBACK_MAX_CONCURRENT: usize = 512; @@ -55,13 +56,17 @@ pub(crate) enum ConntrackClosePolicy { pub(crate) struct HandshakeSharedState { pub(crate) auth_probe: DashMap, + pub(crate) auth_probe_slots: SlotBudget, pub(crate) auth_probe_saturation: Mutex>, pub(crate) auth_probe_eviction_hasher: RandomState, pub(crate) invalid_secret_warned: Mutex>, pub(crate) unknown_sni_warn_next_allowed: Mutex>, pub(crate) sticky_user_by_ip: DashMap, + pub(crate) sticky_user_by_ip_slots: SlotBudget, pub(crate) sticky_user_by_ip_prefix: DashMap, + pub(crate) sticky_user_by_ip_prefix_slots: SlotBudget, pub(crate) sticky_user_by_sni_hash: DashMap, + pub(crate) sticky_user_by_sni_hash_slots: SlotBudget, pub(crate) recent_user_ring: Box<[AtomicU32]>, pub(crate) recent_user_ring_seq: AtomicU64, pub(crate) auth_expensive_checks_total: AtomicU64, @@ -114,13 +119,25 @@ impl ProxySharedState { Arc::new(Self { handshake: HandshakeSharedState { auth_probe: DashMap::new(), + auth_probe_slots: SlotBudget::new( + crate::proxy::handshake::AUTH_PROBE_TRACK_MAX_ENTRIES, + ), auth_probe_saturation: Mutex::new(None), auth_probe_eviction_hasher: RandomState::new(), invalid_secret_warned: Mutex::new(HashSet::new()), unknown_sni_warn_next_allowed: Mutex::new(None), sticky_user_by_ip: DashMap::new(), + sticky_user_by_ip_slots: SlotBudget::new( + crate::proxy::handshake::STICKY_HINT_MAX_ENTRIES, + ), sticky_user_by_ip_prefix: DashMap::new(), + sticky_user_by_ip_prefix_slots: SlotBudget::new( + crate::proxy::handshake::STICKY_HINT_MAX_ENTRIES, + ), sticky_user_by_sni_hash: DashMap::new(), + sticky_user_by_sni_hash_slots: SlotBudget::new( + crate::proxy::handshake::STICKY_HINT_MAX_ENTRIES, + ), recent_user_ring: std::iter::repeat_with(|| AtomicU32::new(0)) .take(HANDSHAKE_RECENT_USER_RING_LEN) .collect::>() diff --git a/src/slot_budget.rs b/src/slot_budget.rs new file mode 100644 index 0000000..005fb56 --- /dev/null +++ b/src/slot_budget.rs @@ -0,0 +1,136 @@ +use std::sync::atomic::{AtomicUsize, Ordering}; + +/// Exact lock-free admission budget for bounded concurrent registries. +pub(crate) struct SlotBudget { + capacity: usize, + used: AtomicUsize, +} + +impl SlotBudget { + /// Creates an empty budget with the given hard capacity. + pub(crate) const fn new(capacity: usize) -> Self { + Self { + capacity, + used: AtomicUsize::new(0), + } + } + + /// Reserves one slot or returns `None` when the hard capacity is exhausted. + pub(crate) fn try_acquire(&self) -> Option> { + let mut current = self.used.load(Ordering::Acquire); + loop { + if current >= self.capacity { + return None; + } + match self.used.compare_exchange_weak( + current, + current + 1, + Ordering::AcqRel, + Ordering::Acquire, + ) { + Ok(_) => { + return Some(SlotLease { + budget: self, + armed: true, + }); + } + Err(actual) => current = actual, + } + } + } + + /// Releases one slot after its registry entry has been removed. + pub(crate) fn release(&self) { + self.release_many(1); + } + + /// Releases a known number of slots after a batch registry removal completes. + pub(crate) fn release_many(&self, amount: usize) { + if amount == 0 { + return; + } + let released = self + .used + .fetch_update(Ordering::AcqRel, Ordering::Acquire, |current| { + current.checked_sub(amount) + }); + #[cfg(not(test))] + debug_assert!(released.is_ok(), "slot budget release must match acquisitions"); + #[cfg(test)] + let _ = released; + } + + /// Returns the exact number of currently committed or reserved slots. + pub(crate) fn used(&self) -> usize { + self.used.load(Ordering::Acquire) + } + + #[cfg(test)] + pub(crate) fn reset_for_testing(&self) { + self.used.store(0, Ordering::Release); + } +} + +/// Provisional slot ownership that rolls back unless committed to a registry entry. +pub(crate) struct SlotLease<'a> { + budget: &'a SlotBudget, + armed: bool, +} + +impl SlotLease<'_> { + /// Transfers the reserved slot to the registry entry being published. + pub(crate) fn commit(mut self) { + self.armed = false; + } +} + +impl Drop for SlotLease<'_> { + fn drop(&mut self) { + if self.armed { + self.budget.release(); + } + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + use std::sync::atomic::{AtomicUsize, Ordering}; + + use super::*; + + #[test] + fn parallel_reservations_never_exceed_the_hard_capacity() { + const CAPACITY: usize = 127; + const ATTEMPTS: usize = 10_000; + + let budget = Arc::new(SlotBudget::new(CAPACITY)); + let admitted = Arc::new(AtomicUsize::new(0)); + std::thread::scope(|scope| { + for worker in 0..16 { + let budget = Arc::clone(&budget); + let admitted = Arc::clone(&admitted); + scope.spawn(move || { + for attempt in (worker..ATTEMPTS).step_by(16) { + if let Some(lease) = budget.try_acquire() { + admitted.fetch_add(1, Ordering::Relaxed); + lease.commit(); + } + std::hint::black_box(attempt); + } + }); + } + }); + + assert_eq!(admitted.load(Ordering::Relaxed), CAPACITY); + assert_eq!(budget.used(), CAPACITY); + } + + #[test] + fn dropped_provisional_lease_returns_its_slot() { + let budget = SlotBudget::new(1); + drop(budget.try_acquire().unwrap()); + + assert!(budget.try_acquire().is_some()); + } +} diff --git a/src/stats/tls_fingerprints.rs b/src/stats/tls_fingerprints.rs index 424aa4e..5655d52 100644 --- a/src/stats/tls_fingerprints.rs +++ b/src/stats/tls_fingerprints.rs @@ -10,6 +10,7 @@ use dashmap::DashMap; use dashmap::mapref::entry::Entry; use crate::protocol::tls_fingerprint::TlsClientFingerprint; +use crate::slot_budget::SlotBudget; use super::Stats; @@ -68,15 +69,33 @@ struct TlsFingerprintEntry { bad_or_probe: AtomicU64, } -#[derive(Default)] pub struct TlsFingerprintCollector { entries: DashMap, + slots: SlotBudget, + capacity: usize, dropped_total: AtomicU64, parse_error_total: AtomicU64, last_cleanup_epoch_secs: AtomicU64, } +impl Default for TlsFingerprintCollector { + fn default() -> Self { + Self::with_capacity(MAX_TLS_FINGERPRINT_BUCKETS) + } +} + impl TlsFingerprintCollector { + fn with_capacity(capacity: usize) -> Self { + Self { + entries: DashMap::new(), + slots: SlotBudget::new(capacity), + capacity, + dropped_total: AtomicU64::new(0), + parse_error_total: AtomicU64::new(0), + last_cleanup_epoch_secs: AtomicU64::new(0), + } + } + pub fn record_observed( &self, fingerprint: &TlsClientFingerprint, @@ -228,7 +247,7 @@ impl TlsFingerprintCollector { TlsFingerprintSnapshot { retention_secs: ttl.as_secs(), - capacity: MAX_TLS_FINGERPRINT_BUCKETS, + capacity: self.capacity, dropped_total: self.dropped_total.load(Ordering::Relaxed), parse_error_total: self.parse_error_total.load(Ordering::Relaxed), by_fingerprint, @@ -286,22 +305,6 @@ impl TlsFingerprintCollector { ja4_raw: fingerprint.ja4_raw.clone(), }; - if let Some(entry) = self.entries.get(&key) { - update_entry( - entry.value(), - now_epoch_secs, - count_total, - count_auth_success, - count_bad_or_probe, - ); - return; - } - - if self.entries.len() >= MAX_TLS_FINGERPRINT_BUCKETS { - self.dropped_total.fetch_add(1, Ordering::Relaxed); - return; - } - match self.entries.entry(key) { Entry::Occupied(entry) => { update_entry( @@ -313,12 +316,17 @@ impl TlsFingerprintCollector { ); } Entry::Vacant(entry) => { + let Some(slot) = self.slots.try_acquire() else { + self.dropped_total.fetch_add(1, Ordering::Relaxed); + return; + }; entry.insert(TlsFingerprintEntry::new( now_epoch_secs, if count_total { 1 } else { 0 }, if count_auth_success { 1 } else { 0 }, if count_bad_or_probe { 1 } else { 0 }, )); + slot.commit(); } } } @@ -339,14 +347,17 @@ impl TlsFingerprintCollector { } fn cleanup(&self, now_epoch_secs: u64, ttl_secs: u64) { - if ttl_secs == 0 { - self.entries.clear(); - return; - } + let mut removed = 0usize; self.entries.retain(|_, entry| { let last_seen = entry.last_seen_epoch_secs.load(Ordering::Relaxed); - now_epoch_secs.saturating_sub(last_seen) <= ttl_secs + let retained = + ttl_secs != 0 && now_epoch_secs.saturating_sub(last_seen) <= ttl_secs; + if !retained { + removed += 1; + } + retained }); + self.slots.release_many(removed); } } @@ -526,31 +537,4 @@ impl Stats { } #[cfg(test)] -mod tests { - use super::*; - - fn fp() -> TlsClientFingerprint { - TlsClientFingerprint { - ja3: "ja3".to_string(), - ja3_raw: "771,4865,,,0".to_string(), - ja4: "t13d010100_hash_hash".to_string(), - ja4_raw: "raw".to_string(), - } - } - - #[test] - fn aggregates_ip_cidr_and_user_scopes() { - let collector = TlsFingerprintCollector::default(); - let ip: IpAddr = "192.0.2.15".parse().expect("test IP parses"); - collector.record_observed(&fp(), ip, Duration::from_secs(60)); - collector.record_auth_success(&fp(), ip, "alice", Duration::from_secs(60)); - let snapshot = collector.snapshot(Duration::from_secs(60), 10); - - assert_eq!(snapshot.by_fingerprint[0].total, 1); - assert_eq!(snapshot.by_fingerprint[0].auth_success, 1); - assert_eq!(snapshot.by_ip[0].scope_key, "192.0.2.15"); - assert_eq!(snapshot.by_cidr[0].scope_key, "192.0.2.0/24"); - assert_eq!(snapshot.by_user[0].scope_key, "alice"); - assert_eq!(snapshot.by_user[0].total, 1); - } -} +mod tests; diff --git a/src/stats/tls_fingerprints/tests.rs b/src/stats/tls_fingerprints/tests.rs new file mode 100644 index 0000000..c0b1edb --- /dev/null +++ b/src/stats/tls_fingerprints/tests.rs @@ -0,0 +1,58 @@ +use super::*; + +fn fp() -> TlsClientFingerprint { + TlsClientFingerprint { + ja3: "ja3".to_string(), + ja3_raw: "771,4865,,,0".to_string(), + ja4: "t13d010100_hash_hash".to_string(), + ja4_raw: "raw".to_string(), + } +} + +#[test] +fn aggregates_ip_cidr_and_user_scopes() { + let collector = TlsFingerprintCollector::default(); + let ip: IpAddr = "192.0.2.15".parse().expect("test IP parses"); + collector.record_observed(&fp(), ip, Duration::from_secs(60)); + collector.record_auth_success(&fp(), ip, "alice", Duration::from_secs(60)); + let snapshot = collector.snapshot(Duration::from_secs(60), 10); + + assert_eq!(snapshot.by_fingerprint[0].total, 1); + assert_eq!(snapshot.by_fingerprint[0].auth_success, 1); + assert_eq!(snapshot.by_ip[0].scope_key, "192.0.2.15"); + assert_eq!(snapshot.by_cidr[0].scope_key, "192.0.2.0/24"); + assert_eq!(snapshot.by_user[0].scope_key, "alice"); + assert_eq!(snapshot.by_user[0].total, 1); +} + +#[test] +fn parallel_distinct_insertions_respect_exact_capacity() { + const CAPACITY: usize = 127; + const ATTEMPTS: usize = 10_000; + + let collector = std::sync::Arc::new(TlsFingerprintCollector::with_capacity(CAPACITY)); + std::thread::scope(|scope| { + for worker in 0..16 { + let collector = std::sync::Arc::clone(&collector); + scope.spawn(move || { + for index in (worker..ATTEMPTS).step_by(16) { + collector.record_scoped( + (TlsFingerprintScopeKind::Ip, format!("192.0.2.{index}")), + &fp(), + 1, + true, + false, + false, + ); + } + }); + } + }); + + assert_eq!(collector.entries.len(), CAPACITY); + assert_eq!(collector.slots.used(), CAPACITY); + assert_eq!( + collector.dropped_total.load(Ordering::Relaxed), + (ATTEMPTS - CAPACITY) as u64 + ); +} diff --git a/src/synlimit_control/command.rs b/src/synlimit_control/command.rs index 19355c0..0376c8a 100644 --- a/src/synlimit_control/command.rs +++ b/src/synlimit_control/command.rs @@ -1,14 +1,14 @@ -use std::path::PathBuf; - use tokio::io::AsyncWriteExt; use tokio::process::Command; +use crate::util::trusted_command::resolve_trusted_helper; + pub(super) async fn run_command( binary: &str, args: &[&str], stdin: Option, ) -> Result<(), String> { - let Some(command_path) = resolve_command(binary) else { + let Some(command_path) = resolve_trusted_helper(binary) else { return Err(format!("{binary} is not available")); }; let mut command = Command::new(command_path); @@ -45,7 +45,7 @@ pub(super) async fn run_command( } pub(super) async fn run_command_stdout(binary: &str, args: &[&str]) -> Result { - let Some(command_path) = resolve_command(binary) else { + let Some(command_path) = resolve_trusted_helper(binary) else { return Err(format!("{binary} is not available")); }; let output = Command::new(command_path) @@ -64,16 +64,6 @@ pub(super) async fn run_command_stdout(binary: &str, args: &[&str]) -> Result Option { - let mut dirs = std::env::var_os("PATH") - .map(|path| std::env::split_paths(&path).collect::>()) - .unwrap_or_default(); - dirs.extend(["/usr/sbin", "/sbin", "/usr/bin", "/bin"].map(PathBuf::from)); - dirs.into_iter() - .map(|dir| dir.join(binary)) - .find(|candidate| candidate.exists() && candidate.is_file()) -} - pub(super) fn has_firewall_privileges() -> bool { #[cfg(target_os = "linux")] { diff --git a/src/tls_front/fetcher.rs b/src/tls_front/fetcher.rs index e49ae86..a831f1f 100644 --- a/src/tls_front/fetcher.rs +++ b/src/tls_front/fetcher.rs @@ -172,6 +172,12 @@ fn sweep_expired_profile_cache(ttl: Duration, now: Instant) { profile_cache().retain(|_, value| now.saturating_duration_since(value.updated_at) <= ttl); } +fn remove_profile_if_unchanged(key: &ProfileCacheKey, observed: ProfileCacheValue) { + profile_cache().remove_if(key, |_, current| { + current.profile == observed.profile && current.updated_at == observed.updated_at + }); +} + /// Current number of adaptive TLS fetch profile-cache entries. pub(crate) fn profile_cache_entries_for_metrics() -> usize { profile_cache().len() @@ -270,8 +276,9 @@ fn order_profiles( if let Some(cached) = profile_cache().get(key) { let age = now.saturating_duration_since(cached.updated_at); if age > strategy.profile_cache_ttl { + let observed = *cached; drop(cached); - profile_cache().remove(key); + remove_profile_if_unchanged(key, observed); return ordered; } diff --git a/src/tls_front/fetcher/tests.rs b/src/tls_front/fetcher/tests.rs index 2309754..63556f3 100644 --- a/src/tls_front/fetcher/tests.rs +++ b/src/tls_front/fetcher/tests.rs @@ -6,6 +6,7 @@ use super::{ TLS_NAMED_GROUP_X25519MLKEM768, TlsFetchStrategy, X25519_KEY_SHARE_LEN, build_client_hello, build_tls_fetch_proxy_header, derive_behavior_profile, encode_tls13_certificate_message, fetch_via_rustls_stream, order_profiles, profile_alpn, profile_cache, profile_cache_key, + remove_profile_if_unchanged, }; use crate::config::TlsFetchProfile; use crate::crypto::SecureRandom; @@ -225,6 +226,31 @@ fn test_order_profiles_drops_expired_cached_winner() { assert!(profile_cache().get(&cache_key).is_none()); } +#[test] +fn expired_profile_removal_preserves_concurrent_refresh() { + let cache_key = profile_cache_key("mask3.example", 443, "tls3.example", None, None, 0, None); + let observed = ProfileCacheValue { + profile: TlsFetchProfile::CompatTls12, + updated_at: Instant::now() - Duration::from_secs(60), + }; + let refreshed = ProfileCacheValue { + profile: TlsFetchProfile::ModernChromeLike, + updated_at: Instant::now(), + }; + profile_cache().insert(cache_key.clone(), observed); + profile_cache().insert(cache_key.clone(), refreshed); + + remove_profile_if_unchanged(&cache_key, observed); + + let current = profile_cache() + .get(&cache_key) + .expect("concurrent refresh must remain cached"); + assert_eq!(current.profile, refreshed.profile); + assert_eq!(current.updated_at, refreshed.updated_at); + drop(current); + profile_cache().remove(&cache_key); +} + #[test] fn test_deterministic_client_hello_is_stable() { let rng = SecureRandom::new(); diff --git a/src/util/mod.rs b/src/util/mod.rs index a81ef10..7698354 100644 --- a/src/util/mod.rs +++ b/src/util/mod.rs @@ -2,6 +2,8 @@ pub mod ip; pub mod time; +#[cfg(unix)] +pub mod trusted_command; #[allow(unused_imports)] pub use ip::*; diff --git a/src/util/trusted_command.rs b/src/util/trusted_command.rs new file mode 100644 index 0000000..676350d --- /dev/null +++ b/src/util/trusted_command.rs @@ -0,0 +1,70 @@ +use std::os::unix::fs::{MetadataExt, PermissionsExt}; +use std::path::{Path, PathBuf}; + +const TRUSTED_HELPER_DIRS: [&str; 4] = ["/usr/sbin", "/usr/bin", "/sbin", "/bin"]; +const TRUSTED_HELPERS: [&str; 5] = ["nft", "iptables", "ip6tables", "conntrack", "pfctl"]; + +/// Resolves a privileged helper only through the fixed system allowlist. +pub(crate) fn resolve_trusted_helper(binary: &str) -> Option { + if !TRUSTED_HELPERS.contains(&binary) { + return None; + } + + TRUSTED_HELPER_DIRS + .iter() + .map(|directory| Path::new(directory).join(binary)) + .find_map(|candidate| trusted_executable(&candidate)) +} + +fn trusted_executable(candidate: &Path) -> Option { + let canonical = std::fs::canonicalize(candidate).ok()?; + let metadata = std::fs::metadata(&canonical).ok()?; + if !metadata.is_file() + || metadata.uid() != 0 + || metadata.permissions().mode() & 0o022 != 0 + || metadata.permissions().mode() & 0o111 == 0 + || !trusted_parent_chain(canonical.parent()?) + { + return None; + } + Some(canonical) +} + +fn trusted_parent_chain(path: &Path) -> bool { + for ancestor in path.ancestors() { + let Ok(metadata) = std::fs::symlink_metadata(ancestor) else { + return false; + }; + if !metadata.is_dir() + || metadata.file_type().is_symlink() + || metadata.uid() != 0 + || metadata.permissions().mode() & 0o022 != 0 + { + return false; + } + } + true +} + +#[cfg(test)] +mod tests { + use std::os::unix::fs::PermissionsExt; + + use super::*; + + #[test] + fn privileged_helper_allowlist_rejects_arbitrary_binary() { + assert!(resolve_trusted_helper("sh").is_none()); + assert!(resolve_trusted_helper("../bin/nft").is_none()); + } + + #[test] + fn writable_executable_is_not_trusted() { + let directory = tempfile::tempdir().unwrap(); + let executable = directory.path().join("nft"); + std::fs::write(&executable, b"#!/bin/sh\nexit 0\n").unwrap(); + std::fs::set_permissions(&executable, std::fs::Permissions::from_mode(0o777)).unwrap(); + + assert!(trusted_executable(&executable).is_none()); + } +} diff --git a/src/web/trace/exchange.rs b/src/web/trace/exchange.rs index 1d95c6c..25ecec9 100644 --- a/src/web/trace/exchange.rs +++ b/src/web/trace/exchange.rs @@ -1,6 +1,5 @@ use std::net::IpAddr; use std::sync::Arc; -use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; use std::time::Instant; use parking_lot::Mutex; @@ -29,6 +28,8 @@ struct BodyCapture { } struct ExchangeState { + phase: ExchangePhase, + reserved: usize, method: String, path: String, route: TraceRoute, @@ -47,6 +48,12 @@ struct ExchangeState { body_capture_blocked: bool, } +#[derive(Clone, Copy, PartialEq, Eq)] +enum ExchangePhase { + Open, + Committed, +} + /// One in-flight request-to-response capture with process-wide byte leases. pub(crate) struct HttpTraceExchange { store: Arc, @@ -55,8 +62,6 @@ pub(crate) struct HttpTraceExchange { started: Instant, started_epoch_millis: u64, state: Mutex, - reserved: AtomicUsize, - committed: AtomicBool, } impl HttpTraceExchange { @@ -104,6 +109,8 @@ impl HttpTraceExchange { started, started_epoch_millis, state: Mutex::new(ExchangeState { + phase: ExchangePhase::Open, + reserved: base_reservation + if dynamic_reserved { dynamic } else { 0 }, method, path, route: TraceRoute::Unknown, @@ -121,21 +128,23 @@ impl HttpTraceExchange { redactions, body_capture_blocked: !dynamic_reserved, }), - reserved: AtomicUsize::new( - base_reservation + if dynamic_reserved { dynamic } else { 0 }, - ), - committed: AtomicBool::new(false), }) } /// Sets the final request route before body polling or decoy forwarding. pub(crate) fn set_route(&self, route: TraceRoute) { - self.state.lock().route = route; + let mut state = self.state.lock(); + if state.phase == ExchangePhase::Open { + state.route = route; + } } /// Sets the trusted effective client address after proxy-header validation. pub(crate) fn set_effective_ip(&self, client_ip: IpAddr) { - self.state.lock().effective_ip = Some(client_ip); + let mut state = self.state.lock(); + if state.phase == ExchangePhase::Open { + state.effective_ip = Some(client_ip); + } } /// Binds non-secret profile and process session identity. @@ -145,8 +154,11 @@ impl HttpTraceExchange { .len() .saturating_add(profile.key_fingerprint.len()); let mut state = self.state.lock(); + if state.phase != ExchangePhase::Open { + return; + } state.identity.session_id = Some(session_id); - if self.reserve(dynamic) { + if self.reserve_locked(&mut state, dynamic) { state.identity.user = Some(profile.user.clone()); state.identity.key_fingerprint = Some(profile.key_fingerprint.clone()); } @@ -160,8 +172,11 @@ impl HttpTraceExchange { .map_or(0, String::len) .saturating_add(identity.key_fingerprint.as_ref().map_or(0, String::len)); let mut state = self.state.lock(); + if state.phase != ExchangePhase::Open { + return; + } state.identity.session_id = identity.session_id; - if self.reserve(dynamic) { + if self.reserve_locked(&mut state, dynamic) { state.identity.user = identity.user; state.identity.key_fingerprint = identity.key_fingerprint; } @@ -172,21 +187,25 @@ impl HttpTraceExchange { if value.is_empty() { return; } - if !self.reserve(value.len()) { - self.block_body_capture(); + let mut state = self.state.lock(); + if state.phase != ExchangePhase::Open { return; } - self.state - .lock() - .redactions - .push(Zeroizing::new(value.to_vec())); + if !self.reserve_locked(&mut state, value.len()) { + Self::block_body_capture_locked(&mut state); + return; + } + state.redactions.push(Zeroizing::new(value.to_vec())); } /// Captures response status and sanitized headers at handler completion. pub(crate) fn response_ready(&self, response: &hyper::Response) { let dynamic = response_dynamic_bytes(response, &self.policy); - let reserved = self.reserve(dynamic); let mut state = self.state.lock(); + if state.phase != ExchangePhase::Open { + return; + } + let reserved = self.reserve_locked(&mut state, dynamic); state.status = Some(response.status().as_u16()); if self.policy.capture_headers && reserved { state.response_headers = sanitized_headers(response.headers()); @@ -210,19 +229,22 @@ impl HttpTraceExchange { /// Appends one body data frame without changing the proxied bytes. pub(crate) fn body_data(&self, direction: TraceDirection, data: &[u8]) { let mut state = self.state.lock(); + if state.phase != ExchangePhase::Open { + return; + } let route = state.route; let body_capture_blocked = state.body_capture_blocked; - let body = match direction { - TraceDirection::Request => &mut state.request_body, - TraceDirection::Response => &mut state.response_body, - }; - body.observed_bytes = body - .observed_bytes - .saturating_add(u64::try_from(data.len()).unwrap_or(u64::MAX)); + { + let body = selected_body(&mut state, direction); + body.observed_bytes = body + .observed_bytes + .saturating_add(u64::try_from(data.len()).unwrap_or(u64::MAX)); + } if data.is_empty() { return; } if body_capture_blocked { + let body = selected_body(&mut state, direction); body.truncated |= !data.is_empty(); return; } @@ -230,6 +252,7 @@ impl HttpTraceExchange { else { return; }; + let body = selected_body(&mut state, direction); if body.captured.len() >= limit { if !data.is_empty() && !body.truncated { self.store.record_truncation(); @@ -237,13 +260,17 @@ impl HttpTraceExchange { body.truncated |= !data.is_empty(); return; } - if body.captured.capacity() == 0 { - if !self.reserve(limit) { + let needs_reservation = body.captured.capacity() == 0; + if needs_reservation { + if !self.reserve_locked(&mut state, limit) { + let body = selected_body(&mut state, direction); body.truncated = true; return; } + let body = selected_body(&mut state, direction); body.captured = Vec::with_capacity(limit); } + let body = selected_body(&mut state, direction); let take = data.len().min(limit - body.captured.len()); body.captured.extend_from_slice(&data[..take]); if take < data.len() { @@ -257,6 +284,9 @@ impl HttpTraceExchange { /// Marks one request or response body terminal state. pub(crate) fn body_finished(&self, direction: TraceDirection, terminal: TraceBodyState) { let mut state = self.state.lock(); + if state.phase != ExchangePhase::Open { + return; + } let body = match direction { TraceDirection::Request => &mut state.request_body, TraceDirection::Response => &mut state.response_body, @@ -293,7 +323,10 @@ impl HttpTraceExchange { .div_ceil(frame::HEADER_BYTES) .clamp(1, limits.max_frames_per_body); let reservation = estimated_frames.saturating_mul(std::mem::size_of::()); - if !self.reserve(reservation) { + let mut state = self.state.lock(); + if state.phase != ExchangePhase::Open + || !self.reserve_locked(&mut state, reservation) + { return; } let frames = match frame::parse_all(body, limits) { @@ -319,35 +352,38 @@ impl HttpTraceExchange { parse_error: Some(frame_error_name(error)), }], }; - self.state.lock().frames.extend(frames); + state.frames.extend(frames); } /// Commits once after response body consumption or drop. pub(crate) fn commit(&self) { - if self.committed.swap(true, Ordering::AcqRel) { - return; - } - let reserved = self.reserved.load(Ordering::Acquire); - let record = self.build_record(); + let (record, reserved) = { + let mut state = self.state.lock(); + if state.phase != ExchangePhase::Open { + return; + } + state.phase = ExchangePhase::Committed; + let reserved = state.reserved; + (self.build_record_locked(&mut state), reserved) + }; if !self.store.try_commit(record, reserved, self.epoch) { self.store.release(reserved); } } - fn reserve(&self, bytes: usize) -> bool { + fn reserve_locked(&self, state: &mut ExchangeState, bytes: usize) -> bool { if bytes == 0 { return true; } if self.store.try_reserve(bytes) { - self.reserved.fetch_add(bytes, Ordering::AcqRel); + state.reserved = state.reserved.saturating_add(bytes); true } else { false } } - fn block_body_capture(&self) { - let mut state = self.state.lock(); + fn block_body_capture_locked(state: &mut ExchangeState) { state.body_capture_blocked = true; state.request_body.captured.clear(); state.response_body.captured.clear(); @@ -355,8 +391,7 @@ impl HttpTraceExchange { state.response_body.truncated = true; } - fn build_record(&self) -> TraceRecord { - let mut state = self.state.lock(); + fn build_record_locked(&self, state: &mut ExchangeState) -> TraceRecord { let redactions = std::mem::take(&mut state.redactions); scrub_body(&mut state.request_body.captured, &redactions); scrub_body(&mut state.response_body.captured, &redactions); @@ -394,7 +429,7 @@ impl HttpTraceExchange { impl Drop for HttpTraceExchange { fn drop(&mut self) { - if !self.committed.load(Ordering::Acquire) { + if self.state.get_mut().phase == ExchangePhase::Open { self.body_finished(TraceDirection::Request, TraceBodyState::Aborted); self.body_finished(TraceDirection::Response, TraceBodyState::Aborted); } @@ -410,83 +445,12 @@ fn body_snapshot(policy: &WebDebugConfig, body: &mut BodyCapture) -> Option &mut BodyCapture { + match direction { + TraceDirection::Request => &mut state.request_body, + TraceDirection::Response => &mut state.response_body, } } + +#[cfg(test)] +mod tests; diff --git a/src/web/trace/exchange/tests.rs b/src/web/trace/exchange/tests.rs new file mode 100644 index 0000000..89aa73c --- /dev/null +++ b/src/web/trace/exchange/tests.rs @@ -0,0 +1,104 @@ +use super::*; + +fn trace_store() -> Arc { + let policy = WebDebugConfig { + enabled: true, + body_capture: WebDebugBodyCapture::Prefix, + body_prefix_bytes: 256, + ..Default::default() + }; + let limits = WebLimitsConfig { + debug_records_capacity: 4, + debug_bytes_global: 16 * 1024, + ..Default::default() + }; + WebTraceStore::new(policy, &limits) +} + +#[test] +fn request_response_capture_redacts_credentials_and_omits_query() { + let request_token = "request-token-0123456789"; + let capability = "capability-0123456789"; + let response_token = "response-token-0123456789"; + let request = hyper::Request::builder() + .uri(format!("/?bridge={capability}")) + .header("authorization", format!("Bearer {request_token}")) + .body(()) + .unwrap(); + let store = trace_store(); + let exchange = store + .begin_http(&request, "192.0.2.30".parse().unwrap()) + .unwrap(); + exchange.set_route(TraceRoute::Bridge); + exchange.body_data( + TraceDirection::Request, + format!("{request_token}:{capability}").as_bytes(), + ); + exchange.body_finished(TraceDirection::Request, TraceBodyState::Complete); + + let response = hyper::Response::builder() + .status(hyper::StatusCode::OK) + .header("x-session-token", response_token) + .body(()) + .unwrap(); + exchange.response_ready(&response); + exchange.body_data(TraceDirection::Response, response_token.as_bytes()); + exchange.body_finished(TraceDirection::Response, TraceBodyState::Complete); + + let records = store.snapshot_matching(|_| true); + assert_eq!(records.len(), 1); + let TraceRecordKind::Http(http) = &records[0].record.kind else { + panic!("expected HTTP debug record"); + }; + assert_eq!(http.path, "/"); + assert_eq!(http.route, TraceRoute::Bridge); + assert!( + http.request_headers + .iter() + .any(|header| header.name == "authorization" && header.value.is_none()) + ); + assert!( + http.response_headers + .iter() + .any(|header| header.name == "x-session-token" && header.value.is_none()) + ); + let request_body = http.request_body.as_ref().unwrap(); + let response_body = http.response_body.as_ref().unwrap(); + for secret in [request_token.as_bytes(), capability.as_bytes()] { + assert!( + !request_body + .captured + .windows(secret.len()) + .any(|value| value == secret) + ); + } + assert!( + !response_body + .captured + .windows(response_token.len()) + .any(|value| value == response_token.as_bytes()) + ); +} + +#[test] +fn late_capture_after_commit_cannot_leak_debug_byte_budget() { + let request = hyper::Request::builder().uri("/").body(()).unwrap(); + let store = trace_store(); + let exchange = store + .begin_http(&request, "192.0.2.31".parse().unwrap()) + .unwrap(); + exchange.commit(); + let committed_bytes = store.status().used_bytes; + + exchange.register_redaction(&[0x41; 2048]); + exchange.body_data(TraceDirection::Request, &[0x42; 2048]); + exchange.record_frames( + TraceDirection::Request, + &[0; frame::HEADER_BYTES], + &WebLimitsConfig::default(), + ); + + assert_eq!(store.status().used_bytes, committed_bytes); + drop(exchange); + assert_eq!(store.clear().leased_bytes, 0); +}