Descriptor-anchored Secured Filesystem Operations

This commit is contained in:
Alexey
2026-09-17 22:45:48 +03:00
parent 9a683d8b3d
commit 02f66c542e
41 changed files with 2997 additions and 1759 deletions
+225 -62
View File
@@ -1,12 +1,24 @@
use std::fs::File;
use std::io::{Read, Write}; use std::io::{Read, Write};
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
#[cfg(unix)] #[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 super::compute_source_revision;
use crate::api::model::ApiFailure; use crate::api::model::ApiFailure;
use crate::config::ProxyConfig; use crate::config::ProxyConfig;
#[cfg(unix)]
use crate::util::secure_fs::AnchoredPath;
const MAX_CONFIG_SOURCE_BYTES: u64 = 8 * 1024 * 1024;
enum AtomicWriteError { enum AtomicWriteError {
Conflict, Conflict,
@@ -19,12 +31,53 @@ struct ExistingTarget {
metadata: std::fs::Metadata, metadata: std::fs::Metadata,
} }
struct ConfigWriteLock {
#[cfg(unix)]
_file: Flock<File>,
}
impl ConfigWriteLock {
fn acquire(path: &Path) -> std::io::Result<Self> {
#[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. /// Replaces one config source through a durable same-directory rename.
pub(in crate::api) async fn write_atomic( pub(in crate::api) async fn write_atomic(
path: PathBuf, path: PathBuf,
contents: String, contents: String,
) -> Result<(), ApiFailure> { ) -> 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 .await
.map_err(|error| ApiFailure::internal(format!("failed to join writer: {error}")))? .map_err(|error| ApiFailure::internal(format!("failed to join writer: {error}")))?
.map_err(|error| ApiFailure::internal(format!("failed to write config: {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, contents: String,
) -> Result<(), ApiFailure> { ) -> Result<(), ApiFailure> {
tokio::task::spawn_blocking(move || { 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) let graph = ProxyConfig::read_source_graph(&config_path)
.map_err(|error| AtomicWriteError::ReadGraph(error.to_string()))?; .map_err(|error| AtomicWriteError::ReadGraph(error.to_string()))?;
if compute_source_revision(&graph) != expected_revision { 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<Option<ExistingTarget>> {
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<Option<ExistingTarget>> { fn open_existing_target(path: &Path) -> std::io::Result<Option<ExistingTarget>> {
let mut options = std::fs::OpenOptions::new(); let mut file = match File::open(path) {
options.read(true);
#[cfg(unix)]
options.custom_flags(libc::O_CLOEXEC | libc::O_NOFOLLOW);
let mut file = match options.open(path) {
Ok(file) => file, Ok(file) => file,
Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None), Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None),
Err(error) => return Err(error), Err(error) => return Err(error),
}; };
let metadata = file.metadata()?; let metadata = file.metadata()?;
if !metadata.is_file() { if !metadata.is_file() || metadata.len() > MAX_CONFIG_SOURCE_BYTES {
return Err(std::io::Error::new( return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput, std::io::ErrorKind::InvalidInput,
"config target must be a regular file", "config target must be a bounded regular file",
)); ));
} }
let mut contents = String::new(); 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::<u64>()
);
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( fn write_atomic_sync(
path: &Path, path: &Path,
expected_contents: Option<&str>, expected_contents: Option<&str>,
@@ -114,72 +299,50 @@ fn write_atomic_sync(
let parent = path.parent().unwrap_or_else(|| Path::new(".")); let parent = path.parent().unwrap_or_else(|| Path::new("."));
std::fs::create_dir_all(parent)?; std::fs::create_dir_all(parent)?;
let existing = open_existing_target(path)?; let existing = open_existing_target(path)?;
validate_expected_contents(existing.as_ref(), expected_contents)?;
let temp = parent.join(format!(".telemt.tmp-{}", rand::random::<u64>()));
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| { if expected_contents.is_some_and(|expected| {
existing existing.is_none_or(|target| target.contents != expected)
.as_ref()
.is_none_or(|target| target.contents != expected)
}) { }) {
return Err(std::io::Error::new( return Err(std::io::Error::new(
std::io::ErrorKind::AlreadyExists, std::io::ErrorKind::AlreadyExists,
"config source changed before persistence", "config source changed before persistence",
)); ));
} }
Ok(())
}
let tmp_name = format!( fn target_unchanged(
".{}.tmp-{}", existing: Option<&ExistingTarget>,
path.file_name() current: Option<&ExistingTarget>,
.and_then(|name| name.to_str()) ) -> bool {
.unwrap_or("config.toml"), match (existing, current) {
rand::random::<u64>()
);
let tmp_path = parent.join(tmp_name);
let write_result = (|| {
let mut options = std::fs::OpenOptions::new();
options.create_new(true).write(true);
#[cfg(unix)]
options.mode(0o600);
let mut file = options.open(&tmp_path)?;
#[cfg(unix)]
if let Some(existing) = existing.as_ref() {
use nix::unistd::{Gid, Uid, fchown};
fchown(
&file,
Some(Uid::from_raw(existing.metadata.uid())),
Some(Gid::from_raw(existing.metadata.gid())),
)
.map_err(|error| std::io::Error::from_raw_os_error(error as i32))?;
file.set_permissions(std::fs::Permissions::from_mode(
existing.metadata.mode() & 0o7777,
))?;
}
file.write_all(contents.as_bytes())?;
file.sync_all()?;
let current = open_existing_target(path)?;
let target_unchanged = match (&existing, &current) {
(Some(expected), Some(current)) => { (Some(expected), Some(current)) => {
same_target(&expected.metadata, &current.metadata) same_target(&expected.metadata, &current.metadata)
&& expected.contents == current.contents && expected.contents == current.contents
} }
(None, None) => true, (None, None) => true,
_ => false, _ => 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(); #[cfg(unix)]
} fn errno_to_io(error: nix::errno::Errno) -> std::io::Error {
Ok(()) std::io::Error::from_raw_os_error(error as i32)
})();
if write_result.is_err() {
let _ = std::fs::remove_file(&tmp_path);
}
write_result
} }
+42
View File
@@ -314,6 +314,48 @@ async fn atomic_write_preserves_existing_file_mode() {
assert_eq!(after.gid(), before.gid()); 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] #[tokio::test]
async fn access_mutation_rejects_sections_with_different_source_owners() { async fn access_mutation_rejects_sections_with_different_source_owners() {
let dir = tempfile::tempdir().unwrap(); let dir = tempfile::tempdir().unwrap();
+23 -17
View File
@@ -1,4 +1,3 @@
use std::fs;
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use std::process::Command; use std::process::Command;
@@ -114,10 +113,9 @@ pub fn run_init(opts: InitOptions) -> Result<(), Box<dyn std::error::Error>> {
eprintln!("[+] Port: {}", opts.port); eprintln!("[+] Port: {}", opts.port);
eprintln!("[+] Domain: {}", opts.domain); eprintln!("[+] Domain: {}", opts.domain);
fs::create_dir_all(&opts.config_dir)?;
let config_path = opts.config_dir.join("config.toml"); let config_path = opts.config_dir.join("config.toml");
let config_content = generate_config(&opts.username, &secret, opts.port, &opts.domain); 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()); eprintln!("[+] Config written to {}", config_path.display());
let exe_path = let exe_path =
@@ -135,22 +133,15 @@ pub fn run_init(opts: InitOptions) -> Result<(), Box<dyn std::error::Error>> {
let service_path = service::service_file_path(init_system); let service_path = service::service_file_path(init_system);
let service_content = service::generate_service_file(init_system, &service_opts); let service_content = service::generate_service_file(init_system, &service_opts);
if let Some(parent) = Path::new(service_path).parent() { let service_mode = if init_system == InitSystem::OpenRC || init_system == InitSystem::FreeBSDRc
let _ = fs::create_dir_all(parent); {
} 0o755
} else {
match fs::write(service_path, &service_content) { 0o644
};
match write_init_file(Path::new(service_path), &service_content, service_mode) {
Ok(()) => { Ok(()) => {
eprintln!("[+] Service file written to {}", service_path); 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) => { Err(e) => {
eprintln!("[!] Cannot write service file (run as root?): {}", e); eprintln!("[!] Cannot write service file (run as root?): {}", e);
@@ -226,6 +217,21 @@ pub fn run_init(opts: InitOptions) -> Result<(), Box<dyn std::error::Error>> {
Ok(()) 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 { fn generate_secret() -> String {
let mut rng = rand::rng(); let mut rng = rand::rng();
let bytes: Vec<u8> = (0..16).map(|_| rng.random::<u8>()).collect(); let bytes: Vec<u8> = (0..16).map(|_| rng.random::<u8>()).collect();
+33 -11
View File
@@ -66,9 +66,9 @@ const MAX_API_REQUEST_BODY_LIMIT_BYTES: usize = 1024 * 1024;
pub(crate) struct LoadedConfig { pub(crate) struct LoadedConfig {
/// Validated and normalized effective configuration. /// Validated and normalized effective configuration.
pub(crate) config: ProxyConfig, 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<PathBuf>, pub(crate) source_files: Vec<PathBuf>,
/// Raw source bytes keyed by canonical source path. /// Raw source bytes keyed by normalized absolute source path.
pub(crate) source_contents: BTreeMap<PathBuf, String>, pub(crate) source_contents: BTreeMap<PathBuf, String>,
/// Legacy hash of the include-expanded rendered snapshot. /// Legacy hash of the include-expanded rendered snapshot.
pub(crate) rendered_hash: u64, pub(crate) rendered_hash: u64,
@@ -77,7 +77,7 @@ pub(crate) struct LoadedConfig {
/// Raw recursive source graph captured before typed deserialization. /// Raw recursive source graph captured before typed deserialization.
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub(crate) struct ConfigSourceGraph { 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<PathBuf, String>, pub(crate) source_contents: BTreeMap<PathBuf, String>,
/// Include-expanded TOML used for typed deserialization. /// Include-expanded TOML used for typed deserialization.
pub(crate) rendered: String, pub(crate) rendered: String,
@@ -177,20 +177,42 @@ impl ProxyConfig {
source_overrides: &BTreeMap<PathBuf, String>, source_overrides: &BTreeMap<PathBuf, String>,
) -> Result<ConfigSourceGraph> { ) -> Result<ConfigSourceGraph> {
let path = path.as_ref(); let path = path.as_ref();
let initial_path = normalize_config_path(path); let mut previous = Self::capture_source_graph(path, source_overrides)?;
let (normalized_path, content) = if let Some(content) = source_overrides.get(&initial_path) { for _ in 0..2 {
(initial_path, content.clone()) let current = Self::capture_source_graph(path, source_overrides)?;
} else { if current.source_contents == previous.source_contents
read_config_source(path)? && current.rendered == previous.rendered
}; {
let base_dir = path.parent().unwrap_or(Path::new(".")); 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<PathBuf, String>,
) -> Result<ConfigSourceGraph> {
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(); let mut source_files = BTreeSet::new();
source_files.insert(normalized_path.clone()); source_files.insert(normalized_path.clone());
let mut source_contents = BTreeMap::new(); let mut source_contents = BTreeMap::new();
source_contents.insert(normalized_path, content.clone()); source_contents.insert(normalized_path, content.clone());
let processed = preprocess_includes( let processed = preprocess_includes(
&content, &content,
base_dir, &base_dir,
0, 0,
&mut source_files, &mut source_files,
&mut source_contents, &mut source_contents,
+35 -76
View File
@@ -1,23 +1,30 @@
use std::collections::{BTreeMap, BTreeSet}; use std::collections::{BTreeMap, BTreeSet};
use std::hash::{DefaultHasher, Hash, Hasher}; use std::hash::{DefaultHasher, Hash, Hasher};
use std::io::Read;
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
#[cfg(unix)]
use std::os::unix::fs::{MetadataExt, OpenOptionsExt};
use crate::error::{ProxyError, Result}; use crate::error::{ProxyError, Result};
const MAX_CONFIG_SOURCE_BYTES: usize = 8 * 1024 * 1024;
pub(super) fn normalize_config_path(path: &Path) -> PathBuf { pub(super) fn normalize_config_path(path: &Path) -> PathBuf {
path.canonicalize().unwrap_or_else(|_| { let absolute = if path.is_absolute() {
if path.is_absolute() {
path.to_path_buf() path.to_path_buf()
} else { } else {
std::env::current_dir() std::env::current_dir()
.map(|cwd| cwd.join(path)) .map(|cwd| cwd.join(path))
.unwrap_or_else(|_| path.to_path_buf()) .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 { 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)> { pub(super) fn read_config_source(path: &Path) -> Result<(PathBuf, String)> {
let mut options = std::fs::OpenOptions::new();
options.read(true);
#[cfg(unix)] #[cfg(unix)]
options.custom_flags(libc::O_CLOEXEC); let bytes = crate::util::secure_fs::read_regular_limited(path, MAX_CONFIG_SOURCE_BYTES)
let mut file = options
.open(path)
.map_err(|error| ProxyError::Config(error.to_string()))?; .map_err(|error| ProxyError::Config(error.to_string()))?;
let opened_metadata = file #[cfg(not(unix))]
.metadata() let bytes = std::fs::read(path).map_err(|error| ProxyError::Config(error.to_string()))?;
.map_err(|error| ProxyError::Config(error.to_string()))?; if bytes.len() > MAX_CONFIG_SOURCE_BYTES {
if !opened_metadata.is_file() {
return Err(ProxyError::Config(format!( return Err(ProxyError::Config(format!(
"config source `{}` must be a regular file", "config source `{}` exceeds {} bytes",
path.display() 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 normalized = normalize_config_path(path);
let current_metadata = std::fs::metadata(&normalized)
.map_err(|error| ProxyError::Config(error.to_string()))?;
if !same_file_identity(&opened_metadata, &current_metadata) {
return Err(ProxyError::Config(format!(
"config source `{}` changed while it was opened",
path.display()
)));
}
let mut contents = String::new();
file.read_to_string(&mut contents)
.map_err(|error| ProxyError::Config(error.to_string()))?;
let completed_metadata = file
.metadata()
.map_err(|error| ProxyError::Config(error.to_string()))?;
if !same_file_version(&opened_metadata, &completed_metadata) {
return Err(ProxyError::Config(format!(
"config source `{}` changed while it was read",
path.display()
)));
}
Ok((normalized, contents)) Ok((normalized, contents))
} }
fn same_file_identity(left: &std::fs::Metadata, right: &std::fs::Metadata) -> bool {
#[cfg(unix)]
{
left.dev() == right.dev() && left.ino() == right.ino()
}
#[cfg(not(unix))]
{
left.len() == right.len() && left.modified().ok() == right.modified().ok()
}
}
fn same_file_version(left: &std::fs::Metadata, right: &std::fs::Metadata) -> bool {
#[cfg(unix)]
{
same_file_identity(left, right)
&& left.len() == right.len()
&& left.mtime() == right.mtime()
&& left.mtime_nsec() == right.mtime_nsec()
&& left.ctime() == right.ctime()
&& left.ctime_nsec() == right.ctime_nsec()
}
#[cfg(not(unix))]
{
same_file_identity(left, right)
}
}
pub(super) fn preprocess_includes( pub(super) fn preprocess_includes(
content: &str, content: &str,
base_dir: &Path, base_dir: &Path,
@@ -113,20 +75,17 @@ pub(super) fn preprocess_includes(
if let Some(rest) = rest.strip_prefix('=') { if let Some(rest) = rest.strip_prefix('=') {
let path_str = rest.trim().trim_matches('"'); let path_str = rest.trim().trim_matches('"');
let resolved = base_dir.join(path_str); let resolved = base_dir.join(path_str);
let normalized = normalize_config_path(&resolved); let (normalized, disk_contents) = read_config_source(&resolved)?;
let cached = source_contents.get(&normalized).cloned(); let included = source_overrides
let (normalized, included) = if let Some(included) = .get(&normalized)
source_overrides.get(&normalized).cloned().or(cached) .cloned()
{ .or_else(|| source_contents.get(&normalized).cloned())
(normalized, included) .unwrap_or(disk_contents);
} else {
read_config_source(&resolved)?
};
source_files.insert(normalized.clone()); source_files.insert(normalized.clone());
source_contents source_contents
.entry(normalized) .entry(normalized.clone())
.or_insert_with(|| included.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( output.push_str(&preprocess_includes(
&included, &included,
included_dir, included_dir,
+72 -1
View File
@@ -12,6 +12,7 @@ const ACCESS_SECRET_BYTES: usize = 16;
pub(crate) struct UserAuthSnapshot { pub(crate) struct UserAuthSnapshot {
entries: Vec<UserAuthEntry>, entries: Vec<UserAuthEntry>,
by_name: HashMap<String, u32>, by_name: HashMap<String, u32>,
by_hint_key: HashMap<u64, Vec<u32>>,
sni_index: HashMap<u64, Vec<u32>>, sni_index: HashMap<u64, Vec<u32>>,
sni_initial_index: HashMap<u8, Vec<u32>>, sni_initial_index: HashMap<u8, Vec<u32>>,
} }
@@ -22,16 +23,21 @@ pub(crate) struct UserAuthEntry {
pub(crate) secret: [u8; ACCESS_SECRET_BYTES], pub(crate) secret: [u8; ACCESS_SECRET_BYTES],
/// Stable secret identity used by process-wide admission fencing. /// Stable secret identity used by process-wide admission fencing.
pub(crate) credential_id: [u8; 16], pub(crate) credential_id: [u8; 16],
/// Stable compact key used only to resolve bounded authentication hints.
pub(crate) hint_key: u64,
} }
impl UserAuthSnapshot { impl UserAuthSnapshot {
pub(super) fn from_users(users: &HashMap<String, String>) -> Result<Self> { pub(super) fn from_users(users: &HashMap<String, String>) -> Result<Self> {
let mut entries = Vec::with_capacity(users.len()); let mut entries = Vec::with_capacity(users.len());
let mut by_name = HashMap::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_index = HashMap::with_capacity(users.len());
let mut sni_initial_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::<Vec<_>>();
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 { let decoded = hex::decode(secret_hex).map_err(|_| ProxyError::InvalidSecret {
user: user.clone(), user: user.clone(),
reason: "Must be 32 hex characters".to_string(), reason: "Must be 32 hex characters".to_string(),
@@ -52,12 +58,27 @@ impl UserAuthSnapshot {
let digest = sha256(&secret); let digest = sha256(&secret);
let mut credential_id = [0; 16]; let mut credential_id = [0; 16];
credential_id.copy_from_slice(&digest[..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 { entries.push(UserAuthEntry {
user: user.clone(), user: user.clone(),
secret, secret,
credential_id, credential_id,
hint_key,
}); });
by_name.insert(user.clone(), user_id); by_name.insert(user.clone(), user_id);
by_hint_key
.entry(hint_key)
.or_insert_with(Vec::new)
.push(user_id);
sni_index sni_index
.entry(Self::sni_lookup_hash(user)) .entry(Self::sni_lookup_hash(user))
.or_insert_with(Vec::new) .or_insert_with(Vec::new)
@@ -77,6 +98,7 @@ impl UserAuthSnapshot {
Ok(Self { Ok(Self {
entries, entries,
by_name, by_name,
by_hint_key,
sni_index, sni_index,
sni_initial_index, sni_initial_index,
}) })
@@ -101,6 +123,10 @@ impl UserAuthSnapshot {
.map(|entry| entry.credential_id) .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]> { pub(crate) fn sni_candidates(&self, sni: &str) -> Option<&[u32]> {
self.sni_index self.sni_index
.get(&Self::sni_lookup_hash(sni)) .get(&Self::sni_lookup_hash(sni))
@@ -123,3 +149,48 @@ impl UserAuthSnapshot {
hasher.finish() 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)
);
}
}
+9 -6
View File
@@ -23,6 +23,8 @@ use hmac::{Hmac, Mac};
use sha2::{Digest, Sha256}; use sha2::{Digest, Sha256};
use super::*; use super::*;
#[cfg(unix)]
use crate::util::secure_fs::open_dir_nofollow;
// Path-based static snapshot fallback for platforms without directory descriptors. // Path-based static snapshot fallback for platforms without directory descriptors.
#[cfg(not(unix))] #[cfg(not(unix))]
@@ -247,16 +249,17 @@ fn load_static_site(
#[cfg(unix)] #[cfg(unix)]
fn open_static_root(root: &Path) -> Result<Dir> { fn open_static_root(root: &Path) -> Result<Dir> {
Dir::open( let descriptor = open_dir_nofollow(root).map_err(|error| {
root,
OFlag::O_RDONLY | OFlag::O_DIRECTORY | OFlag::O_NOFOLLOW | OFlag::O_CLOEXEC,
Mode::empty(),
)
.map_err(|error| {
ProxyError::Config(format!( ProxyError::Config(format!(
"WEB static directory `{}` must be a real directory, not a symlink: {error}", "WEB static directory `{}` must be a real directory, not a symlink: {error}",
root.display() root.display()
)) ))
})?;
Dir::from_fd(descriptor).map_err(|error| {
ProxyError::Config(format!(
"failed to read WEB static directory `{}`: {error}",
root.display()
))
}) })
} }
+2
View File
@@ -46,6 +46,8 @@ mod legacy_policy_tests;
mod me_route_tests; mod me_route_tests;
#[path = "load_basic_tests/me_startup_tests.rs"] #[path = "load_basic_tests/me_startup_tests.rs"]
mod me_startup_tests; mod me_startup_tests;
#[path = "load_basic_tests/source_security_tests.rs"]
mod source_security_tests;
#[path = "load_basic_tests/synlimit_mss_tests.rs"] #[path = "load_basic_tests/synlimit_mss_tests.rs"]
mod synlimit_mss_tests; mod synlimit_mss_tests;
#[path = "load_basic_tests/tls_fetch_tests.rs"] #[path = "load_basic_tests/tls_fetch_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());
}
+10 -414
View File
@@ -1,20 +1,22 @@
use std::collections::BTreeSet;
use std::net::IpAddr;
use std::sync::Arc; use std::sync::Arc;
use std::time::Duration; use std::time::Duration;
use tokio::io::AsyncWriteExt;
use tokio::process::Command;
use tokio::sync::{mpsc, watch}; use tokio::sync::{mpsc, watch};
use tokio_util::sync::CancellationToken; 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::middle_relay::note_global_relay_pressure;
use crate::proxy::shared_state::{ConntrackCloseEvent, ConntrackCloseReason, ProxySharedState}; use crate::proxy::shared_state::{ConntrackCloseEvent, ConntrackCloseReason, ProxySharedState};
use crate::stats::Stats; use crate::stats::Stats;
#[cfg(unix)]
use crate::util::trusted_command::resolve_trusted_helper; // 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 CONNTRACK_EVENT_QUEUE_CAPACITY: usize = 32_768;
const PRESSURE_RELEASE_TICKS: u8 = 3; const PRESSURE_RELEASE_TICKS: u8 = 3;
@@ -311,389 +313,6 @@ fn update_pressure_state(
state.low_streak = 0; 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<NetfilterBackend> {
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<u16> {
let mut ports: BTreeSet<u16> = 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<IpAddr>, u16)>, Vec<(Option<IpAddr>, u16)>) {
let mode = cfg.server.conntrack_control.mode;
let mut v4_targets: BTreeSet<(Option<IpAddr>, u16)> = BTreeSet::new();
let mut v6_targets: BTreeSet<(Option<IpAddr>, 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::<IpAddr>().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::<IpAddr>().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<String>) -> 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<u8> { fn fd_usage_pct() -> Option<u8> {
let soft_limit = nofile_soft_limit()?; let soft_limit = nofile_soft_limit()?;
if soft_limit == 0 { if soft_limit == 0 {
@@ -722,29 +341,6 @@ fn nofile_soft_limit() -> Option<u64> {
} }
} }
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)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
+419
View File
@@ -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<NetfilterBackend> {
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<u16> {
let mut ports: BTreeSet<u16> = 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<IpAddr>, u16)>, Vec<(Option<IpAddr>, u16)>) {
let mode = cfg.server.conntrack_control.mode;
let mut v4_targets: BTreeSet<(Option<IpAddr>, u16)> = BTreeSet::new();
let mut v6_targets: BTreeSet<(Option<IpAddr>, 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::<IpAddr>().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::<IpAddr>().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<String>) -> 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
}
}
+95 -194
View File
@@ -1,6 +1,8 @@
use std::fs::{self, File, OpenOptions}; use std::fs::{self, File, OpenOptions};
use std::io::{ErrorKind, Read, Write}; use std::io::{ErrorKind, Read, Write};
use std::os::unix::fs::{MetadataExt, OpenOptionsExt}; use std::os::unix::fs::{MetadataExt, OpenOptionsExt};
#[cfg(target_os = "linux")]
use std::os::fd::{FromRawFd, OwnedFd};
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use nix::fcntl::{Flock, FlockArg}; use nix::fcntl::{Flock, FlockArg};
@@ -281,7 +283,19 @@ pub fn signal_pid_file<P: AsRef<Path>>(
path: P, path: P,
signal: nix::sys::signal::Signal, signal: nix::sys::signal::Signal,
) -> Result<(), DaemonError> { ) -> 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) nix::sys::signal::kill(Pid::from_raw(pid), signal)
.map_err(|error| DaemonError::PidFile(format!("cannot signal process {}: {}", pid, error))) .map_err(|error| DaemonError::PidFile(format!("cannot signal process {}: {}", pid, error)))
} }
@@ -303,206 +317,93 @@ pub enum DaemonStatus {
pub fn check_status<P: AsRef<Path>>(path: P) -> DaemonStatus { pub fn check_status<P: AsRef<Path>>(path: P) -> DaemonStatus {
let path = path.as_ref(); let path = path.as_ref();
match read_pid_file_if_exists(path) { 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(Some(pid)) => DaemonStatus::Stale(pid),
Ok(None) | Err(_) => DaemonStatus::NotRunning, Ok(None) | Err(_) => DaemonStatus::NotRunning,
} }
} }
fn daemon_lock_is_held(path: &Path) -> Result<bool, DaemonError> {
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<OwnedFd, DaemonError> {
// 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::<libc::siginfo_t>(),
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 { fn is_process_running(pid: i32) -> bool {
nix::sys::signal::kill(Pid::from_raw(pid), None).is_ok() nix::sys::signal::kill(Pid::from_raw(pid), None).is_ok()
} }
#[cfg(test)] #[cfg(test)]
mod tests { 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<std::process::ExitStatus> {
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<T: Send + Sync>() {}
assert_send_sync::<PidFile>();
}
#[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();
}
}
+209
View File
@@ -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<std::process::ExitStatus> {
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<T: Send + Sync>() {}
assert_send_sync::<PidFile>();
}
#[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();
}
+1 -47
View File
@@ -8,8 +8,6 @@
// Infrastructure module used via CLI flags. // Infrastructure module used via CLI flags.
#![allow(dead_code)] #![allow(dead_code)]
use std::path::Path;
use crate::config::{LogRotation, LoggingConfig, LoggingDestination}; use crate::config::{LogRotation, LoggingConfig, LoggingDestination};
use tracing_subscriber::layer::SubscriberExt; use tracing_subscriber::layer::SubscriberExt;
@@ -144,31 +142,9 @@ pub fn init_logging(
} }
LogDestination::File { options } => { 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()) let file_appender = file::BoundedFileAppender::new(options.clone())
.expect("Failed to open log file"); .expect("Failed to open log file");
tracing_appender::non_blocking(file_appender) let (non_blocking, guard) = 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 fmt_layer = fmt::Layer::default() let fmt_layer = fmt::Layer::default()
.with_ansi(false) .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. /// Syslog writer for tracing.
#[cfg(unix)] #[cfg(unix)]
#[derive(Clone, Copy)] #[derive(Clone, Copy)]
+190 -140
View File
@@ -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::io::{self, Write};
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use std::time::{Duration, SystemTime, UNIX_EPOCH}; 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 chrono::{DateTime, Datelike, Duration as ChronoDuration, Utc};
use crate::config::LogRotation; use crate::config::LogRotation;
@@ -20,6 +38,8 @@ pub(crate) struct BoundedFileAppender {
current_size: u64, current_size: u64,
last_cleanup: DateTime<Utc>, last_cleanup: DateTime<Utc>,
file: Option<File>, file: Option<File>,
#[cfg(unix)]
dir_fd: OwnedFd,
now: Box<dyn Fn() -> DateTime<Utc> + Send + Sync>, now: Box<dyn Fn() -> DateTime<Utc> + Send + Sync>,
} }
@@ -46,6 +66,11 @@ impl BoundedFileAppender {
let start = now(); let start = now();
let current_path = active_path_for(&dir, &base_name, options.rotation, &start); 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, &current_path)?;
#[cfg(not(unix))]
let (file, current_size) = open_append_file(&current_path)?; let (file, current_size) = open_append_file(&current_path)?;
let mut appender = Self { let mut appender = Self {
options, options,
@@ -55,6 +80,8 @@ impl BoundedFileAppender {
current_size, current_size,
last_cleanup: start, last_cleanup: start,
file: Some(file), file: Some(file),
#[cfg(unix)]
dir_fd,
now, now,
}; };
appender.cleanup(&start); appender.cleanup(&start);
@@ -79,10 +106,8 @@ impl BoundedFileAppender {
fn rotate_for_size(&mut self, now: &DateTime<Utc>) -> io::Result<()> { fn rotate_for_size(&mut self, now: &DateTime<Utc>) -> io::Result<()> {
self.close_current()?; self.close_current()?;
if self.current_path.exists() {
let archive_path = self.archive_path(now); let archive_path = self.archive_path(now);
fs::rename(&self.current_path, archive_path)?; self.rename_current_if_present(&archive_path)?;
}
self.open_current() self.open_current()
} }
@@ -95,7 +120,7 @@ impl BoundedFileAppender {
let stamp = now.format("%Y%m%d%H%M%S"); let stamp = now.format("%Y%m%d%H%M%S");
for seq in 0..1000 { for seq in 0..1000 {
let candidate = self.dir.join(format!("{file_name}.{stamp}.{seq}")); let candidate = self.dir.join(format!("{file_name}.{stamp}.{seq}"));
if !candidate.exists() { if !self.path_exists(&candidate) {
return candidate; return candidate;
} }
} }
@@ -103,6 +128,9 @@ impl BoundedFileAppender {
} }
fn open_current(&mut self) -> io::Result<()> { 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)?; let (file, current_size) = open_append_file(&self.current_path)?;
self.file = Some(file); self.file = Some(file);
self.current_size = current_size; self.current_size = current_size;
@@ -130,40 +158,10 @@ impl BoundedFileAppender {
fn cleanup(&mut self, now: &DateTime<Utc>) { fn cleanup(&mut self, now: &DateTime<Utc>) {
self.last_cleanup = now.clone(); self.last_cleanup = now.clone();
let Ok(entries) = fs::read_dir(&self.dir) else { let Ok(mut candidates) = self.collect_candidates() else {
return; 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 { if self.options.max_age_secs > 0 {
let cutoff = system_time_from_utc(now) let cutoff = system_time_from_utc(now)
.checked_sub(Duration::from_secs(self.options.max_age_secs)) .checked_sub(Duration::from_secs(self.options.max_age_secs))
@@ -172,7 +170,7 @@ impl BoundedFileAppender {
if candidate.is_current || candidate.modified >= cutoff { if candidate.is_current || candidate.modified >= cutoff {
true true
} else { } else {
let _ = fs::remove_file(&candidate.path); self.remove_candidate(candidate);
false false
} }
}); });
@@ -189,11 +187,153 @@ impl BoundedFileAppender {
if total <= self.options.max_files { if total <= self.options.max_files {
break; break;
} }
let _ = fs::remove_file(candidate.path); self.remove_candidate(&candidate);
total -= 1; 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<Vec<LogFileCandidate>> {
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<Vec<LogFileCandidate>> {
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 { impl Write for BoundedFileAppender {
@@ -233,6 +373,17 @@ struct LogFileCandidate {
is_current: bool, 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)> { fn open_append_file(path: &Path) -> io::Result<(File, u64)> {
let mut options = OpenOptions::new(); let mut options = OpenOptions::new();
options.create(true).append(true); options.create(true).append(true);
@@ -291,105 +442,4 @@ fn system_time_from_utc(now: &DateTime<Utc>) -> SystemTime {
} }
#[cfg(test)] #[cfg(test)]
mod tests { mod tests;
use std::io::Write;
use tempfile::tempdir;
use super::*;
fn fixed_now() -> DateTime<Utc> {
DateTime::<Utc>::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<PathBuf> {
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());
}
}
+125
View File
@@ -0,0 +1,125 @@
use std::io::Write;
use tempfile::tempdir;
use super::*;
fn fixed_now() -> DateTime<Utc> {
DateTime::<Utc>::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<PathBuf> {
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());
}
+28 -41
View File
@@ -74,26 +74,7 @@ pub(super) async fn bootstrap(
data_path.as_deref(), data_path.as_deref(),
); );
if !runtime_base_dir.exists() if let Err(e) = enter_runtime_directory(&runtime_base_dir) {
&& 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) {
eprintln!( eprintln!(
"[telemt] Can't use runtime directory {}: {}", "[telemt] Can't use runtime directory {}: {}",
runtime_base_dir.display(), runtime_base_dir.display(),
@@ -125,7 +106,7 @@ pub(super) async fn bootstrap(
if config_path_explicit { if config_path_explicit {
if let Some(serialized) = serialized.as_ref() { 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!( eprintln!(
"[telemt] Error: failed to create explicit config at {}: {}", "[telemt] Error: failed to create explicit config at {}: {}",
config_path.display(), config_path.display(),
@@ -149,7 +130,7 @@ pub(super) async fn bootstrap(
if let Some(serialized) = serialized.as_ref() { if let Some(serialized) = serialized.as_ref() {
match std::fs::create_dir_all(&runtime_base_dir) { 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(()) => { Ok(()) => {
config_path = runtime_config_path; config_path = runtime_config_path;
eprintln!( eprintln!(
@@ -176,7 +157,7 @@ pub(super) async fn bootstrap(
} }
if !persisted { if !persisted {
match std::fs::write(&fallback_config_path, serialized) { match write_private_file(&fallback_config_path, serialized) {
Ok(()) => { Ok(()) => {
config_path = fallback_config_path; config_path = fallback_config_path;
eprintln!( eprintln!(
@@ -226,24 +207,7 @@ pub(super) async fn bootstrap(
std::process::exit(1); std::process::exit(1);
} }
if data_path.exists() { if let Err(e) = enter_runtime_directory(data_path) {
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) {
eprintln!( eprintln!(
"[telemt] Can't use data_path {}: {}", "[telemt] Can't use data_path {}: {}",
data_path.display(), data_path.display(),
@@ -376,3 +340,26 @@ pub(super) async fn bootstrap(
logging_guard, 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)
}
}
+6 -596
View File
@@ -1,20 +1,8 @@
#![allow(clippy::items_after_test_module)]
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicBool, Ordering}; 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::cli;
use crate::config::ProxyConfig;
use crate::logging::LogCliOptions; 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 MAESTRO_COLOR: &str = "\x1b[92m";
const COLOR_RESET: &str = "\x1b[0m"; 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)] #[cfg(test)]
mod tests { 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>) -> 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<std::sync::Arc<UpstreamManager>>,
) -> Option<ProxyConfigData> {
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;
}
}
}
}
+360
View File
@@ -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>) -> 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<std::sync::Arc<UpstreamManager>>,
) -> Option<ProxyConfigData> {
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;
}
}
}
}
+247
View File
@@ -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));
}
}
+11 -3
View File
@@ -20,6 +20,8 @@ use crate::config::{ListenerTransport, ProxyConfig};
use crate::startup::{COMPONENT_LISTENERS_BIND, StartupTracker}; use crate::startup::{COMPONENT_LISTENERS_BIND, StartupTracker};
use crate::transport::find_listener_processes; use crate::transport::find_listener_processes;
use crate::transport::socket::{activate_listener_socket, bind_listener_socket}; 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 super::plan::{ListenerBindSpec, listener_bind_plan};
use crate::maestro::helpers::{print_proxy_links, print_web_proxy_links}; use crate::maestro::helpers::{print_proxy_links, print_web_proxy_links};
@@ -281,6 +283,7 @@ pub(crate) async fn bind_listeners(
#[cfg(unix)] #[cfg(unix)]
if let Some(unix_path) = &config.server.listen_unix_sock { if let Some(unix_path) = &config.server.listen_unix_sock {
let unix_path = Path::new(unix_path); let unix_path = Path::new(unix_path);
let anchored_path = AnchoredPath::open_trusted_parent(unix_path)?;
remove_stale_unix_socket(unix_path)?; remove_stale_unix_socket(unix_path)?;
let unix_listener = UnixListener::bind(unix_path)?; let unix_listener = UnixListener::bind(unix_path)?;
let socket_metadata = std::fs::symlink_metadata(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 { if let Some(perm_str) = &config.server.listen_unix_sock_perm {
match u32::from_str_radix(perm_str.trim_start_matches('0'), 8) { match u32::from_str_radix(perm_str.trim_start_matches('0'), 8) {
Ok(mode) => { Ok(mode) => {
use std::os::unix::fs::PermissionsExt; use nix::sys::stat::{FchmodatFlags, Mode, fchmodat};
let permissions = std::fs::Permissions::from_mode(mode);
verify_bound_unix_socket(unix_path, socket_identity)?; 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!( error!(
path = %unix_path.display(), path = %unix_path.display(),
permissions = %perm_str, permissions = %perm_str,
+16 -1
View File
@@ -25,6 +25,21 @@ impl RunningClientHandler {
R: AsyncRead + Unpin + Send + 'static, R: AsyncRead + Unpin + Send + 'static,
W: AsyncWrite + 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( Self::handle_authenticated_static_with_shared(
client_reader, client_reader,
client_writer, client_writer,
@@ -40,7 +55,7 @@ impl RunningClientHandler {
local_addr, local_addr,
peer_addr, peer_addr,
ip_tracker, ip_tracker,
ProxySharedState::new(), shared,
) )
.await .await
} }
+1
View File
@@ -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_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, auth_probe_slots_for_testing_in_shared, clear_auth_probe_state_for_testing_in_shared,
clear_unknown_sni_warn_state_for_testing_in_shared, clear_warned_secrets_for_testing_in_shared, clear_unknown_sni_warn_state_for_testing_in_shared, clear_warned_secrets_for_testing_in_shared,
insert_auth_probe_state_for_testing_in_shared,
should_emit_unknown_sni_warn_for_testing_in_shared, warned_secrets_for_testing_in_shared, should_emit_unknown_sni_warn_for_testing_in_shared, warned_secrets_for_testing_in_shared,
}; };
+39 -18
View File
@@ -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<u32> { pub(super) fn sticky_hint_get_by_ip(shared: &ProxySharedState, peer_ip: IpAddr) -> Option<u64> {
shared shared
.handshake .handshake
.sticky_user_by_ip .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( pub(super) fn sticky_hint_get_by_ip_prefix(
shared: &ProxySharedState, shared: &ProxySharedState,
peer_ip: IpAddr, peer_ip: IpAddr,
) -> Option<u32> { ) -> Option<u64> {
shared shared
.handshake .handshake
.sticky_user_by_ip_prefix .sticky_user_by_ip_prefix
@@ -61,7 +61,7 @@ pub(super) fn sticky_hint_get_by_ip_prefix(
.map(|entry| *entry) .map(|entry| *entry)
} }
pub(super) fn sticky_hint_get_by_sni(shared: &ProxySharedState, sni: &str) -> Option<u32> { pub(super) fn sticky_hint_get_by_sni(shared: &ProxySharedState, sni: &str) -> Option<u64> {
let key = sni_hint_hash(sni); let key = sni_hint_hash(sni);
shared shared
.handshake .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( pub(super) fn sticky_hint_record_success_in(
shared: &ProxySharedState, shared: &ProxySharedState,
peer_ip: IpAddr, peer_ip: IpAddr,
user_id: u32, hint_key: u64,
sni: Option<&str>, sni: Option<&str>,
) { ) {
bounded_sticky_hint_upsert( bounded_sticky_hint_upsert(
&shared.handshake.sticky_user_by_ip, &shared.handshake.sticky_user_by_ip,
&shared.handshake.sticky_user_by_ip_slots, &shared.handshake.sticky_user_by_ip_slots,
peer_ip, peer_ip,
user_id, hint_key,
); );
bounded_sticky_hint_upsert( bounded_sticky_hint_upsert(
&shared.handshake.sticky_user_by_ip_prefix, &shared.handshake.sticky_user_by_ip_prefix,
&shared.handshake.sticky_user_by_ip_prefix_slots, &shared.handshake.sticky_user_by_ip_prefix_slots,
ip_prefix_hint_key(peer_ip), ip_prefix_hint_key(peer_ip),
user_id, hint_key,
); );
if let Some(sni) = sni { 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,
&shared.handshake.sticky_user_by_sni_hash_slots, &shared.handshake.sticky_user_by_sni_hash_slots,
sni_hint_hash(sni), sni_hint_hash(sni),
user_id, hint_key,
); );
} }
} }
fn bounded_sticky_hint_upsert<K>( fn bounded_sticky_hint_upsert<K>(
entries: &DashMap<K, u32>, entries: &DashMap<K, u64>,
slots: &crate::slot_budget::SlotBudget, slots: &crate::slot_budget::SlotBudget,
key: K, key: K,
user_id: u32, hint_key: u64,
) where ) where
K: Eq + Hash, K: Clone + Eq + Hash,
{ {
match entries.entry(key) { 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::Occupied(mut entry) => {
entry.insert(user_id); entry.insert(hint_key);
} }
Entry::Vacant(entry) => { Entry::Vacant(entry) => {
let Some(slot) = slots.try_acquire() else { entry.insert(hint_key);
slot.commit();
}
}
return;
}
let Some((victim_key, victim_hint_key)) = entries
.iter()
.next()
.map(|entry| (entry.key().clone(), *entry.value()))
else {
return; return;
}; };
entry.insert(user_id); if entries
slot.commit(); .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; let ring = &shared.handshake.recent_user_ring;
if ring.is_empty() { if ring.is_empty() {
return; return;
@@ -131,7 +152,7 @@ pub(super) fn record_recent_user_success_in(shared: &ProxySharedState, user_id:
.recent_user_ring_seq .recent_user_ring_seq
.fetch_add(1, Ordering::Relaxed); .fetch_add(1, Ordering::Relaxed);
let idx = (seq as usize) % ring.len(); 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( pub(super) fn mark_candidate_if_new(
@@ -387,7 +408,7 @@ mod bounded_registry_tests {
sticky_hint_record_success_in( sticky_hint_record_success_in(
shared.as_ref(), shared.as_ref(),
peer_ip, peer_ip,
index as u32, index as u64 | 1,
Some(&format!("host-{index}.example")), Some(&format!("host-{index}.example")),
); );
} }
+25 -1
View File
@@ -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) { 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.clear();
shared.handshake.auth_probe_slots.reset_for_testing(); shared.handshake.auth_probe_slots.release_many(removed);
match shared.handshake.auth_probe_saturation.lock() { match shared.handshake.auth_probe_saturation.lock() {
Ok(mut saturation) => { Ok(mut saturation) => {
*saturation = None; *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( pub(crate) fn auth_probe_state_for_testing_in_shared(
shared: &ProxySharedState, shared: &ProxySharedState,
) -> &DashMap<IpAddr, AuthProbeState> { ) -> &DashMap<IpAddr, AuthProbeState> {
+40 -14
View File
@@ -151,10 +151,14 @@ where
if let Some(snapshot) = config.runtime_user_auth() { if let Some(snapshot) = config.runtime_user_auth() {
let sticky_ip_hint = sticky_hint_get_by_ip(shared, peer.ip()); 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_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 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 exact_user_id = exact_user.and_then(|user| snapshot.user_id_by_name(user));
let has_hint = sticky_ip_hint.is_some() let has_hint = sticky_ip_candidates.is_some_and(|ids| !ids.is_empty())
|| sticky_prefix_hint.is_some() || sticky_prefix_candidates.is_some_and(|ids| !ids.is_empty())
|| preferred_user_id.is_some() || preferred_user_id.is_some()
|| exact_user_id.is_some(); || exact_user_id.is_some();
let overload = auth_probe_saturation_is_throttled_in(shared, Instant::now()); 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)); let mut matched = exact_user_id.is_some_and(|user_id| try_user_id!(user_id));
if exact_user.is_none() 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() if exact_user.is_none()
@@ -218,9 +230,17 @@ where
if exact_user.is_none() if exact_user.is_none()
&& !matched && !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 { if exact_user.is_none() && !matched && !budget_exhausted {
@@ -231,18 +251,22 @@ where
.recent_user_ring_seq .recent_user_ring_seq
.load(Ordering::Relaxed); .load(Ordering::Relaxed);
let scan_limit = ring.len().min(RECENT_USER_RING_SCAN_LIMIT); 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 idx = (next_seq as usize + ring.len() - 1 - offset) % ring.len();
let encoded_user_id = ring[idx].load(Ordering::Relaxed); let hint_key = ring[idx].load(Ordering::Relaxed);
if encoded_user_id == 0 { if hint_key == 0 {
continue; continue;
} }
if try_user_id!(encoded_user_id - 1) { 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; matched = true;
break; break 'recent_hints;
} }
if budget_exhausted { if budget_exhausted {
break; break 'recent_hints;
}
}
} }
} }
} }
@@ -357,8 +381,10 @@ where
auth_probe_record_success_in(shared, peer.ip()); auth_probe_record_success_in(shared, peer.ip());
if let Some(user_id) = matched_user_id { if let Some(user_id) = matched_user_id {
sticky_hint_record_success_in(shared, peer.ip(), user_id, None); if let Some(entry) = snapshot.entry_by_id(user_id) {
record_recent_user_success_in(shared, 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; let max_pending = config.general.crypto_pending_buffer;
+12 -2
View File
@@ -396,8 +396,18 @@ where
auth_probe_record_success_in(shared, peer.ip()); auth_probe_record_success_in(shared, peer.ip());
if let Some(user_id) = validated_user_id { if let Some(user_id) = validated_user_id {
sticky_hint_record_success_in(shared, peer.ip(), user_id, client_sni.as_deref()); if let Some(entry) = config
record_recent_user_success_in(shared, user_id); .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(( HandshakeResult::Success((
+49 -15
View File
@@ -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_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 preferred_user_id = preferred_user_hint.and_then(|user| snapshot.user_id_by_name(user));
let sticky_sni_hint = client_sni let sticky_sni_hint = client_sni
.as_deref() .as_deref()
.and_then(|sni| sticky_hint_get_by_sni(shared, sni)); .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_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 let sni_candidates = client_sni
.as_deref() .as_deref()
.and_then(|sni| snapshot.sni_candidates(sni)); .and_then(|sni| snapshot.sni_candidates(sni));
@@ -52,10 +58,10 @@ pub(super) async fn validate_tls_client(
.as_deref() .as_deref()
.and_then(|sni| snapshot.sni_initial_candidates(sni)); .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() || preferred_user_id.is_some()
|| sticky_sni_hint.is_some() || sticky_sni_candidates.is_some_and(|ids| !ids.is_empty())
|| sticky_prefix_hint.is_some() || sticky_prefix_candidates.is_some_and(|ids| !ids.is_empty())
|| sni_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()); || sni_initial_candidates.is_some_and(|ids| !ids.is_empty());
let overload = auth_probe_saturation_is_throttled_in(shared, Instant::now()); 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; let mut matched = false;
if let Some(user_id) = sticky_ip_hint { if 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 !matched && let Some(user_id) = preferred_user_id { if !matched && let Some(user_id) = preferred_user_id {
matched = try_user_id!(user_id); matched = try_user_id!(user_id);
} }
if !matched && let Some(user_id) = sticky_sni_hint { if !matched && let Some(candidate_ids) = sticky_sni_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 !matched && let Some(user_id) = sticky_prefix_hint { if !matched && 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 !matched if !matched
@@ -149,18 +179,22 @@ pub(super) async fn validate_tls_client(
.recent_user_ring_seq .recent_user_ring_seq
.load(Ordering::Relaxed); .load(Ordering::Relaxed);
let scan_limit = ring.len().min(RECENT_USER_RING_SCAN_LIMIT); 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 idx = (next_seq as usize + ring.len() - 1 - offset) % ring.len();
let encoded_user_id = ring[idx].load(Ordering::Relaxed); let hint_key = ring[idx].load(Ordering::Relaxed);
if encoded_user_id == 0 { if hint_key == 0 {
continue; continue;
} }
if try_user_id!(encoded_user_id - 1) { 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; matched = true;
break; break 'recent_hints;
} }
if budget_exhausted { if budget_exhausted {
break; break 'recent_hints;
}
}
} }
} }
} }
+6 -6
View File
@@ -1,7 +1,7 @@
use std::collections::hash_map::RandomState; use std::collections::hash_map::RandomState;
use std::collections::{HashMap, HashSet}; use std::collections::{HashMap, HashSet};
use std::net::{IpAddr, SocketAddr}; 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::sync::{Arc, Mutex};
use std::time::Instant; use std::time::Instant;
@@ -61,13 +61,13 @@ pub(crate) struct HandshakeSharedState {
pub(crate) auth_probe_eviction_hasher: RandomState, pub(crate) auth_probe_eviction_hasher: RandomState,
pub(crate) invalid_secret_warned: Mutex<HashSet<(String, String)>>, pub(crate) invalid_secret_warned: Mutex<HashSet<(String, String)>>,
pub(crate) unknown_sni_warn_next_allowed: Mutex<Option<Instant>>, pub(crate) unknown_sni_warn_next_allowed: Mutex<Option<Instant>>,
pub(crate) sticky_user_by_ip: DashMap<IpAddr, u32>, pub(crate) sticky_user_by_ip: DashMap<IpAddr, u64>,
pub(crate) sticky_user_by_ip_slots: SlotBudget, pub(crate) sticky_user_by_ip_slots: SlotBudget,
pub(crate) sticky_user_by_ip_prefix: DashMap<u64, u32>, pub(crate) sticky_user_by_ip_prefix: DashMap<u64, u64>,
pub(crate) sticky_user_by_ip_prefix_slots: SlotBudget, pub(crate) sticky_user_by_ip_prefix_slots: SlotBudget,
pub(crate) sticky_user_by_sni_hash: DashMap<u64, u32>, pub(crate) sticky_user_by_sni_hash: DashMap<u64, u64>,
pub(crate) sticky_user_by_sni_hash_slots: SlotBudget, 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) recent_user_ring_seq: AtomicU64,
pub(crate) auth_expensive_checks_total: AtomicU64, pub(crate) auth_expensive_checks_total: AtomicU64,
pub(crate) auth_budget_exhausted_total: AtomicU64, pub(crate) auth_budget_exhausted_total: AtomicU64,
@@ -138,7 +138,7 @@ impl ProxySharedState {
sticky_user_by_sni_hash_slots: SlotBudget::new( sticky_user_by_sni_hash_slots: SlotBudget::new(
crate::proxy::handshake::STICKY_HINT_MAX_ENTRIES, 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) .take(HANDSHAKE_RECENT_USER_RING_LEN)
.collect::<Vec<_>>() .collect::<Vec<_>>()
.into_boxed_slice(), .into_boxed_slice(),
+30 -14
View File
@@ -1241,7 +1241,10 @@ async fn tls_runtime_snapshot_updates_sticky_and_recent_hints() {
.sticky_user_by_ip .sticky_user_by_ip
.get(&peer.ip()) .get(&peer.ip())
.map(|entry| *entry), .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" "successful runtime-snapshot auth must seed sticky ip cache"
); );
assert_eq!( assert_eq!(
@@ -3047,7 +3050,8 @@ async fn valid_tls_is_blocked_by_per_ip_preauth_throttle_without_saturation() {
let rng = SecureRandom::new(); let rng = SecureRandom::new();
let peer: SocketAddr = "198.51.100.103:45103".parse().unwrap(); 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()), normalize_auth_probe_ip(peer.ip()),
AuthProbeState { AuthProbeState {
fail_streak: AUTH_PROBE_BACKOFF_START_FAILS, 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 peer: SocketAddr = "198.51.100.104:45104".parse().unwrap();
let now = Instant::now(); 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()), normalize_auth_probe_ip(peer.ip()),
AuthProbeState { AuthProbeState {
fail_streak: AUTH_PROBE_BACKOFF_START_FAILS, 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 rng = SecureRandom::new();
let peer: SocketAddr = "198.51.100.205:45205".parse().unwrap(); let peer: SocketAddr = "198.51.100.205:45205".parse().unwrap();
let now = Instant::now(); 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()), normalize_auth_probe_ip(peer.ip()),
AuthProbeState { AuthProbeState {
fail_streak: AUTH_PROBE_BACKOFF_START_FAILS + AUTH_PROBE_SATURATION_GRACE_FAILS, 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 peer: SocketAddr = "198.51.100.106:45106".parse().unwrap();
let now = Instant::now(); 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()), normalize_auth_probe_ip(peer.ip()),
AuthProbeState { AuthProbeState {
fail_streak: AUTH_PROBE_BACKOFF_START_FAILS, 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 replay_checker = ReplayChecker::new(128, Duration::from_secs(60));
let peer: SocketAddr = "198.51.100.206:45206".parse().unwrap(); let peer: SocketAddr = "198.51.100.206:45206".parse().unwrap();
let now = Instant::now(); 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()), normalize_auth_probe_ip(peer.ip()),
AuthProbeState { AuthProbeState {
fail_streak: AUTH_PROBE_BACKOFF_START_FAILS + AUTH_PROBE_SATURATION_GRACE_FAILS, 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 rng = SecureRandom::new();
let peer: SocketAddr = "198.51.100.207:45207".parse().unwrap(); let peer: SocketAddr = "198.51.100.207:45207".parse().unwrap();
let now = Instant::now(); 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()), normalize_auth_probe_ip(peer.ip()),
AuthProbeState { AuthProbeState {
fail_streak: AUTH_PROBE_BACKOFF_START_FAILS, 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 replay_checker = ReplayChecker::new(128, Duration::from_secs(60));
let peer: SocketAddr = "198.51.100.208:45208".parse().unwrap(); let peer: SocketAddr = "198.51.100.208:45208".parse().unwrap();
let now = Instant::now(); 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()), normalize_auth_probe_ip(peer.ip()),
AuthProbeState { AuthProbeState {
fail_streak: AUTH_PROBE_BACKOFF_START_FAILS, 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 rng = SecureRandom::new();
let peer: SocketAddr = "198.51.100.209:45209".parse().unwrap(); let peer: SocketAddr = "198.51.100.209:45209".parse().unwrap();
let now = Instant::now(); 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()), normalize_auth_probe_ip(peer.ip()),
AuthProbeState { AuthProbeState {
fail_streak: AUTH_PROBE_BACKOFF_START_FAILS + AUTH_PROBE_SATURATION_GRACE_FAILS - 1, 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 rng = SecureRandom::new();
let peer: SocketAddr = "198.51.100.210:45210".parse().unwrap(); let peer: SocketAddr = "198.51.100.210:45210".parse().unwrap();
let now = Instant::now(); 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()), normalize_auth_probe_ip(peer.ip()),
AuthProbeState { AuthProbeState {
fail_streak: AUTH_PROBE_BACKOFF_START_FAILS + AUTH_PROBE_SATURATION_GRACE_FAILS, 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 rng = SecureRandom::new();
let peer: SocketAddr = "198.51.100.211:45211".parse().unwrap(); let peer: SocketAddr = "198.51.100.211:45211".parse().unwrap();
let now = Instant::now(); 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()), normalize_auth_probe_ip(peer.ip()),
AuthProbeState { AuthProbeState {
fail_streak: AUTH_PROBE_BACKOFF_START_FAILS + AUTH_PROBE_SATURATION_GRACE_FAILS, 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 rng = Arc::new(SecureRandom::new());
let peer: SocketAddr = "198.51.100.212:45212".parse().unwrap(); let peer: SocketAddr = "198.51.100.212:45212".parse().unwrap();
let now = Instant::now(); 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()), normalize_auth_probe_ip(peer.ip()),
AuthProbeState { AuthProbeState {
fail_streak: AUTH_PROBE_BACKOFF_START_FAILS + AUTH_PROBE_SATURATION_GRACE_FAILS, 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 rng = SecureRandom::new();
let peer: SocketAddr = "198.51.100.213:45213".parse().unwrap(); let peer: SocketAddr = "198.51.100.213:45213".parse().unwrap();
let now = Instant::now(); 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()), normalize_auth_probe_ip(peer.ip()),
AuthProbeState { AuthProbeState {
fail_streak: AUTH_PROBE_BACKOFF_START_FAILS + AUTH_PROBE_SATURATION_GRACE_FAILS, 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 peer: SocketAddr = "198.51.100.110:45110".parse().unwrap();
let now = Instant::now(); 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()), normalize_auth_probe_ip(peer.ip()),
AuthProbeState { AuthProbeState {
fail_streak: AUTH_PROBE_BACKOFF_START_FAILS, fail_streak: AUTH_PROBE_BACKOFF_START_FAILS,
+30 -7
View File
@@ -1,11 +1,9 @@
use std::collections::{BTreeMap, BTreeSet}; use std::collections::{BTreeMap, BTreeSet};
use std::io::Write;
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use std::sync::Arc; use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH}; use std::time::{SystemTime, UNIX_EPOCH};
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use tokio::io::AsyncReadExt;
use tokio::sync::Mutex; use tokio::sync::Mutex;
use tracing::{info, warn}; use tracing::{info, warn};
@@ -173,6 +171,19 @@ fn now_epoch_secs() -> u64 {
} }
async fn read_state_file(path: &Path) -> std::io::Result<Option<QuotaStateFile>> { async fn read_state_file(path: &Path) -> std::io::Result<Option<QuotaStateFile>> {
#[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 { let file = match tokio::fs::File::open(path).await {
Ok(file) => file, Ok(file) => file,
Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None), 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<Option<QuotaStateFile>>
"quota state file grew beyond the 16 MiB limit while reading", "quota state file grew beyond the 16 MiB limit while reading",
)); ));
} }
payload
};
let state = serde_json::from_slice(&payload).map_err(|error| { let state = serde_json::from_slice(&payload).map_err(|error| {
std::io::Error::new( std::io::Error::new(
std::io::ErrorKind::InvalidData, std::io::ErrorKind::InvalidData,
@@ -217,11 +230,6 @@ async fn wait_for_blocking_io<T>(
} }
fn write_state_file_blocking(path: &Path, state: &QuotaStateFile) -> std::io::Result<()> { 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)?; let mut payload = serde_json::to_vec_pretty(state)?;
payload.push(b'\n'); payload.push(b'\n');
if payload.len() as u64 > QUOTA_STATE_MAX_BYTES { 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; let mut last_collision = None;
for _ in 0..8 { for _ in 0..8 {
let tmp_path = path.with_extension(format!( 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", "failed to allocate a unique quota checkpoint temporary file",
) )
})) }))
}
} }
fn quota_user_state(quota: UserQuotaSnapshot) -> QuotaUserState { fn quota_user_state(quota: UserQuotaSnapshot) -> QuotaUserState {
-7
View File
@@ -54,10 +54,7 @@ impl SlotBudget {
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |current| { .fetch_update(Ordering::AcqRel, Ordering::Acquire, |current| {
current.checked_sub(amount) current.checked_sub(amount)
}); });
#[cfg(not(test))]
debug_assert!(released.is_ok(), "slot budget release must match acquisitions"); 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. /// Returns the exact number of currently committed or reserved slots.
@@ -65,10 +62,6 @@ impl SlotBudget {
self.used.load(Ordering::Acquire) 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. /// Provisional slot ownership that rolls back unless committed to a registry entry.
+13 -14
View File
@@ -1,27 +1,26 @@
use std::path::Path; use std::path::Path;
use tokio::io::AsyncReadExt;
use super::*; use super::*;
pub(super) async fn read_disk_entry_bounded(path: &Path) -> std::io::Result<Vec<u8>> { pub(super) async fn read_disk_entry_bounded(path: &Path) -> std::io::Result<Vec<u8>> {
let file = tokio::fs::File::open(path).await?; #[cfg(unix)]
if file.metadata().await?.len() > TLS_FRONT_DISK_ENTRY_MAX_BYTES { {
crate::util::secure_fs::read_regular_limited_async(
path.to_path_buf(),
TLS_FRONT_DISK_ENTRY_MAX_BYTES as usize,
)
.await
}
#[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( return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData, std::io::ErrorKind::InvalidData,
"TLS cache entry exceeds the 1 MiB limit", "TLS cache entry exceeds the 1 MiB limit",
)); ));
} }
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",
));
}
Ok(bytes) Ok(bytes)
}
} }
pub(super) fn cert_info_matches_domain(cached: &CachedTlsData) -> bool { pub(super) fn cert_info_matches_domain(cached: &CachedTlsData) -> bool {
+3 -7
View File
@@ -231,9 +231,6 @@ impl TlsFrontCache {
pub async fn load_from_disk(&self) { pub async fn load_from_disk(&self) {
let path = self.disk_path.clone(); let path = self.disk_path.clone();
if tokio::fs::create_dir_all(&path).await.is_err() {
return;
}
let mut loaded = 0usize; let mut loaded = 0usize;
for name in &self.disk_entry_names { for name in &self.disk_entry_names {
let entry_path = path.join(name); let entry_path = path.join(name);
@@ -297,9 +294,6 @@ impl TlsFrontCache {
} }
async fn persist(&self, domain: &str, data: &CachedTlsData) { 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 fname = format!("{}.json", domain.replace(['/', '\\'], "_"));
let path = self.disk_path.join(fname); let path = self.disk_path.join(fname);
if let Ok(json) = serde_json::to_vec_pretty(data) { if let Ok(json) = serde_json::to_vec_pretty(data) {
@@ -311,7 +305,9 @@ impl TlsFrontCache {
); );
return; 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; let _ = tokio::fs::write(path, json).await;
} }
} }
+26 -13
View File
@@ -12,6 +12,8 @@ use tracing::{debug, info, warn};
use crate::config::ProxyConfig; use crate::config::ProxyConfig;
use crate::error::Result; use crate::error::Result;
use crate::transport::UpstreamManager; use crate::transport::UpstreamManager;
#[cfg(unix)]
use crate::util::secure_fs::{atomic_replace_async, read_regular_limited_async};
use super::MePool; use super::MePool;
use super::http_fetch::{HTTPS_RESPONSE_BODY_MAX_BYTES, https_get}; 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<ProxyConfigData> { pub async fn load_proxy_config_cache(path: &str) -> Result<ProxyConfigData> {
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}")) 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)) Ok(parse_proxy_config_text(&text, 200))
} }
pub async fn save_proxy_config_cache(path: &str, raw_text: &str) -> Result<()> { pub async fn save_proxy_config_cache(path: &str, raw_text: &str) -> Result<()> {
if let Some(parent) = Path::new(path).parent() #[cfg(unix)]
&& !parent.as_os_str().is_empty() let write = atomic_replace_async(
{ Path::new(path).to_path_buf(),
tokio::fs::create_dir_all(parent).await.map_err(|e| { raw_text.as_bytes().to_vec(),
crate::error::ProxyError::Proxy(format!( 0o640,
"create proxy-config cache dir '{}' failed: {e}", )
parent.display() .await;
)) #[cfg(not(unix))]
})?; let write = tokio::fs::write(path, raw_text).await;
} write.map_err(|e| {
tokio::fs::write(path, raw_text).await.map_err(|e| {
crate::error::ProxyError::Proxy(format!("write proxy-config cache '{path}' failed: {e}")) crate::error::ProxyError::Proxy(format!("write proxy-config cache '{path}' failed: {e}"))
})?; })?;
Ok(()) Ok(())
+14 -2
View File
@@ -1,4 +1,5 @@
use httpdate; use httpdate;
use std::path::PathBuf;
use std::sync::Arc; use std::sync::Arc;
use std::time::SystemTime; use std::time::SystemTime;
use tracing::{debug, info, warn}; use tracing::{debug, info, warn};
@@ -7,6 +8,8 @@ use super::http_fetch::https_get;
use super::selftest::record_timeskew_sample; use super::selftest::record_timeskew_sample;
use crate::error::{ProxyError, Result}; use crate::error::{ProxyError, Result};
use crate::transport::UpstreamManager; 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; 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 match download_proxy_secret_with_max_len_via_upstream(max_len, upstream, proxy_secret_url).await
{ {
Ok(data) => { 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)"); warn!(error = %e, "Failed to cache proxy-secret (non-fatal)");
} else { } else {
debug!(path = cache, len = data.len(), "Cached proxy-secret"); 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. // 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() => { Ok(data) if validate_proxy_secret_len(data.len(), max_len).is_ok() => {
let age_hours = tokio::fs::metadata(cache) let age_hours = tokio::fs::metadata(cache)
.await .await
+2
View File
@@ -1,6 +1,8 @@
//! Utils //! Utils
pub mod ip; pub mod ip;
#[cfg(unix)]
pub mod secure_fs;
pub mod time; pub mod time;
#[cfg(unix)] #[cfg(unix)]
pub mod trusted_command; pub mod trusted_command;
+20
View File
@@ -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;
+159
View File
@@ -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> {
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> {
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<Self> {
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<u32>) -> io::Result<Self> {
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<OwnedFd> {
open_dir_components(path, None, false)
}
fn open_or_create_dir_nofollow(path: &Path, mode: u32) -> io::Result<OwnedFd> {
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<OwnedFd> {
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<OwnedFd> {
open_dir_components(path, None, true)
}
fn open_dir_components(
path: &Path,
create_mode: Option<u32>,
require_trusted: bool,
) -> io::Result<OwnedFd> {
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(&current)?;
}
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(&current, 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(&current, name, mode) {
Ok(()) | Err(nix::errno::Errno::EEXIST) => {}
Err(error) => return Err(errno_to_io(error)),
}
openat(&current, 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)
}
+78
View File
@@ -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);
}
+220
View File
@@ -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<std::fs::File> {
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<Vec<u8>> {
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<Vec<u8>> {
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<std::fs::File> {
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<Fd: AsFd>(
parent: Fd,
name: &OsStr,
mode: u32,
) -> io::Result<std::fs::File> {
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<u8>,
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::<u64>());
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)
}