Slot Budget + Config Store Atomic Writer fixes + Trusted Command

This commit is contained in:
Alexey
2026-09-16 22:06:59 +03:00
parent 55f3d19ee0
commit 9a683d8b3d
33 changed files with 1680 additions and 687 deletions
+27 -3
View File
@@ -9,8 +9,11 @@ use super::ApiShared;
use super::config_store::{ use super::config_store::{
EDITABLE_SECTIONS, EDITABLE_SERVER_FIELDS, compute_snapshot_revision, is_editable_section, EDITABLE_SECTIONS, EDITABLE_SERVER_FIELDS, compute_snapshot_revision, is_editable_section,
load_candidate_snapshot, load_config_snapshot, render_server_listeners, 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 super::model::ApiFailure;
use crate::config::ProxyConfig; use crate::config::ProxyConfig;
use crate::config::hot_reload::classify_config_changes; use crate::config::hot_reload::classify_config_changes;
@@ -43,7 +46,10 @@ pub(super) struct PatchConfigResponse {
} }
struct PreparedConfigPatch { struct PreparedConfigPatch {
config_path: PathBuf,
expected_revision: String,
owner_path: PathBuf, owner_path: PathBuf,
expected_owner_contents: String,
owner_contents: String, owner_contents: String,
desired_config: Arc<ProxyConfig>, desired_config: Arc<ProxyConfig>,
response: PatchConfigResponse, response: PatchConfigResponse,
@@ -77,7 +83,14 @@ pub(super) async fn patch_config(
} else { } else {
None 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 { if let Some(reservation) = reservation {
prepared.response.reload = Some(reservation.enqueue(prepared.desired_config)); prepared.response.reload = Some(reservation.enqueue(prepared.desired_config));
} }
@@ -111,7 +124,14 @@ pub(super) async fn apply_patch_to_path(
expected_revision: Option<String>, expected_revision: Option<String>,
) -> Result<PatchConfigResponse, ApiFailure> { ) -> Result<PatchConfigResponse, ApiFailure> {
let prepared = prepare_patch_to_path(config_path, patch_json, expected_revision).await?; 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) Ok(prepared.response)
} }
@@ -197,6 +217,7 @@ async fn prepare_patch_to_path(
.get(&owner_path) .get(&owner_path)
.cloned() .cloned()
.ok_or_else(|| ApiFailure::internal("config source owner is missing from snapshot"))?; .ok_or_else(|| ApiFailure::internal("config source owner is missing from snapshot"))?;
let expected_owner_contents = owner_contents.clone();
for section in &touched { for section in &touched {
if *section == "server" { if *section == "server" {
let rendered = render_server_listeners(&requested_cfg)?; 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)?; deferred_process_fields(&old_cfg, &new_cfg).map_err(ApiFailure::bad_request)?;
Ok(PreparedConfigPatch { Ok(PreparedConfigPatch {
config_path: config_path.to_path_buf(),
expected_revision: current,
owner_path, owner_path,
expected_owner_contents,
owner_contents, owner_contents,
desired_config: Arc::new(new_cfg), desired_config: Arc::new(new_cfg),
response: PatchConfigResponse { response: PatchConfigResponse {
+24
View File
@@ -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] #[tokio::test]
async fn patch_rejects_multiple_source_owners_without_writing() { async fn patch_rejects_multiple_source_owners_without_writing() {
let dir = tempfile::tempdir().unwrap(); let dir = tempfile::tempdir().unwrap();
+11 -9
View File
@@ -10,13 +10,16 @@ use super::model::ApiFailure;
// Source-preserving TOML rendering and atomic persistence helpers. // Source-preserving TOML rendering and atomic persistence helpers.
mod persistence; mod persistence;
// Compare-and-replace file persistence and metadata preservation.
mod atomic;
#[cfg(test)] #[cfg(test)]
use persistence::{find_toml_table_bounds, render_access_section, save_sections_to_disk}; use persistence::{find_toml_table_bounds, render_access_section, save_sections_to_disk};
pub(in crate::api) use persistence::{ pub(in crate::api) use persistence::{
render_server_listeners, render_top_level_section, save_access_sections_to_disk, 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)] #[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(super) enum AccessSection { pub(super) enum AccessSection {
@@ -54,22 +57,21 @@ pub(super) fn parse_if_match(headers: &hyper::HeaderMap) -> Option<String> {
.map(|value| value.trim_matches('"').to_string()) .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, config_path: &Path,
expected_revision: Option<&str>, expected_revision: Option<&str>,
) -> Result<(), ApiFailure> { ) -> Result<(ProxyConfig, String), ApiFailure> {
let Some(expected) = expected_revision else { let loaded = load_config_snapshot(config_path, false).await?;
return Ok(()); let revision = compute_snapshot_revision(&loaded);
}; if expected_revision.is_some_and(|expected| expected != revision) {
let current = current_revision(config_path).await?;
if current != expected {
return Err(ApiFailure::new( return Err(ApiFailure::new(
hyper::StatusCode::CONFLICT, hyper::StatusCode::CONFLICT,
"revision_conflict", "revision_conflict",
"Config revision mismatch", "Config revision mismatch",
)); ));
} }
Ok(()) Ok((loaded.config, revision))
} }
pub(super) async fn current_revision(config_path: &Path) -> Result<String, ApiFailure> { pub(super) async fn current_revision(config_path: &Path) -> Result<String, ApiFailure> {
+185
View File
@@ -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<Option<ExistingTarget>> {
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::<u64>()
);
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, &current) {
(Some(expected), Some(current)) => {
same_target(&expected.metadata, &current.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
}
+6 -171
View File
@@ -1,9 +1,5 @@
use std::collections::BTreeMap; use std::collections::BTreeMap;
use std::io::{Read, Write}; use std::path::Path;
use std::path::{Path, PathBuf};
#[cfg(unix)]
use std::os::unix::fs::{MetadataExt, OpenOptionsExt, PermissionsExt};
use chrono::{DateTime, Utc}; use chrono::{DateTime, Utc};
use serde::Serialize; use serde::Serialize;
@@ -13,9 +9,12 @@ use crate::config::{ProxyConfig, RateLimitBps};
#[cfg(test)] #[cfg(test)]
use super::compute_revision; use super::compute_revision;
use super::{ use super::{
AccessSection, compute_snapshot_revision, compute_source_revision, load_candidate_snapshot, AccessSection, compute_snapshot_revision, load_candidate_snapshot, load_config_snapshot,
load_config_snapshot, resolve_single_source_owner, toml_path_exists, 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; use crate::api::model::ApiFailure;
/// Re-render the given top-level tables from `cfg` and upsert each into the /// 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 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 { fn revision_conflict() -> ApiFailure {
ApiFailure::new( ApiFailure::new(
hyper::StatusCode::CONFLICT, hyper::StatusCode::CONFLICT,
@@ -457,115 +404,3 @@ fn revision_conflict() -> ApiFailure {
"Config revision changed before persistence", "Config revision changed before persistence",
) )
} }
fn open_existing_target(path: &Path) -> std::io::Result<Option<ExistingTarget>> {
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::<u64>()
);
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, &current) {
(Some(expected), Some(current)) => {
same_target(&expected.metadata, &current.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
}
+54
View File
@@ -260,6 +260,60 @@ async fn access_mutation_writes_only_the_single_included_owner() {
assert_eq!(revision, current_revision(&root).await.unwrap()); 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] #[tokio::test]
async fn access_mutation_rejects_sections_with_different_source_owners() { async fn access_mutation_rejects_sections_with_different_source_owners() {
let dir = tempfile::tempdir().unwrap(); let dir = tempfile::tempdir().unwrap();
+2 -2
View File
@@ -134,8 +134,8 @@ pub(super) async fn handle(
} }
let expected_revision = parse_if_match(req.headers()); let expected_revision = parse_if_match(req.headers());
let _mutation_guard = shared.mutation_lock.lock().await; let _mutation_guard = shared.mutation_lock.lock().await;
let disk_cfg = load_config_from_disk(&shared.config_path).await?; let (disk_cfg, _) =
ensure_expected_revision(&shared.config_path, expected_revision.as_deref()).await?; load_config_for_mutation(&shared.config_path, expected_revision.as_deref()).await?;
if !disk_cfg.access.users.contains_key(user) { if !disk_cfg.access.users.contains_key(user) {
return Ok(error_response( return Ok(error_response(
request_id, request_id,
+1 -1
View File
@@ -59,7 +59,7 @@ mod web_runtime;
mod web_status; mod web_status;
use config_store::{ 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, parse_if_match,
}; };
use events::ApiEventStore; use events::ApiEventStore;
+2 -2
View File
@@ -9,8 +9,8 @@ use crate::stats::Stats;
use super::ApiShared; use super::ApiShared;
use super::config_store::{ use super::config_store::{
AccessSection, current_revision, ensure_expected_revision, load_config_from_disk, AccessSection, current_revision, load_config_for_mutation,
save_access_sections_to_disk, save_access_sections_to_disk_if_revision,
}; };
use super::model::{ use super::model::{
ApiFailure, CreateUserRequest, CreateUserResponse, PatchUserRequest, RotateSecretRequest, ApiFailure, CreateUserRequest, CreateUserResponse, PatchUserRequest, RotateSecretRequest,
+9 -4
View File
@@ -42,8 +42,8 @@ pub(in crate::api) async fn create_user(
let expiration = parse_optional_expiration(body.expiration_rfc3339.as_deref())?; let expiration = parse_optional_expiration(body.expiration_rfc3339.as_deref())?;
let _guard = shared.mutation_lock.lock().await; let _guard = shared.mutation_lock.lock().await;
let mut cfg = load_config_from_disk(&shared.config_path).await?; let (mut cfg, base_revision) =
ensure_expected_revision(&shared.config_path, expected_revision.as_deref()).await?; load_config_for_mutation(&shared.config_path, expected_revision.as_deref()).await?;
if cfg.access.users.contains_key(&body.username) { if cfg.access.users.contains_key(&body.username) {
return Err(ApiFailure::new( return Err(ApiFailure::new(
@@ -122,8 +122,13 @@ pub(in crate::api) async fn create_user(
touched_sections.push(AccessSection::UserEnabled); touched_sections.push(AccessSection::UserEnabled);
} }
let revision = let revision = save_access_sections_to_disk_if_revision(
save_access_sections_to_disk(&shared.config_path, &cfg, &touched_sections).await?; &shared.config_path,
&cfg,
&touched_sections,
Some(&base_revision),
)
.await?;
shared shared
.proxy_shared .proxy_shared
.stage_user( .stage_user(
+18 -8
View File
@@ -15,8 +15,8 @@ pub(in crate::api) async fn rotate_secret(
} }
let _guard = shared.mutation_lock.lock().await; let _guard = shared.mutation_lock.lock().await;
let mut cfg = load_config_from_disk(&shared.config_path).await?; let (mut cfg, base_revision) =
ensure_expected_revision(&shared.config_path, expected_revision.as_deref()).await?; load_config_for_mutation(&shared.config_path, expected_revision.as_deref()).await?;
if !cfg.access.users.contains_key(user) { if !cfg.access.users.contains_key(user) {
return Err(ApiFailure::new( 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.access.users.insert(user.to_string(), secret.clone());
cfg.validate() cfg.validate()
.map_err(|e| ApiFailure::bad_request(format!("config validation failed: {}", e)))?; .map_err(|e| ApiFailure::bad_request(format!("config validation failed: {}", e)))?;
let revision = let revision = save_access_sections_to_disk_if_revision(
save_access_sections_to_disk(&shared.config_path, &cfg, &[AccessSection::Users]).await?; &shared.config_path,
&cfg,
&[AccessSection::Users],
Some(&base_revision),
)
.await?;
shared shared
.proxy_shared .proxy_shared
.stage_user(user, &secret, cfg.access.is_user_enabled(user)) .stage_user(user, &secret, cfg.access.is_user_enabled(user))
@@ -67,8 +72,8 @@ pub(in crate::api) async fn delete_user(
shared: &ApiShared, shared: &ApiShared,
) -> Result<(String, String), ApiFailure> { ) -> Result<(String, String), ApiFailure> {
let _guard = shared.mutation_lock.lock().await; let _guard = shared.mutation_lock.lock().await;
let mut cfg = load_config_from_disk(&shared.config_path).await?; let (mut cfg, base_revision) =
ensure_expected_revision(&shared.config_path, expected_revision.as_deref()).await?; load_config_for_mutation(&shared.config_path, expected_revision.as_deref()).await?;
if !cfg.access.users.contains_key(user) { if !cfg.access.users.contains_key(user) {
return Err(ApiFailure::new( return Err(ApiFailure::new(
@@ -111,8 +116,13 @@ pub(in crate::api) async fn delete_user(
cfg.validate() cfg.validate()
.map_err(|e| ApiFailure::bad_request(format!("config validation failed: {}", e)))?; .map_err(|e| ApiFailure::bad_request(format!("config validation failed: {}", e)))?;
let revision = let revision = save_access_sections_to_disk_if_revision(
save_access_sections_to_disk(&shared.config_path, &cfg, &touched_sections).await?; &shared.config_path,
&cfg,
&touched_sections,
Some(&base_revision),
)
.await?;
let deleted_incarnation = shared.proxy_shared.delete_user(user).incarnation; let deleted_incarnation = shared.proxy_shared.delete_user(user).incarnation;
let configured_users = cfg.access.users.keys().cloned().collect(); let configured_users = cfg.access.users.keys().cloned().collect();
if let Err(error) = shared if let Err(error) = shared
+18 -8
View File
@@ -32,8 +32,8 @@ pub(in crate::api) async fn patch_user(
} }
let expiration = parse_patch_expiration(&body.expiration_rfc3339)?; let expiration = parse_patch_expiration(&body.expiration_rfc3339)?;
let _guard = shared.mutation_lock.lock().await; let _guard = shared.mutation_lock.lock().await;
let mut cfg = load_config_from_disk(&shared.config_path).await?; let (mut cfg, base_revision) =
ensure_expected_revision(&shared.config_path, expected_revision.as_deref()).await?; load_config_for_mutation(&shared.config_path, expected_revision.as_deref()).await?;
if !cfg.access.users.contains_key(user) { if !cfg.access.users.contains_key(user) {
return Err(ApiFailure::new( return Err(ApiFailure::new(
@@ -168,7 +168,13 @@ pub(in crate::api) async fn patch_user(
let revision = if touched_sections.is_empty() { let revision = if touched_sections.is_empty() {
current_revision(&shared.config_path).await? current_revision(&shared.config_path).await?
} else { } 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 { if touches_users || touches_user_enabled {
let secret = cfg let secret = cfg
@@ -212,8 +218,8 @@ pub(in crate::api) async fn set_user_enabled(
shared: &ApiShared, shared: &ApiShared,
) -> Result<(UserInfo, String), ApiFailure> { ) -> Result<(UserInfo, String), ApiFailure> {
let _guard = shared.mutation_lock.lock().await; let _guard = shared.mutation_lock.lock().await;
let mut cfg = load_config_from_disk(&shared.config_path).await?; let (mut cfg, base_revision) =
ensure_expected_revision(&shared.config_path, expected_revision.as_deref()).await?; load_config_for_mutation(&shared.config_path, expected_revision.as_deref()).await?;
if !cfg.access.users.contains_key(user) { if !cfg.access.users.contains_key(user) {
return Err(ApiFailure::new( return Err(ApiFailure::new(
@@ -231,9 +237,13 @@ pub(in crate::api) async fn set_user_enabled(
cfg.validate() cfg.validate()
.map_err(|e| ApiFailure::bad_request(format!("config validation failed: {}", e)))?; .map_err(|e| ApiFailure::bad_request(format!("config validation failed: {}", e)))?;
let revision = let revision = save_access_sections_to_disk_if_revision(
save_access_sections_to_disk(&shared.config_path, &cfg, &[AccessSection::UserEnabled]) &shared.config_path,
.await?; &cfg,
&[AccessSection::UserEnabled],
Some(&base_revision),
)
.await?;
let secret = cfg let secret = cfg
.access .access
.users .users
+9 -8
View File
@@ -36,7 +36,9 @@ mod validate_server;
mod validate_web; mod validate_web;
mod validation; 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::{ use self::normalize::{
is_valid_ad_tag, is_valid_tls_domain_name, normalize_domain_to_ascii, 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, normalize_exclusive_mask_target, normalize_mask_host_to_ascii, parse_exclusive_mask_target,
@@ -175,13 +177,12 @@ impl ProxyConfig {
source_overrides: &BTreeMap<PathBuf, String>, source_overrides: &BTreeMap<PathBuf, String>,
) -> Result<ConfigSourceGraph> { ) -> Result<ConfigSourceGraph> {
let path = path.as_ref(); let path = path.as_ref();
let normalized_path = normalize_config_path(path); let initial_path = normalize_config_path(path);
let content = source_overrides let (normalized_path, content) = if let Some(content) = source_overrides.get(&initial_path) {
.get(&normalized_path) (initial_path, content.clone())
.cloned() } else {
.map(Ok) read_config_source(path)?
.unwrap_or_else(|| std::fs::read_to_string(path)) };
.map_err(|e| ProxyError::Config(e.to_string()))?;
let base_dir = path.parent().unwrap_or(Path::new(".")); let base_dir = path.parent().unwrap_or(Path::new("."));
let mut source_files = BTreeSet::new(); let mut source_files = BTreeSet::new();
source_files.insert(normalized_path.clone()); source_files.insert(normalized_path.clone());
+83 -7
View File
@@ -1,7 +1,11 @@
use std::collections::{BTreeMap, BTreeSet}; use std::collections::{BTreeMap, BTreeSet};
use std::hash::{DefaultHasher, Hash, Hasher}; use std::hash::{DefaultHasher, Hash, Hasher};
use std::io::Read;
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
#[cfg(unix)]
use std::os::unix::fs::{MetadataExt, OpenOptionsExt};
use crate::error::{ProxyError, Result}; use crate::error::{ProxyError, Result};
pub(super) fn normalize_config_path(path: &Path) -> PathBuf { pub(super) fn normalize_config_path(path: &Path) -> PathBuf {
@@ -22,6 +26,74 @@ pub(super) fn hash_rendered_snapshot(rendered: &str) -> u64 {
hasher.finish() 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, &current_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( pub(super) fn preprocess_includes(
content: &str, content: &str,
base_dir: &Path, base_dir: &Path,
@@ -42,14 +114,18 @@ pub(super) fn preprocess_includes(
let path_str = rest.trim().trim_matches('"'); let path_str = rest.trim().trim_matches('"');
let resolved = base_dir.join(path_str); let resolved = base_dir.join(path_str);
let normalized = normalize_config_path(&resolved); 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()); source_files.insert(normalized.clone());
let included = source_overrides source_contents
.get(&normalized) .entry(normalized)
.cloned() .or_insert_with(|| included.clone());
.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());
let included_dir = resolved.parent().unwrap_or(base_dir); let included_dir = resolved.parent().unwrap_or(base_dir);
output.push_str(&preprocess_includes( output.push_str(&preprocess_includes(
&included, &included,
+231 -119
View File
@@ -5,7 +5,18 @@ use std::path::Path;
use std::sync::Arc; use std::sync::Arc;
#[cfg(unix)] #[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 bytes::Bytes;
use hmac::{Hmac, Mac}; use hmac::{Hmac, Mac};
@@ -13,6 +24,10 @@ use sha2::{Digest, Sha256};
use super::*; 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_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 WEB_DEBUG_FINGERPRINT_CONTEXT: &[u8] = b"telemt-web-debug-key-fingerprint-v1\0";
const MAX_WEB_STATIC_DEPTH: usize = 64; const MAX_WEB_STATIC_DEPTH: usize = 64;
@@ -193,34 +208,31 @@ fn load_static_site(
total_files: &mut usize, total_files: &mut usize,
total_bytes: &mut usize, total_bytes: &mut usize,
) -> Result<WebStaticSite> { ) -> Result<WebStaticSite> {
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(); let mut assets = BTreeMap::new();
load_static_directory( #[cfg(unix)]
&canonical_root, {
&canonical_root, let directory = open_static_root(root)?;
&mut assets, load_static_directory(
total_files, directory,
total_bytes, Path::new(""),
limits, root,
0, &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}")) { if !assets.contains_key(&format!("/{index}")) {
return Err(ProxyError::Config(format!( return Err(ProxyError::Config(format!(
"WEB static directory `{}` does not contain index `{index}`", "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> {
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( fn load_static_directory(
mut directory: Dir,
relative: &Path,
root: &Path, root: &Path,
directory: &Path,
assets: &mut BTreeMap<String, WebStaticAsset>, assets: &mut BTreeMap<String, WebStaticAsset>,
total_files: &mut usize, total_files: &mut usize,
total_bytes: &mut usize, total_bytes: &mut usize,
limits: &WebLimitsConfig, limits: &WebLimitsConfig,
depth: usize, depth: usize,
) -> Result<()> { ) -> Result<()> {
let entries = fs::read_dir(directory).map_err(|error| { let mut entries = Vec::new();
ProxyError::Config(format!( for entry in directory.iter() {
"failed to read WEB static directory `{}`: {error}",
directory.display()
))
})?;
for entry in entries {
let entry = entry.map_err(|error| { 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 { if *total_files >= limits.max_static_files {
return Err(ProxyError::Config( return Err(ProxyError::Config(
"WEB static entries exceed process-wide web.limits.max_static_files".to_string(), "WEB static entries exceed process-wide web.limits.max_static_files".to_string(),
)); ));
} }
*total_files += 1; *total_files += 1;
let path = entry.path(); entries.push(OsString::from_vec(name.to_vec()));
let file_type = entry.file_type().map_err(|error| { }
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!( ProxyError::Config(format!(
"failed to inspect WEB static entry `{}`: {error}", "failed to open WEB static entry `{}` without following symlinks: {error}",
path.display() display_path.display()
)) ))
})?; })?;
if file_type.is_symlink() { let file = fs::File::from(descriptor);
return Err(ProxyError::Config(format!( let metadata = file.metadata().map_err(|error| {
"WEB static entry `{}` must not be a symlink", ProxyError::Config(format!(
path.display() "failed to inspect WEB static entry `{}`: {error}",
))); display_path.display()
} ))
if file_type.is_dir() { })?;
if metadata.is_dir() {
if depth >= MAX_WEB_STATIC_DEPTH { if depth >= MAX_WEB_STATIC_DEPTH {
return Err(ProxyError::Config(format!( return Err(ProxyError::Config(format!(
"WEB static directory `{}` exceeds the maximum nesting depth", "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( load_static_directory(
child,
&relative_path,
root, root,
&path,
assets, assets,
total_files, total_files,
total_bytes, total_bytes,
@@ -289,83 +341,107 @@ fn load_static_directory(
)?; )?;
continue; 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() { if !metadata.is_file() {
return Err(ProxyError::Config(format!( return Err(ProxyError::Config(format!(
"WEB static entry `{}` changed before it was opened", "WEB static entry `{}` must be a regular file",
path.display() display_path.display()
))); )));
} }
let file_len = usize::try_from(metadata.len()).map_err(|_| { load_static_file(
ProxyError::Config(format!("WEB static file `{}` is too large", path.display())) file,
})?; &metadata,
if file_len > limits.max_static_file_bytes { &relative_path,
return Err(ProxyError::Config(format!( &display_path,
"WEB static file `{}` exceeds web.limits.max_static_file_bytes", assets,
path.display() total_bytes,
))); limits,
} )?;
*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,
},
);
} }
Ok(()) Ok(())
} }
fn load_static_file(
mut file: fs::File,
metadata: &fs::Metadata,
relative: &Path,
display_path: &Path,
assets: &mut BTreeMap<String, WebStaticAsset>,
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<String> { fn static_route(relative: &Path) -> Result<String> {
let mut route = String::new(); let mut route = String::new();
for component in relative.components() { for component in relative.components() {
@@ -425,4 +501,40 @@ mod tests {
"IpJrt3e7sKtzPyoXy6w-Zj6GGEvsvclN66JzQEfPYLA" "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");
}
} }
@@ -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<String, WebStaticAsset>,
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<String, WebStaticAsset>,
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(())
}
+18 -11
View File
@@ -1,6 +1,5 @@
use std::collections::BTreeSet; use std::collections::BTreeSet;
use std::net::IpAddr; use std::net::IpAddr;
use std::path::PathBuf;
use std::sync::Arc; use std::sync::Arc;
use std::time::Duration; 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::middle_relay::note_global_relay_pressure;
use crate::proxy::shared_state::{ConntrackCloseEvent, ConntrackCloseReason, ProxySharedState}; use crate::proxy::shared_state::{ConntrackCloseEvent, ConntrackCloseReason, ProxySharedState};
use crate::stats::Stats; use crate::stats::Stats;
#[cfg(unix)]
use crate::util::trusted_command::resolve_trusted_helper;
const CONNTRACK_EVENT_QUEUE_CAPACITY: usize = 32_768; const CONNTRACK_EVENT_QUEUE_CAPACITY: usize = 32_768;
const PRESSURE_RELEASE_TICKS: u8 = 3; const PRESSURE_RELEASE_TICKS: u8 = 3;
@@ -381,13 +382,15 @@ fn pick_backend(configured: ConntrackBackend) -> Option<NetfilterBackend> {
} }
fn command_exists(binary: &str) -> bool { fn command_exists(binary: &str) -> bool {
let Some(path_var) = std::env::var_os("PATH") else { #[cfg(unix)]
return false; {
}; resolve_trusted_helper(binary).is_some()
std::env::split_paths(&path_var).any(|dir| { }
let candidate: PathBuf = dir.join(binary); #[cfg(not(unix))]
candidate.exists() && candidate.is_file() {
}) let _ = binary;
false
}
} }
fn listener_port_set(cfg: &ProxyConfig) -> Vec<u16> { fn listener_port_set(cfg: &ProxyConfig) -> Vec<u16> {
@@ -651,10 +654,14 @@ async fn delete_conntrack_entry(event: ConntrackCloseEvent) -> DeleteOutcome {
} }
async fn run_command(binary: &str, args: &[&str], stdin: Option<String>) -> Result<(), String> { async fn run_command(binary: &str, args: &[&str], stdin: Option<String>) -> Result<(), String> {
if !command_exists(binary) { #[cfg(unix)]
let Some(command_path) = resolve_trusted_helper(binary) else {
return Err(format!("{binary} is not available")); 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); command.args(args);
if stdin.is_some() { if stdin.is_some() {
command.stdin(std::process::Stdio::piped()); command.stdin(std::process::Stdio::piped());
+1
View File
@@ -27,6 +27,7 @@ mod protocol;
mod proxy; mod proxy;
mod quota_state; mod quota_state;
mod service; mod service;
mod slot_budget;
mod startup; mod startup;
mod stats; mod stats;
mod stream; mod stream;
+7 -4
View File
@@ -75,7 +75,7 @@ pub(crate) use self::auth_probe::{
auth_probe_saturation_is_throttled_for_testing_in_shared, auth_probe_saturation_is_throttled_for_testing_in_shared,
auth_probe_saturation_state_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, 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, 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, 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; const AUTH_PROBE_TRACK_RETENTION_SECS: u64 = 10 * 60;
#[cfg(test)] #[cfg(test)]
const AUTH_PROBE_TRACK_MAX_ENTRIES: usize = 256; pub(super) const AUTH_PROBE_TRACK_MAX_ENTRIES: usize = 256;
#[cfg(not(test))] #[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_PRUNE_SCAN_LIMIT: usize = 1_024;
const AUTH_PROBE_BACKOFF_START_FAILS: u32 = 4; const AUTH_PROBE_BACKOFF_START_FAILS: u32 = 4;
const AUTH_PROBE_SATURATION_GRACE_FAILS: u32 = 2; 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 CANDIDATE_HINT_TRACK_CAP: usize = 64;
const OVERLOAD_CANDIDATE_BUDGET_HINTED: usize = 16; const OVERLOAD_CANDIDATE_BUDGET_HINTED: usize = 16;
const OVERLOAD_CANDIDATE_BUDGET_UNHINTED: usize = 8; const OVERLOAD_CANDIDATE_BUDGET_UNHINTED: usize = 8;
+94 -18
View File
@@ -76,27 +76,48 @@ pub(super) fn sticky_hint_record_success_in(
user_id: u32, user_id: u32,
sni: Option<&str>, sni: Option<&str>,
) { ) {
if shared.handshake.sticky_user_by_ip.len() > STICKY_HINT_MAX_ENTRIES { bounded_sticky_hint_upsert(
shared.handshake.sticky_user_by_ip.clear(); &shared.handshake.sticky_user_by_ip,
} &shared.handshake.sticky_user_by_ip_slots,
shared.handshake.sticky_user_by_ip.insert(peer_ip, user_id); 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(); bounded_sticky_hint_upsert(
} &shared.handshake.sticky_user_by_ip_prefix,
shared &shared.handshake.sticky_user_by_ip_prefix_slots,
.handshake ip_prefix_hint_key(peer_ip),
.sticky_user_by_ip_prefix user_id,
.insert(ip_prefix_hint_key(peer_ip), user_id); );
if let Some(sni) = sni { if let Some(sni) = sni {
if shared.handshake.sticky_user_by_sni_hash.len() > STICKY_HINT_MAX_ENTRIES { bounded_sticky_hint_upsert(
shared.handshake.sticky_user_by_sni_hash.clear(); &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<K>(
entries: &DashMap<K, u32>,
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( pub(super) fn decode_user_secrets_in(
shared: &ProxySharedState, shared: &ProxySharedState,
config: &ProxyConfig, config: &ProxyConfig,
+71 -127
View File
@@ -98,9 +98,14 @@ pub(super) fn auth_probe_is_throttled_in(
}; };
if auth_probe_state_expired(&entry, now) { if auth_probe_state_expired(&entry, now) {
drop(entry); drop(entry);
state.remove_if(&peer_ip, |_, current| { if state
auth_probe_state_expired(current, now) .remove_if(&peer_ip, |_, current| {
}); auth_probe_state_expired(current, now)
})
.is_some()
{
shared.handshake.auth_probe_slots.release();
}
return false; return false;
} }
now < entry.blocked_until now < entry.blocked_until
@@ -118,9 +123,14 @@ pub(super) fn auth_probe_saturation_grace_exhausted_in(
}; };
if auth_probe_state_expired(&entry, now) { if auth_probe_state_expired(&entry, now) {
drop(entry); drop(entry);
state.remove_if(&peer_ip, |_, current| { if state
auth_probe_state_expired(current, now) .remove_if(&peer_ip, |_, current| {
}); auth_probe_state_expired(current, now)
})
.is_some()
{
shared.handshake.auth_probe_slots.release();
}
return false; return false;
} }
@@ -216,7 +226,13 @@ pub(super) fn auth_probe_record_failure_in(
) { ) {
let peer_ip = normalize_auth_probe_ip(peer_ip); let peer_ip = normalize_auth_probe_ip(peer_ip);
let state = &shared.handshake.auth_probe; 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( 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<IpAddr, AuthProbeState>, state: &DashMap<IpAddr, AuthProbeState>,
peer_ip: IpAddr, peer_ip: IpAddr,
now: Instant, 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<IpAddr, AuthProbeState>,
slots: Option<&crate::slot_budget::SlotBudget>,
peer_ip: IpAddr,
now: Instant,
) { ) {
let make_new_state = || AuthProbeState { let make_new_state = || AuthProbeState {
fail_streak: 1, fail_streak: 1,
@@ -279,6 +305,9 @@ pub(super) fn auth_probe_record_failure_with_state_in(
}) })
.is_some() .is_some()
{ {
if let Some(slots) = slots {
slots.release();
}
break; break;
} }
continue; continue;
@@ -347,9 +376,15 @@ pub(super) fn auth_probe_record_failure_with_state_in(
} }
for stale_key in stale_keys { for stale_key in stale_keys {
state.remove_if(&stale_key, |_, current| { if state
auth_probe_state_expired(current, now) .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 { 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); auth_probe_note_saturation_in(shared, now);
return; return;
}; };
state.remove_if(&evict_key, |_, current| { if state
current.fail_streak == evict_fail_streak && current.last_seen == evict_last_seen .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); 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) { match state.entry(peer_ip) {
Entry::Occupied(mut entry) => { Entry::Occupied(mut entry) => {
update_existing(entry.get_mut()); update_existing(entry.get_mut());
} }
Entry::Vacant(entry) => { Entry::Vacant(entry) => {
entry.insert(make_new_state()); 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) { pub(super) fn auth_probe_record_success_in(shared: &ProxySharedState, peer_ip: IpAddr) {
let peer_ip = normalize_auth_probe_ip(peer_ip); let peer_ip = normalize_auth_probe_ip(peer_ip);
let state = &shared.handshake.auth_probe; let state = &shared.handshake.auth_probe;
state.remove(&peer_ip); if state.remove(&peer_ip).is_some() {
} shared.handshake.auth_probe_slots.release();
#[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<u32> {
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();
}
} }
} }
#[cfg(test)] #[cfg(test)]
pub(crate) fn auth_probe_state_for_testing_in_shared( mod testing;
shared: &ProxySharedState,
) -> &DashMap<IpAddr, AuthProbeState> {
&shared.handshake.auth_probe
}
#[cfg(test)] #[cfg(test)]
pub(crate) fn auth_probe_saturation_state_for_testing_in_shared( pub(crate) use testing::*;
shared: &ProxySharedState,
) -> &Mutex<Option<AuthProbeSaturationState>> {
&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<AuthProbeSaturationState>> {
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<HashSet<(String, String)>> {
&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)
}
#[inline] #[inline]
pub(super) fn find_matching_tls_domain<'a>(config: &'a ProxyConfig, sni: &str) -> Option<&'a str> { pub(super) fn find_matching_tls_domain<'a>(config: &'a ProxyConfig, sni: &str) -> Option<&'a str> {
+137
View File
@@ -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<u32> {
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<IpAddr, AuthProbeState> {
&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<Option<AuthProbeSaturationState>> {
&shared.handshake.auth_probe_saturation
}
pub(crate) fn auth_probe_saturation_state_lock_for_testing_in_shared(
shared: &ProxySharedState,
) -> std::sync::MutexGuard<'_, Option<AuthProbeSaturationState>> {
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<HashSet<(String, String)>> {
&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
);
}
+17
View File
@@ -16,6 +16,7 @@ use crate::proxy::user_admission::{
UserAdmissionAuthority, UserAdmissionPublication, UserCredentialId, UserIncarnation, UserAdmissionAuthority, UserAdmissionPublication, UserCredentialId, UserIncarnation,
UserMutationResult, UserSessionRegistration, UserMutationResult, UserSessionRegistration,
}; };
use crate::slot_budget::SlotBudget;
const HANDSHAKE_RECENT_USER_RING_LEN: usize = 64; const HANDSHAKE_RECENT_USER_RING_LEN: usize = 64;
const MASKING_FALLBACK_MAX_CONCURRENT: usize = 512; const MASKING_FALLBACK_MAX_CONCURRENT: usize = 512;
@@ -55,13 +56,17 @@ pub(crate) enum ConntrackClosePolicy {
pub(crate) struct HandshakeSharedState { pub(crate) struct HandshakeSharedState {
pub(crate) auth_probe: DashMap<IpAddr, AuthProbeState>, pub(crate) auth_probe: DashMap<IpAddr, AuthProbeState>,
pub(crate) auth_probe_slots: SlotBudget,
pub(crate) auth_probe_saturation: Mutex<Option<AuthProbeSaturationState>>, pub(crate) auth_probe_saturation: Mutex<Option<AuthProbeSaturationState>>,
pub(crate) auth_probe_eviction_hasher: RandomState, pub(crate) auth_probe_eviction_hasher: RandomState,
pub(crate) invalid_secret_warned: Mutex<HashSet<(String, String)>>, pub(crate) invalid_secret_warned: Mutex<HashSet<(String, String)>>,
pub(crate) unknown_sni_warn_next_allowed: Mutex<Option<Instant>>, pub(crate) unknown_sni_warn_next_allowed: Mutex<Option<Instant>>,
pub(crate) sticky_user_by_ip: DashMap<IpAddr, u32>, pub(crate) sticky_user_by_ip: DashMap<IpAddr, u32>,
pub(crate) sticky_user_by_ip_slots: SlotBudget,
pub(crate) sticky_user_by_ip_prefix: DashMap<u64, u32>, pub(crate) sticky_user_by_ip_prefix: DashMap<u64, u32>,
pub(crate) sticky_user_by_ip_prefix_slots: SlotBudget,
pub(crate) sticky_user_by_sni_hash: DashMap<u64, u32>, pub(crate) sticky_user_by_sni_hash: DashMap<u64, u32>,
pub(crate) sticky_user_by_sni_hash_slots: SlotBudget,
pub(crate) recent_user_ring: Box<[AtomicU32]>, pub(crate) recent_user_ring: Box<[AtomicU32]>,
pub(crate) recent_user_ring_seq: AtomicU64, pub(crate) recent_user_ring_seq: AtomicU64,
pub(crate) auth_expensive_checks_total: AtomicU64, pub(crate) auth_expensive_checks_total: AtomicU64,
@@ -114,13 +119,25 @@ impl ProxySharedState {
Arc::new(Self { Arc::new(Self {
handshake: HandshakeSharedState { handshake: HandshakeSharedState {
auth_probe: DashMap::new(), 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_saturation: Mutex::new(None),
auth_probe_eviction_hasher: RandomState::new(), auth_probe_eviction_hasher: RandomState::new(),
invalid_secret_warned: Mutex::new(HashSet::new()), invalid_secret_warned: Mutex::new(HashSet::new()),
unknown_sni_warn_next_allowed: Mutex::new(None), unknown_sni_warn_next_allowed: Mutex::new(None),
sticky_user_by_ip: DashMap::new(), 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: 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: 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)) recent_user_ring: std::iter::repeat_with(|| AtomicU32::new(0))
.take(HANDSHAKE_RECENT_USER_RING_LEN) .take(HANDSHAKE_RECENT_USER_RING_LEN)
.collect::<Vec<_>>() .collect::<Vec<_>>()
+136
View File
@@ -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<SlotLease<'_>> {
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());
}
}
+35 -51
View File
@@ -10,6 +10,7 @@ use dashmap::DashMap;
use dashmap::mapref::entry::Entry; use dashmap::mapref::entry::Entry;
use crate::protocol::tls_fingerprint::TlsClientFingerprint; use crate::protocol::tls_fingerprint::TlsClientFingerprint;
use crate::slot_budget::SlotBudget;
use super::Stats; use super::Stats;
@@ -68,15 +69,33 @@ struct TlsFingerprintEntry {
bad_or_probe: AtomicU64, bad_or_probe: AtomicU64,
} }
#[derive(Default)]
pub struct TlsFingerprintCollector { pub struct TlsFingerprintCollector {
entries: DashMap<TlsFingerprintKey, TlsFingerprintEntry>, entries: DashMap<TlsFingerprintKey, TlsFingerprintEntry>,
slots: SlotBudget,
capacity: usize,
dropped_total: AtomicU64, dropped_total: AtomicU64,
parse_error_total: AtomicU64, parse_error_total: AtomicU64,
last_cleanup_epoch_secs: AtomicU64, last_cleanup_epoch_secs: AtomicU64,
} }
impl Default for TlsFingerprintCollector {
fn default() -> Self {
Self::with_capacity(MAX_TLS_FINGERPRINT_BUCKETS)
}
}
impl TlsFingerprintCollector { 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( pub fn record_observed(
&self, &self,
fingerprint: &TlsClientFingerprint, fingerprint: &TlsClientFingerprint,
@@ -228,7 +247,7 @@ impl TlsFingerprintCollector {
TlsFingerprintSnapshot { TlsFingerprintSnapshot {
retention_secs: ttl.as_secs(), retention_secs: ttl.as_secs(),
capacity: MAX_TLS_FINGERPRINT_BUCKETS, capacity: self.capacity,
dropped_total: self.dropped_total.load(Ordering::Relaxed), dropped_total: self.dropped_total.load(Ordering::Relaxed),
parse_error_total: self.parse_error_total.load(Ordering::Relaxed), parse_error_total: self.parse_error_total.load(Ordering::Relaxed),
by_fingerprint, by_fingerprint,
@@ -286,22 +305,6 @@ impl TlsFingerprintCollector {
ja4_raw: fingerprint.ja4_raw.clone(), 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) { match self.entries.entry(key) {
Entry::Occupied(entry) => { Entry::Occupied(entry) => {
update_entry( update_entry(
@@ -313,12 +316,17 @@ impl TlsFingerprintCollector {
); );
} }
Entry::Vacant(entry) => { Entry::Vacant(entry) => {
let Some(slot) = self.slots.try_acquire() else {
self.dropped_total.fetch_add(1, Ordering::Relaxed);
return;
};
entry.insert(TlsFingerprintEntry::new( entry.insert(TlsFingerprintEntry::new(
now_epoch_secs, now_epoch_secs,
if count_total { 1 } else { 0 }, if count_total { 1 } else { 0 },
if count_auth_success { 1 } else { 0 }, if count_auth_success { 1 } else { 0 },
if count_bad_or_probe { 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) { fn cleanup(&self, now_epoch_secs: u64, ttl_secs: u64) {
if ttl_secs == 0 { let mut removed = 0usize;
self.entries.clear();
return;
}
self.entries.retain(|_, entry| { self.entries.retain(|_, entry| {
let last_seen = entry.last_seen_epoch_secs.load(Ordering::Relaxed); 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)] #[cfg(test)]
mod tests { 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);
}
}
+58
View File
@@ -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
);
}
+4 -14
View File
@@ -1,14 +1,14 @@
use std::path::PathBuf;
use tokio::io::AsyncWriteExt; use tokio::io::AsyncWriteExt;
use tokio::process::Command; use tokio::process::Command;
use crate::util::trusted_command::resolve_trusted_helper;
pub(super) async fn run_command( pub(super) async fn run_command(
binary: &str, binary: &str,
args: &[&str], args: &[&str],
stdin: Option<String>, stdin: Option<String>,
) -> Result<(), String> { ) -> 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")); return Err(format!("{binary} is not available"));
}; };
let mut command = Command::new(command_path); 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<String, String> { pub(super) async fn run_command_stdout(binary: &str, args: &[&str]) -> Result<String, 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")); return Err(format!("{binary} is not available"));
}; };
let output = Command::new(command_path) let output = Command::new(command_path)
@@ -64,16 +64,6 @@ pub(super) async fn run_command_stdout(binary: &str, args: &[&str]) -> Result<St
}) })
} }
fn resolve_command(binary: &str) -> Option<PathBuf> {
let mut dirs = std::env::var_os("PATH")
.map(|path| std::env::split_paths(&path).collect::<Vec<_>>())
.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 { pub(super) fn has_firewall_privileges() -> bool {
#[cfg(target_os = "linux")] #[cfg(target_os = "linux")]
{ {
+8 -1
View File
@@ -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); 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. /// Current number of adaptive TLS fetch profile-cache entries.
pub(crate) fn profile_cache_entries_for_metrics() -> usize { pub(crate) fn profile_cache_entries_for_metrics() -> usize {
profile_cache().len() profile_cache().len()
@@ -270,8 +276,9 @@ fn order_profiles(
if let Some(cached) = profile_cache().get(key) { if let Some(cached) = profile_cache().get(key) {
let age = now.saturating_duration_since(cached.updated_at); let age = now.saturating_duration_since(cached.updated_at);
if age > strategy.profile_cache_ttl { if age > strategy.profile_cache_ttl {
let observed = *cached;
drop(cached); drop(cached);
profile_cache().remove(key); remove_profile_if_unchanged(key, observed);
return ordered; return ordered;
} }
+26
View File
@@ -6,6 +6,7 @@ use super::{
TLS_NAMED_GROUP_X25519MLKEM768, TlsFetchStrategy, X25519_KEY_SHARE_LEN, build_client_hello, TLS_NAMED_GROUP_X25519MLKEM768, TlsFetchStrategy, X25519_KEY_SHARE_LEN, build_client_hello,
build_tls_fetch_proxy_header, derive_behavior_profile, encode_tls13_certificate_message, 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, fetch_via_rustls_stream, order_profiles, profile_alpn, profile_cache, profile_cache_key,
remove_profile_if_unchanged,
}; };
use crate::config::TlsFetchProfile; use crate::config::TlsFetchProfile;
use crate::crypto::SecureRandom; use crate::crypto::SecureRandom;
@@ -225,6 +226,31 @@ fn test_order_profiles_drops_expired_cached_winner() {
assert!(profile_cache().get(&cache_key).is_none()); 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] #[test]
fn test_deterministic_client_hello_is_stable() { fn test_deterministic_client_hello_is_stable() {
let rng = SecureRandom::new(); let rng = SecureRandom::new();
+2
View File
@@ -2,6 +2,8 @@
pub mod ip; pub mod ip;
pub mod time; pub mod time;
#[cfg(unix)]
pub mod trusted_command;
#[allow(unused_imports)] #[allow(unused_imports)]
pub use ip::*; pub use ip::*;
+70
View File
@@ -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<PathBuf> {
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<PathBuf> {
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());
}
}
+83 -119
View File
@@ -1,6 +1,5 @@
use std::net::IpAddr; use std::net::IpAddr;
use std::sync::Arc; use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::time::Instant; use std::time::Instant;
use parking_lot::Mutex; use parking_lot::Mutex;
@@ -29,6 +28,8 @@ struct BodyCapture {
} }
struct ExchangeState { struct ExchangeState {
phase: ExchangePhase,
reserved: usize,
method: String, method: String,
path: String, path: String,
route: TraceRoute, route: TraceRoute,
@@ -47,6 +48,12 @@ struct ExchangeState {
body_capture_blocked: bool, 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. /// One in-flight request-to-response capture with process-wide byte leases.
pub(crate) struct HttpTraceExchange { pub(crate) struct HttpTraceExchange {
store: Arc<WebTraceStore>, store: Arc<WebTraceStore>,
@@ -55,8 +62,6 @@ pub(crate) struct HttpTraceExchange {
started: Instant, started: Instant,
started_epoch_millis: u64, started_epoch_millis: u64,
state: Mutex<ExchangeState>, state: Mutex<ExchangeState>,
reserved: AtomicUsize,
committed: AtomicBool,
} }
impl HttpTraceExchange { impl HttpTraceExchange {
@@ -104,6 +109,8 @@ impl HttpTraceExchange {
started, started,
started_epoch_millis, started_epoch_millis,
state: Mutex::new(ExchangeState { state: Mutex::new(ExchangeState {
phase: ExchangePhase::Open,
reserved: base_reservation + if dynamic_reserved { dynamic } else { 0 },
method, method,
path, path,
route: TraceRoute::Unknown, route: TraceRoute::Unknown,
@@ -121,21 +128,23 @@ impl HttpTraceExchange {
redactions, redactions,
body_capture_blocked: !dynamic_reserved, 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. /// Sets the final request route before body polling or decoy forwarding.
pub(crate) fn set_route(&self, route: TraceRoute) { 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. /// Sets the trusted effective client address after proxy-header validation.
pub(crate) fn set_effective_ip(&self, client_ip: IpAddr) { 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. /// Binds non-secret profile and process session identity.
@@ -145,8 +154,11 @@ impl HttpTraceExchange {
.len() .len()
.saturating_add(profile.key_fingerprint.len()); .saturating_add(profile.key_fingerprint.len());
let mut state = self.state.lock(); let mut state = self.state.lock();
if state.phase != ExchangePhase::Open {
return;
}
state.identity.session_id = Some(session_id); 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.user = Some(profile.user.clone());
state.identity.key_fingerprint = Some(profile.key_fingerprint.clone()); state.identity.key_fingerprint = Some(profile.key_fingerprint.clone());
} }
@@ -160,8 +172,11 @@ impl HttpTraceExchange {
.map_or(0, String::len) .map_or(0, String::len)
.saturating_add(identity.key_fingerprint.as_ref().map_or(0, String::len)); .saturating_add(identity.key_fingerprint.as_ref().map_or(0, String::len));
let mut state = self.state.lock(); let mut state = self.state.lock();
if state.phase != ExchangePhase::Open {
return;
}
state.identity.session_id = identity.session_id; 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.user = identity.user;
state.identity.key_fingerprint = identity.key_fingerprint; state.identity.key_fingerprint = identity.key_fingerprint;
} }
@@ -172,21 +187,25 @@ impl HttpTraceExchange {
if value.is_empty() { if value.is_empty() {
return; return;
} }
if !self.reserve(value.len()) { let mut state = self.state.lock();
self.block_body_capture(); if state.phase != ExchangePhase::Open {
return; return;
} }
self.state if !self.reserve_locked(&mut state, value.len()) {
.lock() Self::block_body_capture_locked(&mut state);
.redactions return;
.push(Zeroizing::new(value.to_vec())); }
state.redactions.push(Zeroizing::new(value.to_vec()));
} }
/// Captures response status and sanitized headers at handler completion. /// Captures response status and sanitized headers at handler completion.
pub(crate) fn response_ready<B>(&self, response: &hyper::Response<B>) { pub(crate) fn response_ready<B>(&self, response: &hyper::Response<B>) {
let dynamic = response_dynamic_bytes(response, &self.policy); let dynamic = response_dynamic_bytes(response, &self.policy);
let reserved = self.reserve(dynamic);
let mut state = self.state.lock(); 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()); state.status = Some(response.status().as_u16());
if self.policy.capture_headers && reserved { if self.policy.capture_headers && reserved {
state.response_headers = sanitized_headers(response.headers()); state.response_headers = sanitized_headers(response.headers());
@@ -210,19 +229,22 @@ impl HttpTraceExchange {
/// Appends one body data frame without changing the proxied bytes. /// Appends one body data frame without changing the proxied bytes.
pub(crate) fn body_data(&self, direction: TraceDirection, data: &[u8]) { pub(crate) fn body_data(&self, direction: TraceDirection, data: &[u8]) {
let mut state = self.state.lock(); let mut state = self.state.lock();
if state.phase != ExchangePhase::Open {
return;
}
let route = state.route; let route = state.route;
let body_capture_blocked = state.body_capture_blocked; let body_capture_blocked = state.body_capture_blocked;
let body = match direction { {
TraceDirection::Request => &mut state.request_body, let body = selected_body(&mut state, direction);
TraceDirection::Response => &mut state.response_body, body.observed_bytes = body
}; .observed_bytes
body.observed_bytes = body .saturating_add(u64::try_from(data.len()).unwrap_or(u64::MAX));
.observed_bytes }
.saturating_add(u64::try_from(data.len()).unwrap_or(u64::MAX));
if data.is_empty() { if data.is_empty() {
return; return;
} }
if body_capture_blocked { if body_capture_blocked {
let body = selected_body(&mut state, direction);
body.truncated |= !data.is_empty(); body.truncated |= !data.is_empty();
return; return;
} }
@@ -230,6 +252,7 @@ impl HttpTraceExchange {
else { else {
return; return;
}; };
let body = selected_body(&mut state, direction);
if body.captured.len() >= limit { if body.captured.len() >= limit {
if !data.is_empty() && !body.truncated { if !data.is_empty() && !body.truncated {
self.store.record_truncation(); self.store.record_truncation();
@@ -237,13 +260,17 @@ impl HttpTraceExchange {
body.truncated |= !data.is_empty(); body.truncated |= !data.is_empty();
return; return;
} }
if body.captured.capacity() == 0 { let needs_reservation = body.captured.capacity() == 0;
if !self.reserve(limit) { if needs_reservation {
if !self.reserve_locked(&mut state, limit) {
let body = selected_body(&mut state, direction);
body.truncated = true; body.truncated = true;
return; return;
} }
let body = selected_body(&mut state, direction);
body.captured = Vec::with_capacity(limit); body.captured = Vec::with_capacity(limit);
} }
let body = selected_body(&mut state, direction);
let take = data.len().min(limit - body.captured.len()); let take = data.len().min(limit - body.captured.len());
body.captured.extend_from_slice(&data[..take]); body.captured.extend_from_slice(&data[..take]);
if take < data.len() { if take < data.len() {
@@ -257,6 +284,9 @@ impl HttpTraceExchange {
/// Marks one request or response body terminal state. /// Marks one request or response body terminal state.
pub(crate) fn body_finished(&self, direction: TraceDirection, terminal: TraceBodyState) { pub(crate) fn body_finished(&self, direction: TraceDirection, terminal: TraceBodyState) {
let mut state = self.state.lock(); let mut state = self.state.lock();
if state.phase != ExchangePhase::Open {
return;
}
let body = match direction { let body = match direction {
TraceDirection::Request => &mut state.request_body, TraceDirection::Request => &mut state.request_body,
TraceDirection::Response => &mut state.response_body, TraceDirection::Response => &mut state.response_body,
@@ -293,7 +323,10 @@ impl HttpTraceExchange {
.div_ceil(frame::HEADER_BYTES) .div_ceil(frame::HEADER_BYTES)
.clamp(1, limits.max_frames_per_body); .clamp(1, limits.max_frames_per_body);
let reservation = estimated_frames.saturating_mul(std::mem::size_of::<TraceFrame>()); let reservation = estimated_frames.saturating_mul(std::mem::size_of::<TraceFrame>());
if !self.reserve(reservation) { let mut state = self.state.lock();
if state.phase != ExchangePhase::Open
|| !self.reserve_locked(&mut state, reservation)
{
return; return;
} }
let frames = match frame::parse_all(body, limits) { let frames = match frame::parse_all(body, limits) {
@@ -319,35 +352,38 @@ impl HttpTraceExchange {
parse_error: Some(frame_error_name(error)), 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. /// Commits once after response body consumption or drop.
pub(crate) fn commit(&self) { pub(crate) fn commit(&self) {
if self.committed.swap(true, Ordering::AcqRel) { let (record, reserved) = {
return; let mut state = self.state.lock();
} if state.phase != ExchangePhase::Open {
let reserved = self.reserved.load(Ordering::Acquire); return;
let record = self.build_record(); }
state.phase = ExchangePhase::Committed;
let reserved = state.reserved;
(self.build_record_locked(&mut state), reserved)
};
if !self.store.try_commit(record, reserved, self.epoch) { if !self.store.try_commit(record, reserved, self.epoch) {
self.store.release(reserved); self.store.release(reserved);
} }
} }
fn reserve(&self, bytes: usize) -> bool { fn reserve_locked(&self, state: &mut ExchangeState, bytes: usize) -> bool {
if bytes == 0 { if bytes == 0 {
return true; return true;
} }
if self.store.try_reserve(bytes) { if self.store.try_reserve(bytes) {
self.reserved.fetch_add(bytes, Ordering::AcqRel); state.reserved = state.reserved.saturating_add(bytes);
true true
} else { } else {
false false
} }
} }
fn block_body_capture(&self) { fn block_body_capture_locked(state: &mut ExchangeState) {
let mut state = self.state.lock();
state.body_capture_blocked = true; state.body_capture_blocked = true;
state.request_body.captured.clear(); state.request_body.captured.clear();
state.response_body.captured.clear(); state.response_body.captured.clear();
@@ -355,8 +391,7 @@ impl HttpTraceExchange {
state.response_body.truncated = true; state.response_body.truncated = true;
} }
fn build_record(&self) -> TraceRecord { fn build_record_locked(&self, state: &mut ExchangeState) -> TraceRecord {
let mut state = self.state.lock();
let redactions = std::mem::take(&mut state.redactions); let redactions = std::mem::take(&mut state.redactions);
scrub_body(&mut state.request_body.captured, &redactions); scrub_body(&mut state.request_body.captured, &redactions);
scrub_body(&mut state.response_body.captured, &redactions); scrub_body(&mut state.response_body.captured, &redactions);
@@ -394,7 +429,7 @@ impl HttpTraceExchange {
impl Drop for HttpTraceExchange { impl Drop for HttpTraceExchange {
fn drop(&mut self) { 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::Request, TraceBodyState::Aborted);
self.body_finished(TraceDirection::Response, TraceBodyState::Aborted); self.body_finished(TraceDirection::Response, TraceBodyState::Aborted);
} }
@@ -410,83 +445,12 @@ fn body_snapshot(policy: &WebDebugConfig, body: &mut BodyCapture) -> Option<Trac
}) })
} }
#[cfg(test)] fn selected_body(state: &mut ExchangeState, direction: TraceDirection) -> &mut BodyCapture {
mod tests { match direction {
use super::*; TraceDirection::Request => &mut state.request_body,
TraceDirection::Response => &mut state.response_body,
#[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 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()
};
let store = WebTraceStore::new(policy, &limits);
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())
);
} }
} }
#[cfg(test)]
mod tests;
+104
View File
@@ -0,0 +1,104 @@
use super::*;
fn trace_store() -> Arc<WebTraceStore> {
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);
}