diff --git a/src/api/config_store/atomic.rs b/src/api/config_store/atomic.rs index d13aa23..c5e7864 100644 --- a/src/api/config_store/atomic.rs +++ b/src/api/config_store/atomic.rs @@ -1,12 +1,24 @@ +use std::fs::File; use std::io::{Read, Write}; use std::path::{Path, PathBuf}; #[cfg(unix)] -use std::os::unix::fs::{MetadataExt, OpenOptionsExt, PermissionsExt}; +use std::os::unix::fs::{MetadataExt, PermissionsExt}; + +#[cfg(unix)] +use nix::fcntl::{Flock, FlockArg, OFlag, openat, renameat}; +#[cfg(unix)] +use nix::sys::stat::Mode; +#[cfg(unix)] +use nix::unistd::{UnlinkatFlags, fsync, unlinkat}; use super::compute_source_revision; use crate::api::model::ApiFailure; use crate::config::ProxyConfig; +#[cfg(unix)] +use crate::util::secure_fs::AnchoredPath; + +const MAX_CONFIG_SOURCE_BYTES: u64 = 8 * 1024 * 1024; enum AtomicWriteError { Conflict, @@ -19,12 +31,53 @@ struct ExistingTarget { metadata: std::fs::Metadata, } +struct ConfigWriteLock { + #[cfg(unix)] + _file: Flock, +} + +impl ConfigWriteLock { + fn acquire(path: &Path) -> std::io::Result { + #[cfg(unix)] + { + let lock_path = sibling_lock_path(path); + let anchored = AnchoredPath::open_creating_parents(&lock_path, 0o750)?; + let descriptor = openat( + anchored.parent(), + anchored.name(), + OFlag::O_RDWR | OFlag::O_CREAT | OFlag::O_NOFOLLOW | OFlag::O_CLOEXEC, + Mode::from_bits_truncate(0o600), + ) + .map_err(errno_to_io)?; + let file = File::from(descriptor); + let metadata = file.metadata()?; + if !metadata.is_file() || metadata.nlink() != 1 { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidInput, + "config lock must be a regular file with one directory entry", + )); + } + let file = Flock::lock(file, FlockArg::LockExclusive) + .map_err(|(_, error)| errno_to_io(error))?; + Ok(Self { _file: file }) + } + #[cfg(not(unix))] + { + let _ = path; + Ok(Self {}) + } + } +} + /// 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)) + tokio::task::spawn_blocking(move || { + let _lock = ConfigWriteLock::acquire(&path)?; + 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}"))) @@ -39,6 +92,8 @@ pub(in crate::api) async fn write_atomic_if_unchanged( contents: String, ) -> Result<(), ApiFailure> { tokio::task::spawn_blocking(move || { + // Every API mutation locks the root source so writes to different includes serialize. + let _lock = ConfigWriteLock::acquire(&config_path).map_err(AtomicWriteError::Io)?; let graph = ProxyConfig::read_source_graph(&config_path) .map_err(|error| AtomicWriteError::ReadGraph(error.to_string()))?; if compute_source_revision(&graph) != expected_revision { @@ -73,21 +128,65 @@ fn revision_conflict() -> ApiFailure { ) } +fn sibling_lock_path(path: &Path) -> PathBuf { + let mut name = path + .file_name() + .unwrap_or_else(|| std::ffi::OsStr::new("config.toml")) + .to_os_string(); + name.push(".lock"); + path.parent().unwrap_or_else(|| Path::new(".")).join(name) +} + +#[cfg(unix)] +fn open_existing_target(anchored: &AnchoredPath) -> std::io::Result> { + let descriptor = match openat( + anchored.parent(), + anchored.name(), + OFlag::O_RDONLY | OFlag::O_NONBLOCK | OFlag::O_NOFOLLOW | OFlag::O_CLOEXEC, + Mode::empty(), + ) { + Ok(descriptor) => descriptor, + Err(nix::errno::Errno::ENOENT) => return Ok(None), + Err(error) => return Err(errno_to_io(error)), + }; + let mut file = File::from(descriptor); + let metadata = file.metadata()?; + if !metadata.is_file() || metadata.nlink() != 1 || metadata.len() > MAX_CONFIG_SOURCE_BYTES { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidInput, + "config target must be a bounded regular file with one directory entry", + )); + } + let mut contents = String::with_capacity(metadata.len() as usize); + Read::take(&mut file, MAX_CONFIG_SOURCE_BYTES + 1).read_to_string(&mut contents)?; + if contents.len() as u64 > MAX_CONFIG_SOURCE_BYTES { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + "config target exceeds the source size limit", + )); + } + let completed = file.metadata()?; + if !same_target(&metadata, &completed) || metadata.len() != completed.len() { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + "config target changed while it was read", + )); + } + Ok(Some(ExistingTarget { contents, metadata })) +} + +#[cfg(not(unix))] fn open_existing_target(path: &Path) -> std::io::Result> { - let mut options = std::fs::OpenOptions::new(); - options.read(true); - #[cfg(unix)] - options.custom_flags(libc::O_CLOEXEC | libc::O_NOFOLLOW); - let mut file = match options.open(path) { + let mut file = match File::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() { + if !metadata.is_file() || metadata.len() > MAX_CONFIG_SOURCE_BYTES { return Err(std::io::Error::new( std::io::ErrorKind::InvalidInput, - "config target must be a regular file", + "config target must be a bounded regular file", )); } let mut contents = String::new(); @@ -106,6 +205,92 @@ fn same_target(left: &std::fs::Metadata, right: &std::fs::Metadata) -> bool { } } +#[cfg(unix)] +fn write_atomic_sync( + path: &Path, + expected_contents: Option<&str>, + contents: &str, +) -> std::io::Result<()> { + let anchored = AnchoredPath::open_creating_parents(path, 0o750)?; + let existing = open_existing_target(&anchored)?; + validate_expected_contents(existing.as_ref(), expected_contents)?; + let temp_name = format!( + ".{}.tmp-{}", + path.file_name() + .and_then(|name| name.to_str()) + .unwrap_or("config.toml"), + rand::random::() + ); + let descriptor = openat( + anchored.parent(), + temp_name.as_str(), + OFlag::O_WRONLY + | OFlag::O_CREAT + | OFlag::O_EXCL + | OFlag::O_NOFOLLOW + | OFlag::O_CLOEXEC, + Mode::from_bits_truncate(0o600), + ) + .map_err(errno_to_io)?; + let write_result = write_and_publish( + descriptor, + &anchored, + &temp_name, + existing.as_ref(), + contents, + ); + if write_result.is_err() { + let _ = unlinkat( + anchored.parent(), + temp_name.as_str(), + UnlinkatFlags::NoRemoveDir, + ); + } + write_result +} + +#[cfg(unix)] +fn write_and_publish( + descriptor: std::os::fd::OwnedFd, + anchored: &AnchoredPath, + temp_name: &str, + existing: Option<&ExistingTarget>, + contents: &str, +) -> std::io::Result<()> { + let mut file = File::from(descriptor); + if let Some(existing) = existing { + use nix::unistd::{Gid, Uid, fchown}; + + fchown( + &file, + Some(Uid::from_raw(existing.metadata.uid())), + Some(Gid::from_raw(existing.metadata.gid())), + ) + .map_err(errno_to_io)?; + 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(anchored)?; + if !target_unchanged(existing, current.as_ref()) { + return Err(std::io::Error::new( + std::io::ErrorKind::AlreadyExists, + "config target changed during persistence", + )); + } + renameat( + anchored.parent(), + temp_name, + anchored.parent(), + anchored.name(), + ) + .map_err(errno_to_io)?; + fsync(anchored.parent()).map_err(errno_to_io) +} + +#[cfg(not(unix))] fn write_atomic_sync( path: &Path, expected_contents: Option<&str>, @@ -114,72 +299,50 @@ fn write_atomic_sync( let parent = path.parent().unwrap_or_else(|| Path::new(".")); std::fs::create_dir_all(parent)?; let existing = open_existing_target(path)?; + validate_expected_contents(existing.as_ref(), expected_contents)?; + let temp = parent.join(format!(".telemt.tmp-{}", rand::random::())); + std::fs::write(&temp, contents)?; + let current = open_existing_target(path)?; + if !target_unchanged(existing.as_ref(), current.as_ref()) { + let _ = std::fs::remove_file(&temp); + return Err(std::io::Error::new( + std::io::ErrorKind::AlreadyExists, + "config target changed during persistence", + )); + } + std::fs::rename(temp, path) +} + +fn validate_expected_contents( + existing: Option<&ExistingTarget>, + expected_contents: Option<&str>, +) -> std::io::Result<()> { if expected_contents.is_some_and(|expected| { - existing - .as_ref() - .is_none_or(|target| target.contents != expected) + existing.is_none_or(|target| target.contents != expected) }) { return Err(std::io::Error::new( std::io::ErrorKind::AlreadyExists, "config source changed before persistence", )); } - - let tmp_name = format!( - ".{}.tmp-{}", - path.file_name() - .and_then(|name| name.to_str()) - .unwrap_or("config.toml"), - rand::random::() - ); - let tmp_path = parent.join(tmp_name); - - let write_result = (|| { - let mut options = std::fs::OpenOptions::new(); - options.create_new(true).write(true); - #[cfg(unix)] - options.mode(0o600); - let mut file = options.open(&tmp_path)?; - #[cfg(unix)] - if let Some(existing) = existing.as_ref() { - use nix::unistd::{Gid, Uid, fchown}; - - fchown( - &file, - Some(Uid::from_raw(existing.metadata.uid())), - Some(Gid::from_raw(existing.metadata.gid())), - ) - .map_err(|error| std::io::Error::from_raw_os_error(error as i32))?; - file.set_permissions(std::fs::Permissions::from_mode( - existing.metadata.mode() & 0o7777, - ))?; - } - file.write_all(contents.as_bytes())?; - file.sync_all()?; - let current = open_existing_target(path)?; - let target_unchanged = match (&existing, ¤t) { - (Some(expected), Some(current)) => { - same_target(&expected.metadata, ¤t.metadata) - && expected.contents == current.contents - } - (None, None) => true, - _ => false, - }; - if !target_unchanged { - return Err(std::io::Error::new( - std::io::ErrorKind::AlreadyExists, - "config target changed during persistence", - )); - } - std::fs::rename(&tmp_path, path)?; - if let Ok(dir) = std::fs::File::open(parent) { - let _ = dir.sync_all(); - } - Ok(()) - })(); - - if write_result.is_err() { - let _ = std::fs::remove_file(&tmp_path); - } - write_result + Ok(()) +} + +fn target_unchanged( + existing: Option<&ExistingTarget>, + current: Option<&ExistingTarget>, +) -> bool { + match (existing, current) { + (Some(expected), Some(current)) => { + same_target(&expected.metadata, ¤t.metadata) + && expected.contents == current.contents + } + (None, None) => true, + _ => false, + } +} + +#[cfg(unix)] +fn errno_to_io(error: nix::errno::Errno) -> std::io::Error { + std::io::Error::from_raw_os_error(error as i32) } diff --git a/src/api/config_store/tests.rs b/src/api/config_store/tests.rs index ccbcc12..7d09de4 100644 --- a/src/api/config_store/tests.rs +++ b/src/api/config_store/tests.rs @@ -314,6 +314,48 @@ async fn atomic_write_preserves_existing_file_mode() { assert_eq!(after.gid(), before.gid()); } +#[tokio::test] +async fn config_sidecar_lock_serializes_competing_revision_writers() { + let dir = tempfile::tempdir().unwrap(); + let path = dir.path().join("config.toml"); + let original = concat!( + "[censorship]\n", + "tls_domain = \"original.example\"\n", + "[access.users]\n", + "alice = \"00000000000000000000000000000000\"\n" + ); + tokio::fs::write(&path, original).await.unwrap(); + let graph = ProxyConfig::read_source_graph(&path).unwrap(); + let revision = compute_source_revision(&graph); + + let first = tokio::spawn(write_atomic_if_unchanged( + path.clone(), + revision.clone(), + path.clone(), + original.to_string(), + original.replace("original.example", "first.example"), + )); + let second = tokio::spawn(write_atomic_if_unchanged( + path.clone(), + revision, + path.clone(), + original.to_string(), + original.replace("original.example", "second.example"), + )); + let first = first.await.unwrap(); + let second = second.await.unwrap(); + + assert_ne!(first.is_ok(), second.is_ok()); + let conflict = if let Err(error) = first { + error + } else { + second.unwrap_err() + }; + assert_eq!(conflict.code, "revision_conflict"); + let persisted = tokio::fs::read_to_string(path).await.unwrap(); + assert!(persisted.contains("first.example") || persisted.contains("second.example")); +} + #[tokio::test] async fn access_mutation_rejects_sections_with_different_source_owners() { let dir = tempfile::tempdir().unwrap(); diff --git a/src/cli/init.rs b/src/cli/init.rs index fa36218..f53fa5f 100644 --- a/src/cli/init.rs +++ b/src/cli/init.rs @@ -1,4 +1,3 @@ -use std::fs; use std::path::{Path, PathBuf}; use std::process::Command; @@ -114,10 +113,9 @@ pub fn run_init(opts: InitOptions) -> Result<(), Box> { eprintln!("[+] Port: {}", opts.port); eprintln!("[+] Domain: {}", opts.domain); - fs::create_dir_all(&opts.config_dir)?; let config_path = opts.config_dir.join("config.toml"); let config_content = generate_config(&opts.username, &secret, opts.port, &opts.domain); - fs::write(&config_path, &config_content)?; + write_init_file(&config_path, &config_content, 0o600)?; eprintln!("[+] Config written to {}", config_path.display()); let exe_path = @@ -135,22 +133,15 @@ pub fn run_init(opts: InitOptions) -> Result<(), Box> { let service_path = service::service_file_path(init_system); let service_content = service::generate_service_file(init_system, &service_opts); - if let Some(parent) = Path::new(service_path).parent() { - let _ = fs::create_dir_all(parent); - } - - match fs::write(service_path, &service_content) { + let service_mode = if init_system == InitSystem::OpenRC || init_system == InitSystem::FreeBSDRc + { + 0o755 + } else { + 0o644 + }; + match write_init_file(Path::new(service_path), &service_content, service_mode) { Ok(()) => { eprintln!("[+] Service file written to {}", service_path); - - // OpenRC and FreeBSD service scripts must be executable. - #[cfg(unix)] - if init_system == InitSystem::OpenRC || init_system == InitSystem::FreeBSDRc { - use std::os::unix::fs::PermissionsExt; - let mut perms = fs::metadata(service_path)?.permissions(); - perms.set_mode(0o755); - fs::set_permissions(service_path, perms)?; - } } Err(e) => { eprintln!("[!] Cannot write service file (run as root?): {}", e); @@ -226,6 +217,21 @@ pub fn run_init(opts: InitOptions) -> Result<(), Box> { Ok(()) } +fn write_init_file(path: &Path, contents: &str, mode: u32) -> std::io::Result<()> { + #[cfg(unix)] + { + crate::util::secure_fs::atomic_replace(path, contents.as_bytes(), mode) + } + #[cfg(not(unix))] + { + if let Some(parent) = path.parent() { + std::fs::create_dir_all(parent)?; + } + let _ = mode; + std::fs::write(path, contents) + } +} + fn generate_secret() -> String { let mut rng = rand::rng(); let bytes: Vec = (0..16).map(|_| rng.random::()).collect(); diff --git a/src/config/load.rs b/src/config/load.rs index c497a4b..2ec5bd5 100644 --- a/src/config/load.rs +++ b/src/config/load.rs @@ -66,9 +66,9 @@ const MAX_API_REQUEST_BODY_LIMIT_BYTES: usize = 1024 * 1024; pub(crate) struct LoadedConfig { /// Validated and normalized effective configuration. pub(crate) config: ProxyConfig, - /// Canonical paths participating in the recursive include graph. + /// Normalized absolute paths participating in the recursive include graph. pub(crate) source_files: Vec, - /// Raw source bytes keyed by canonical source path. + /// Raw source bytes keyed by normalized absolute source path. pub(crate) source_contents: BTreeMap, /// Legacy hash of the include-expanded rendered snapshot. pub(crate) rendered_hash: u64, @@ -77,7 +77,7 @@ pub(crate) struct LoadedConfig { /// Raw recursive source graph captured before typed deserialization. #[derive(Debug, Clone)] pub(crate) struct ConfigSourceGraph { - /// Raw source bytes keyed by canonical source path. + /// Raw source bytes keyed by normalized absolute source path. pub(crate) source_contents: BTreeMap, /// Include-expanded TOML used for typed deserialization. pub(crate) rendered: String, @@ -177,20 +177,42 @@ impl ProxyConfig { source_overrides: &BTreeMap, ) -> Result { let path = path.as_ref(); - let initial_path = normalize_config_path(path); - let (normalized_path, content) = if let Some(content) = source_overrides.get(&initial_path) { - (initial_path, content.clone()) - } else { - read_config_source(path)? - }; - let base_dir = path.parent().unwrap_or(Path::new(".")); + let mut previous = Self::capture_source_graph(path, source_overrides)?; + for _ in 0..2 { + let current = Self::capture_source_graph(path, source_overrides)?; + if current.source_contents == previous.source_contents + && current.rendered == previous.rendered + { + return Ok(current); + } + previous = current; + } + Err(ProxyError::Config( + "config source graph changed repeatedly while it was read".to_string(), + )) + } + + fn capture_source_graph( + path: &Path, + source_overrides: &BTreeMap, + ) -> Result { + let path = path.as_ref(); + let (normalized_path, disk_content) = read_config_source(path)?; + let content = source_overrides + .get(&normalized_path) + .cloned() + .unwrap_or(disk_content); + let base_dir = normalized_path + .parent() + .unwrap_or(Path::new(".")) + .to_path_buf(); let mut source_files = BTreeSet::new(); source_files.insert(normalized_path.clone()); let mut source_contents = BTreeMap::new(); source_contents.insert(normalized_path, content.clone()); let processed = preprocess_includes( &content, - base_dir, + &base_dir, 0, &mut source_files, &mut source_contents, diff --git a/src/config/load/includes.rs b/src/config/load/includes.rs index aef22ef..b964157 100644 --- a/src/config/load/includes.rs +++ b/src/config/load/includes.rs @@ -1,23 +1,30 @@ use std::collections::{BTreeMap, BTreeSet}; use std::hash::{DefaultHasher, Hash, Hasher}; -use std::io::Read; use std::path::{Path, PathBuf}; -#[cfg(unix)] -use std::os::unix::fs::{MetadataExt, OpenOptionsExt}; - use crate::error::{ProxyError, Result}; +const MAX_CONFIG_SOURCE_BYTES: usize = 8 * 1024 * 1024; + pub(super) fn normalize_config_path(path: &Path) -> PathBuf { - path.canonicalize().unwrap_or_else(|_| { - if path.is_absolute() { - path.to_path_buf() - } else { - std::env::current_dir() - .map(|cwd| cwd.join(path)) - .unwrap_or_else(|_| path.to_path_buf()) + let absolute = if path.is_absolute() { + path.to_path_buf() + } else { + std::env::current_dir() + .map(|cwd| cwd.join(path)) + .unwrap_or_else(|_| path.to_path_buf()) + }; + let mut normalized = PathBuf::new(); + for component in absolute.components() { + match component { + std::path::Component::CurDir => {} + std::path::Component::ParentDir => { + normalized.pop(); + } + component => normalized.push(component.as_os_str()), } - }) + } + normalized } pub(super) fn hash_rendered_snapshot(rendered: &str) -> u64 { @@ -27,73 +34,28 @@ pub(super) fn hash_rendered_snapshot(rendered: &str) -> u64 { } 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) + let bytes = crate::util::secure_fs::read_regular_limited(path, MAX_CONFIG_SOURCE_BYTES) .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() { + #[cfg(not(unix))] + let bytes = std::fs::read(path).map_err(|error| ProxyError::Config(error.to_string()))?; + if bytes.len() > MAX_CONFIG_SOURCE_BYTES { return Err(ProxyError::Config(format!( - "config source `{}` must be a regular file", - path.display() + "config source `{}` exceeds {} bytes", + path.display(), + MAX_CONFIG_SOURCE_BYTES ))); } + let contents = String::from_utf8(bytes).map_err(|error| { + ProxyError::Config(format!( + "config source `{}` is not valid UTF-8: {error}", + path.display() + )) + })?; let normalized = normalize_config_path(path); - let current_metadata = std::fs::metadata(&normalized) - .map_err(|error| ProxyError::Config(error.to_string()))?; - if !same_file_identity(&opened_metadata, ¤t_metadata) { - return Err(ProxyError::Config(format!( - "config source `{}` changed while it was opened", - path.display() - ))); - } - let mut contents = String::new(); - file.read_to_string(&mut contents) - .map_err(|error| ProxyError::Config(error.to_string()))?; - let completed_metadata = file - .metadata() - .map_err(|error| ProxyError::Config(error.to_string()))?; - if !same_file_version(&opened_metadata, &completed_metadata) { - return Err(ProxyError::Config(format!( - "config source `{}` changed while it was read", - path.display() - ))); - } Ok((normalized, contents)) } -fn same_file_identity(left: &std::fs::Metadata, right: &std::fs::Metadata) -> bool { - #[cfg(unix)] - { - left.dev() == right.dev() && left.ino() == right.ino() - } - #[cfg(not(unix))] - { - left.len() == right.len() && left.modified().ok() == right.modified().ok() - } -} - -fn same_file_version(left: &std::fs::Metadata, right: &std::fs::Metadata) -> bool { - #[cfg(unix)] - { - same_file_identity(left, right) - && left.len() == right.len() - && left.mtime() == right.mtime() - && left.mtime_nsec() == right.mtime_nsec() - && left.ctime() == right.ctime() - && left.ctime_nsec() == right.ctime_nsec() - } - #[cfg(not(unix))] - { - same_file_identity(left, right) - } -} - pub(super) fn preprocess_includes( content: &str, base_dir: &Path, @@ -113,20 +75,17 @@ pub(super) fn preprocess_includes( if let Some(rest) = rest.strip_prefix('=') { let path_str = rest.trim().trim_matches('"'); let resolved = base_dir.join(path_str); - let normalized = normalize_config_path(&resolved); - let cached = source_contents.get(&normalized).cloned(); - let (normalized, included) = if let Some(included) = - source_overrides.get(&normalized).cloned().or(cached) - { - (normalized, included) - } else { - read_config_source(&resolved)? - }; + let (normalized, disk_contents) = read_config_source(&resolved)?; + let included = source_overrides + .get(&normalized) + .cloned() + .or_else(|| source_contents.get(&normalized).cloned()) + .unwrap_or(disk_contents); source_files.insert(normalized.clone()); source_contents - .entry(normalized) + .entry(normalized.clone()) .or_insert_with(|| included.clone()); - let included_dir = resolved.parent().unwrap_or(base_dir); + let included_dir = normalized.parent().unwrap_or(base_dir); output.push_str(&preprocess_includes( &included, included_dir, diff --git a/src/config/load/runtime_auth.rs b/src/config/load/runtime_auth.rs index 3fc914b..8dc3964 100644 --- a/src/config/load/runtime_auth.rs +++ b/src/config/load/runtime_auth.rs @@ -12,6 +12,7 @@ const ACCESS_SECRET_BYTES: usize = 16; pub(crate) struct UserAuthSnapshot { entries: Vec, by_name: HashMap, + by_hint_key: HashMap>, sni_index: HashMap>, sni_initial_index: HashMap>, } @@ -22,16 +23,21 @@ pub(crate) struct UserAuthEntry { pub(crate) secret: [u8; ACCESS_SECRET_BYTES], /// Stable secret identity used by process-wide admission fencing. pub(crate) credential_id: [u8; 16], + /// Stable compact key used only to resolve bounded authentication hints. + pub(crate) hint_key: u64, } impl UserAuthSnapshot { pub(super) fn from_users(users: &HashMap) -> Result { let mut entries = Vec::with_capacity(users.len()); let mut by_name = HashMap::with_capacity(users.len()); + let mut by_hint_key = HashMap::with_capacity(users.len()); let mut sni_index = HashMap::with_capacity(users.len()); let mut sni_initial_index = HashMap::with_capacity(users.len()); - for (user, secret_hex) in users { + let mut ordered_users = users.iter().collect::>(); + ordered_users.sort_unstable_by(|(left, _), (right, _)| left.cmp(right)); + for (user, secret_hex) in ordered_users { let decoded = hex::decode(secret_hex).map_err(|_| ProxyError::InvalidSecret { user: user.clone(), reason: "Must be 32 hex characters".to_string(), @@ -52,12 +58,27 @@ impl UserAuthSnapshot { let digest = sha256(&secret); let mut credential_id = [0; 16]; credential_id.copy_from_slice(&digest[..16]); + let hint_key = u64::from_le_bytes([ + credential_id[0], + credential_id[1], + credential_id[2], + credential_id[3], + credential_id[4], + credential_id[5], + credential_id[6], + credential_id[7], + ]) | 1; entries.push(UserAuthEntry { user: user.clone(), secret, credential_id, + hint_key, }); by_name.insert(user.clone(), user_id); + by_hint_key + .entry(hint_key) + .or_insert_with(Vec::new) + .push(user_id); sni_index .entry(Self::sni_lookup_hash(user)) .or_insert_with(Vec::new) @@ -77,6 +98,7 @@ impl UserAuthSnapshot { Ok(Self { entries, by_name, + by_hint_key, sni_index, sni_initial_index, }) @@ -101,6 +123,10 @@ impl UserAuthSnapshot { .map(|entry| entry.credential_id) } + pub(crate) fn candidate_ids_by_hint_key(&self, hint_key: u64) -> Option<&[u32]> { + self.by_hint_key.get(&hint_key).map(Vec::as_slice) + } + pub(crate) fn sni_candidates(&self, sni: &str) -> Option<&[u32]> { self.sni_index .get(&Self::sni_lookup_hash(sni)) @@ -123,3 +149,48 @@ impl UserAuthSnapshot { hasher.finish() } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn credential_hint_survives_positional_id_shift() { + let mut initial = HashMap::new(); + initial.insert( + "alice".to_string(), + "11111111111111111111111111111111".to_string(), + ); + initial.insert( + "bob".to_string(), + "22222222222222222222222222222222".to_string(), + ); + let initial = UserAuthSnapshot::from_users(&initial).unwrap(); + let initial_id = initial.user_id_by_name("alice").unwrap(); + let hint_key = initial.entry_by_id(initial_id).unwrap().hint_key; + + let mut reloaded = HashMap::new(); + reloaded.insert( + "aaron".to_string(), + "33333333333333333333333333333333".to_string(), + ); + reloaded.insert( + "alice".to_string(), + "11111111111111111111111111111111".to_string(), + ); + reloaded.insert( + "bob".to_string(), + "22222222222222222222222222222222".to_string(), + ); + let reloaded = UserAuthSnapshot::from_users(&reloaded).unwrap(); + let reloaded_id = reloaded.user_id_by_name("alice").unwrap(); + + assert_ne!(initial_id, reloaded_id); + assert!( + reloaded + .candidate_ids_by_hint_key(hint_key) + .unwrap() + .contains(&reloaded_id) + ); + } +} diff --git a/src/config/load/runtime_web.rs b/src/config/load/runtime_web.rs index 3b701c1..877aa38 100644 --- a/src/config/load/runtime_web.rs +++ b/src/config/load/runtime_web.rs @@ -23,6 +23,8 @@ use hmac::{Hmac, Mac}; use sha2::{Digest, Sha256}; use super::*; +#[cfg(unix)] +use crate::util::secure_fs::open_dir_nofollow; // Path-based static snapshot fallback for platforms without directory descriptors. #[cfg(not(unix))] @@ -247,16 +249,17 @@ fn load_static_site( #[cfg(unix)] fn open_static_root(root: &Path) -> Result { - Dir::open( - root, - OFlag::O_RDONLY | OFlag::O_DIRECTORY | OFlag::O_NOFOLLOW | OFlag::O_CLOEXEC, - Mode::empty(), - ) - .map_err(|error| { + let descriptor = open_dir_nofollow(root).map_err(|error| { ProxyError::Config(format!( "WEB static directory `{}` must be a real directory, not a symlink: {error}", root.display() )) + })?; + Dir::from_fd(descriptor).map_err(|error| { + ProxyError::Config(format!( + "failed to read WEB static directory `{}`: {error}", + root.display() + )) }) } diff --git a/src/config/tests/load_basic_tests.rs b/src/config/tests/load_basic_tests.rs index e24285f..82c7da4 100644 --- a/src/config/tests/load_basic_tests.rs +++ b/src/config/tests/load_basic_tests.rs @@ -46,6 +46,8 @@ mod legacy_policy_tests; mod me_route_tests; #[path = "load_basic_tests/me_startup_tests.rs"] mod me_startup_tests; +#[path = "load_basic_tests/source_security_tests.rs"] +mod source_security_tests; #[path = "load_basic_tests/synlimit_mss_tests.rs"] mod synlimit_mss_tests; #[path = "load_basic_tests/tls_fetch_tests.rs"] diff --git a/src/config/tests/load_basic_tests/source_security_tests.rs b/src/config/tests/load_basic_tests/source_security_tests.rs new file mode 100644 index 0000000..86706eb --- /dev/null +++ b/src/config/tests/load_basic_tests/source_security_tests.rs @@ -0,0 +1,42 @@ +#[cfg(unix)] +use std::os::unix::fs::symlink; + +use super::*; + +#[cfg(unix)] +#[test] +fn config_loader_rejects_final_and_intermediate_symlinks() { + let directory = tempfile::tempdir().unwrap(); + let real_directory = directory.path().join("real"); + let linked_directory = directory.path().join("linked"); + std::fs::create_dir(&real_directory).unwrap(); + let real_config = real_directory.join("config.toml"); + let final_link = directory.path().join("config.toml"); + std::fs::write(&real_config, "[general]\n").unwrap(); + symlink(&real_config, &final_link).unwrap(); + symlink(&real_directory, &linked_directory).unwrap(); + + assert!(ProxyConfig::load(&final_link).is_err()); + assert!(ProxyConfig::load(linked_directory.join("config.toml")).is_err()); +} + +#[test] +fn config_loader_rejects_oversized_source() { + let directory = tempfile::tempdir().unwrap(); + let path = directory.path().join("config.toml"); + std::fs::write(&path, vec![b' '; 8 * 1024 * 1024 + 1]).unwrap(); + + let error = ProxyConfig::load(&path).unwrap_err().to_string(); + + assert!(error.contains("size limit") || error.contains("exceeds")); +} + +#[cfg(unix)] +#[test] +fn config_loader_rejects_fifo_without_blocking() { + let directory = tempfile::tempdir().unwrap(); + let path = directory.path().join("config.toml"); + nix::unistd::mkfifo(&path, nix::sys::stat::Mode::S_IRUSR).unwrap(); + + assert!(ProxyConfig::load(&path).is_err()); +} diff --git a/src/conntrack_control.rs b/src/conntrack_control.rs index 8024cb2..275e75d 100644 --- a/src/conntrack_control.rs +++ b/src/conntrack_control.rs @@ -1,20 +1,22 @@ -use std::collections::BTreeSet; -use std::net::IpAddr; use std::sync::Arc; use std::time::Duration; -use tokio::io::AsyncWriteExt; -use tokio::process::Command; use tokio::sync::{mpsc, watch}; use tokio_util::sync::CancellationToken; -use tracing::{debug, info, warn}; +use tracing::{info, warn}; -use crate::config::{ConntrackBackend, ConntrackMode, ProxyConfig}; +use crate::config::ProxyConfig; use crate::proxy::middle_relay::note_global_relay_pressure; use crate::proxy::shared_state::{ConntrackCloseEvent, ConntrackCloseReason, ProxySharedState}; use crate::stats::Stats; -#[cfg(unix)] -use crate::util::trusted_command::resolve_trusted_helper; + +// Privileged netfilter rule and conntrack helper execution. +mod firewall; + +use firewall::{ + DeleteOutcome, delete_conntrack_entry, effective_conntrack_enabled, probe_runtime_support, + reconcile_rules, +}; const CONNTRACK_EVENT_QUEUE_CAPACITY: usize = 32_768; const PRESSURE_RELEASE_TICKS: u8 = 3; @@ -311,389 +313,6 @@ fn update_pressure_state( state.low_streak = 0; } -async fn reconcile_rules( - cfg: &ProxyConfig, - runtime_support: ConntrackRuntimeSupport, - stats: &Stats, -) { - if !cfg.server.conntrack_control.inline_conntrack_control { - clear_notrack_rules_all_backends().await; - stats.set_conntrack_rule_apply_ok(true); - return; - } - - if !effective_conntrack_enabled(cfg, runtime_support) { - clear_notrack_rules_all_backends().await; - stats.set_conntrack_rule_apply_ok(false); - return; - } - - let backend = runtime_support - .netfilter_backend - .expect("netfilter backend must be available for effective conntrack control"); - - let apply_result = match backend { - NetfilterBackend::Nftables => apply_nft_rules(cfg).await, - NetfilterBackend::Iptables => apply_iptables_rules(cfg).await, - }; - - if let Err(error) = apply_result { - warn!(error = %error, "Failed to reconcile conntrack/notrack rules"); - stats.set_conntrack_rule_apply_ok(false); - } else { - stats.set_conntrack_rule_apply_ok(true); - } -} - -fn probe_runtime_support(configured_backend: ConntrackBackend) -> ConntrackRuntimeSupport { - ConntrackRuntimeSupport { - netfilter_backend: pick_backend(configured_backend), - has_cap_net_admin: has_cap_net_admin(), - has_conntrack_binary: command_exists("conntrack"), - } -} - -fn effective_conntrack_enabled( - cfg: &ProxyConfig, - runtime_support: ConntrackRuntimeSupport, -) -> bool { - cfg.server.conntrack_control.inline_conntrack_control - && runtime_support.has_cap_net_admin - && runtime_support.netfilter_backend.is_some() - && runtime_support.has_conntrack_binary -} - -fn pick_backend(configured: ConntrackBackend) -> Option { - match configured { - ConntrackBackend::Auto => { - if command_exists("nft") { - Some(NetfilterBackend::Nftables) - } else if command_exists("iptables") { - Some(NetfilterBackend::Iptables) - } else { - None - } - } - ConntrackBackend::Nftables => command_exists("nft").then_some(NetfilterBackend::Nftables), - ConntrackBackend::Iptables => { - command_exists("iptables").then_some(NetfilterBackend::Iptables) - } - } -} - -fn command_exists(binary: &str) -> bool { - #[cfg(unix)] - { - resolve_trusted_helper(binary).is_some() - } - #[cfg(not(unix))] - { - let _ = binary; - false - } -} - -fn listener_port_set(cfg: &ProxyConfig) -> Vec { - let mut ports: BTreeSet = BTreeSet::new(); - if cfg.server.listeners.is_empty() { - ports.insert(cfg.server.port); - } else { - for listener in &cfg.server.listeners { - ports.insert(listener.port.unwrap_or(cfg.server.port)); - } - } - ports.into_iter().collect() -} - -fn notrack_targets(cfg: &ProxyConfig) -> (Vec<(Option, u16)>, Vec<(Option, u16)>) { - let mode = cfg.server.conntrack_control.mode; - let mut v4_targets: BTreeSet<(Option, u16)> = BTreeSet::new(); - let mut v6_targets: BTreeSet<(Option, u16)> = BTreeSet::new(); - - match mode { - ConntrackMode::Tracked => {} - ConntrackMode::Notrack => { - if cfg.server.listeners.is_empty() { - let port = cfg.server.port; - if let Some(ipv4) = cfg - .server - .listen_addr_ipv4 - .as_ref() - .and_then(|s| s.parse::().ok()) - { - if ipv4.is_unspecified() { - v4_targets.insert((None, port)); - } else { - v4_targets.insert((Some(ipv4), port)); - } - } - if let Some(ipv6) = cfg - .server - .listen_addr_ipv6 - .as_ref() - .and_then(|s| s.parse::().ok()) - { - if ipv6.is_unspecified() { - v6_targets.insert((None, port)); - } else { - v6_targets.insert((Some(ipv6), port)); - } - } - } else { - for listener in &cfg.server.listeners { - let port = listener.port.unwrap_or(cfg.server.port); - if listener.ip.is_ipv4() { - if listener.ip.is_unspecified() { - v4_targets.insert((None, port)); - } else { - v4_targets.insert((Some(listener.ip), port)); - } - } else if listener.ip.is_unspecified() { - v6_targets.insert((None, port)); - } else { - v6_targets.insert((Some(listener.ip), port)); - } - } - } - } - ConntrackMode::Hybrid => { - let ports = listener_port_set(cfg); - for ip in &cfg.server.conntrack_control.hybrid_listener_ips { - if ip.is_ipv4() { - for port in &ports { - v4_targets.insert((Some(*ip), *port)); - } - } else { - for port in &ports { - v6_targets.insert((Some(*ip), *port)); - } - } - } - } - } - - ( - v4_targets.into_iter().collect(), - v6_targets.into_iter().collect(), - ) -} - -async fn apply_nft_rules(cfg: &ProxyConfig) -> Result<(), String> { - let _ = run_command( - "nft", - &["delete", "table", "inet", "telemt_conntrack"], - None, - ) - .await; - if matches!(cfg.server.conntrack_control.mode, ConntrackMode::Tracked) { - return Ok(()); - } - - let (v4_targets, v6_targets) = notrack_targets(cfg); - let mut rules = Vec::new(); - for (ip, port) in v4_targets { - let rule = if let Some(ip) = ip { - format!("tcp dport {} ip daddr {} notrack", port, ip) - } else { - format!("tcp dport {} notrack", port) - }; - rules.push(rule); - } - for (ip, port) in v6_targets { - let rule = if let Some(ip) = ip { - format!("tcp dport {} ip6 daddr {} notrack", port, ip) - } else { - format!("tcp dport {} notrack", port) - }; - rules.push(rule); - } - - let rule_blob = if rules.is_empty() { - String::new() - } else { - format!(" {}\n", rules.join("\n ")) - }; - let script = format!( - "table inet telemt_conntrack {{\n chain preraw {{\n type filter hook prerouting priority raw; policy accept;\n{rule_blob} }}\n}}\n" - ); - run_command("nft", &["-f", "-"], Some(script)).await -} - -async fn apply_iptables_rules(cfg: &ProxyConfig) -> Result<(), String> { - apply_iptables_rules_for_binary("iptables", cfg, true).await?; - apply_iptables_rules_for_binary("ip6tables", cfg, false).await?; - Ok(()) -} - -async fn apply_iptables_rules_for_binary( - binary: &str, - cfg: &ProxyConfig, - ipv4: bool, -) -> Result<(), String> { - if !command_exists(binary) { - return Ok(()); - } - let chain = "TELEMT_NOTRACK"; - let _ = run_command( - binary, - &["-t", "raw", "-D", "PREROUTING", "-j", chain], - None, - ) - .await; - let _ = run_command(binary, &["-t", "raw", "-F", chain], None).await; - let _ = run_command(binary, &["-t", "raw", "-X", chain], None).await; - - if matches!(cfg.server.conntrack_control.mode, ConntrackMode::Tracked) { - return Ok(()); - } - - run_command(binary, &["-t", "raw", "-N", chain], None).await?; - run_command(binary, &["-t", "raw", "-F", chain], None).await?; - if run_command( - binary, - &["-t", "raw", "-C", "PREROUTING", "-j", chain], - None, - ) - .await - .is_err() - { - run_command( - binary, - &["-t", "raw", "-I", "PREROUTING", "1", "-j", chain], - None, - ) - .await?; - } - - let (v4_targets, v6_targets) = notrack_targets(cfg); - let selected = if ipv4 { v4_targets } else { v6_targets }; - for (ip, port) in selected { - let mut args = vec![ - "-t".to_string(), - "raw".to_string(), - "-A".to_string(), - chain.to_string(), - "-p".to_string(), - "tcp".to_string(), - "--dport".to_string(), - port.to_string(), - ]; - if let Some(ip) = ip { - args.push("-d".to_string()); - args.push(ip.to_string()); - } - args.push("-j".to_string()); - args.push("CT".to_string()); - args.push("--notrack".to_string()); - let arg_refs: Vec<&str> = args.iter().map(String::as_str).collect(); - run_command(binary, &arg_refs, None).await?; - } - Ok(()) -} - -async fn clear_notrack_rules_all_backends() { - let _ = run_command( - "nft", - &["delete", "table", "inet", "telemt_conntrack"], - None, - ) - .await; - let _ = run_command( - "iptables", - &["-t", "raw", "-D", "PREROUTING", "-j", "TELEMT_NOTRACK"], - None, - ) - .await; - let _ = run_command("iptables", &["-t", "raw", "-F", "TELEMT_NOTRACK"], None).await; - let _ = run_command("iptables", &["-t", "raw", "-X", "TELEMT_NOTRACK"], None).await; - let _ = run_command( - "ip6tables", - &["-t", "raw", "-D", "PREROUTING", "-j", "TELEMT_NOTRACK"], - None, - ) - .await; - let _ = run_command("ip6tables", &["-t", "raw", "-F", "TELEMT_NOTRACK"], None).await; - let _ = run_command("ip6tables", &["-t", "raw", "-X", "TELEMT_NOTRACK"], None).await; -} - -enum DeleteOutcome { - Deleted, - NotFound, - Error, -} - -async fn delete_conntrack_entry(event: ConntrackCloseEvent) -> DeleteOutcome { - if !command_exists("conntrack") { - return DeleteOutcome::Error; - } - let args = vec![ - "-D".to_string(), - "-p".to_string(), - "tcp".to_string(), - "-s".to_string(), - event.src.ip().to_string(), - "--sport".to_string(), - event.src.port().to_string(), - "-d".to_string(), - event.dst.ip().to_string(), - "--dport".to_string(), - event.dst.port().to_string(), - ]; - let arg_refs: Vec<&str> = args.iter().map(String::as_str).collect(); - match run_command("conntrack", &arg_refs, None).await { - Ok(()) => DeleteOutcome::Deleted, - Err(error) => { - if error.contains("0 flow entries have been deleted") { - DeleteOutcome::NotFound - } else { - debug!(error = %error, "conntrack delete failed"); - DeleteOutcome::Error - } - } - } -} - -async fn run_command(binary: &str, args: &[&str], stdin: Option) -> Result<(), String> { - #[cfg(unix)] - let Some(command_path) = resolve_trusted_helper(binary) else { - return Err(format!("{binary} is not available")); - }; - #[cfg(not(unix))] - return Err(format!("{binary} is not available")); - #[cfg(unix)] - let mut command = Command::new(command_path); - command.args(args); - if stdin.is_some() { - command.stdin(std::process::Stdio::piped()); - } - command.stdout(std::process::Stdio::null()); - command.stderr(std::process::Stdio::piped()); - let mut child = command - .spawn() - .map_err(|e| format!("spawn {binary} failed: {e}"))?; - if let Some(blob) = stdin - && let Some(mut writer) = child.stdin.take() - { - writer - .write_all(blob.as_bytes()) - .await - .map_err(|e| format!("stdin write {binary} failed: {e}"))?; - } - let output = child - .wait_with_output() - .await - .map_err(|e| format!("wait {binary} failed: {e}"))?; - if output.status.success() { - return Ok(()); - } - let stderr = String::from_utf8_lossy(&output.stderr).trim().to_string(); - Err(if stderr.is_empty() { - format!("{binary} exited with status {}", output.status) - } else { - stderr - }) -} - fn fd_usage_pct() -> Option { let soft_limit = nofile_soft_limit()?; if soft_limit == 0 { @@ -722,29 +341,6 @@ fn nofile_soft_limit() -> Option { } } -fn has_cap_net_admin() -> bool { - #[cfg(target_os = "linux")] - { - let Ok(status) = std::fs::read_to_string("/proc/self/status") else { - return false; - }; - for line in status.lines() { - if let Some(raw) = line.strip_prefix("CapEff:") { - let caps = raw.trim(); - if let Ok(bits) = u64::from_str_radix(caps, 16) { - const CAP_NET_ADMIN_BIT: u64 = 12; - return (bits & (1u64 << CAP_NET_ADMIN_BIT)) != 0; - } - } - } - false - } - #[cfg(not(target_os = "linux"))] - { - false - } -} - #[cfg(test)] mod tests { use super::*; diff --git a/src/conntrack_control/firewall.rs b/src/conntrack_control/firewall.rs new file mode 100644 index 0000000..73110cf --- /dev/null +++ b/src/conntrack_control/firewall.rs @@ -0,0 +1,419 @@ +use std::collections::BTreeSet; +use std::net::IpAddr; + +use tokio::io::AsyncWriteExt; +use tokio::process::Command; +use tracing::{debug, warn}; + +use crate::config::{ConntrackBackend, ConntrackMode, ProxyConfig}; +use crate::proxy::shared_state::ConntrackCloseEvent; +use crate::stats::Stats; +#[cfg(unix)] +use crate::util::trusted_command::resolve_trusted_helper; + +use super::{ConntrackRuntimeSupport, NetfilterBackend}; + +pub(super) async fn reconcile_rules( + cfg: &ProxyConfig, + runtime_support: ConntrackRuntimeSupport, + stats: &Stats, +) { + if !cfg.server.conntrack_control.inline_conntrack_control { + clear_notrack_rules_all_backends().await; + stats.set_conntrack_rule_apply_ok(true); + return; + } + + if !effective_conntrack_enabled(cfg, runtime_support) { + clear_notrack_rules_all_backends().await; + stats.set_conntrack_rule_apply_ok(false); + return; + } + + let backend = runtime_support + .netfilter_backend + .expect("netfilter backend must be available for effective conntrack control"); + let apply_result = match backend { + NetfilterBackend::Nftables => apply_nft_rules(cfg).await, + NetfilterBackend::Iptables => apply_iptables_rules(cfg).await, + }; + if let Err(error) = apply_result { + warn!(error = %error, "Failed to reconcile conntrack/notrack rules"); + stats.set_conntrack_rule_apply_ok(false); + } else { + stats.set_conntrack_rule_apply_ok(true); + } +} + +pub(super) fn probe_runtime_support( + configured_backend: ConntrackBackend, +) -> ConntrackRuntimeSupport { + ConntrackRuntimeSupport { + netfilter_backend: pick_backend(configured_backend), + has_cap_net_admin: has_cap_net_admin(), + has_conntrack_binary: command_exists("conntrack"), + } +} + +pub(super) fn effective_conntrack_enabled( + cfg: &ProxyConfig, + runtime_support: ConntrackRuntimeSupport, +) -> bool { + cfg.server.conntrack_control.inline_conntrack_control + && runtime_support.has_cap_net_admin + && runtime_support.netfilter_backend.is_some() + && runtime_support.has_conntrack_binary +} + +fn pick_backend(configured: ConntrackBackend) -> Option { + match configured { + ConntrackBackend::Auto => { + if command_exists("nft") { + Some(NetfilterBackend::Nftables) + } else if command_exists("iptables") { + Some(NetfilterBackend::Iptables) + } else { + None + } + } + ConntrackBackend::Nftables => command_exists("nft").then_some(NetfilterBackend::Nftables), + ConntrackBackend::Iptables => { + command_exists("iptables").then_some(NetfilterBackend::Iptables) + } + } +} + +fn command_exists(binary: &str) -> bool { + #[cfg(unix)] + { + resolve_trusted_helper(binary).is_some() + } + #[cfg(not(unix))] + { + let _ = binary; + false + } +} + +fn listener_port_set(cfg: &ProxyConfig) -> Vec { + let mut ports: BTreeSet = BTreeSet::new(); + if cfg.server.listeners.is_empty() { + ports.insert(cfg.server.port); + } else { + for listener in &cfg.server.listeners { + ports.insert(listener.port.unwrap_or(cfg.server.port)); + } + } + ports.into_iter().collect() +} + +fn notrack_targets(cfg: &ProxyConfig) -> (Vec<(Option, u16)>, Vec<(Option, u16)>) { + let mode = cfg.server.conntrack_control.mode; + let mut v4_targets: BTreeSet<(Option, u16)> = BTreeSet::new(); + let mut v6_targets: BTreeSet<(Option, u16)> = BTreeSet::new(); + + match mode { + ConntrackMode::Tracked => {} + ConntrackMode::Notrack => { + if cfg.server.listeners.is_empty() { + let port = cfg.server.port; + if let Some(ipv4) = cfg + .server + .listen_addr_ipv4 + .as_ref() + .and_then(|value| value.parse::().ok()) + { + if ipv4.is_unspecified() { + v4_targets.insert((None, port)); + } else { + v4_targets.insert((Some(ipv4), port)); + } + } + if let Some(ipv6) = cfg + .server + .listen_addr_ipv6 + .as_ref() + .and_then(|value| value.parse::().ok()) + { + if ipv6.is_unspecified() { + v6_targets.insert((None, port)); + } else { + v6_targets.insert((Some(ipv6), port)); + } + } + } else { + for listener in &cfg.server.listeners { + let port = listener.port.unwrap_or(cfg.server.port); + if listener.ip.is_ipv4() { + if listener.ip.is_unspecified() { + v4_targets.insert((None, port)); + } else { + v4_targets.insert((Some(listener.ip), port)); + } + } else if listener.ip.is_unspecified() { + v6_targets.insert((None, port)); + } else { + v6_targets.insert((Some(listener.ip), port)); + } + } + } + } + ConntrackMode::Hybrid => { + let ports = listener_port_set(cfg); + for ip in &cfg.server.conntrack_control.hybrid_listener_ips { + if ip.is_ipv4() { + for port in &ports { + v4_targets.insert((Some(*ip), *port)); + } + } else { + for port in &ports { + v6_targets.insert((Some(*ip), *port)); + } + } + } + } + } + + ( + v4_targets.into_iter().collect(), + v6_targets.into_iter().collect(), + ) +} + +async fn apply_nft_rules(cfg: &ProxyConfig) -> Result<(), String> { + let _ = run_command( + "nft", + &["delete", "table", "inet", "telemt_conntrack"], + None, + ) + .await; + if matches!(cfg.server.conntrack_control.mode, ConntrackMode::Tracked) { + return Ok(()); + } + + let (v4_targets, v6_targets) = notrack_targets(cfg); + let mut rules = Vec::new(); + for (ip, port) in v4_targets { + let rule = if let Some(ip) = ip { + format!("tcp dport {} ip daddr {} notrack", port, ip) + } else { + format!("tcp dport {} notrack", port) + }; + rules.push(rule); + } + for (ip, port) in v6_targets { + let rule = if let Some(ip) = ip { + format!("tcp dport {} ip6 daddr {} notrack", port, ip) + } else { + format!("tcp dport {} notrack", port) + }; + rules.push(rule); + } + + let rule_blob = if rules.is_empty() { + String::new() + } else { + format!(" {}\n", rules.join("\n ")) + }; + let script = format!( + "table inet telemt_conntrack {{\n chain preraw {{\n type filter hook prerouting priority raw; policy accept;\n{rule_blob} }}\n}}\n" + ); + run_command("nft", &["-f", "-"], Some(script)).await +} + +async fn apply_iptables_rules(cfg: &ProxyConfig) -> Result<(), String> { + apply_iptables_rules_for_binary("iptables", cfg, true).await?; + apply_iptables_rules_for_binary("ip6tables", cfg, false).await?; + Ok(()) +} + +async fn apply_iptables_rules_for_binary( + binary: &str, + cfg: &ProxyConfig, + ipv4: bool, +) -> Result<(), String> { + if !command_exists(binary) { + return Ok(()); + } + let chain = "TELEMT_NOTRACK"; + let _ = run_command( + binary, + &["-t", "raw", "-D", "PREROUTING", "-j", chain], + None, + ) + .await; + let _ = run_command(binary, &["-t", "raw", "-F", chain], None).await; + let _ = run_command(binary, &["-t", "raw", "-X", chain], None).await; + if matches!(cfg.server.conntrack_control.mode, ConntrackMode::Tracked) { + return Ok(()); + } + + run_command(binary, &["-t", "raw", "-N", chain], None).await?; + run_command(binary, &["-t", "raw", "-F", chain], None).await?; + if run_command( + binary, + &["-t", "raw", "-C", "PREROUTING", "-j", chain], + None, + ) + .await + .is_err() + { + run_command( + binary, + &["-t", "raw", "-I", "PREROUTING", "1", "-j", chain], + None, + ) + .await?; + } + + let (v4_targets, v6_targets) = notrack_targets(cfg); + let selected = if ipv4 { v4_targets } else { v6_targets }; + for (ip, port) in selected { + let mut args = vec![ + "-t".to_string(), + "raw".to_string(), + "-A".to_string(), + chain.to_string(), + "-p".to_string(), + "tcp".to_string(), + "--dport".to_string(), + port.to_string(), + ]; + if let Some(ip) = ip { + args.push("-d".to_string()); + args.push(ip.to_string()); + } + args.push("-j".to_string()); + args.push("CT".to_string()); + args.push("--notrack".to_string()); + let arg_refs: Vec<&str> = args.iter().map(String::as_str).collect(); + run_command(binary, &arg_refs, None).await?; + } + Ok(()) +} + +async fn clear_notrack_rules_all_backends() { + let _ = run_command( + "nft", + &["delete", "table", "inet", "telemt_conntrack"], + None, + ) + .await; + let _ = run_command( + "iptables", + &["-t", "raw", "-D", "PREROUTING", "-j", "TELEMT_NOTRACK"], + None, + ) + .await; + let _ = run_command("iptables", &["-t", "raw", "-F", "TELEMT_NOTRACK"], None).await; + let _ = run_command("iptables", &["-t", "raw", "-X", "TELEMT_NOTRACK"], None).await; + let _ = run_command( + "ip6tables", + &["-t", "raw", "-D", "PREROUTING", "-j", "TELEMT_NOTRACK"], + None, + ) + .await; + let _ = run_command("ip6tables", &["-t", "raw", "-F", "TELEMT_NOTRACK"], None).await; + let _ = run_command("ip6tables", &["-t", "raw", "-X", "TELEMT_NOTRACK"], None).await; +} + +pub(super) enum DeleteOutcome { + Deleted, + NotFound, + Error, +} + +pub(super) async fn delete_conntrack_entry(event: ConntrackCloseEvent) -> DeleteOutcome { + if !command_exists("conntrack") { + return DeleteOutcome::Error; + } + let args = vec![ + "-D".to_string(), + "-p".to_string(), + "tcp".to_string(), + "-s".to_string(), + event.src.ip().to_string(), + "--sport".to_string(), + event.src.port().to_string(), + "-d".to_string(), + event.dst.ip().to_string(), + "--dport".to_string(), + event.dst.port().to_string(), + ]; + let arg_refs: Vec<&str> = args.iter().map(String::as_str).collect(); + match run_command("conntrack", &arg_refs, None).await { + Ok(()) => DeleteOutcome::Deleted, + Err(error) => { + if error.contains("0 flow entries have been deleted") { + DeleteOutcome::NotFound + } else { + debug!(error = %error, "conntrack delete failed"); + DeleteOutcome::Error + } + } + } +} + +async fn run_command(binary: &str, args: &[&str], stdin: Option) -> Result<(), String> { + #[cfg(unix)] + let Some(command_path) = resolve_trusted_helper(binary) else { + return Err(format!("{binary} is not available")); + }; + #[cfg(not(unix))] + return Err(format!("{binary} is not available")); + #[cfg(unix)] + let mut command = Command::new(command_path); + command.args(args); + if stdin.is_some() { + command.stdin(std::process::Stdio::piped()); + } + command.stdout(std::process::Stdio::null()); + command.stderr(std::process::Stdio::piped()); + let mut child = command + .spawn() + .map_err(|error| format!("spawn {binary} failed: {error}"))?; + if let Some(blob) = stdin + && let Some(mut writer) = child.stdin.take() + { + writer + .write_all(blob.as_bytes()) + .await + .map_err(|error| format!("stdin write {binary} failed: {error}"))?; + } + let output = child + .wait_with_output() + .await + .map_err(|error| format!("wait {binary} failed: {error}"))?; + if output.status.success() { + return Ok(()); + } + let stderr = String::from_utf8_lossy(&output.stderr).trim().to_string(); + Err(if stderr.is_empty() { + format!("{binary} exited with status {}", output.status) + } else { + stderr + }) +} + +fn has_cap_net_admin() -> bool { + #[cfg(target_os = "linux")] + { + let Ok(status) = std::fs::read_to_string("/proc/self/status") else { + return false; + }; + for line in status.lines() { + if let Some(raw) = line.strip_prefix("CapEff:") { + let caps = raw.trim(); + if let Ok(bits) = u64::from_str_radix(caps, 16) { + const CAP_NET_ADMIN_BIT: u64 = 12; + return (bits & (1u64 << CAP_NET_ADMIN_BIT)) != 0; + } + } + } + false + } + #[cfg(not(target_os = "linux"))] + { + false + } +} diff --git a/src/daemon/pid_file.rs b/src/daemon/pid_file.rs index ae1c3e1..3536e4d 100644 --- a/src/daemon/pid_file.rs +++ b/src/daemon/pid_file.rs @@ -1,6 +1,8 @@ use std::fs::{self, File, OpenOptions}; use std::io::{ErrorKind, Read, Write}; use std::os::unix::fs::{MetadataExt, OpenOptionsExt}; +#[cfg(target_os = "linux")] +use std::os::fd::{FromRawFd, OwnedFd}; use std::path::{Path, PathBuf}; use nix::fcntl::{Flock, FlockArg}; @@ -281,7 +283,19 @@ pub fn signal_pid_file>( path: P, signal: nix::sys::signal::Signal, ) -> Result<(), DaemonError> { - let pid = read_pid_file(&path)?; + let path = path.as_ref(); + let pid = read_pid_file(path)?; + #[cfg(target_os = "linux")] + let pidfd = open_pidfd(pid)?; + if !daemon_lock_is_held(path)? { + return Err(DaemonError::PidFile(format!( + "refusing to signal unlocked or stale PID file {}", + path.display() + ))); + } + #[cfg(target_os = "linux")] + return signal_pidfd(&pidfd, pid, signal); + #[cfg(not(target_os = "linux"))] nix::sys::signal::kill(Pid::from_raw(pid), signal) .map_err(|error| DaemonError::PidFile(format!("cannot signal process {}: {}", pid, error))) } @@ -303,206 +317,93 @@ pub enum DaemonStatus { pub fn check_status>(path: P) -> DaemonStatus { let path = path.as_ref(); match read_pid_file_if_exists(path) { - Ok(Some(pid)) if is_process_running(pid) => DaemonStatus::Running(pid), + Ok(Some(pid)) + if daemon_lock_is_held(path).unwrap_or(false) && is_process_running(pid) => + { + DaemonStatus::Running(pid) + } Ok(Some(pid)) => DaemonStatus::Stale(pid), Ok(None) | Err(_) => DaemonStatus::NotRunning, } } +fn daemon_lock_is_held(path: &Path) -> Result { + let lock_path = sibling_lock_path(path); + let file = match OpenOptions::new() + .read(true) + .write(true) + .custom_flags(libc::O_CLOEXEC | libc::O_NOFOLLOW) + .open(&lock_path) + { + Ok(file) => file, + Err(error) if error.kind() == ErrorKind::NotFound => return Ok(false), + Err(error) => { + return Err(DaemonError::PidFile(format!( + "cannot inspect lock {}: {}", + lock_path.display(), + error + ))); + } + }; + validate_regular_single_link(&file, &lock_path)?; + match Flock::lock(file, FlockArg::LockExclusiveNonblock) { + Ok(_available) => Ok(false), + Err((_file, nix::errno::Errno::EWOULDBLOCK)) => Ok(true), + Err((_file, error)) => Err(DaemonError::PidFile(format!( + "cannot inspect lock ownership for {}: {}", + lock_path.display(), + error + ))), + } +} + +#[cfg(target_os = "linux")] +fn open_pidfd(pid: i32) -> Result { + // SAFETY: `pidfd_open` receives a validated positive PID and no pointer arguments. + let descriptor = unsafe { libc::syscall(libc::SYS_pidfd_open, pid, 0) }; + if descriptor < 0 { + return Err(DaemonError::PidFile(format!( + "cannot open stable process handle for {}: {}", + pid, + std::io::Error::last_os_error() + ))); + } + // SAFETY: a successful `pidfd_open` returns one newly owned descriptor. + Ok(unsafe { OwnedFd::from_raw_fd(descriptor as i32) }) +} + +#[cfg(target_os = "linux")] +fn signal_pidfd( + pidfd: &OwnedFd, + pid: i32, + signal: nix::sys::signal::Signal, +) -> Result<(), DaemonError> { + use std::os::fd::AsRawFd; + + // SAFETY: the pidfd is owned and valid, and both optional pointer arguments are null. + let result = unsafe { + libc::syscall( + libc::SYS_pidfd_send_signal, + pidfd.as_raw_fd(), + signal as libc::c_int, + std::ptr::null::(), + 0, + ) + }; + if result == 0 { + Ok(()) + } else { + Err(DaemonError::PidFile(format!( + "cannot signal process {} through stable handle: {}", + pid, + std::io::Error::last_os_error() + ))) + } +} + fn is_process_running(pid: i32) -> bool { nix::sys::signal::kill(Pid::from_raw(pid), None).is_ok() } #[cfg(test)] -mod tests { - use std::os::unix::fs::MetadataExt; - use std::os::unix::fs::symlink; - use std::process::{Child, Command, Stdio}; - use std::thread; - use std::time::{Duration, Instant}; - - use super::*; - - const HELPER_PID_PATH: &str = "TELEMT_PID_LOCK_HELPER_PATH"; - const HELPER_READY_PATH: &str = "TELEMT_PID_LOCK_HELPER_READY"; - const HELPER_STOP_PATH: &str = "TELEMT_PID_LOCK_HELPER_STOP"; - - fn wait_for_path(path: &Path, timeout: Duration) -> bool { - let deadline = Instant::now() + timeout; - while Instant::now() < deadline { - if path.exists() { - return true; - } - thread::sleep(Duration::from_millis(10)); - } - false - } - - fn wait_for_child(child: &mut Child, timeout: Duration) -> Option { - let deadline = Instant::now() + timeout; - while Instant::now() < deadline { - if let Some(status) = child.try_wait().unwrap() { - return Some(status); - } - thread::sleep(Duration::from_millis(10)); - } - None - } - - #[test] - fn pid_file_remains_send_and_sync() { - fn assert_send_sync() {} - - assert_send_sync::(); - } - - #[test] - fn lock_holder_subprocess() { - let Some(pid_path) = std::env::var_os(HELPER_PID_PATH) else { - return; - }; - let ready_path = PathBuf::from(std::env::var_os(HELPER_READY_PATH).unwrap()); - let stop_path = PathBuf::from(std::env::var_os(HELPER_STOP_PATH).unwrap()); - let mut pid_file = PidFile::new(PathBuf::from(pid_path)); - pid_file.acquire().unwrap(); - fs::write(&ready_path, b"ready").unwrap(); - assert!(wait_for_path(&stop_path, Duration::from_secs(10))); - pid_file.release().unwrap(); - } - - #[test] - fn persistent_sibling_lock_serializes_processes_after_pid_unlink() { - let directory = tempfile::tempdir().unwrap(); - let pid_path = directory.path().join("telemt.pid"); - let lock_path = sibling_lock_path(&pid_path); - let ready_path = directory.path().join("ready"); - let stop_path = directory.path().join("stop"); - let mut child = Command::new(std::env::current_exe().unwrap()) - .args([ - "--exact", - "daemon::pid_file::tests::lock_holder_subprocess", - "--nocapture", - ]) - .env(HELPER_PID_PATH, &pid_path) - .env(HELPER_READY_PATH, &ready_path) - .env(HELPER_STOP_PATH, &stop_path) - .stdout(Stdio::null()) - .stderr(Stdio::null()) - .spawn() - .unwrap(); - - if !wait_for_path(&ready_path, Duration::from_secs(5)) { - let _ = child.kill(); - let _ = child.wait(); - panic!("PID lock holder did not become ready"); - } - let lock_inode = fs::metadata(&lock_path).unwrap().ino(); - fs::remove_file(&pid_path).unwrap(); - - let mut contender = PidFile::new(&pid_path); - assert!(contender.acquire().is_err()); - - fs::write(&stop_path, b"stop").unwrap(); - let status = wait_for_child(&mut child, Duration::from_secs(5)).unwrap_or_else(|| { - let _ = child.kill(); - child.wait().unwrap() - }); - assert!(status.success()); - assert_eq!(fs::metadata(&lock_path).unwrap().ino(), lock_inode); - - contender.acquire().unwrap(); - assert_eq!(fs::metadata(&lock_path).unwrap().ino(), lock_inode); - contender.release().unwrap(); - assert!(!pid_path.exists()); - assert!(lock_path.exists()); - } - - #[test] - fn stale_pid_checks_are_read_only() { - let directory = tempfile::tempdir().unwrap(); - let pid_path = directory.path().join("telemt.pid"); - fs::write(&pid_path, b"2000000000\n").unwrap(); - let pid_file = PidFile::new(&pid_path); - - assert_eq!(pid_file.check_running().unwrap(), None); - assert_eq!(check_status(&pid_path), DaemonStatus::Stale(2_000_000_000)); - assert!(pid_path.exists()); - } - - #[test] - fn unowned_release_does_not_remove_pid_file() { - let directory = tempfile::tempdir().unwrap(); - let pid_path = directory.path().join("telemt.pid"); - fs::write(&pid_path, b"2000000000\n").unwrap(); - let mut pid_file = PidFile::new(&pid_path); - - pid_file.release().unwrap(); - - assert!(pid_path.exists()); - } - - #[test] - fn acquire_rejects_pid_symlink_without_truncating_target() { - let directory = tempfile::tempdir().unwrap(); - let pid_path = directory.path().join("telemt.pid"); - let target_path = directory.path().join("target"); - fs::write(&target_path, b"preserve\n").unwrap(); - symlink(&target_path, &pid_path).unwrap(); - let mut pid_file = PidFile::new(&pid_path); - - assert!(pid_file.acquire().is_err()); - assert_eq!(fs::read(&target_path).unwrap(), b"preserve\n"); - } - - #[test] - fn release_does_not_remove_replacement_path() { - let directory = tempfile::tempdir().unwrap(); - let pid_path = directory.path().join("telemt.pid"); - let owned_path = directory.path().join("owned.pid"); - let mut pid_file = PidFile::new(&pid_path); - pid_file.acquire().unwrap(); - fs::rename(&pid_path, &owned_path).unwrap(); - fs::write(&pid_path, b"replacement\n").unwrap(); - - let error = pid_file.release().unwrap_err(); - - assert!(error.to_string().contains("refusing to remove replaced PID file")); - assert_eq!(fs::read(&pid_path).unwrap(), b"replacement\n"); - } - - #[test] - fn pid_parser_rejects_process_group_values() { - let directory = tempfile::tempdir().unwrap(); - let pid_path = directory.path().join("telemt.pid"); - - for value in ["-1\n", "0\n", "1\n"] { - fs::write(&pid_path, value).unwrap(); - assert!(read_pid_file(&pid_path).is_err()); - } - } - - #[test] - fn pid_file_release_keeps_lock_inode() { - let directory = tempfile::tempdir().unwrap(); - let pid_path = directory.path().join("telemt.pid"); - let lock_path = sibling_lock_path(&pid_path); - let mut pid_file = PidFile::new(&pid_path); - - pid_file.acquire().unwrap(); - assert!( - pid_file - .ownership_file_handles() - .into_iter() - .all(|file| file.is_some()) - ); - assert_eq!(read_pid_file(&pid_path).unwrap(), std::process::id() as i32); - let lock_inode = fs::metadata(&lock_path).unwrap().ino(); - pid_file.release().unwrap(); - - assert!(!pid_path.exists()); - assert!(lock_path.exists()); - pid_file.acquire().unwrap(); - assert_eq!(fs::metadata(&lock_path).unwrap().ino(), lock_inode); - pid_file.release().unwrap(); - } -} +mod tests; diff --git a/src/daemon/pid_file/tests.rs b/src/daemon/pid_file/tests.rs new file mode 100644 index 0000000..9aad337 --- /dev/null +++ b/src/daemon/pid_file/tests.rs @@ -0,0 +1,209 @@ +use std::os::unix::fs::{MetadataExt, symlink}; +use std::process::{Child, Command, Stdio}; +use std::thread; +use std::time::{Duration, Instant}; + +use super::*; + +const HELPER_PID_PATH: &str = "TELEMT_PID_LOCK_HELPER_PATH"; +const HELPER_READY_PATH: &str = "TELEMT_PID_LOCK_HELPER_READY"; +const HELPER_STOP_PATH: &str = "TELEMT_PID_LOCK_HELPER_STOP"; + +fn wait_for_path(path: &Path, timeout: Duration) -> bool { + let deadline = Instant::now() + timeout; + while Instant::now() < deadline { + if path.exists() { + return true; + } + thread::sleep(Duration::from_millis(10)); + } + false +} + +fn wait_for_child(child: &mut Child, timeout: Duration) -> Option { + let deadline = Instant::now() + timeout; + while Instant::now() < deadline { + if let Some(status) = child.try_wait().unwrap() { + return Some(status); + } + thread::sleep(Duration::from_millis(10)); + } + None +} + +#[test] +fn pid_file_remains_send_and_sync() { + fn assert_send_sync() {} + + assert_send_sync::(); +} + +#[test] +fn lock_holder_subprocess() { + let Some(pid_path) = std::env::var_os(HELPER_PID_PATH) else { + return; + }; + let ready_path = PathBuf::from(std::env::var_os(HELPER_READY_PATH).unwrap()); + let stop_path = PathBuf::from(std::env::var_os(HELPER_STOP_PATH).unwrap()); + let mut pid_file = PidFile::new(PathBuf::from(pid_path)); + pid_file.acquire().unwrap(); + fs::write(&ready_path, b"ready").unwrap(); + assert!(wait_for_path(&stop_path, Duration::from_secs(10))); + pid_file.release().unwrap(); +} + +#[test] +fn persistent_sibling_lock_serializes_processes_after_pid_unlink() { + let directory = tempfile::tempdir().unwrap(); + let pid_path = directory.path().join("telemt.pid"); + let lock_path = sibling_lock_path(&pid_path); + let ready_path = directory.path().join("ready"); + let stop_path = directory.path().join("stop"); + let mut child = Command::new(std::env::current_exe().unwrap()) + .args([ + "--exact", + "daemon::pid_file::tests::lock_holder_subprocess", + "--nocapture", + ]) + .env(HELPER_PID_PATH, &pid_path) + .env(HELPER_READY_PATH, &ready_path) + .env(HELPER_STOP_PATH, &stop_path) + .stdout(Stdio::null()) + .stderr(Stdio::null()) + .spawn() + .unwrap(); + + if !wait_for_path(&ready_path, Duration::from_secs(5)) { + let _ = child.kill(); + let _ = child.wait(); + panic!("PID lock holder did not become ready"); + } + let lock_inode = fs::metadata(&lock_path).unwrap().ino(); + fs::remove_file(&pid_path).unwrap(); + + let mut contender = PidFile::new(&pid_path); + assert!(contender.acquire().is_err()); + + fs::write(&stop_path, b"stop").unwrap(); + let status = wait_for_child(&mut child, Duration::from_secs(5)).unwrap_or_else(|| { + let _ = child.kill(); + child.wait().unwrap() + }); + assert!(status.success()); + assert_eq!(fs::metadata(&lock_path).unwrap().ino(), lock_inode); + + contender.acquire().unwrap(); + assert_eq!(fs::metadata(&lock_path).unwrap().ino(), lock_inode); + contender.release().unwrap(); + assert!(!pid_path.exists()); + assert!(lock_path.exists()); +} + +#[test] +fn stale_pid_checks_are_read_only() { + let directory = tempfile::tempdir().unwrap(); + let pid_path = directory.path().join("telemt.pid"); + fs::write(&pid_path, b"2000000000\n").unwrap(); + let pid_file = PidFile::new(&pid_path); + + assert_eq!(pid_file.check_running().unwrap(), None); + assert_eq!(check_status(&pid_path), DaemonStatus::Stale(2_000_000_000)); + assert!(pid_path.exists()); +} + +#[test] +fn status_requires_live_lock_ownership() { + let directory = tempfile::tempdir().unwrap(); + let pid_path = directory.path().join("telemt.pid"); + fs::write(&pid_path, format!("{}\n", std::process::id())).unwrap(); + assert_eq!( + check_status(&pid_path), + DaemonStatus::Stale(std::process::id() as i32) + ); + + fs::remove_file(&pid_path).unwrap(); + let mut owner = PidFile::new(&pid_path); + owner.acquire().unwrap(); + assert_eq!( + check_status(&pid_path), + DaemonStatus::Running(std::process::id() as i32) + ); + owner.release().unwrap(); +} + +#[test] +fn unowned_release_does_not_remove_pid_file() { + let directory = tempfile::tempdir().unwrap(); + let pid_path = directory.path().join("telemt.pid"); + fs::write(&pid_path, b"2000000000\n").unwrap(); + let mut pid_file = PidFile::new(&pid_path); + + pid_file.release().unwrap(); + + assert!(pid_path.exists()); +} + +#[test] +fn acquire_rejects_pid_symlink_without_truncating_target() { + let directory = tempfile::tempdir().unwrap(); + let pid_path = directory.path().join("telemt.pid"); + let target_path = directory.path().join("target"); + fs::write(&target_path, b"preserve\n").unwrap(); + symlink(&target_path, &pid_path).unwrap(); + let mut pid_file = PidFile::new(&pid_path); + + assert!(pid_file.acquire().is_err()); + assert_eq!(fs::read(&target_path).unwrap(), b"preserve\n"); +} + +#[test] +fn release_does_not_remove_replacement_path() { + let directory = tempfile::tempdir().unwrap(); + let pid_path = directory.path().join("telemt.pid"); + let owned_path = directory.path().join("owned.pid"); + let mut pid_file = PidFile::new(&pid_path); + pid_file.acquire().unwrap(); + fs::rename(&pid_path, &owned_path).unwrap(); + fs::write(&pid_path, b"replacement\n").unwrap(); + + let error = pid_file.release().unwrap_err(); + + assert!(error.to_string().contains("refusing to remove replaced PID file")); + assert_eq!(fs::read(&pid_path).unwrap(), b"replacement\n"); +} + +#[test] +fn pid_parser_rejects_process_group_values() { + let directory = tempfile::tempdir().unwrap(); + let pid_path = directory.path().join("telemt.pid"); + + for value in ["-1\n", "0\n", "1\n"] { + fs::write(&pid_path, value).unwrap(); + assert!(read_pid_file(&pid_path).is_err()); + } +} + +#[test] +fn pid_file_release_keeps_lock_inode() { + let directory = tempfile::tempdir().unwrap(); + let pid_path = directory.path().join("telemt.pid"); + let lock_path = sibling_lock_path(&pid_path); + let mut pid_file = PidFile::new(&pid_path); + + pid_file.acquire().unwrap(); + assert!( + pid_file + .ownership_file_handles() + .into_iter() + .all(|file| file.is_some()) + ); + assert_eq!(read_pid_file(&pid_path).unwrap(), std::process::id() as i32); + let lock_inode = fs::metadata(&lock_path).unwrap().ino(); + pid_file.release().unwrap(); + + assert!(!pid_path.exists()); + assert!(lock_path.exists()); + pid_file.acquire().unwrap(); + assert_eq!(fs::metadata(&lock_path).unwrap().ino(), lock_inode); + pid_file.release().unwrap(); +} diff --git a/src/logging.rs b/src/logging.rs index 08f1131..8ed5ea7 100644 --- a/src/logging.rs +++ b/src/logging.rs @@ -8,8 +8,6 @@ // Infrastructure module used via CLI flags. #![allow(dead_code)] -use std::path::Path; - use crate::config::{LogRotation, LoggingConfig, LoggingDestination}; use tracing_subscriber::layer::SubscriberExt; @@ -144,31 +142,9 @@ pub fn init_logging( } LogDestination::File { options } => { - let (non_blocking, guard) = if options.max_size_bytes > 0 - || options.max_files > 0 - || options.max_age_secs > 0 - { - let file_appender = file::BoundedFileAppender::new(options.clone()) - .expect("Failed to open log file"); - tracing_appender::non_blocking(file_appender) - } else if !matches!(options.rotation, LogRotation::Never) { - let path = Path::new(&options.path); - let dir = log_file_dir(path); - let prefix = log_file_name(path); - let file_appender = tracing_appender::rolling::RollingFileAppender::builder() - .rotation(to_tracing_rotation(options.rotation)) - .filename_prefix(prefix) - .build(dir) - .expect("Failed to open log file"); - tracing_appender::non_blocking(file_appender) - } else { - let file = std::fs::OpenOptions::new() - .create(true) - .append(true) - .open(&options.path) - .expect("Failed to open log file"); - tracing_appender::non_blocking(file) - }; + let file_appender = file::BoundedFileAppender::new(options.clone()) + .expect("Failed to open log file"); + let (non_blocking, guard) = tracing_appender::non_blocking(file_appender); let fmt_layer = fmt::Layer::default() .with_ansi(false) @@ -185,28 +161,6 @@ pub fn init_logging( } } -fn log_file_dir(path: &Path) -> &Path { - path.parent() - .filter(|parent| !parent.as_os_str().is_empty()) - .unwrap_or_else(|| Path::new(".")) -} - -fn log_file_name(path: &Path) -> &str { - path.file_name() - .and_then(|s| s.to_str()) - .unwrap_or("telemt") -} - -fn to_tracing_rotation(rotation: LogRotation) -> tracing_appender::rolling::Rotation { - match rotation { - LogRotation::Never => tracing_appender::rolling::Rotation::NEVER, - LogRotation::Minutely => tracing_appender::rolling::Rotation::MINUTELY, - LogRotation::Hourly => tracing_appender::rolling::Rotation::HOURLY, - LogRotation::Daily => tracing_appender::rolling::Rotation::DAILY, - LogRotation::Weekly => tracing_appender::rolling::Rotation::WEEKLY, - } -} - /// Syslog writer for tracing. #[cfg(unix)] #[derive(Clone, Copy)] diff --git a/src/logging/file.rs b/src/logging/file.rs index 3b96903..f59c8a1 100644 --- a/src/logging/file.rs +++ b/src/logging/file.rs @@ -1,8 +1,26 @@ -use std::fs::{self, File, OpenOptions}; +use std::fs::{self, File}; +#[cfg(not(unix))] +use std::fs::OpenOptions; use std::io::{self, Write}; use std::path::{Path, PathBuf}; use std::time::{Duration, SystemTime, UNIX_EPOCH}; +#[cfg(unix)] +use std::ffi::OsString; +#[cfg(unix)] +use std::os::fd::OwnedFd; +#[cfg(unix)] +use std::os::unix::ffi::OsStringExt; + +#[cfg(unix)] +use nix::dir::Dir; +#[cfg(unix)] +use nix::fcntl::{OFlag, openat, renameat}; +#[cfg(unix)] +use nix::sys::stat::Mode; +#[cfg(unix)] +use nix::unistd::{UnlinkatFlags, dup, unlinkat}; + use chrono::{DateTime, Datelike, Duration as ChronoDuration, Utc}; use crate::config::LogRotation; @@ -20,6 +38,8 @@ pub(crate) struct BoundedFileAppender { current_size: u64, last_cleanup: DateTime, file: Option, + #[cfg(unix)] + dir_fd: OwnedFd, now: Box DateTime + Send + Sync>, } @@ -46,6 +66,11 @@ impl BoundedFileAppender { let start = now(); let current_path = active_path_for(&dir, &base_name, options.rotation, &start); + #[cfg(unix)] + let dir_fd = crate::util::secure_fs::open_dir_nofollow_or_create(&dir, 0o750)?; + #[cfg(unix)] + let (file, current_size) = open_append_file(&dir_fd, ¤t_path)?; + #[cfg(not(unix))] let (file, current_size) = open_append_file(¤t_path)?; let mut appender = Self { options, @@ -55,6 +80,8 @@ impl BoundedFileAppender { current_size, last_cleanup: start, file: Some(file), + #[cfg(unix)] + dir_fd, now, }; appender.cleanup(&start); @@ -79,10 +106,8 @@ impl BoundedFileAppender { fn rotate_for_size(&mut self, now: &DateTime) -> io::Result<()> { self.close_current()?; - if self.current_path.exists() { - let archive_path = self.archive_path(now); - fs::rename(&self.current_path, archive_path)?; - } + let archive_path = self.archive_path(now); + self.rename_current_if_present(&archive_path)?; self.open_current() } @@ -95,7 +120,7 @@ impl BoundedFileAppender { let stamp = now.format("%Y%m%d%H%M%S"); for seq in 0..1000 { let candidate = self.dir.join(format!("{file_name}.{stamp}.{seq}")); - if !candidate.exists() { + if !self.path_exists(&candidate) { return candidate; } } @@ -103,6 +128,9 @@ impl BoundedFileAppender { } fn open_current(&mut self) -> io::Result<()> { + #[cfg(unix)] + let (file, current_size) = open_append_file(&self.dir_fd, &self.current_path)?; + #[cfg(not(unix))] let (file, current_size) = open_append_file(&self.current_path)?; self.file = Some(file); self.current_size = current_size; @@ -130,40 +158,10 @@ impl BoundedFileAppender { fn cleanup(&mut self, now: &DateTime) { self.last_cleanup = now.clone(); - let Ok(entries) = fs::read_dir(&self.dir) else { + let Ok(mut candidates) = self.collect_candidates() else { return; }; - let mut candidates = Vec::new(); - let prefix = format!("{}.", self.base_name); - for entry in entries.flatten() { - let path = entry.path(); - let Ok(file_type) = entry.file_type() else { - continue; - }; - if !file_type.is_file() { - continue; - } - - let is_current = path == self.current_path; - let Some(name) = entry.file_name().to_str().map(|name| name.to_string()) else { - continue; - }; - if !is_current && !name.starts_with(&prefix) { - continue; - } - - let Ok(metadata) = entry.metadata() else { - continue; - }; - let modified = metadata.modified().unwrap_or(UNIX_EPOCH); - candidates.push(LogFileCandidate { - path, - modified, - is_current, - }); - } - if self.options.max_age_secs > 0 { let cutoff = system_time_from_utc(now) .checked_sub(Duration::from_secs(self.options.max_age_secs)) @@ -172,7 +170,7 @@ impl BoundedFileAppender { if candidate.is_current || candidate.modified >= cutoff { true } else { - let _ = fs::remove_file(&candidate.path); + self.remove_candidate(candidate); false } }); @@ -189,11 +187,153 @@ impl BoundedFileAppender { if total <= self.options.max_files { break; } - let _ = fs::remove_file(candidate.path); + self.remove_candidate(&candidate); total -= 1; } } } + + #[cfg(unix)] + fn path_exists(&self, path: &Path) -> bool { + let Some(name) = path.file_name() else { + return true; + }; + match openat( + &self.dir_fd, + name, + OFlag::O_RDONLY | OFlag::O_NONBLOCK | OFlag::O_NOFOLLOW | OFlag::O_CLOEXEC, + Mode::empty(), + ) { + Ok(_) => true, + Err(nix::errno::Errno::ENOENT) => false, + Err(_) => true, + } + } + + #[cfg(not(unix))] + fn path_exists(&self, path: &Path) -> bool { + path.exists() + } + + #[cfg(unix)] + fn rename_current_if_present(&self, archive_path: &Path) -> io::Result<()> { + let current_name = self.current_path.file_name().ok_or_else(|| { + io::Error::new(io::ErrorKind::InvalidInput, "log path has no file name") + })?; + let archive_name = archive_path.file_name().ok_or_else(|| { + io::Error::new(io::ErrorKind::InvalidInput, "archive path has no file name") + })?; + match renameat( + &self.dir_fd, + current_name, + &self.dir_fd, + archive_name, + ) { + Ok(()) => Ok(()), + Err(nix::errno::Errno::ENOENT) => Ok(()), + Err(error) => Err(io::Error::from_raw_os_error(error as i32)), + } + } + + #[cfg(not(unix))] + fn rename_current_if_present(&self, archive_path: &Path) -> io::Result<()> { + if self.current_path.exists() { + fs::rename(&self.current_path, archive_path)?; + } + Ok(()) + } + + #[cfg(unix)] + fn collect_candidates(&self) -> io::Result> { + use std::os::unix::fs::MetadataExt; + + let descriptor = dup(&self.dir_fd).map_err(|error| { + io::Error::from_raw_os_error(error as i32) + })?; + let mut directory = Dir::from_fd(descriptor) + .map_err(|error| io::Error::from_raw_os_error(error as i32))?; + let mut candidates = Vec::new(); + let prefix = format!("{}.", self.base_name); + for entry in directory.iter().flatten() { + let bytes = entry.file_name().to_bytes(); + if bytes == b"." || bytes == b".." { + continue; + } + let name = OsString::from_vec(bytes.to_vec()); + let path = self.dir.join(&name); + let is_current = path == self.current_path; + let Some(name_text) = name.to_str() else { + continue; + }; + if !is_current && !name_text.starts_with(&prefix) { + continue; + } + let Ok(descriptor) = openat( + &self.dir_fd, + name.as_os_str(), + OFlag::O_RDONLY | OFlag::O_NONBLOCK | OFlag::O_NOFOLLOW | OFlag::O_CLOEXEC, + Mode::empty(), + ) else { + continue; + }; + let file = File::from(descriptor); + let Ok(metadata) = file.metadata() else { + continue; + }; + if !metadata.is_file() || metadata.nlink() != 1 { + continue; + } + candidates.push(LogFileCandidate { + path, + modified: metadata.modified().unwrap_or(UNIX_EPOCH), + is_current, + }); + } + Ok(candidates) + } + + #[cfg(not(unix))] + fn collect_candidates(&self) -> io::Result> { + let mut candidates = Vec::new(); + let prefix = format!("{}.", self.base_name); + for entry in fs::read_dir(&self.dir)?.flatten() { + let path = entry.path(); + let Ok(file_type) = entry.file_type() else { + continue; + }; + if !file_type.is_file() { + continue; + } + let is_current = path == self.current_path; + let Some(name) = entry.file_name().to_str().map(str::to_string) else { + continue; + }; + if !is_current && !name.starts_with(&prefix) { + continue; + } + let Ok(metadata) = entry.metadata() else { + continue; + }; + candidates.push(LogFileCandidate { + path, + modified: metadata.modified().unwrap_or(UNIX_EPOCH), + is_current, + }); + } + Ok(candidates) + } + + #[cfg(unix)] + fn remove_candidate(&self, candidate: &LogFileCandidate) { + if let Some(name) = candidate.path.file_name() { + let _ = unlinkat(&self.dir_fd, name, UnlinkatFlags::NoRemoveDir); + } + } + + #[cfg(not(unix))] + fn remove_candidate(&self, candidate: &LogFileCandidate) { + let _ = fs::remove_file(&candidate.path); + } } impl Write for BoundedFileAppender { @@ -233,6 +373,17 @@ struct LogFileCandidate { is_current: bool, } +#[cfg(unix)] +fn open_append_file(dir_fd: &OwnedFd, path: &Path) -> io::Result<(File, u64)> { + let name = path.file_name().ok_or_else(|| { + io::Error::new(io::ErrorKind::InvalidInput, "log path has no file name") + })?; + let file = crate::util::secure_fs::open_append_regular_at(dir_fd, name, 0o640)?; + let current_size = file.metadata()?.len(); + Ok((file, current_size)) +} + +#[cfg(not(unix))] fn open_append_file(path: &Path) -> io::Result<(File, u64)> { let mut options = OpenOptions::new(); options.create(true).append(true); @@ -291,105 +442,4 @@ fn system_time_from_utc(now: &DateTime) -> SystemTime { } #[cfg(test)] -mod tests { - use std::io::Write; - - use tempfile::tempdir; - - use super::*; - - fn fixed_now() -> DateTime { - DateTime::::from(UNIX_EPOCH + Duration::from_secs(10)) - } - - fn options(path: PathBuf) -> FileLogOptions { - FileLogOptions { - path: path.to_string_lossy().to_string(), - rotation: LogRotation::Never, - max_size_bytes: 0, - max_files: 0, - max_age_secs: 0, - } - } - - fn matching_logs(dir: &Path) -> Vec { - let mut files: Vec<_> = fs::read_dir(dir) - .unwrap() - .flatten() - .map(|entry| entry.path()) - .filter(|path| { - path.file_name() - .and_then(|name| name.to_str()) - .map(|name| name.starts_with("telemt.log")) - .unwrap_or(false) - }) - .collect(); - files.sort(); - files - } - - #[test] - fn size_rotation_keeps_latest_write_in_active_file() { - let dir = tempdir().unwrap(); - let path = dir.path().join("telemt.log"); - let mut options = options(path.clone()); - options.max_size_bytes = 6; - - let mut appender = BoundedFileAppender::with_now(options, Box::new(fixed_now)).unwrap(); - appender.write_all(b"abc\n").unwrap(); - appender.write_all(b"def\n").unwrap(); - appender.flush().unwrap(); - - assert_eq!(fs::read_to_string(path).unwrap(), "def\n"); - assert_eq!(matching_logs(dir.path()).len(), 2); - } - - #[test] - fn max_files_retention_removes_oldest_archives() { - let dir = tempdir().unwrap(); - let path = dir.path().join("telemt.log"); - let mut options = options(path); - options.max_size_bytes = 4; - options.max_files = 2; - - let mut appender = BoundedFileAppender::with_now(options, Box::new(fixed_now)).unwrap(); - for line in [b"aa\n", b"bb\n", b"cc\n", b"dd\n"] { - appender.write_all(line).unwrap(); - } - appender.flush().unwrap(); - - assert!(matching_logs(dir.path()).len() <= 2); - } - - #[cfg(unix)] - #[test] - fn max_age_retention_removes_old_archives() { - use std::ffi::CString; - use std::os::unix::ffi::OsStrExt; - - let dir = tempdir().unwrap(); - let path = dir.path().join("telemt.log"); - let old_archive = dir.path().join("telemt.log.20000101000000.0"); - fs::write(&old_archive, "old").unwrap(); - - let c_path = CString::new(old_archive.as_os_str().as_bytes()).unwrap(); - let times = [ - libc::timespec { - tv_sec: 0, - tv_nsec: 0, - }, - libc::timespec { - tv_sec: 0, - tv_nsec: 0, - }, - ]; - let rc = unsafe { libc::utimensat(libc::AT_FDCWD, c_path.as_ptr(), times.as_ptr(), 0) }; - assert_eq!(rc, 0); - - let mut options = options(path); - options.max_age_secs = 1; - let _appender = BoundedFileAppender::with_now(options, Box::new(fixed_now)).unwrap(); - - assert!(!old_archive.exists()); - } -} +mod tests; diff --git a/src/logging/file/tests.rs b/src/logging/file/tests.rs new file mode 100644 index 0000000..edaf984 --- /dev/null +++ b/src/logging/file/tests.rs @@ -0,0 +1,125 @@ +use std::io::Write; + +use tempfile::tempdir; + +use super::*; + +fn fixed_now() -> DateTime { + DateTime::::from(UNIX_EPOCH + Duration::from_secs(10)) +} + +fn options(path: PathBuf) -> FileLogOptions { + FileLogOptions { + path: path.to_string_lossy().to_string(), + rotation: LogRotation::Never, + max_size_bytes: 0, + max_files: 0, + max_age_secs: 0, + } +} + +fn matching_logs(dir: &Path) -> Vec { + let mut files: Vec<_> = fs::read_dir(dir) + .unwrap() + .flatten() + .map(|entry| entry.path()) + .filter(|path| { + path.file_name() + .and_then(|name| name.to_str()) + .map(|name| name.starts_with("telemt.log")) + .unwrap_or(false) + }) + .collect(); + files.sort(); + files +} + +#[test] +fn size_rotation_keeps_latest_write_in_active_file() { + let dir = tempdir().unwrap(); + let path = dir.path().join("telemt.log"); + let mut options = options(path.clone()); + options.max_size_bytes = 6; + + let mut appender = BoundedFileAppender::with_now(options, Box::new(fixed_now)).unwrap(); + appender.write_all(b"abc\n").unwrap(); + appender.write_all(b"def\n").unwrap(); + appender.flush().unwrap(); + + assert_eq!(fs::read_to_string(path).unwrap(), "def\n"); + assert_eq!(matching_logs(dir.path()).len(), 2); +} + +#[test] +fn max_files_retention_removes_oldest_archives() { + let dir = tempdir().unwrap(); + let path = dir.path().join("telemt.log"); + let mut options = options(path); + options.max_size_bytes = 4; + options.max_files = 2; + + let mut appender = BoundedFileAppender::with_now(options, Box::new(fixed_now)).unwrap(); + for line in [b"aa\n", b"bb\n", b"cc\n", b"dd\n"] { + appender.write_all(line).unwrap(); + } + appender.flush().unwrap(); + + assert!(matching_logs(dir.path()).len() <= 2); +} + +#[cfg(unix)] +#[test] +fn max_age_retention_removes_old_archives() { + use std::ffi::CString; + use std::os::unix::ffi::OsStrExt; + + let dir = tempdir().unwrap(); + let path = dir.path().join("telemt.log"); + let old_archive = dir.path().join("telemt.log.20000101000000.0"); + fs::write(&old_archive, "old").unwrap(); + + let c_path = CString::new(old_archive.as_os_str().as_bytes()).unwrap(); + let times = [ + libc::timespec { + tv_sec: 0, + tv_nsec: 0, + }, + libc::timespec { + tv_sec: 0, + tv_nsec: 0, + }, + ]; + let rc = unsafe { libc::utimensat(libc::AT_FDCWD, c_path.as_ptr(), times.as_ptr(), 0) }; + assert_eq!(rc, 0); + + let mut options = options(path); + options.max_age_secs = 1; + let _appender = BoundedFileAppender::with_now(options, Box::new(fixed_now)).unwrap(); + + assert!(!old_archive.exists()); +} + +#[cfg(unix)] +#[test] +fn rotation_stays_bound_to_opened_directory_after_path_replacement() { + use std::os::unix::fs::symlink; + + let root = tempdir().unwrap(); + let original = root.path().join("logs"); + let moved = root.path().join("logs-moved"); + let redirect = root.path().join("redirect"); + fs::create_dir(&original).unwrap(); + fs::create_dir(&redirect).unwrap(); + let mut options = options(original.join("telemt.log")); + options.max_size_bytes = 4; + let mut appender = BoundedFileAppender::with_now(options, Box::new(fixed_now)).unwrap(); + appender.write_all(b"aa\n").unwrap(); + fs::rename(&original, &moved).unwrap(); + symlink(&redirect, &original).unwrap(); + + appender.write_all(b"bb\n").unwrap(); + appender.flush().unwrap(); + + assert!(!matching_logs(&moved).is_empty()); + assert!(matching_logs(&redirect).is_empty()); +} diff --git a/src/maestro/bootstrap.rs b/src/maestro/bootstrap.rs index a89ee39..fd9abb0 100644 --- a/src/maestro/bootstrap.rs +++ b/src/maestro/bootstrap.rs @@ -74,26 +74,7 @@ pub(super) async fn bootstrap( data_path.as_deref(), ); - if !runtime_base_dir.exists() - && let Err(e) = std::fs::create_dir_all(&runtime_base_dir) - { - eprintln!( - "[telemt] Can't create runtime directory {}: {}", - runtime_base_dir.display(), - e - ); - std::process::exit(1); - } - - if !runtime_base_dir.is_dir() { - eprintln!( - "[telemt] Runtime path exists but is not a directory: {}", - runtime_base_dir.display() - ); - std::process::exit(1); - } - - if let Err(e) = std::env::set_current_dir(&runtime_base_dir) { + if let Err(e) = enter_runtime_directory(&runtime_base_dir) { eprintln!( "[telemt] Can't use runtime directory {}: {}", runtime_base_dir.display(), @@ -125,7 +106,7 @@ pub(super) async fn bootstrap( if config_path_explicit { if let Some(serialized) = serialized.as_ref() { - if let Err(write_error) = std::fs::write(&config_path, serialized) { + if let Err(write_error) = write_private_file(&config_path, serialized) { eprintln!( "[telemt] Error: failed to create explicit config at {}: {}", config_path.display(), @@ -149,7 +130,7 @@ pub(super) async fn bootstrap( if let Some(serialized) = serialized.as_ref() { match std::fs::create_dir_all(&runtime_base_dir) { - Ok(()) => match std::fs::write(&runtime_config_path, serialized) { + Ok(()) => match write_private_file(&runtime_config_path, serialized) { Ok(()) => { config_path = runtime_config_path; eprintln!( @@ -176,7 +157,7 @@ pub(super) async fn bootstrap( } if !persisted { - match std::fs::write(&fallback_config_path, serialized) { + match write_private_file(&fallback_config_path, serialized) { Ok(()) => { config_path = fallback_config_path; eprintln!( @@ -226,24 +207,7 @@ pub(super) async fn bootstrap( std::process::exit(1); } - if data_path.exists() { - if !data_path.is_dir() { - eprintln!( - "[telemt] data_path exists but is not a directory: {}", - data_path.display() - ); - std::process::exit(1); - } - } else if let Err(e) = std::fs::create_dir_all(data_path) { - eprintln!( - "[telemt] Can't create data_path {}: {}", - data_path.display(), - e - ); - std::process::exit(1); - } - - if let Err(e) = std::env::set_current_dir(data_path) { + if let Err(e) = enter_runtime_directory(data_path) { eprintln!( "[telemt] Can't use data_path {}: {}", data_path.display(), @@ -376,3 +340,26 @@ pub(super) async fn bootstrap( logging_guard, }) } + +fn enter_runtime_directory(path: &std::path::Path) -> std::io::Result<()> { + #[cfg(unix)] + { + crate::util::secure_fs::chdir_nofollow_or_create(path, 0o750) + } + #[cfg(not(unix))] + { + std::fs::create_dir_all(path)?; + std::env::set_current_dir(path) + } +} + +fn write_private_file(path: &std::path::Path, contents: &str) -> std::io::Result<()> { + #[cfg(unix)] + { + crate::util::secure_fs::atomic_replace(path, contents.as_bytes(), 0o600) + } + #[cfg(not(unix))] + { + std::fs::write(path, contents) + } +} diff --git a/src/maestro/helpers.rs b/src/maestro/helpers.rs index 9a50f16..3245674 100644 --- a/src/maestro/helpers.rs +++ b/src/maestro/helpers.rs @@ -1,20 +1,8 @@ -#![allow(clippy::items_after_test_module)] - use std::path::{Path, PathBuf}; use std::sync::atomic::{AtomicBool, Ordering}; -use std::time::Duration; - -use tokio::sync::watch; -use tracing::{debug, error, info, warn}; use crate::cli; -use crate::config::ProxyConfig; use crate::logging::LogCliOptions; -use crate::transport::UpstreamManager; -use crate::transport::middle_proxy::{ - ProxyConfigData, fetch_proxy_config_with_raw_via_upstream, load_proxy_config_cache, - save_proxy_config_cache, -}; const MAESTRO_COLOR: &str = "\x1b[92m"; const COLOR_RESET: &str = "\x1b[0m"; @@ -314,588 +302,10 @@ fn print_help() { } } +// Runtime reporting and startup snapshot helpers. +mod runtime; + +pub(crate) use runtime::*; + #[cfg(test)] -mod tests { - use std::path::{Path, PathBuf}; - - use super::{ - expected_handshake_close_description, format_maestro_line, is_expected_handshake_eof, - peer_close_description, resolve_runtime_base_dir, resolve_runtime_config_path, - }; - use crate::error::{ProxyError, StreamError}; - - #[test] - fn maestro_line_formatter_respects_disabled_colors() { - let plain = format_maestro_line("boot", false); - assert_eq!(plain, "MAESTRO: boot"); - assert!(!plain.contains('\x1b')); - } - - #[test] - fn maestro_line_formatter_keeps_color_when_enabled() { - let colored = format_maestro_line("boot", true); - assert!(colored.contains("\x1b[92mMAESTRO\x1b[0m")); - } - - #[test] - fn resolve_runtime_config_path_anchors_relative_to_startup_cwd() { - let nonce = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_nanos(); - let startup_cwd = std::env::temp_dir().join(format!("telemt_cfg_path_{nonce}")); - std::fs::create_dir_all(&startup_cwd).unwrap(); - let target = startup_cwd.join("config.toml"); - std::fs::write(&target, " ").unwrap(); - - let resolved = resolve_runtime_config_path("config.toml", &startup_cwd, true); - assert_eq!(resolved, target.canonicalize().unwrap()); - - let _ = std::fs::remove_file(&target); - let _ = std::fs::remove_dir(&startup_cwd); - } - - #[test] - fn resolve_runtime_config_path_keeps_absolute_for_missing_file() { - let nonce = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_nanos(); - let startup_cwd = std::env::temp_dir().join(format!("telemt_cfg_path_missing_{nonce}")); - std::fs::create_dir_all(&startup_cwd).unwrap(); - - let resolved = resolve_runtime_config_path("missing.toml", &startup_cwd, true); - assert_eq!(resolved, startup_cwd.join("missing.toml")); - - let _ = std::fs::remove_dir(&startup_cwd); - } - - #[test] - fn resolve_runtime_config_path_uses_startup_candidates_when_not_explicit() { - let nonce = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_nanos(); - let startup_cwd = - std::env::temp_dir().join(format!("telemt_cfg_startup_candidates_{nonce}")); - std::fs::create_dir_all(&startup_cwd).unwrap(); - let telemt = startup_cwd.join("telemt.toml"); - std::fs::write(&telemt, " ").unwrap(); - - let resolved = resolve_runtime_config_path("config.toml", &startup_cwd, false); - assert_eq!(resolved, telemt.canonicalize().unwrap()); - - let _ = std::fs::remove_file(&telemt); - let _ = std::fs::remove_dir(&startup_cwd); - } - - #[test] - fn resolve_runtime_config_path_defaults_to_startup_config_when_none_found() { - let nonce = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_nanos(); - let startup_cwd = std::env::temp_dir().join(format!("telemt_cfg_startup_default_{nonce}")); - std::fs::create_dir_all(&startup_cwd).unwrap(); - - let resolved = resolve_runtime_config_path("config.toml", &startup_cwd, false); - assert_eq!(resolved, startup_cwd.join("config.toml")); - - let _ = std::fs::remove_dir(&startup_cwd); - } - - #[test] - fn resolve_runtime_base_dir_prefers_cli_data_path() { - let nonce = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_nanos(); - let startup_cwd = std::env::temp_dir().join(format!("telemt_runtime_base_cwd_{nonce}")); - let data_path = std::env::temp_dir().join(format!("telemt_runtime_base_data_{nonce}")); - std::fs::create_dir_all(&startup_cwd).unwrap(); - std::fs::create_dir_all(&data_path).unwrap(); - - let resolved = resolve_runtime_base_dir( - &startup_cwd.join("config.toml"), - &startup_cwd, - true, - Some(&data_path), - ); - assert_eq!(resolved, data_path.canonicalize().unwrap()); - - let _ = std::fs::remove_dir(&data_path); - let _ = std::fs::remove_dir(&startup_cwd); - } - - #[test] - fn resolve_runtime_base_dir_uses_working_directory_before_explicit_config_parent() { - let nonce = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_nanos(); - let startup_cwd = std::env::temp_dir().join(format!("telemt_runtime_base_start_{nonce}")); - let config_dir = std::env::temp_dir().join(format!("telemt_runtime_base_cfg_{nonce}")); - std::fs::create_dir_all(&startup_cwd).unwrap(); - std::fs::create_dir_all(&config_dir).unwrap(); - - let resolved = - resolve_runtime_base_dir(&config_dir.join("telemt.toml"), &startup_cwd, true, None); - assert_eq!(resolved, startup_cwd.canonicalize().unwrap()); - - let _ = std::fs::remove_dir(&config_dir); - let _ = std::fs::remove_dir(&startup_cwd); - } - - #[test] - fn resolve_runtime_base_dir_uses_explicit_config_parent_from_root() { - let nonce = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_nanos(); - let config_dir = std::env::temp_dir().join(format!("telemt_runtime_base_root_cfg_{nonce}")); - std::fs::create_dir_all(&config_dir).unwrap(); - - let resolved = - resolve_runtime_base_dir(&config_dir.join("telemt.toml"), Path::new("/"), true, None); - assert_eq!(resolved, config_dir.canonicalize().unwrap()); - - let _ = std::fs::remove_dir(&config_dir); - } - - #[test] - fn resolve_runtime_base_dir_uses_systemd_working_directory_before_etc() { - let nonce = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_nanos(); - let startup_cwd = std::env::temp_dir().join(format!("telemt_runtime_base_systemd_{nonce}")); - std::fs::create_dir_all(&startup_cwd).unwrap(); - - let resolved = - resolve_runtime_base_dir(&startup_cwd.join("config.toml"), &startup_cwd, false, None); - assert_eq!(resolved, startup_cwd.canonicalize().unwrap()); - - let _ = std::fs::remove_dir(&startup_cwd); - } - - #[test] - fn resolve_runtime_base_dir_falls_back_to_etc_from_root() { - let resolved = resolve_runtime_base_dir( - Path::new("/etc/telemt/config.toml"), - Path::new("/"), - false, - None, - ); - assert_eq!(resolved, PathBuf::from("/etc/telemt")); - } - - #[test] - fn expected_handshake_eof_matches_connection_reset() { - let err = ProxyError::Io(std::io::Error::from(std::io::ErrorKind::ConnectionReset)); - assert!(is_expected_handshake_eof(&err)); - } - - #[test] - fn expected_handshake_eof_matches_stream_io_unexpected_eof() { - let err = ProxyError::Stream(StreamError::Io(std::io::Error::from( - std::io::ErrorKind::UnexpectedEof, - ))); - assert!(is_expected_handshake_eof(&err)); - } - - #[test] - fn peer_close_description_is_human_readable_for_all_peer_close_kinds() { - let cases = [ - ( - std::io::ErrorKind::ConnectionReset, - "Peer reset TCP connection (RST)", - ), - ( - std::io::ErrorKind::ConnectionAborted, - "Peer aborted TCP connection during transport", - ), - ( - std::io::ErrorKind::BrokenPipe, - "Peer closed write side (broken pipe)", - ), - ( - std::io::ErrorKind::NotConnected, - "Socket was already closed by peer", - ), - ]; - - for (kind, expected) in cases { - let err = ProxyError::Io(std::io::Error::from(kind)); - assert_eq!(peer_close_description(&err), Some(expected)); - } - } - - #[test] - fn handshake_close_description_is_human_readable_for_all_expected_kinds() { - let cases = [ - ( - ProxyError::Io(std::io::Error::from(std::io::ErrorKind::UnexpectedEof)), - "Peer closed before sending full 64-byte MTProto handshake", - ), - ( - ProxyError::Io(std::io::Error::from(std::io::ErrorKind::ConnectionReset)), - "Peer reset TCP connection during initial MTProto handshake", - ), - ( - ProxyError::Io(std::io::Error::from(std::io::ErrorKind::ConnectionAborted)), - "Peer aborted TCP connection during initial MTProto handshake", - ), - ( - ProxyError::Io(std::io::Error::from(std::io::ErrorKind::BrokenPipe)), - "Peer closed write side before MTProto handshake completed", - ), - ( - ProxyError::Io(std::io::Error::from(std::io::ErrorKind::NotConnected)), - "Handshake socket was already closed by peer", - ), - ( - ProxyError::Stream(StreamError::UnexpectedEof), - "Peer closed before sending full 64-byte MTProto handshake", - ), - ]; - - for (err, expected) in cases { - assert_eq!(expected_handshake_close_description(&err), Some(expected)); - } - } -} - -pub(crate) fn print_proxy_links(host: &str, port: u16, config: &ProxyConfig) { - print_maestro_line(format!("Proxy links ({host})")); - for user_name in config - .general - .links - .show - .resolve_users(&config.access.users) - { - if let Some(secret) = config.access.users.get(user_name) { - print_maestro_line(format!("User: {user_name}")); - if config.general.modes.classic { - print_maestro_line(format!( - "Classic: tg://proxy?server={host}&port={port}&secret={secret}" - )); - } - if config.general.modes.secure { - print_maestro_line(format!( - "DD: tg://proxy?server={host}&port={port}&secret=dd{secret}" - )); - } - if config.general.modes.tls { - let mut domains = Vec::with_capacity(1 + config.censorship.tls_domains.len()); - domains.push(config.censorship.tls_domain.clone()); - for d in &config.censorship.tls_domains { - if !domains.contains(d) { - domains.push(d.clone()); - } - } - - for domain in domains { - let domain_hex = hex::encode(&domain); - print_maestro_line(format!( - "EE-TLS: tg://proxy?server={host}&port={port}&secret=ee{secret}{domain_hex}" - )); - } - } - } else { - warn!(target: "telemt::links", "User '{}' in show_link not found", user_name); - } - } -} - -/// Prints WEB links only for profiles selected by the existing link policy. -pub(crate) fn print_web_proxy_links(config: &ProxyConfig) { - if !config.web.enabled || config.general.links.show.is_empty() { - return; - } - let Some(runtime) = config.web.runtime.as_ref() else { - return; - }; - let shown = config - .general - .links - .show - .resolve_users(&config.access.users); - let mut heading_printed = false; - for profile in &runtime.profiles { - if !shown.iter().any(|user| user.as_str() == profile.user) { - continue; - } - if !heading_printed { - print_maestro_line("WEB proxy links"); - heading_printed = true; - } - let Some(secret) = config.access.users.get(&profile.user) else { - continue; - }; - let prefix = match profile.secret_mode { - crate::config::WebSecretMode::Plain => "", - crate::config::WebSecretMode::Dd => "dd", - }; - print_maestro_line(format!( - "User: {} ({:?})", - profile.user, profile.secret_mode - )); - print_maestro_line(format!( - "WEB: tg://webproxy?server={}&secret={prefix}{secret}", - profile.host, - )); - } -} - -pub(crate) async fn write_beobachten_snapshot(path: &str, payload: &str) -> std::io::Result<()> { - if let Some(parent) = std::path::Path::new(path).parent() - && !parent.as_os_str().is_empty() - { - tokio::fs::create_dir_all(parent).await?; - } - tokio::fs::write(path, payload).await -} - -pub(crate) fn unit_label(value: u64, singular: &'static str, plural: &'static str) -> &'static str { - if value == 1 { singular } else { plural } -} - -pub(crate) fn format_uptime(total_secs: u64) -> String { - const SECS_PER_MINUTE: u64 = 60; - const SECS_PER_HOUR: u64 = 60 * SECS_PER_MINUTE; - const SECS_PER_DAY: u64 = 24 * SECS_PER_HOUR; - const SECS_PER_MONTH: u64 = 30 * SECS_PER_DAY; - const SECS_PER_YEAR: u64 = 12 * SECS_PER_MONTH; - - let mut remaining = total_secs; - let years = remaining / SECS_PER_YEAR; - remaining %= SECS_PER_YEAR; - let months = remaining / SECS_PER_MONTH; - remaining %= SECS_PER_MONTH; - let days = remaining / SECS_PER_DAY; - remaining %= SECS_PER_DAY; - let hours = remaining / SECS_PER_HOUR; - remaining %= SECS_PER_HOUR; - let minutes = remaining / SECS_PER_MINUTE; - let seconds = remaining % SECS_PER_MINUTE; - - let mut parts = Vec::new(); - if total_secs > SECS_PER_YEAR { - parts.push(format!("{} {}", years, unit_label(years, "year", "years"))); - } - if total_secs > SECS_PER_MONTH { - parts.push(format!( - "{} {}", - months, - unit_label(months, "month", "months") - )); - } - if total_secs > SECS_PER_DAY { - parts.push(format!("{} {}", days, unit_label(days, "day", "days"))); - } - if total_secs > SECS_PER_HOUR { - parts.push(format!("{} {}", hours, unit_label(hours, "hour", "hours"))); - } - if total_secs > SECS_PER_MINUTE { - parts.push(format!( - "{} {}", - minutes, - unit_label(minutes, "minute", "minutes") - )); - } - parts.push(format!( - "{} {}", - seconds, - unit_label(seconds, "second", "seconds") - )); - - format!("{} / {} seconds", parts.join(", "), total_secs) -} - -#[allow(dead_code)] -pub(crate) async fn wait_until_admission_open(admission_rx: &mut watch::Receiver) -> bool { - loop { - if *admission_rx.borrow() { - return true; - } - if admission_rx.changed().await.is_err() { - return *admission_rx.borrow(); - } - } -} - -pub(crate) fn is_expected_handshake_eof(err: &crate::error::ProxyError) -> bool { - expected_handshake_close_description(err).is_some() -} - -pub(crate) fn peer_close_description(err: &crate::error::ProxyError) -> Option<&'static str> { - fn from_kind(kind: std::io::ErrorKind) -> Option<&'static str> { - match kind { - std::io::ErrorKind::ConnectionReset => Some("Peer reset TCP connection (RST)"), - std::io::ErrorKind::ConnectionAborted => { - Some("Peer aborted TCP connection during transport") - } - std::io::ErrorKind::BrokenPipe => Some("Peer closed write side (broken pipe)"), - std::io::ErrorKind::NotConnected => Some("Socket was already closed by peer"), - _ => None, - } - } - - match err { - crate::error::ProxyError::Io(ioe) => from_kind(ioe.kind()), - crate::error::ProxyError::Stream(crate::error::StreamError::Io(ioe)) => { - from_kind(ioe.kind()) - } - _ => None, - } -} - -pub(crate) fn expected_handshake_close_description( - err: &crate::error::ProxyError, -) -> Option<&'static str> { - fn from_kind(kind: std::io::ErrorKind) -> Option<&'static str> { - match kind { - std::io::ErrorKind::UnexpectedEof => { - Some("Peer closed before sending full 64-byte MTProto handshake") - } - std::io::ErrorKind::ConnectionReset => { - Some("Peer reset TCP connection during initial MTProto handshake") - } - std::io::ErrorKind::ConnectionAborted => { - Some("Peer aborted TCP connection during initial MTProto handshake") - } - std::io::ErrorKind::BrokenPipe => { - Some("Peer closed write side before MTProto handshake completed") - } - std::io::ErrorKind::NotConnected => Some("Handshake socket was already closed by peer"), - _ => None, - } - } - - match err { - crate::error::ProxyError::Io(ioe) => from_kind(ioe.kind()), - crate::error::ProxyError::Stream(crate::error::StreamError::UnexpectedEof) => { - Some("Peer closed before sending full 64-byte MTProto handshake") - } - crate::error::ProxyError::Stream(crate::error::StreamError::Io(ioe)) => { - from_kind(ioe.kind()) - } - _ => None, - } -} - -pub(crate) async fn load_startup_proxy_config_snapshot( - url: &str, - cache_path: Option<&str>, - me2dc_fallback: bool, - label: &'static str, - upstream: Option>, -) -> Option { - loop { - match fetch_proxy_config_with_raw_via_upstream(url, upstream.clone()).await { - Ok((cfg, raw)) => { - if !cfg.map.is_empty() { - if let Some(path) = cache_path - && let Err(e) = save_proxy_config_cache(path, &raw).await - { - warn!(error = %e, path, snapshot = label, "Failed to store startup proxy-config cache"); - } - return Some(cfg); - } - - warn!( - snapshot = label, - url, "Startup proxy-config is empty; trying disk cache" - ); - if let Some(path) = cache_path { - match load_proxy_config_cache(path).await { - Ok(cached) if !cached.map.is_empty() => { - info!( - snapshot = label, - path, - proxy_for_lines = cached.proxy_for_lines, - "Loaded startup proxy-config from disk cache" - ); - return Some(cached); - } - Ok(_) => { - warn!( - snapshot = label, - path, "Startup proxy-config cache is empty; ignoring cache file" - ); - } - Err(cache_err) => { - debug!( - snapshot = label, - path, - error = %cache_err, - "Startup proxy-config cache unavailable" - ); - } - } - } - - if me2dc_fallback { - error!( - snapshot = label, - "Startup proxy-config unavailable and no saved config found; falling back to direct mode" - ); - return None; - } - - warn!( - snapshot = label, - retry_in_secs = 2, - "Startup proxy-config unavailable and no saved config found; retrying because me2dc_fallback=false" - ); - tokio::time::sleep(Duration::from_secs(2)).await; - } - Err(fetch_err) => { - if let Some(path) = cache_path { - match load_proxy_config_cache(path).await { - Ok(cached) if !cached.map.is_empty() => { - info!( - snapshot = label, - path, - proxy_for_lines = cached.proxy_for_lines, - "Loaded startup proxy-config from disk cache" - ); - return Some(cached); - } - Ok(_) => { - warn!( - snapshot = label, - path, "Startup proxy-config cache is empty; ignoring cache file" - ); - } - Err(cache_err) => { - debug!( - snapshot = label, - path, - error = %cache_err, - "Startup proxy-config cache unavailable" - ); - } - } - } - - if me2dc_fallback { - error!( - snapshot = label, - error = %fetch_err, - "Startup proxy-config unavailable and no cached data; falling back to direct mode" - ); - return None; - } - - warn!( - snapshot = label, - error = %fetch_err, - retry_in_secs = 2, - "Startup proxy-config unavailable; retrying because me2dc_fallback=false" - ); - tokio::time::sleep(Duration::from_secs(2)).await; - } - } - } -} +mod tests; diff --git a/src/maestro/helpers/runtime.rs b/src/maestro/helpers/runtime.rs new file mode 100644 index 0000000..245650f --- /dev/null +++ b/src/maestro/helpers/runtime.rs @@ -0,0 +1,360 @@ +use std::time::Duration; + +use tokio::sync::watch; +use tracing::{debug, error, info, warn}; + +use crate::config::ProxyConfig; +use crate::transport::UpstreamManager; +use crate::transport::middle_proxy::{ + ProxyConfigData, fetch_proxy_config_with_raw_via_upstream, load_proxy_config_cache, + save_proxy_config_cache, +}; + +use super::print_maestro_line; + +pub(crate) fn print_proxy_links(host: &str, port: u16, config: &ProxyConfig) { + print_maestro_line(format!("Proxy links ({host})")); + for user_name in config + .general + .links + .show + .resolve_users(&config.access.users) + { + if let Some(secret) = config.access.users.get(user_name) { + print_maestro_line(format!("User: {user_name}")); + if config.general.modes.classic { + print_maestro_line(format!( + "Classic: tg://proxy?server={host}&port={port}&secret={secret}" + )); + } + if config.general.modes.secure { + print_maestro_line(format!( + "DD: tg://proxy?server={host}&port={port}&secret=dd{secret}" + )); + } + if config.general.modes.tls { + let mut domains = Vec::with_capacity(1 + config.censorship.tls_domains.len()); + domains.push(config.censorship.tls_domain.clone()); + for d in &config.censorship.tls_domains { + if !domains.contains(d) { + domains.push(d.clone()); + } + } + + for domain in domains { + let domain_hex = hex::encode(&domain); + print_maestro_line(format!( + "EE-TLS: tg://proxy?server={host}&port={port}&secret=ee{secret}{domain_hex}" + )); + } + } + } else { + warn!(target: "telemt::links", "User '{}' in show_link not found", user_name); + } + } +} + +/// Prints WEB links only for profiles selected by the existing link policy. +pub(crate) fn print_web_proxy_links(config: &ProxyConfig) { + if !config.web.enabled || config.general.links.show.is_empty() { + return; + } + let Some(runtime) = config.web.runtime.as_ref() else { + return; + }; + let shown = config + .general + .links + .show + .resolve_users(&config.access.users); + let mut heading_printed = false; + for profile in &runtime.profiles { + if !shown.iter().any(|user| user.as_str() == profile.user) { + continue; + } + if !heading_printed { + print_maestro_line("WEB proxy links"); + heading_printed = true; + } + let Some(secret) = config.access.users.get(&profile.user) else { + continue; + }; + let prefix = match profile.secret_mode { + crate::config::WebSecretMode::Plain => "", + crate::config::WebSecretMode::Dd => "dd", + }; + print_maestro_line(format!( + "User: {} ({:?})", + profile.user, profile.secret_mode + )); + print_maestro_line(format!( + "WEB: tg://webproxy?server={}&secret={prefix}{secret}", + profile.host, + )); + } +} + +pub(crate) async fn write_beobachten_snapshot(path: &str, payload: &str) -> std::io::Result<()> { + #[cfg(unix)] + { + crate::util::secure_fs::atomic_replace_async( + std::path::PathBuf::from(path), + payload.as_bytes().to_vec(), + 0o600, + ) + .await + } + #[cfg(not(unix))] + { + if let Some(parent) = std::path::Path::new(path).parent() + && !parent.as_os_str().is_empty() + { + tokio::fs::create_dir_all(parent).await?; + } + tokio::fs::write(path, payload).await + } +} + +pub(crate) fn unit_label(value: u64, singular: &'static str, plural: &'static str) -> &'static str { + if value == 1 { singular } else { plural } +} + +pub(crate) fn format_uptime(total_secs: u64) -> String { + const SECS_PER_MINUTE: u64 = 60; + const SECS_PER_HOUR: u64 = 60 * SECS_PER_MINUTE; + const SECS_PER_DAY: u64 = 24 * SECS_PER_HOUR; + const SECS_PER_MONTH: u64 = 30 * SECS_PER_DAY; + const SECS_PER_YEAR: u64 = 12 * SECS_PER_MONTH; + + let mut remaining = total_secs; + let years = remaining / SECS_PER_YEAR; + remaining %= SECS_PER_YEAR; + let months = remaining / SECS_PER_MONTH; + remaining %= SECS_PER_MONTH; + let days = remaining / SECS_PER_DAY; + remaining %= SECS_PER_DAY; + let hours = remaining / SECS_PER_HOUR; + remaining %= SECS_PER_HOUR; + let minutes = remaining / SECS_PER_MINUTE; + let seconds = remaining % SECS_PER_MINUTE; + + let mut parts = Vec::new(); + if total_secs > SECS_PER_YEAR { + parts.push(format!("{} {}", years, unit_label(years, "year", "years"))); + } + if total_secs > SECS_PER_MONTH { + parts.push(format!( + "{} {}", + months, + unit_label(months, "month", "months") + )); + } + if total_secs > SECS_PER_DAY { + parts.push(format!("{} {}", days, unit_label(days, "day", "days"))); + } + if total_secs > SECS_PER_HOUR { + parts.push(format!("{} {}", hours, unit_label(hours, "hour", "hours"))); + } + if total_secs > SECS_PER_MINUTE { + parts.push(format!( + "{} {}", + minutes, + unit_label(minutes, "minute", "minutes") + )); + } + parts.push(format!( + "{} {}", + seconds, + unit_label(seconds, "second", "seconds") + )); + + format!("{} / {} seconds", parts.join(", "), total_secs) +} + +#[allow(dead_code)] +pub(crate) async fn wait_until_admission_open(admission_rx: &mut watch::Receiver) -> bool { + loop { + if *admission_rx.borrow() { + return true; + } + if admission_rx.changed().await.is_err() { + return *admission_rx.borrow(); + } + } +} + +pub(crate) fn is_expected_handshake_eof(err: &crate::error::ProxyError) -> bool { + expected_handshake_close_description(err).is_some() +} + +pub(crate) fn peer_close_description(err: &crate::error::ProxyError) -> Option<&'static str> { + fn from_kind(kind: std::io::ErrorKind) -> Option<&'static str> { + match kind { + std::io::ErrorKind::ConnectionReset => Some("Peer reset TCP connection (RST)"), + std::io::ErrorKind::ConnectionAborted => { + Some("Peer aborted TCP connection during transport") + } + std::io::ErrorKind::BrokenPipe => Some("Peer closed write side (broken pipe)"), + std::io::ErrorKind::NotConnected => Some("Socket was already closed by peer"), + _ => None, + } + } + + match err { + crate::error::ProxyError::Io(ioe) => from_kind(ioe.kind()), + crate::error::ProxyError::Stream(crate::error::StreamError::Io(ioe)) => { + from_kind(ioe.kind()) + } + _ => None, + } +} + +pub(crate) fn expected_handshake_close_description( + err: &crate::error::ProxyError, +) -> Option<&'static str> { + fn from_kind(kind: std::io::ErrorKind) -> Option<&'static str> { + match kind { + std::io::ErrorKind::UnexpectedEof => { + Some("Peer closed before sending full 64-byte MTProto handshake") + } + std::io::ErrorKind::ConnectionReset => { + Some("Peer reset TCP connection during initial MTProto handshake") + } + std::io::ErrorKind::ConnectionAborted => { + Some("Peer aborted TCP connection during initial MTProto handshake") + } + std::io::ErrorKind::BrokenPipe => { + Some("Peer closed write side before MTProto handshake completed") + } + std::io::ErrorKind::NotConnected => Some("Handshake socket was already closed by peer"), + _ => None, + } + } + + match err { + crate::error::ProxyError::Io(ioe) => from_kind(ioe.kind()), + crate::error::ProxyError::Stream(crate::error::StreamError::UnexpectedEof) => { + Some("Peer closed before sending full 64-byte MTProto handshake") + } + crate::error::ProxyError::Stream(crate::error::StreamError::Io(ioe)) => { + from_kind(ioe.kind()) + } + _ => None, + } +} + +pub(crate) async fn load_startup_proxy_config_snapshot( + url: &str, + cache_path: Option<&str>, + me2dc_fallback: bool, + label: &'static str, + upstream: Option>, +) -> Option { + loop { + match fetch_proxy_config_with_raw_via_upstream(url, upstream.clone()).await { + Ok((cfg, raw)) => { + if !cfg.map.is_empty() { + if let Some(path) = cache_path + && let Err(e) = save_proxy_config_cache(path, &raw).await + { + warn!(error = %e, path, snapshot = label, "Failed to store startup proxy-config cache"); + } + return Some(cfg); + } + + warn!( + snapshot = label, + url, "Startup proxy-config is empty; trying disk cache" + ); + if let Some(path) = cache_path { + match load_proxy_config_cache(path).await { + Ok(cached) if !cached.map.is_empty() => { + info!( + snapshot = label, + path, + proxy_for_lines = cached.proxy_for_lines, + "Loaded startup proxy-config from disk cache" + ); + return Some(cached); + } + Ok(_) => { + warn!( + snapshot = label, + path, "Startup proxy-config cache is empty; ignoring cache file" + ); + } + Err(cache_err) => { + debug!( + snapshot = label, + path, + error = %cache_err, + "Startup proxy-config cache unavailable" + ); + } + } + } + + if me2dc_fallback { + error!( + snapshot = label, + "Startup proxy-config unavailable and no saved config found; falling back to direct mode" + ); + return None; + } + + warn!( + snapshot = label, + retry_in_secs = 2, + "Startup proxy-config unavailable and no saved config found; retrying because me2dc_fallback=false" + ); + tokio::time::sleep(Duration::from_secs(2)).await; + } + Err(fetch_err) => { + if let Some(path) = cache_path { + match load_proxy_config_cache(path).await { + Ok(cached) if !cached.map.is_empty() => { + info!( + snapshot = label, + path, + proxy_for_lines = cached.proxy_for_lines, + "Loaded startup proxy-config from disk cache" + ); + return Some(cached); + } + Ok(_) => { + warn!( + snapshot = label, + path, "Startup proxy-config cache is empty; ignoring cache file" + ); + } + Err(cache_err) => { + debug!( + snapshot = label, + path, + error = %cache_err, + "Startup proxy-config cache unavailable" + ); + } + } + } + + if me2dc_fallback { + error!( + snapshot = label, + error = %fetch_err, + "Startup proxy-config unavailable and no cached data; falling back to direct mode" + ); + return None; + } + + warn!( + snapshot = label, + error = %fetch_err, + retry_in_secs = 2, + "Startup proxy-config unavailable; retrying because me2dc_fallback=false" + ); + tokio::time::sleep(Duration::from_secs(2)).await; + } + } + } +} diff --git a/src/maestro/helpers/tests.rs b/src/maestro/helpers/tests.rs new file mode 100644 index 0000000..de28726 --- /dev/null +++ b/src/maestro/helpers/tests.rs @@ -0,0 +1,247 @@ + use std::path::{Path, PathBuf}; + + use super::{ + expected_handshake_close_description, format_maestro_line, is_expected_handshake_eof, + peer_close_description, resolve_runtime_base_dir, resolve_runtime_config_path, + }; + use crate::error::{ProxyError, StreamError}; + + #[test] + fn maestro_line_formatter_respects_disabled_colors() { + let plain = format_maestro_line("boot", false); + assert_eq!(plain, "MAESTRO: boot"); + assert!(!plain.contains('\x1b')); + } + + #[test] + fn maestro_line_formatter_keeps_color_when_enabled() { + let colored = format_maestro_line("boot", true); + assert!(colored.contains("\x1b[92mMAESTRO\x1b[0m")); + } + + #[test] + fn resolve_runtime_config_path_anchors_relative_to_startup_cwd() { + let nonce = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_nanos(); + let startup_cwd = std::env::temp_dir().join(format!("telemt_cfg_path_{nonce}")); + std::fs::create_dir_all(&startup_cwd).unwrap(); + let target = startup_cwd.join("config.toml"); + std::fs::write(&target, " ").unwrap(); + + let resolved = resolve_runtime_config_path("config.toml", &startup_cwd, true); + assert_eq!(resolved, target.canonicalize().unwrap()); + + let _ = std::fs::remove_file(&target); + let _ = std::fs::remove_dir(&startup_cwd); + } + + #[test] + fn resolve_runtime_config_path_keeps_absolute_for_missing_file() { + let nonce = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_nanos(); + let startup_cwd = std::env::temp_dir().join(format!("telemt_cfg_path_missing_{nonce}")); + std::fs::create_dir_all(&startup_cwd).unwrap(); + + let resolved = resolve_runtime_config_path("missing.toml", &startup_cwd, true); + assert_eq!(resolved, startup_cwd.join("missing.toml")); + + let _ = std::fs::remove_dir(&startup_cwd); + } + + #[test] + fn resolve_runtime_config_path_uses_startup_candidates_when_not_explicit() { + let nonce = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_nanos(); + let startup_cwd = + std::env::temp_dir().join(format!("telemt_cfg_startup_candidates_{nonce}")); + std::fs::create_dir_all(&startup_cwd).unwrap(); + let telemt = startup_cwd.join("telemt.toml"); + std::fs::write(&telemt, " ").unwrap(); + + let resolved = resolve_runtime_config_path("config.toml", &startup_cwd, false); + assert_eq!(resolved, telemt.canonicalize().unwrap()); + + let _ = std::fs::remove_file(&telemt); + let _ = std::fs::remove_dir(&startup_cwd); + } + + #[test] + fn resolve_runtime_config_path_defaults_to_startup_config_when_none_found() { + let nonce = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_nanos(); + let startup_cwd = std::env::temp_dir().join(format!("telemt_cfg_startup_default_{nonce}")); + std::fs::create_dir_all(&startup_cwd).unwrap(); + + let resolved = resolve_runtime_config_path("config.toml", &startup_cwd, false); + assert_eq!(resolved, startup_cwd.join("config.toml")); + + let _ = std::fs::remove_dir(&startup_cwd); + } + + #[test] + fn resolve_runtime_base_dir_prefers_cli_data_path() { + let nonce = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_nanos(); + let startup_cwd = std::env::temp_dir().join(format!("telemt_runtime_base_cwd_{nonce}")); + let data_path = std::env::temp_dir().join(format!("telemt_runtime_base_data_{nonce}")); + std::fs::create_dir_all(&startup_cwd).unwrap(); + std::fs::create_dir_all(&data_path).unwrap(); + + let resolved = resolve_runtime_base_dir( + &startup_cwd.join("config.toml"), + &startup_cwd, + true, + Some(&data_path), + ); + assert_eq!(resolved, data_path.canonicalize().unwrap()); + + let _ = std::fs::remove_dir(&data_path); + let _ = std::fs::remove_dir(&startup_cwd); + } + + #[test] + fn resolve_runtime_base_dir_uses_working_directory_before_explicit_config_parent() { + let nonce = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_nanos(); + let startup_cwd = std::env::temp_dir().join(format!("telemt_runtime_base_start_{nonce}")); + let config_dir = std::env::temp_dir().join(format!("telemt_runtime_base_cfg_{nonce}")); + std::fs::create_dir_all(&startup_cwd).unwrap(); + std::fs::create_dir_all(&config_dir).unwrap(); + + let resolved = + resolve_runtime_base_dir(&config_dir.join("telemt.toml"), &startup_cwd, true, None); + assert_eq!(resolved, startup_cwd.canonicalize().unwrap()); + + let _ = std::fs::remove_dir(&config_dir); + let _ = std::fs::remove_dir(&startup_cwd); + } + + #[test] + fn resolve_runtime_base_dir_uses_explicit_config_parent_from_root() { + let nonce = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_nanos(); + let config_dir = std::env::temp_dir().join(format!("telemt_runtime_base_root_cfg_{nonce}")); + std::fs::create_dir_all(&config_dir).unwrap(); + + let resolved = + resolve_runtime_base_dir(&config_dir.join("telemt.toml"), Path::new("/"), true, None); + assert_eq!(resolved, config_dir.canonicalize().unwrap()); + + let _ = std::fs::remove_dir(&config_dir); + } + + #[test] + fn resolve_runtime_base_dir_uses_systemd_working_directory_before_etc() { + let nonce = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_nanos(); + let startup_cwd = std::env::temp_dir().join(format!("telemt_runtime_base_systemd_{nonce}")); + std::fs::create_dir_all(&startup_cwd).unwrap(); + + let resolved = + resolve_runtime_base_dir(&startup_cwd.join("config.toml"), &startup_cwd, false, None); + assert_eq!(resolved, startup_cwd.canonicalize().unwrap()); + + let _ = std::fs::remove_dir(&startup_cwd); + } + + #[test] + fn resolve_runtime_base_dir_falls_back_to_etc_from_root() { + let resolved = resolve_runtime_base_dir( + Path::new("/etc/telemt/config.toml"), + Path::new("/"), + false, + None, + ); + assert_eq!(resolved, PathBuf::from("/etc/telemt")); + } + + #[test] + fn expected_handshake_eof_matches_connection_reset() { + let err = ProxyError::Io(std::io::Error::from(std::io::ErrorKind::ConnectionReset)); + assert!(is_expected_handshake_eof(&err)); + } + + #[test] + fn expected_handshake_eof_matches_stream_io_unexpected_eof() { + let err = ProxyError::Stream(StreamError::Io(std::io::Error::from( + std::io::ErrorKind::UnexpectedEof, + ))); + assert!(is_expected_handshake_eof(&err)); + } + + #[test] + fn peer_close_description_is_human_readable_for_all_peer_close_kinds() { + let cases = [ + ( + std::io::ErrorKind::ConnectionReset, + "Peer reset TCP connection (RST)", + ), + ( + std::io::ErrorKind::ConnectionAborted, + "Peer aborted TCP connection during transport", + ), + ( + std::io::ErrorKind::BrokenPipe, + "Peer closed write side (broken pipe)", + ), + ( + std::io::ErrorKind::NotConnected, + "Socket was already closed by peer", + ), + ]; + + for (kind, expected) in cases { + let err = ProxyError::Io(std::io::Error::from(kind)); + assert_eq!(peer_close_description(&err), Some(expected)); + } + } + + #[test] + fn handshake_close_description_is_human_readable_for_all_expected_kinds() { + let cases = [ + ( + ProxyError::Io(std::io::Error::from(std::io::ErrorKind::UnexpectedEof)), + "Peer closed before sending full 64-byte MTProto handshake", + ), + ( + ProxyError::Io(std::io::Error::from(std::io::ErrorKind::ConnectionReset)), + "Peer reset TCP connection during initial MTProto handshake", + ), + ( + ProxyError::Io(std::io::Error::from(std::io::ErrorKind::ConnectionAborted)), + "Peer aborted TCP connection during initial MTProto handshake", + ), + ( + ProxyError::Io(std::io::Error::from(std::io::ErrorKind::BrokenPipe)), + "Peer closed write side before MTProto handshake completed", + ), + ( + ProxyError::Io(std::io::Error::from(std::io::ErrorKind::NotConnected)), + "Handshake socket was already closed by peer", + ), + ( + ProxyError::Stream(StreamError::UnexpectedEof), + "Peer closed before sending full 64-byte MTProto handshake", + ), + ]; + + for (err, expected) in cases { + assert_eq!(expected_handshake_close_description(&err), Some(expected)); + } + } diff --git a/src/maestro/listeners/bind.rs b/src/maestro/listeners/bind.rs index 7b325f8..38033b3 100644 --- a/src/maestro/listeners/bind.rs +++ b/src/maestro/listeners/bind.rs @@ -20,6 +20,8 @@ use crate::config::{ListenerTransport, ProxyConfig}; use crate::startup::{COMPONENT_LISTENERS_BIND, StartupTracker}; use crate::transport::find_listener_processes; use crate::transport::socket::{activate_listener_socket, bind_listener_socket}; +#[cfg(unix)] +use crate::util::secure_fs::AnchoredPath; use super::plan::{ListenerBindSpec, listener_bind_plan}; use crate::maestro::helpers::{print_proxy_links, print_web_proxy_links}; @@ -281,6 +283,7 @@ pub(crate) async fn bind_listeners( #[cfg(unix)] if let Some(unix_path) = &config.server.listen_unix_sock { let unix_path = Path::new(unix_path); + let anchored_path = AnchoredPath::open_trusted_parent(unix_path)?; remove_stale_unix_socket(unix_path)?; let unix_listener = UnixListener::bind(unix_path)?; let socket_metadata = std::fs::symlink_metadata(unix_path)?; @@ -295,10 +298,15 @@ pub(crate) async fn bind_listeners( if let Some(perm_str) = &config.server.listen_unix_sock_perm { match u32::from_str_radix(perm_str.trim_start_matches('0'), 8) { Ok(mode) => { - use std::os::unix::fs::PermissionsExt; - let permissions = std::fs::Permissions::from_mode(mode); + use nix::sys::stat::{FchmodatFlags, Mode, fchmodat}; + verify_bound_unix_socket(unix_path, socket_identity)?; - if let Err(error_value) = std::fs::set_permissions(unix_path, permissions) { + if let Err(error_value) = fchmodat( + anchored_path.parent(), + anchored_path.name(), + Mode::from_bits_truncate(mode), + FchmodatFlags::NoFollowSymlink, + ) { error!( path = %unix_path.display(), permissions = %perm_str, diff --git a/src/proxy/client/authenticated.rs b/src/proxy/client/authenticated.rs index d131863..66f385f 100644 --- a/src/proxy/client/authenticated.rs +++ b/src/proxy/client/authenticated.rs @@ -25,6 +25,21 @@ impl RunningClientHandler { R: AsyncRead + Unpin + Send + 'static, W: AsyncWrite + Unpin + Send + 'static, { + // Manually constructed handshake fixtures bypass credential validation, so + // materialize the process-authority state that a real handshake generation owns. + let config = if config.runtime_user_credential_id(&success.user).is_none() { + let mut test_config = (*config).clone(); + test_config.access.users.insert( + success.user.clone(), + "00000000000000000000000000000000".to_string(), + ); + test_config.rebuild_runtime_user_auth()?; + Arc::new(test_config) + } else { + config + }; + let shared = ProxySharedState::new(); + shared.apply_user_config(&config.access.users, &config.access.user_enabled); Self::handle_authenticated_static_with_shared( client_reader, client_writer, @@ -40,7 +55,7 @@ impl RunningClientHandler { local_addr, peer_addr, ip_tracker, - ProxySharedState::new(), + shared, ) .await } diff --git a/src/proxy/handshake.rs b/src/proxy/handshake.rs index e0d76c6..ca2fa51 100644 --- a/src/proxy/handshake.rs +++ b/src/proxy/handshake.rs @@ -77,6 +77,7 @@ pub(crate) use self::auth_probe::{ auth_probe_saturation_state_lock_for_testing_in_shared, 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, + insert_auth_probe_state_for_testing_in_shared, should_emit_unknown_sni_warn_for_testing_in_shared, warned_secrets_for_testing_in_shared, }; diff --git a/src/proxy/handshake/auth_candidates.rs b/src/proxy/handshake/auth_candidates.rs index fb7fe1f..5445ac8 100644 --- a/src/proxy/handshake/auth_candidates.rs +++ b/src/proxy/handshake/auth_candidates.rs @@ -42,7 +42,7 @@ pub(super) fn ip_prefix_hint_key(peer_ip: IpAddr) -> u64 { } } -pub(super) fn sticky_hint_get_by_ip(shared: &ProxySharedState, peer_ip: IpAddr) -> Option { +pub(super) fn sticky_hint_get_by_ip(shared: &ProxySharedState, peer_ip: IpAddr) -> Option { shared .handshake .sticky_user_by_ip @@ -53,7 +53,7 @@ pub(super) fn sticky_hint_get_by_ip(shared: &ProxySharedState, peer_ip: IpAddr) pub(super) fn sticky_hint_get_by_ip_prefix( shared: &ProxySharedState, peer_ip: IpAddr, -) -> Option { +) -> Option { shared .handshake .sticky_user_by_ip_prefix @@ -61,7 +61,7 @@ pub(super) fn sticky_hint_get_by_ip_prefix( .map(|entry| *entry) } -pub(super) fn sticky_hint_get_by_sni(shared: &ProxySharedState, sni: &str) -> Option { +pub(super) fn sticky_hint_get_by_sni(shared: &ProxySharedState, sni: &str) -> Option { let key = sni_hint_hash(sni); shared .handshake @@ -73,20 +73,20 @@ pub(super) fn sticky_hint_get_by_sni(shared: &ProxySharedState, sni: &str) -> Op pub(super) fn sticky_hint_record_success_in( shared: &ProxySharedState, peer_ip: IpAddr, - user_id: u32, + hint_key: u64, sni: Option<&str>, ) { bounded_sticky_hint_upsert( &shared.handshake.sticky_user_by_ip, &shared.handshake.sticky_user_by_ip_slots, peer_ip, - user_id, + hint_key, ); bounded_sticky_hint_upsert( &shared.handshake.sticky_user_by_ip_prefix, &shared.handshake.sticky_user_by_ip_prefix_slots, ip_prefix_hint_key(peer_ip), - user_id, + hint_key, ); if let Some(sni) = sni { @@ -94,34 +94,55 @@ pub(super) fn sticky_hint_record_success_in( &shared.handshake.sticky_user_by_sni_hash, &shared.handshake.sticky_user_by_sni_hash_slots, sni_hint_hash(sni), - user_id, + hint_key, ); } } fn bounded_sticky_hint_upsert( - entries: &DashMap, + entries: &DashMap, slots: &crate::slot_budget::SlotBudget, key: K, - user_id: u32, + hint_key: u64, ) where - K: Eq + Hash, + K: Clone + Eq + Hash, { - match entries.entry(key) { - Entry::Occupied(mut entry) => { - entry.insert(user_id); + if let Some(mut existing) = entries.get_mut(&key) { + *existing = hint_key; + return; + } + + for _ in 0..2 { + if let Some(slot) = slots.try_acquire() { + match entries.entry(key.clone()) { + Entry::Occupied(mut entry) => { + entry.insert(hint_key); + } + Entry::Vacant(entry) => { + entry.insert(hint_key); + slot.commit(); + } + } + return; } - Entry::Vacant(entry) => { - let Some(slot) = slots.try_acquire() else { - return; - }; - entry.insert(user_id); - slot.commit(); + + let Some((victim_key, victim_hint_key)) = entries + .iter() + .next() + .map(|entry| (entry.key().clone(), *entry.value())) + else { + return; + }; + if entries + .remove_if(&victim_key, |_, current| *current == victim_hint_key) + .is_some() + { + slots.release(); } } } -pub(super) fn record_recent_user_success_in(shared: &ProxySharedState, user_id: u32) { +pub(super) fn record_recent_user_success_in(shared: &ProxySharedState, hint_key: u64) { let ring = &shared.handshake.recent_user_ring; if ring.is_empty() { return; @@ -131,7 +152,7 @@ pub(super) fn record_recent_user_success_in(shared: &ProxySharedState, user_id: .recent_user_ring_seq .fetch_add(1, Ordering::Relaxed); let idx = (seq as usize) % ring.len(); - ring[idx].store(user_id.saturating_add(1), Ordering::Relaxed); + ring[idx].store(hint_key, Ordering::Relaxed); } pub(super) fn mark_candidate_if_new( @@ -387,7 +408,7 @@ mod bounded_registry_tests { sticky_hint_record_success_in( shared.as_ref(), peer_ip, - index as u32, + index as u64 | 1, Some(&format!("host-{index}.example")), ); } diff --git a/src/proxy/handshake/auth_probe/testing.rs b/src/proxy/handshake/auth_probe/testing.rs index 9d5c4e8..18329f4 100644 --- a/src/proxy/handshake/auth_probe/testing.rs +++ b/src/proxy/handshake/auth_probe/testing.rs @@ -21,8 +21,10 @@ pub(crate) fn auth_probe_fail_streak_for_testing_in_shared( } pub(crate) fn clear_auth_probe_state_for_testing_in_shared(shared: &ProxySharedState) { + let removed = shared.handshake.auth_probe.len(); + assert_eq!(shared.handshake.auth_probe_slots.used(), removed); shared.handshake.auth_probe.clear(); - shared.handshake.auth_probe_slots.reset_for_testing(); + shared.handshake.auth_probe_slots.release_many(removed); match shared.handshake.auth_probe_saturation.lock() { Ok(mut saturation) => { *saturation = None; @@ -35,6 +37,28 @@ pub(crate) fn clear_auth_probe_state_for_testing_in_shared(shared: &ProxySharedS } } +pub(crate) fn insert_auth_probe_state_for_testing_in_shared( + shared: &ProxySharedState, + peer_ip: IpAddr, + state: AuthProbeState, +) { + let peer_ip = normalize_auth_probe_ip(peer_ip); + let slot = shared + .handshake + .auth_probe_slots + .try_acquire() + .expect("test auth-probe registry capacity must be available"); + match shared.handshake.auth_probe.entry(peer_ip) { + Entry::Occupied(mut entry) => { + entry.insert(state); + } + Entry::Vacant(entry) => { + entry.insert(state); + slot.commit(); + } + } +} + pub(crate) fn auth_probe_state_for_testing_in_shared( shared: &ProxySharedState, ) -> &DashMap { diff --git a/src/proxy/handshake/mtproto.rs b/src/proxy/handshake/mtproto.rs index da9b351..883e311 100644 --- a/src/proxy/handshake/mtproto.rs +++ b/src/proxy/handshake/mtproto.rs @@ -151,10 +151,14 @@ where if let Some(snapshot) = config.runtime_user_auth() { let sticky_ip_hint = sticky_hint_get_by_ip(shared, peer.ip()); let sticky_prefix_hint = sticky_hint_get_by_ip_prefix(shared, peer.ip()); + let sticky_ip_candidates = sticky_ip_hint + .and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key)); + let sticky_prefix_candidates = sticky_prefix_hint + .and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key)); let preferred_user_id = preferred_user.and_then(|user| snapshot.user_id_by_name(user)); let exact_user_id = exact_user.and_then(|user| snapshot.user_id_by_name(user)); - let has_hint = sticky_ip_hint.is_some() - || sticky_prefix_hint.is_some() + let has_hint = sticky_ip_candidates.is_some_and(|ids| !ids.is_empty()) + || sticky_prefix_candidates.is_some_and(|ids| !ids.is_empty()) || preferred_user_id.is_some() || exact_user_id.is_some(); let overload = auth_probe_saturation_is_throttled_in(shared, Instant::now()); @@ -204,9 +208,17 @@ where let mut matched = exact_user_id.is_some_and(|user_id| try_user_id!(user_id)); if exact_user.is_none() - && let Some(user_id) = sticky_ip_hint + && let Some(candidate_ids) = sticky_ip_candidates { - matched = try_user_id!(user_id); + for &user_id in candidate_ids { + if try_user_id!(user_id) { + matched = true; + break; + } + if budget_exhausted { + break; + } + } } if exact_user.is_none() @@ -218,9 +230,17 @@ where if exact_user.is_none() && !matched - && let Some(user_id) = sticky_prefix_hint + && let Some(candidate_ids) = sticky_prefix_candidates { - matched = try_user_id!(user_id); + for &user_id in candidate_ids { + if try_user_id!(user_id) { + matched = true; + break; + } + if budget_exhausted { + break; + } + } } if exact_user.is_none() && !matched && !budget_exhausted { @@ -231,18 +251,22 @@ where .recent_user_ring_seq .load(Ordering::Relaxed); let scan_limit = ring.len().min(RECENT_USER_RING_SCAN_LIMIT); - for offset in 0..scan_limit { + 'recent_hints: for offset in 0..scan_limit { let idx = (next_seq as usize + ring.len() - 1 - offset) % ring.len(); - let encoded_user_id = ring[idx].load(Ordering::Relaxed); - if encoded_user_id == 0 { + let hint_key = ring[idx].load(Ordering::Relaxed); + if hint_key == 0 { continue; } - if try_user_id!(encoded_user_id - 1) { - matched = true; - break; - } - if budget_exhausted { - break; + if let Some(candidate_ids) = snapshot.candidate_ids_by_hint_key(hint_key) { + for &user_id in candidate_ids { + if try_user_id!(user_id) { + matched = true; + break 'recent_hints; + } + if budget_exhausted { + break 'recent_hints; + } + } } } } @@ -357,8 +381,10 @@ where auth_probe_record_success_in(shared, peer.ip()); if let Some(user_id) = matched_user_id { - sticky_hint_record_success_in(shared, peer.ip(), user_id, None); - record_recent_user_success_in(shared, user_id); + if let Some(entry) = snapshot.entry_by_id(user_id) { + sticky_hint_record_success_in(shared, peer.ip(), entry.hint_key, None); + record_recent_user_success_in(shared, entry.hint_key); + } } let max_pending = config.general.crypto_pending_buffer; diff --git a/src/proxy/handshake/tls_handshake.rs b/src/proxy/handshake/tls_handshake.rs index f1838ea..90cd95c 100644 --- a/src/proxy/handshake/tls_handshake.rs +++ b/src/proxy/handshake/tls_handshake.rs @@ -396,8 +396,18 @@ where auth_probe_record_success_in(shared, peer.ip()); if let Some(user_id) = validated_user_id { - sticky_hint_record_success_in(shared, peer.ip(), user_id, client_sni.as_deref()); - record_recent_user_success_in(shared, user_id); + if let Some(entry) = config + .runtime_user_auth() + .and_then(|snapshot| snapshot.entry_by_id(user_id)) + { + sticky_hint_record_success_in( + shared, + peer.ip(), + entry.hint_key, + client_sni.as_deref(), + ); + record_recent_user_success_in(shared, entry.hint_key); + } } HandshakeResult::Success(( diff --git a/src/proxy/handshake/tls_validation.rs b/src/proxy/handshake/tls_validation.rs index 4b83c08..f2d74c7 100644 --- a/src/proxy/handshake/tls_validation.rs +++ b/src/proxy/handshake/tls_validation.rs @@ -40,11 +40,17 @@ pub(super) async fn validate_tls_client( }; let sticky_ip_hint = sticky_hint_get_by_ip(shared, peer.ip()); + let sticky_ip_candidates = sticky_ip_hint + .and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key)); let preferred_user_id = preferred_user_hint.and_then(|user| snapshot.user_id_by_name(user)); let sticky_sni_hint = client_sni .as_deref() .and_then(|sni| sticky_hint_get_by_sni(shared, sni)); + let sticky_sni_candidates = sticky_sni_hint + .and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key)); let sticky_prefix_hint = sticky_hint_get_by_ip_prefix(shared, peer.ip()); + let sticky_prefix_candidates = sticky_prefix_hint + .and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key)); let sni_candidates = client_sni .as_deref() .and_then(|sni| snapshot.sni_candidates(sni)); @@ -52,10 +58,10 @@ pub(super) async fn validate_tls_client( .as_deref() .and_then(|sni| snapshot.sni_initial_candidates(sni)); - let has_hint = sticky_ip_hint.is_some() + let has_hint = sticky_ip_candidates.is_some_and(|ids| !ids.is_empty()) || preferred_user_id.is_some() - || sticky_sni_hint.is_some() - || sticky_prefix_hint.is_some() + || sticky_sni_candidates.is_some_and(|ids| !ids.is_empty()) + || sticky_prefix_candidates.is_some_and(|ids| !ids.is_empty()) || sni_candidates.is_some_and(|ids| !ids.is_empty()) || sni_initial_candidates.is_some_and(|ids| !ids.is_empty()); let overload = auth_probe_saturation_is_throttled_in(shared, Instant::now()); @@ -95,20 +101,44 @@ pub(super) async fn validate_tls_client( } let mut matched = false; - if let Some(user_id) = sticky_ip_hint { - matched = try_user_id!(user_id); + if let Some(candidate_ids) = sticky_ip_candidates { + for &user_id in candidate_ids { + if try_user_id!(user_id) { + matched = true; + break; + } + if budget_exhausted { + break; + } + } } if !matched && let Some(user_id) = preferred_user_id { matched = try_user_id!(user_id); } - if !matched && let Some(user_id) = sticky_sni_hint { - matched = try_user_id!(user_id); + if !matched && let Some(candidate_ids) = sticky_sni_candidates { + for &user_id in candidate_ids { + if try_user_id!(user_id) { + matched = true; + break; + } + if budget_exhausted { + break; + } + } } - if !matched && let Some(user_id) = sticky_prefix_hint { - matched = try_user_id!(user_id); + if !matched && let Some(candidate_ids) = sticky_prefix_candidates { + for &user_id in candidate_ids { + if try_user_id!(user_id) { + matched = true; + break; + } + if budget_exhausted { + break; + } + } } if !matched @@ -149,18 +179,22 @@ pub(super) async fn validate_tls_client( .recent_user_ring_seq .load(Ordering::Relaxed); let scan_limit = ring.len().min(RECENT_USER_RING_SCAN_LIMIT); - for offset in 0..scan_limit { + 'recent_hints: for offset in 0..scan_limit { let idx = (next_seq as usize + ring.len() - 1 - offset) % ring.len(); - let encoded_user_id = ring[idx].load(Ordering::Relaxed); - if encoded_user_id == 0 { + let hint_key = ring[idx].load(Ordering::Relaxed); + if hint_key == 0 { continue; } - if try_user_id!(encoded_user_id - 1) { - matched = true; - break; - } - if budget_exhausted { - break; + if let Some(candidate_ids) = snapshot.candidate_ids_by_hint_key(hint_key) { + for &user_id in candidate_ids { + if try_user_id!(user_id) { + matched = true; + break 'recent_hints; + } + if budget_exhausted { + break 'recent_hints; + } + } } } } diff --git a/src/proxy/shared_state.rs b/src/proxy/shared_state.rs index a3c7e67..2749d7a 100644 --- a/src/proxy/shared_state.rs +++ b/src/proxy/shared_state.rs @@ -1,7 +1,7 @@ use std::collections::hash_map::RandomState; use std::collections::{HashMap, HashSet}; use std::net::{IpAddr, SocketAddr}; -use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, Ordering}; +use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; use std::sync::{Arc, Mutex}; use std::time::Instant; @@ -61,13 +61,13 @@ pub(crate) struct HandshakeSharedState { pub(crate) auth_probe_eviction_hasher: RandomState, pub(crate) invalid_secret_warned: Mutex>, pub(crate) unknown_sni_warn_next_allowed: Mutex>, - pub(crate) sticky_user_by_ip: DashMap, + pub(crate) sticky_user_by_ip: DashMap, pub(crate) sticky_user_by_ip_slots: SlotBudget, - pub(crate) sticky_user_by_ip_prefix: DashMap, + pub(crate) sticky_user_by_ip_prefix: DashMap, pub(crate) sticky_user_by_ip_prefix_slots: SlotBudget, - pub(crate) sticky_user_by_sni_hash: DashMap, + pub(crate) sticky_user_by_sni_hash: DashMap, pub(crate) sticky_user_by_sni_hash_slots: SlotBudget, - pub(crate) recent_user_ring: Box<[AtomicU32]>, + pub(crate) recent_user_ring: Box<[AtomicU64]>, pub(crate) recent_user_ring_seq: AtomicU64, pub(crate) auth_expensive_checks_total: AtomicU64, pub(crate) auth_budget_exhausted_total: AtomicU64, @@ -138,7 +138,7 @@ impl ProxySharedState { 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(|| AtomicU64::new(0)) .take(HANDSHAKE_RECENT_USER_RING_LEN) .collect::>() .into_boxed_slice(), diff --git a/src/proxy/tests/handshake_security_tests.rs b/src/proxy/tests/handshake_security_tests.rs index 8a21e58..010e56f 100644 --- a/src/proxy/tests/handshake_security_tests.rs +++ b/src/proxy/tests/handshake_security_tests.rs @@ -1241,7 +1241,10 @@ async fn tls_runtime_snapshot_updates_sticky_and_recent_hints() { .sticky_user_by_ip .get(&peer.ip()) .map(|entry| *entry), - Some(0), + config + .runtime_user_auth() + .and_then(|snapshot| snapshot.entry_by_id(0)) + .map(|entry| entry.hint_key), "successful runtime-snapshot auth must seed sticky ip cache" ); assert_eq!( @@ -3047,7 +3050,8 @@ async fn valid_tls_is_blocked_by_per_ip_preauth_throttle_without_saturation() { let rng = SecureRandom::new(); let peer: SocketAddr = "198.51.100.103:45103".parse().unwrap(); - auth_probe_state_for_testing_in_shared(shared.as_ref()).insert( + insert_auth_probe_state_for_testing_in_shared( + shared.as_ref(), normalize_auth_probe_ip(peer.ip()), AuthProbeState { fail_streak: AUTH_PROBE_BACKOFF_START_FAILS, @@ -3085,7 +3089,8 @@ async fn saturation_allows_valid_tls_even_when_peer_ip_is_currently_throttled() let peer: SocketAddr = "198.51.100.104:45104".parse().unwrap(); let now = Instant::now(); - auth_probe_state_for_testing_in_shared(shared.as_ref()).insert( + insert_auth_probe_state_for_testing_in_shared( + shared.as_ref(), normalize_auth_probe_ip(peer.ip()), AuthProbeState { fail_streak: AUTH_PROBE_BACKOFF_START_FAILS, @@ -3181,7 +3186,8 @@ async fn saturation_grace_exhaustion_preauth_throttles_repeated_invalid_tls_prob let rng = SecureRandom::new(); let peer: SocketAddr = "198.51.100.205:45205".parse().unwrap(); let now = Instant::now(); - auth_probe_state_for_testing_in_shared(shared.as_ref()).insert( + insert_auth_probe_state_for_testing_in_shared( + shared.as_ref(), normalize_auth_probe_ip(peer.ip()), AuthProbeState { fail_streak: AUTH_PROBE_BACKOFF_START_FAILS + AUTH_PROBE_SATURATION_GRACE_FAILS, @@ -3235,7 +3241,8 @@ async fn saturation_allows_valid_mtproto_even_when_peer_ip_is_currently_throttle let peer: SocketAddr = "198.51.100.106:45106".parse().unwrap(); let now = Instant::now(); - auth_probe_state_for_testing_in_shared(shared.as_ref()).insert( + insert_auth_probe_state_for_testing_in_shared( + shared.as_ref(), normalize_auth_probe_ip(peer.ip()), AuthProbeState { fail_streak: AUTH_PROBE_BACKOFF_START_FAILS, @@ -3328,7 +3335,8 @@ async fn saturation_grace_exhaustion_preauth_throttles_repeated_invalid_mtproto_ let replay_checker = ReplayChecker::new(128, Duration::from_secs(60)); let peer: SocketAddr = "198.51.100.206:45206".parse().unwrap(); let now = Instant::now(); - auth_probe_state_for_testing_in_shared(shared.as_ref()).insert( + insert_auth_probe_state_for_testing_in_shared( + shared.as_ref(), normalize_auth_probe_ip(peer.ip()), AuthProbeState { fail_streak: AUTH_PROBE_BACKOFF_START_FAILS + AUTH_PROBE_SATURATION_GRACE_FAILS, @@ -3378,7 +3386,8 @@ async fn saturation_grace_progression_tls_reaches_cap_then_stops_incrementing() let rng = SecureRandom::new(); let peer: SocketAddr = "198.51.100.207:45207".parse().unwrap(); let now = Instant::now(); - auth_probe_state_for_testing_in_shared(shared.as_ref()).insert( + insert_auth_probe_state_for_testing_in_shared( + shared.as_ref(), normalize_auth_probe_ip(peer.ip()), AuthProbeState { fail_streak: AUTH_PROBE_BACKOFF_START_FAILS, @@ -3461,7 +3470,8 @@ async fn saturation_grace_progression_mtproto_reaches_cap_then_stops_incrementin let replay_checker = ReplayChecker::new(128, Duration::from_secs(60)); let peer: SocketAddr = "198.51.100.208:45208".parse().unwrap(); let now = Instant::now(); - auth_probe_state_for_testing_in_shared(shared.as_ref()).insert( + insert_auth_probe_state_for_testing_in_shared( + shared.as_ref(), normalize_auth_probe_ip(peer.ip()), AuthProbeState { fail_streak: AUTH_PROBE_BACKOFF_START_FAILS, @@ -3545,7 +3555,8 @@ async fn saturation_grace_boundary_still_admits_valid_tls_before_exhaustion() { let rng = SecureRandom::new(); let peer: SocketAddr = "198.51.100.209:45209".parse().unwrap(); let now = Instant::now(); - auth_probe_state_for_testing_in_shared(shared.as_ref()).insert( + insert_auth_probe_state_for_testing_in_shared( + shared.as_ref(), normalize_auth_probe_ip(peer.ip()), AuthProbeState { fail_streak: AUTH_PROBE_BACKOFF_START_FAILS + AUTH_PROBE_SATURATION_GRACE_FAILS - 1, @@ -3599,7 +3610,8 @@ async fn saturation_grace_exhaustion_blocks_valid_tls_until_backoff_expires() { let rng = SecureRandom::new(); let peer: SocketAddr = "198.51.100.210:45210".parse().unwrap(); let now = Instant::now(); - auth_probe_state_for_testing_in_shared(shared.as_ref()).insert( + insert_auth_probe_state_for_testing_in_shared( + shared.as_ref(), normalize_auth_probe_ip(peer.ip()), AuthProbeState { fail_streak: AUTH_PROBE_BACKOFF_START_FAILS + AUTH_PROBE_SATURATION_GRACE_FAILS, @@ -3667,7 +3679,8 @@ async fn saturation_grace_exhaustion_is_shared_across_tls_and_mtproto_for_same_p let rng = SecureRandom::new(); let peer: SocketAddr = "198.51.100.211:45211".parse().unwrap(); let now = Instant::now(); - auth_probe_state_for_testing_in_shared(shared.as_ref()).insert( + insert_auth_probe_state_for_testing_in_shared( + shared.as_ref(), normalize_auth_probe_ip(peer.ip()), AuthProbeState { fail_streak: AUTH_PROBE_BACKOFF_START_FAILS + AUTH_PROBE_SATURATION_GRACE_FAILS, @@ -3735,7 +3748,8 @@ async fn adversarial_same_peer_invalid_tls_storm_does_not_bypass_saturation_grac let rng = Arc::new(SecureRandom::new()); let peer: SocketAddr = "198.51.100.212:45212".parse().unwrap(); let now = Instant::now(); - auth_probe_state_for_testing_in_shared(shared.as_ref()).insert( + insert_auth_probe_state_for_testing_in_shared( + shared.as_ref(), normalize_auth_probe_ip(peer.ip()), AuthProbeState { fail_streak: AUTH_PROBE_BACKOFF_START_FAILS + AUTH_PROBE_SATURATION_GRACE_FAILS, @@ -3801,7 +3815,8 @@ async fn light_fuzz_saturation_grace_tls_invalid_inputs_never_authenticate_or_pa let rng = SecureRandom::new(); let peer: SocketAddr = "198.51.100.213:45213".parse().unwrap(); let now = Instant::now(); - auth_probe_state_for_testing_in_shared(shared.as_ref()).insert( + insert_auth_probe_state_for_testing_in_shared( + shared.as_ref(), normalize_auth_probe_ip(peer.ip()), AuthProbeState { fail_streak: AUTH_PROBE_BACKOFF_START_FAILS + AUTH_PROBE_SATURATION_GRACE_FAILS, @@ -3986,7 +4001,8 @@ async fn expired_saturation_keeps_per_ip_throttle_enforced_for_valid_tls() { let peer: SocketAddr = "198.51.100.110:45110".parse().unwrap(); let now = Instant::now(); - auth_probe_state_for_testing_in_shared(shared.as_ref()).insert( + insert_auth_probe_state_for_testing_in_shared( + shared.as_ref(), normalize_auth_probe_ip(peer.ip()), AuthProbeState { fail_streak: AUTH_PROBE_BACKOFF_START_FAILS, diff --git a/src/quota_state.rs b/src/quota_state.rs index bf3149b..5c65b62 100644 --- a/src/quota_state.rs +++ b/src/quota_state.rs @@ -1,11 +1,9 @@ use std::collections::{BTreeMap, BTreeSet}; -use std::io::Write; use std::path::{Path, PathBuf}; use std::sync::Arc; use std::time::{SystemTime, UNIX_EPOCH}; use serde::{Deserialize, Serialize}; -use tokio::io::AsyncReadExt; use tokio::sync::Mutex; use tracing::{info, warn}; @@ -173,6 +171,19 @@ fn now_epoch_secs() -> u64 { } async fn read_state_file(path: &Path) -> std::io::Result> { + #[cfg(unix)] + let payload = match crate::util::secure_fs::read_regular_limited_async( + path.to_path_buf(), + QUOTA_STATE_MAX_BYTES as usize, + ) + .await + { + Ok(payload) => payload, + Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None), + Err(error) => return Err(error), + }; + #[cfg(not(unix))] + let payload = { let file = match tokio::fs::File::open(path).await { Ok(file) => file, Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None), @@ -194,6 +205,8 @@ async fn read_state_file(path: &Path) -> std::io::Result> "quota state file grew beyond the 16 MiB limit while reading", )); } + payload + }; let state = serde_json::from_slice(&payload).map_err(|error| { std::io::Error::new( std::io::ErrorKind::InvalidData, @@ -217,11 +230,6 @@ async fn wait_for_blocking_io( } fn write_state_file_blocking(path: &Path, state: &QuotaStateFile) -> std::io::Result<()> { - let parent = path - .parent() - .filter(|parent| !parent.as_os_str().is_empty()) - .unwrap_or_else(|| Path::new(".")); - std::fs::create_dir_all(parent)?; let mut payload = serde_json::to_vec_pretty(state)?; payload.push(b'\n'); if payload.len() as u64 > QUOTA_STATE_MAX_BYTES { @@ -231,6 +239,20 @@ fn write_state_file_blocking(path: &Path, state: &QuotaStateFile) -> std::io::Re )); } + #[cfg(unix)] + { + return crate::util::secure_fs::atomic_replace(path, &payload, 0o600); + } + #[cfg(not(unix))] + { + use std::io::Write; + + let parent = path + .parent() + .filter(|parent| !parent.as_os_str().is_empty()) + .unwrap_or_else(|| Path::new(".")); + std::fs::create_dir_all(parent)?; + let mut last_collision = None; for _ in 0..8 { let tmp_path = path.with_extension(format!( @@ -270,6 +292,7 @@ fn write_state_file_blocking(path: &Path, state: &QuotaStateFile) -> std::io::Re "failed to allocate a unique quota checkpoint temporary file", ) })) + } } fn quota_user_state(quota: UserQuotaSnapshot) -> QuotaUserState { diff --git a/src/slot_budget.rs b/src/slot_budget.rs index 005fb56..cf1784c 100644 --- a/src/slot_budget.rs +++ b/src/slot_budget.rs @@ -54,10 +54,7 @@ impl SlotBudget { .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. @@ -65,10 +62,6 @@ impl SlotBudget { 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. diff --git a/src/tls_front/cache/disk.rs b/src/tls_front/cache/disk.rs index 3a73263..c516ee0 100644 --- a/src/tls_front/cache/disk.rs +++ b/src/tls_front/cache/disk.rs @@ -1,27 +1,26 @@ use std::path::Path; -use tokio::io::AsyncReadExt; - use super::*; pub(super) async fn read_disk_entry_bounded(path: &Path) -> std::io::Result> { - let file = tokio::fs::File::open(path).await?; - if file.metadata().await?.len() > TLS_FRONT_DISK_ENTRY_MAX_BYTES { - return Err(std::io::Error::new( - std::io::ErrorKind::InvalidData, - "TLS cache entry exceeds the 1 MiB limit", - )); + #[cfg(unix)] + { + crate::util::secure_fs::read_regular_limited_async( + path.to_path_buf(), + TLS_FRONT_DISK_ENTRY_MAX_BYTES as usize, + ) + .await } - let mut bytes = Vec::new(); - file.take(TLS_FRONT_DISK_ENTRY_MAX_BYTES.saturating_add(1)) - .read_to_end(&mut bytes) - .await?; - if bytes.len() as u64 > TLS_FRONT_DISK_ENTRY_MAX_BYTES { - return Err(std::io::Error::new( - std::io::ErrorKind::InvalidData, - "TLS cache entry grew beyond the 1 MiB limit while reading", - )); + #[cfg(not(unix))] + { + let bytes = tokio::fs::read(path).await?; + if bytes.len() as u64 > TLS_FRONT_DISK_ENTRY_MAX_BYTES { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + "TLS cache entry exceeds the 1 MiB limit", + )); + } + Ok(bytes) } - Ok(bytes) } pub(super) fn cert_info_matches_domain(cached: &CachedTlsData) -> bool { diff --git a/src/tls_front/cache/runtime.rs b/src/tls_front/cache/runtime.rs index 91132e4..fca4039 100644 --- a/src/tls_front/cache/runtime.rs +++ b/src/tls_front/cache/runtime.rs @@ -231,9 +231,6 @@ impl TlsFrontCache { pub async fn load_from_disk(&self) { let path = self.disk_path.clone(); - if tokio::fs::create_dir_all(&path).await.is_err() { - return; - } let mut loaded = 0usize; for name in &self.disk_entry_names { let entry_path = path.join(name); @@ -297,9 +294,6 @@ impl TlsFrontCache { } async fn persist(&self, domain: &str, data: &CachedTlsData) { - if tokio::fs::create_dir_all(&self.disk_path).await.is_err() { - return; - } let fname = format!("{}.json", domain.replace(['/', '\\'], "_")); let path = self.disk_path.join(fname); if let Ok(json) = serde_json::to_vec_pretty(data) { @@ -311,7 +305,9 @@ impl TlsFrontCache { ); return; } - // best-effort write + #[cfg(unix)] + let _ = crate::util::secure_fs::atomic_replace_async(path, json, 0o600).await; + #[cfg(not(unix))] let _ = tokio::fs::write(path, json).await; } } diff --git a/src/transport/middle_proxy/config_updater.rs b/src/transport/middle_proxy/config_updater.rs index 50278b5..351fa61 100644 --- a/src/transport/middle_proxy/config_updater.rs +++ b/src/transport/middle_proxy/config_updater.rs @@ -12,6 +12,8 @@ use tracing::{debug, info, warn}; use crate::config::ProxyConfig; use crate::error::Result; use crate::transport::UpstreamManager; +#[cfg(unix)] +use crate::util::secure_fs::{atomic_replace_async, read_regular_limited_async}; use super::MePool; use super::http_fetch::{HTTPS_RESPONSE_BODY_MAX_BYTES, https_get}; @@ -73,25 +75,36 @@ pub fn parse_proxy_config_text(text: &str, http_status: u16) -> ProxyConfigData } pub async fn load_proxy_config_cache(path: &str) -> Result { - let text = tokio::fs::read_to_string(path).await.map_err(|e| { + #[cfg(unix)] + let bytes = read_regular_limited_async( + Path::new(path).to_path_buf(), + HTTPS_RESPONSE_BODY_MAX_BYTES, + ) + .await; + #[cfg(not(unix))] + let bytes = tokio::fs::read(path).await; + let bytes = bytes.map_err(|e| { crate::error::ProxyError::Proxy(format!("read proxy-config cache '{path}' failed: {e}")) })?; + let text = String::from_utf8(bytes).map_err(|e| { + crate::error::ProxyError::Proxy(format!( + "proxy-config cache '{path}' is not valid UTF-8: {e}" + )) + })?; Ok(parse_proxy_config_text(&text, 200)) } pub async fn save_proxy_config_cache(path: &str, raw_text: &str) -> Result<()> { - if let Some(parent) = Path::new(path).parent() - && !parent.as_os_str().is_empty() - { - tokio::fs::create_dir_all(parent).await.map_err(|e| { - crate::error::ProxyError::Proxy(format!( - "create proxy-config cache dir '{}' failed: {e}", - parent.display() - )) - })?; - } - - tokio::fs::write(path, raw_text).await.map_err(|e| { + #[cfg(unix)] + let write = atomic_replace_async( + Path::new(path).to_path_buf(), + raw_text.as_bytes().to_vec(), + 0o640, + ) + .await; + #[cfg(not(unix))] + let write = tokio::fs::write(path, raw_text).await; + write.map_err(|e| { crate::error::ProxyError::Proxy(format!("write proxy-config cache '{path}' failed: {e}")) })?; Ok(()) diff --git a/src/transport/middle_proxy/secret.rs b/src/transport/middle_proxy/secret.rs index 6bae755..442ff50 100644 --- a/src/transport/middle_proxy/secret.rs +++ b/src/transport/middle_proxy/secret.rs @@ -1,4 +1,5 @@ use httpdate; +use std::path::PathBuf; use std::sync::Arc; use std::time::SystemTime; use tracing::{debug, info, warn}; @@ -7,6 +8,8 @@ use super::http_fetch::https_get; use super::selftest::record_timeskew_sample; use crate::error::{ProxyError, Result}; use crate::transport::UpstreamManager; +#[cfg(unix)] +use crate::util::secure_fs::{atomic_replace_async, read_regular_limited_async}; pub const PROXY_SECRET_MIN_LEN: usize = 32; @@ -58,7 +61,12 @@ pub async fn fetch_proxy_secret_with_upstream( match download_proxy_secret_with_max_len_via_upstream(max_len, upstream, proxy_secret_url).await { Ok(data) => { - if let Err(e) = tokio::fs::write(cache, &data).await { + #[cfg(unix)] + let cache_result = + atomic_replace_async(PathBuf::from(cache), data.clone(), 0o600).await; + #[cfg(not(unix))] + let cache_result = tokio::fs::write(cache, &data).await; + if let Err(e) = cache_result { warn!(error = %e, "Failed to cache proxy-secret (non-fatal)"); } else { debug!(path = cache, len = data.len(), "Cached proxy-secret"); @@ -72,7 +80,11 @@ pub async fn fetch_proxy_secret_with_upstream( } // 2) Fallback to cache/file regardless of age; require len in bounds. - match tokio::fs::read(cache).await { + #[cfg(unix)] + let cached = read_regular_limited_async(PathBuf::from(cache), max_len).await; + #[cfg(not(unix))] + let cached = tokio::fs::read(cache).await; + match cached { Ok(data) if validate_proxy_secret_len(data.len(), max_len).is_ok() => { let age_hours = tokio::fs::metadata(cache) .await diff --git a/src/util/mod.rs b/src/util/mod.rs index 7698354..cf6c160 100644 --- a/src/util/mod.rs +++ b/src/util/mod.rs @@ -1,6 +1,8 @@ //! Utils pub mod ip; +#[cfg(unix)] +pub mod secure_fs; pub mod time; #[cfg(unix)] pub mod trusted_command; diff --git a/src/util/secure_fs/mod.rs b/src/util/secure_fs/mod.rs new file mode 100644 index 0000000..1806107 --- /dev/null +++ b/src/util/secure_fs/mod.rs @@ -0,0 +1,20 @@ +//! Descriptor-anchored filesystem operations for privileged runtime paths. +//! +//! Submodules: +//! - `path`: symlink-free directory traversal and anchored path ownership +//! - `write`: regular-file opening and durable atomic replacement + +mod path; +mod write; + +pub(crate) use path::{ + AnchoredPath, chdir_nofollow_or_create, open_dir_nofollow, + open_dir_nofollow_or_create, +}; +pub(crate) use write::{ + atomic_replace, atomic_replace_async, open_append_regular, open_append_regular_at, + read_regular_limited, read_regular_limited_async, +}; + +#[cfg(test)] +mod tests; diff --git a/src/util/secure_fs/path.rs b/src/util/secure_fs/path.rs new file mode 100644 index 0000000..866fc78 --- /dev/null +++ b/src/util/secure_fs/path.rs @@ -0,0 +1,159 @@ +use std::ffi::{OsStr, OsString}; +use std::io; +use std::os::fd::OwnedFd; +use std::os::unix::fs::{MetadataExt, PermissionsExt}; +use std::path::{Component, Path}; + +use nix::fcntl::{OFlag, open, openat}; +use nix::sys::stat::{Mode, mkdirat}; + +const DIRECTORY_FLAGS: OFlag = OFlag::O_RDONLY + .union(OFlag::O_DIRECTORY) + .union(OFlag::O_NOFOLLOW) + .union(OFlag::O_CLOEXEC); + +/// An immutable parent-directory descriptor paired with one final path component. +pub(crate) struct AnchoredPath { + parent: OwnedFd, + name: OsString, +} + +impl AnchoredPath { + /// Opens every parent component without following symbolic links. + pub(crate) fn open(path: &Path) -> io::Result { + Self::open_with_parent_creation(path, None) + } + + /// Creates missing parent directories while traversing without symbolic links. + pub(crate) fn open_creating_parents(path: &Path, mode: u32) -> io::Result { + Self::open_with_parent_creation(path, Some(mode)) + } + + /// Opens a parent chain that cannot be renamed by group or world users. + pub(crate) fn open_trusted_parent(path: &Path) -> io::Result { + let name = path + .file_name() + .filter(|name| !name.is_empty()) + .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "path has no file name"))? + .to_os_string(); + let parent_path = path.parent().unwrap_or_else(|| Path::new(".")); + let parent = open_trusted_dir_nofollow(parent_path)?; + Ok(Self { parent, name }) + } + + fn open_with_parent_creation(path: &Path, create_mode: Option) -> io::Result { + let name = path + .file_name() + .filter(|name| !name.is_empty()) + .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "path has no file name"))? + .to_os_string(); + let parent = path.parent().unwrap_or_else(|| Path::new(".")); + let parent = match create_mode { + Some(mode) => open_or_create_dir_nofollow(parent, mode)?, + None => open_dir_nofollow(parent)?, + }; + Ok(Self { parent, name }) + } + + /// Returns the anchored parent descriptor. + pub(crate) fn parent(&self) -> &OwnedFd { + &self.parent + } + + /// Returns the single final component resolved relative to the parent descriptor. + pub(crate) fn name(&self) -> &OsStr { + &self.name + } +} + +/// Opens a directory by walking every component with `O_NOFOLLOW`. +pub(crate) fn open_dir_nofollow(path: &Path) -> io::Result { + open_dir_components(path, None, false) +} + +fn open_or_create_dir_nofollow(path: &Path, mode: u32) -> io::Result { + open_dir_components(path, Some(mode), false) +} + +/// Opens a directory after securely creating any missing components. +pub(crate) fn open_dir_nofollow_or_create(path: &Path, mode: u32) -> io::Result { + open_or_create_dir_nofollow(path, mode) +} + +/// Opens a directory only when its entire path is owned by root or the effective user. +fn open_trusted_dir_nofollow(path: &Path) -> io::Result { + open_dir_components(path, None, true) +} + +fn open_dir_components( + path: &Path, + create_mode: Option, + require_trusted: bool, +) -> io::Result { + let start = if path.is_absolute() { + Path::new("/") + } else { + Path::new(".") + }; + let mut current = open(start, DIRECTORY_FLAGS, Mode::empty()).map_err(errno_to_io)?; + if require_trusted { + validate_trusted_directory(¤t)?; + } + + for component in path.components() { + let name = match component { + Component::RootDir | Component::CurDir => continue, + Component::Normal(name) => name, + Component::ParentDir => OsStr::new(".."), + Component::Prefix(_) => { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "unsupported path prefix", + )); + } + }; + let next = match openat(¤t, name, DIRECTORY_FLAGS, Mode::empty()) { + Ok(descriptor) => descriptor, + Err(nix::errno::Errno::ENOENT) if create_mode.is_some() => { + let mode = Mode::from_bits_truncate(create_mode.unwrap_or(0o750)); + match mkdirat(¤t, name, mode) { + Ok(()) | Err(nix::errno::Errno::EEXIST) => {} + Err(error) => return Err(errno_to_io(error)), + } + openat(¤t, name, DIRECTORY_FLAGS, Mode::empty()).map_err(errno_to_io)? + } + Err(error) => return Err(errno_to_io(error)), + }; + if require_trusted { + validate_trusted_directory(&next)?; + } + current = next; + } + Ok(current) +} + +fn validate_trusted_directory(descriptor: &OwnedFd) -> io::Result<()> { + let file = std::fs::File::from(descriptor.try_clone()?); + let metadata = file.metadata()?; + let effective_uid = nix::unistd::Uid::effective().as_raw(); + if !metadata.is_dir() + || (metadata.uid() != 0 && metadata.uid() != effective_uid) + || metadata.permissions().mode() & 0o022 != 0 + { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "directory path is not owned and protected from group/world writes", + )); + } + Ok(()) +} + +/// Creates missing components and changes cwd to the exact opened directory inode. +pub(crate) fn chdir_nofollow_or_create(path: &Path, mode: u32) -> io::Result<()> { + let descriptor = open_or_create_dir_nofollow(path, mode)?; + nix::unistd::fchdir(&descriptor).map_err(errno_to_io) +} + +pub(super) fn errno_to_io(error: nix::errno::Errno) -> io::Error { + io::Error::from_raw_os_error(error as i32) +} diff --git a/src/util/secure_fs/tests.rs b/src/util/secure_fs/tests.rs new file mode 100644 index 0000000..c0b80fc --- /dev/null +++ b/src/util/secure_fs/tests.rs @@ -0,0 +1,78 @@ +use std::os::unix::fs::{PermissionsExt, symlink}; + +use super::path::AnchoredPath; +use super::write::atomic_replace_after_anchor; +use super::*; + +#[test] +fn directory_walk_rejects_intermediate_symlink() { + let directory = tempfile::tempdir().unwrap(); + let real = directory.path().join("real"); + let link = directory.path().join("link"); + std::fs::create_dir(&real).unwrap(); + symlink(&real, &link).unwrap(); + + assert!(open_dir_nofollow(&link).is_err()); +} + +#[test] +fn atomic_replace_does_not_follow_final_symlink() { + let directory = tempfile::tempdir().unwrap(); + let sentinel = directory.path().join("sentinel"); + let target = directory.path().join("target"); + std::fs::write(&sentinel, b"preserve").unwrap(); + symlink(&sentinel, &target).unwrap(); + + atomic_replace(&target, b"replacement", 0o600).unwrap(); + + assert_eq!(std::fs::read(&sentinel).unwrap(), b"preserve"); + assert_eq!(std::fs::read(&target).unwrap(), b"replacement"); +} + +#[test] +fn append_open_rejects_final_symlink() { + let directory = tempfile::tempdir().unwrap(); + let sentinel = directory.path().join("sentinel"); + let target = directory.path().join("target"); + std::fs::write(&sentinel, b"preserve").unwrap(); + symlink(&sentinel, &target).unwrap(); + + assert!(open_append_regular(&target, 0o640).is_err()); + assert_eq!(std::fs::read(&sentinel).unwrap(), b"preserve"); +} + +#[test] +fn trusted_parent_rejects_group_writable_directory() { + let current = std::env::current_dir().unwrap(); + let directory = tempfile::Builder::new() + .prefix("telemt-insecure-parent-") + .tempdir_in(current) + .unwrap(); + std::fs::set_permissions(directory.path(), std::fs::Permissions::from_mode(0o770)).unwrap(); + + assert!(AnchoredPath::open_trusted_parent(&directory.path().join("listener.sock")).is_err()); +} + +#[test] +fn anchored_replace_survives_parent_path_substitution() { + let directory = tempfile::tempdir().unwrap(); + let original = directory.path().join("original"); + let moved = directory.path().join("moved"); + let redirect = directory.path().join("redirect"); + std::fs::create_dir(&original).unwrap(); + std::fs::create_dir(&redirect).unwrap(); + let target = original.join("state"); + let anchored = AnchoredPath::open(&target).unwrap(); + std::fs::rename(&original, &moved).unwrap(); + symlink(&redirect, &original).unwrap(); + + atomic_replace_after_anchor(&anchored, b"anchored", 0o600).unwrap(); + + assert_eq!(std::fs::read(moved.join("state")).unwrap(), b"anchored"); + assert!(!redirect.join("state").exists()); + let mode = std::fs::metadata(moved.join("state")) + .unwrap() + .permissions() + .mode(); + assert_eq!(mode & 0o777, 0o600); +} diff --git a/src/util/secure_fs/write.rs b/src/util/secure_fs/write.rs new file mode 100644 index 0000000..c6ffa67 --- /dev/null +++ b/src/util/secure_fs/write.rs @@ -0,0 +1,220 @@ +use std::io::{self, Read, Write}; +use std::ffi::OsStr; +use std::os::fd::{AsFd, OwnedFd}; +use std::path::Path; + +use nix::fcntl::{OFlag, openat, renameat}; +use nix::sys::stat::Mode; +use nix::unistd::{UnlinkatFlags, fsync, unlinkat}; + +use super::path::{AnchoredPath, errno_to_io}; + +fn open_regular_at( + anchored: &AnchoredPath, + flags: OFlag, + mode: u32, +) -> io::Result { + let descriptor = openat( + anchored.parent(), + anchored.name(), + flags | OFlag::O_NOFOLLOW | OFlag::O_CLOEXEC, + Mode::from_bits_truncate(mode), + ) + .map_err(errno_to_io)?; + let file = std::fs::File::from(descriptor); + let metadata = file.metadata()?; + if !metadata.is_file() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "target must be a regular file", + )); + } + use std::os::unix::fs::MetadataExt; + if metadata.nlink() != 1 { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "target must have exactly one directory entry", + )); + } + Ok(file) +} + +/// Reads a regular file through an anchored parent with an allocation bound. +pub(crate) fn read_regular_limited(path: &Path, max_bytes: usize) -> io::Result> { + let anchored = AnchoredPath::open(path)?; + let mut file = open_regular_at( + &anchored, + OFlag::O_RDONLY | OFlag::O_NONBLOCK, + Mode::empty().bits(), + )?; + let before = file.metadata()?; + if before.len() > max_bytes as u64 { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "file exceeds configured size limit", + )); + } + let mut bytes = Vec::with_capacity(before.len() as usize); + Read::take(&mut file, max_bytes.saturating_add(1) as u64) + .read_to_end(&mut bytes)?; + if bytes.len() > max_bytes { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "file exceeds configured size limit", + )); + } + let after = file.metadata()?; + if !same_file_version(&before, &after) { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "file changed while it was read", + )); + } + Ok(bytes) +} + +/// Reads a bounded regular file on the blocking pool. +pub(crate) async fn read_regular_limited_async( + path: std::path::PathBuf, + max_bytes: usize, +) -> io::Result> { + tokio::task::spawn_blocking(move || read_regular_limited(&path, max_bytes)) + .await + .map_err(|error| io::Error::other(format!("secure reader task failed: {error}")))? +} + +fn same_file_version(left: &std::fs::Metadata, right: &std::fs::Metadata) -> bool { + use std::os::unix::fs::MetadataExt; + + left.dev() == right.dev() + && left.ino() == right.ino() + && left.len() == right.len() + && left.mtime() == right.mtime() + && left.mtime_nsec() == right.mtime_nsec() + && left.ctime() == right.ctime() + && left.ctime_nsec() == right.ctime_nsec() +} + +/// Opens an append-only regular file without following path components or hard links. +pub(crate) fn open_append_regular(path: &Path, mode: u32) -> io::Result { + let anchored = AnchoredPath::open_creating_parents(path, 0o750)?; + open_regular_at( + &anchored, + OFlag::O_WRONLY | OFlag::O_APPEND | OFlag::O_CREAT, + mode, + ) +} + +/// Opens one append-only file relative to an already anchored directory. +pub(crate) fn open_append_regular_at( + parent: Fd, + name: &OsStr, + mode: u32, +) -> io::Result { + if Path::new(name).components().count() != 1 { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "anchored file name must contain one component", + )); + } + let descriptor = openat( + parent, + name, + OFlag::O_WRONLY + | OFlag::O_APPEND + | OFlag::O_CREAT + | OFlag::O_NOFOLLOW + | OFlag::O_CLOEXEC, + Mode::from_bits_truncate(mode), + ) + .map_err(errno_to_io)?; + let file = std::fs::File::from(descriptor); + let metadata = file.metadata()?; + if !metadata.is_file() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "target must be a regular file", + )); + } + use std::os::unix::fs::MetadataExt; + if metadata.nlink() != 1 { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "target must have exactly one directory entry", + )); + } + Ok(file) +} + +/// Durably replaces a file through a same-directory descriptor-anchored rename. +pub(crate) fn atomic_replace(path: &Path, contents: &[u8], mode: u32) -> io::Result<()> { + let anchored = AnchoredPath::open_creating_parents(path, 0o750)?; + atomic_replace_anchored(&anchored, contents, mode) +} + +/// Durably replaces a file on the blocking pool. +pub(crate) async fn atomic_replace_async( + path: std::path::PathBuf, + contents: Vec, + mode: u32, +) -> io::Result<()> { + tokio::task::spawn_blocking(move || atomic_replace(&path, &contents, mode)) + .await + .map_err(|error| io::Error::other(format!("secure writer task failed: {error}")))? +} + +fn atomic_replace_anchored( + anchored: &AnchoredPath, + contents: &[u8], + mode: u32, +) -> io::Result<()> { + let temp_name = format!(".telemt.tmp-{}", rand::random::()); + let descriptor = openat( + anchored.parent(), + temp_name.as_str(), + OFlag::O_WRONLY + | OFlag::O_CREAT + | OFlag::O_EXCL + | OFlag::O_NOFOLLOW + | OFlag::O_CLOEXEC, + Mode::from_bits_truncate(mode), + ) + .map_err(errno_to_io)?; + let result = write_and_publish(descriptor, anchored, &temp_name, contents); + if result.is_err() { + let _ = unlinkat( + anchored.parent(), + temp_name.as_str(), + UnlinkatFlags::NoRemoveDir, + ); + } + result +} + +fn write_and_publish( + descriptor: OwnedFd, + anchored: &AnchoredPath, + temp_name: &str, + contents: &[u8], +) -> io::Result<()> { + let mut file = std::fs::File::from(descriptor); + file.write_all(contents)?; + file.sync_all()?; + renameat( + anchored.parent(), + temp_name, + anchored.parent(), + anchored.name(), + ) + .map_err(errno_to_io)?; + fsync(anchored.parent()).map_err(errno_to_io) +} + +#[cfg(test)] +pub(super) fn atomic_replace_after_anchor( + anchored: &AnchoredPath, + contents: &[u8], + mode: u32, +) -> io::Result<()> { + atomic_replace_anchored(anchored, contents, mode) +}