Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
Co-Authored-By: John Preston <17900494+john-preston@users.noreply.github.com>
This commit is contained in:
Alexey
2026-08-23 03:12:11 +03:00
parent 8dbd24b11b
commit 1029703c2c
58 changed files with 7460 additions and 2646 deletions
Generated
+1 -1
View File
@@ -2900,7 +2900,7 @@ checksum = "7b2093cf4c8eb1e67749a6762251bc9cd836b6fc171623bd0a9d324d37af2417"
[[package]]
name = "telemt"
version = "3.5.0"
version = "3.5.1"
dependencies = [
"aes",
"anyhow",
+2 -2
View File
@@ -1,6 +1,6 @@
[package]
name = "telemt"
version = "3.5.0"
version = "3.5.1"
edition = "2024"
[features]
@@ -72,7 +72,7 @@ anyhow = "1.0.102"
reqwest = { version = "0.13.4", features = ["rustls"], default-features = false }
notify = "8.2.0"
ipnetwork = { version = "0.21.1", features = ["serde"] }
hyper = { version = "1.10.1", features = ["server", "http1"] }
hyper = { version = "1.10.1", features = ["client", "server", "http1"] }
hyper-util = { version = "0.1.20", features = ["tokio", "server-auto"] }
http-body-util = "0.1.3"
httpdate = "1.0.3"
File diff suppressed because it is too large Load Diff
-34
View File
@@ -1,34 +0,0 @@
### 3.0.0 Anschluss
- **Middle Proxy now is stable**, confirmed on canary-deploy over ~20 users
- Ad-tag now is working
- DC=203/CDN now is working over ME
- `getProxyConfig` and `ProxySecret` are automated
- Version order is now in format `3.0.0` - without Windows-style "microfixes"
### 3.0.1 Kabelsammler
- Handshake timeouts fixed
- Connectivity logging refactored
- Docker: tmpfs for ProxyConfig and ProxySecret
- Public Host and Port in config
- ME Relays Head-of-Line Blocking fixed
- ME Ping
### 3.0.2 Microtrencher
- New [network] section
- ME Fixes
- Small bugs coverage
### 3.0.3 Ausrutscher
- ME as stateful, no conn-id migration
- No `flush()` on datapath after RpcWriter
- Hightech parser for IPv6 without regexp
- `nat_probe = true` by default
- Timeout for `recv()` in STUN-client
- ConnRegistry review
- Dualstack emergency reconnect
### 3.0.4 Schneeflecken
- Only WARN and Links in Normal log
- Consistent IP-family detection
- Includes for config
- `nonce_frame_hex` in log only with `DEBUG`
+6
View File
@@ -337,9 +337,15 @@ pub(super) fn overlay_hot_fields(old: &ProxyConfig, new: &ProxyConfig) -> ProxyC
cfg.access.user_max_unique_ips_global_each = new.access.user_max_unique_ips_global_each;
cfg.access.user_max_unique_ips_mode = new.access.user_max_unique_ips_mode;
cfg.access.user_max_unique_ips_window_secs = new.access.user_max_unique_ips_window_secs;
let process_limits = cfg.web.limits.clone();
cfg.web = new.web.clone();
cfg.web.limits = process_limits;
if cfg.rebuild_runtime_user_auth().is_err() {
cfg.runtime_user_auth = None;
}
if cfg.rebuild_runtime_web().is_err() {
cfg.web = old.web.clone();
}
cfg
}
+3
View File
@@ -123,6 +123,7 @@ fn listener_synlimit_fields_are_process_owned() {
let mut old = sample_config();
old.server.listeners.push(ListenerConfig {
ip: "0.0.0.0".parse().unwrap(),
transport: crate::config::ListenerTransport::Mtproxy,
port: Some(443),
client_mss: None,
synlimit: SynLimitMode::Iptables,
@@ -138,6 +139,8 @@ fn listener_synlimit_fields_are_process_owned() {
announce_ip: None,
proxy_protocol: None,
reuse_allow: false,
web_client_ip_source: crate::config::WebClientIpSource::XForwardedFor,
web_trusted_proxy_cidrs: Vec::new(),
});
let mut new = old.clone();
new.server.port = 8443;
+12
View File
@@ -22,6 +22,8 @@ mod includes;
mod strict_keys;
// Precomputed user authentication data for handshake hot paths.
mod runtime_auth;
// Validated immutable WEB configuration and static-site snapshots.
mod runtime_web;
// Post-deserialization validation helpers.
mod decode;
mod effective;
@@ -30,6 +32,7 @@ mod validate_core;
mod validate_me;
mod validate_runtime;
mod validate_server;
mod validate_web;
mod validation;
use self::includes::{hash_rendered_snapshot, normalize_config_path, preprocess_includes};
@@ -96,6 +99,10 @@ pub struct ProxyConfig {
#[serde(default)]
pub server: ServerConfig,
/// WEB carrier ingress and public-site fallback configuration.
#[serde(default)]
pub web: WebConfig,
/// Timeout values used by client, fallback, and upstream operations.
#[serde(default)]
pub timeouts: TimeoutsConfig,
@@ -204,6 +211,11 @@ impl ProxyConfig {
Ok(())
}
/// Rebuilds validated WEB capabilities and immutable decoy snapshots.
pub(crate) fn rebuild_runtime_web(&mut self) -> Result<()> {
runtime_web::rebuild(self)
}
pub(crate) fn runtime_user_auth(&self) -> Option<&UserAuthSnapshot> {
self.runtime_user_auth.as_deref()
}
+7
View File
@@ -120,6 +120,7 @@ pub(super) fn apply(config: &mut ProxyConfig) -> Result<()> {
if let Ok(ipv4) = ipv4_str.parse::<IpAddr>() {
config.server.listeners.push(ListenerConfig {
ip: ipv4,
transport: ListenerTransport::Mtproxy,
port: Some(config.server.port),
client_mss: None,
synlimit: SynLimitMode::default(),
@@ -135,6 +136,8 @@ pub(super) fn apply(config: &mut ProxyConfig) -> Result<()> {
announce_ip: None,
proxy_protocol: None,
reuse_allow: false,
web_client_ip_source: WebClientIpSource::XForwardedFor,
web_trusted_proxy_cidrs: Vec::new(),
});
}
if let Some(ipv6_str) = &config.server.listen_addr_ipv6
@@ -142,6 +145,7 @@ pub(super) fn apply(config: &mut ProxyConfig) -> Result<()> {
{
config.server.listeners.push(ListenerConfig {
ip: ipv6,
transport: ListenerTransport::Mtproxy,
port: Some(config.server.port),
client_mss: None,
synlimit: SynLimitMode::default(),
@@ -157,6 +161,8 @@ pub(super) fn apply(config: &mut ProxyConfig) -> Result<()> {
announce_ip: None,
proxy_protocol: None,
reuse_allow: false,
web_client_ip_source: WebClientIpSource::XForwardedFor,
web_trusted_proxy_cidrs: Vec::new(),
});
}
}
@@ -211,5 +217,6 @@ pub(super) fn apply(config: &mut ProxyConfig) -> Result<()> {
validate_logging_config(&config.logging)?;
validate_upstreams(config)?;
config.rebuild_runtime_user_auth()?;
config.rebuild_runtime_web()?;
Ok(())
}
+1
View File
@@ -7,6 +7,7 @@ pub(super) fn load_source_graph(graph: ConfigSourceGraph) -> Result<LoadedConfig
validate_runtime::validate(&mut config)?;
validate_me::validate(&mut config)?;
validate_server::validate(&mut config)?;
validate_web::validate(&mut config)?;
effective::apply(&mut config)?;
Ok(LoadedConfig {
config,
+414
View File
@@ -0,0 +1,414 @@
use std::collections::{BTreeMap, HashSet};
use std::fs;
use std::io::Read;
use std::path::Path;
use std::sync::Arc;
#[cfg(unix)]
use std::os::unix::fs::OpenOptionsExt;
use bytes::Bytes;
use hmac::{Hmac, Mac};
use sha2::{Digest, Sha256};
use super::*;
const WEB_CAPABILITY_CONTEXT: &[u8] = b"tdesktop-web-proxy-bridge-v1\n";
const MAX_WEB_STATIC_DEPTH: usize = 64;
/// Builds the immutable WEB routing and decoy snapshot for one generation.
pub(super) fn rebuild(config: &mut ProxyConfig) -> Result<()> {
let auth = config.runtime_user_auth().ok_or_else(|| {
ProxyError::Config("WEB runtime requires the user authentication snapshot".to_string())
})?;
let mut runtime_vhosts = BTreeMap::new();
let mut runtime_profiles = Vec::new();
let mut static_files = 0usize;
let mut static_bytes = 0usize;
for vhost in &config.web.vhosts {
let decoy = build_decoy(
vhost,
&config.web.limits,
&mut static_files,
&mut static_bytes,
)?;
let mut profiles = Vec::with_capacity(vhost.profiles.len());
let mut capabilities = HashSet::with_capacity(vhost.profiles.len());
for profile in &vhost.profiles {
let user_id = auth.user_id_by_name(&profile.user).ok_or_else(|| {
ProxyError::Config(format!(
"WEB profile references unknown access user `{}`",
profile.user
))
})?;
let auth_entry = auth.entry_by_id(user_id).ok_or_else(|| {
ProxyError::Config("WEB profile user snapshot is inconsistent".to_string())
})?;
let (client_secret, client_secret_len) =
client_secret(auth_entry.secret, profile.secret_mode);
let capability = derive_web_capability(
&client_secret[..client_secret_len],
vhost.host.as_bytes(),
)?;
if !capabilities.insert(capability) {
return Err(ProxyError::Config(format!(
"WEB vhost `{}` contains profiles with the same client capability",
vhost.host
)));
}
let runtime_profile = Arc::new(WebRuntimeProfile {
host: vhost.host.clone(),
public_addr: vhost.public_addr,
user: profile.user.clone(),
secret_mode: profile.secret_mode,
capability,
max_sessions: profile
.max_sessions
.unwrap_or(config.web.limits.max_sessions_global),
max_streams: profile
.max_streams
.unwrap_or(config.web.limits.max_streams_global),
max_streams_per_session: profile
.max_streams_per_session
.unwrap_or(config.web.limits.max_streams_per_session),
});
profiles.push(Arc::clone(&runtime_profile));
runtime_profiles.push(runtime_profile);
}
runtime_vhosts.insert(
vhost.host.clone(),
Arc::new(WebRuntimeVhost {
host: vhost.host.clone(),
decoy,
decoy_header_secs: config.web.timeouts.decoy_header_secs,
profiles,
}),
);
}
config.web.runtime = Some(Arc::new(WebRuntimeConfig {
vhosts: runtime_vhosts,
profiles: runtime_profiles,
}));
Ok(())
}
/// Derives the Telegram Desktop WEB capability for one exact secret and host.
pub(crate) fn derive_web_capability(secret: &[u8], host: &[u8]) -> Result<[u8; 32]> {
let mut mac = Hmac::<Sha256>::new_from_slice(secret).map_err(|_| {
ProxyError::Config("WEB capability secret must not be empty".to_string())
})?;
mac.update(WEB_CAPABILITY_CONTEXT);
mac.update(host);
Ok(mac.finalize().into_bytes().into())
}
fn client_secret(secret: [u8; 16], mode: WebSecretMode) -> ([u8; 17], usize) {
let mut client_secret = [0u8; 17];
match mode {
WebSecretMode::Plain => {
client_secret[..16].copy_from_slice(&secret);
(client_secret, 16)
}
WebSecretMode::Dd => {
client_secret[0] = 0xdd;
client_secret[1..].copy_from_slice(&secret);
(client_secret, 17)
}
}
}
fn build_decoy(
vhost: &WebVhostConfig,
limits: &WebLimitsConfig,
static_files: &mut usize,
static_bytes: &mut usize,
) -> Result<WebRuntimeDecoy> {
match &vhost.decoy {
WebDecoyConfig::HttpUpstream { upstream } => {
let parsed = url::Url::parse(upstream).map_err(|error| {
ProxyError::Config(format!(
"WEB decoy upstream for `{}` is invalid: {error}",
vhost.host
))
})?;
let ip = match parsed.host() {
Some(url::Host::Ipv4(ip)) => std::net::IpAddr::V4(ip),
Some(url::Host::Ipv6(ip)) => std::net::IpAddr::V6(ip),
_ => {
return Err(ProxyError::Config(
"WEB decoy host must be an IP literal".to_string(),
));
}
};
let host = ip.to_string();
let port = parsed.port_or_known_default().ok_or_else(|| {
ProxyError::Config("WEB decoy port cannot be resolved".to_string())
})?;
let authority = match (ip, parsed.port()) {
(std::net::IpAddr::V6(_), Some(_)) => format!("[{host}]:{port}"),
(std::net::IpAddr::V6(_), None) => format!("[{host}]"),
(std::net::IpAddr::V4(_), Some(_)) => format!("{host}:{port}"),
(std::net::IpAddr::V4(_), None) => host.clone(),
};
Ok(WebRuntimeDecoy::HttpUpstream {
addr: SocketAddr::new(ip, port),
authority,
})
}
WebDecoyConfig::StaticDirectory { directory, index } => {
let site = load_static_site(
directory,
index,
limits,
static_files,
static_bytes,
)?;
Ok(WebRuntimeDecoy::StaticDirectory(Arc::new(site)))
}
}
}
fn load_static_site(
root: &Path,
index: &str,
limits: &WebLimitsConfig,
total_files: &mut usize,
total_bytes: &mut usize,
) -> Result<WebStaticSite> {
let root_metadata = fs::symlink_metadata(root).map_err(|error| {
ProxyError::Config(format!(
"failed to inspect WEB static directory `{}`: {error}",
root.display()
))
})?;
if root_metadata.file_type().is_symlink() || !root_metadata.is_dir() {
return Err(ProxyError::Config(format!(
"WEB static directory `{}` must be a real directory, not a symlink",
root.display()
)));
}
let canonical_root = fs::canonicalize(root).map_err(|error| {
ProxyError::Config(format!(
"failed to canonicalize WEB static directory `{}`: {error}",
root.display()
))
})?;
let mut assets = BTreeMap::new();
load_static_directory(
&canonical_root,
&canonical_root,
&mut assets,
total_files,
total_bytes,
limits,
0,
)?;
if !assets.contains_key(&format!("/{index}")) {
return Err(ProxyError::Config(format!(
"WEB static directory `{}` does not contain index `{index}`",
root.display()
)));
}
Ok(WebStaticSite {
assets,
index: index.to_string(),
})
}
fn load_static_directory(
root: &Path,
directory: &Path,
assets: &mut BTreeMap<String, WebStaticAsset>,
total_files: &mut usize,
total_bytes: &mut usize,
limits: &WebLimitsConfig,
depth: usize,
) -> Result<()> {
let entries = fs::read_dir(directory).map_err(|error| {
ProxyError::Config(format!(
"failed to read WEB static directory `{}`: {error}",
directory.display()
))
})?;
for entry in entries {
let entry = entry.map_err(|error| {
ProxyError::Config(format!("failed to read WEB static entry: {error}"))
})?;
if *total_files >= limits.max_static_files {
return Err(ProxyError::Config(
"WEB static entries exceed process-wide web.limits.max_static_files"
.to_string(),
));
}
*total_files += 1;
let path = entry.path();
let file_type = entry.file_type().map_err(|error| {
ProxyError::Config(format!(
"failed to inspect WEB static entry `{}`: {error}",
path.display()
))
})?;
if file_type.is_symlink() {
return Err(ProxyError::Config(format!(
"WEB static entry `{}` must not be a symlink",
path.display()
)));
}
if file_type.is_dir() {
if depth >= MAX_WEB_STATIC_DEPTH {
return Err(ProxyError::Config(format!(
"WEB static directory `{}` exceeds the maximum nesting depth",
path.display()
)));
}
load_static_directory(
root,
&path,
assets,
total_files,
total_bytes,
limits,
depth + 1,
)?;
continue;
}
if !file_type.is_file() {
return Err(ProxyError::Config(format!(
"WEB static entry `{}` must be a regular file",
path.display()
)));
}
let mut options = fs::OpenOptions::new();
options.read(true);
#[cfg(unix)]
options.custom_flags(libc::O_CLOEXEC | libc::O_NOFOLLOW);
let file = options.open(&path).map_err(|error| {
ProxyError::Config(format!(
"failed to open WEB static file `{}`: {error}",
path.display()
))
})?;
let metadata = file.metadata().map_err(|error| {
ProxyError::Config(format!(
"failed to inspect WEB static file `{}`: {error}",
path.display()
))
})?;
if !metadata.is_file() {
return Err(ProxyError::Config(format!(
"WEB static entry `{}` changed before it was opened",
path.display()
)));
}
let file_len = usize::try_from(metadata.len()).map_err(|_| {
ProxyError::Config(format!("WEB static file `{}` is too large", path.display()))
})?;
if file_len > limits.max_static_file_bytes {
return Err(ProxyError::Config(format!(
"WEB static file `{}` exceeds web.limits.max_static_file_bytes",
path.display()
)));
}
*total_bytes = total_bytes.checked_add(file_len).ok_or_else(|| {
ProxyError::Config("WEB static snapshot byte count overflowed usize".to_string())
})?;
if *total_bytes > limits.max_static_bytes {
return Err(ProxyError::Config(
"WEB static snapshots exceed process-wide web.limits.max_static_bytes"
.to_string(),
));
}
let relative = path.strip_prefix(root).map_err(|_| {
ProxyError::Config("WEB static path escaped its configured root".to_string())
})?;
let route = static_route(relative)?;
let mut body = Vec::with_capacity(file_len);
file.take(limits.max_static_file_bytes as u64 + 1)
.read_to_end(&mut body)
.map_err(|error| {
ProxyError::Config(format!(
"failed to read WEB static file `{}`: {error}",
path.display()
))
})?;
if body.len() != file_len {
return Err(ProxyError::Config(format!(
"WEB static file `{}` changed while its snapshot was built",
path.display()
)));
}
let etag = format!("\"{}\"", hex::encode(Sha256::digest(&body)));
assets.insert(
route,
WebStaticAsset {
body: Bytes::from(body),
content_type: static_content_type(&path),
etag,
},
);
}
Ok(())
}
fn static_route(relative: &Path) -> Result<String> {
let mut route = String::new();
for component in relative.components() {
let std::path::Component::Normal(component) = component else {
return Err(ProxyError::Config(
"WEB static path contains an unsafe component".to_string(),
));
};
let component = component.to_str().ok_or_else(|| {
ProxyError::Config("WEB static file names must be valid UTF-8".to_string())
})?;
route.push('/');
route.push_str(component);
}
Ok(route)
}
fn static_content_type(path: &Path) -> &'static str {
match path.extension().and_then(|extension| extension.to_str()) {
Some("html") | Some("htm") => "text/html; charset=utf-8",
Some("css") => "text/css; charset=utf-8",
Some("js") | Some("mjs") => "text/javascript; charset=utf-8",
Some("json") => "application/json",
Some("txt") => "text/plain; charset=utf-8",
Some("svg") => "image/svg+xml",
Some("png") => "image/png",
Some("jpg") | Some("jpeg") => "image/jpeg",
Some("gif") => "image/gif",
Some("webp") => "image/webp",
Some("ico") => "image/x-icon",
Some("woff") => "font/woff",
Some("woff2") => "font/woff2",
Some("wasm") => "application/wasm",
_ => "application/octet-stream",
}
}
#[cfg(test)]
mod tests {
use base64::Engine as _;
use super::*;
#[test]
fn capability_matches_reference_vectors() {
let secret = hex::decode("000102030405060708090a0b0c0d0e0f").unwrap();
let plain = derive_web_capability(&secret, b"proxy.example.com").unwrap();
assert_eq!(
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(plain),
"MHLEY5PmW1GWqJkSrlmJpvJUiLhBH_QKy6yKg8a0JPk"
);
let mut dd_secret = vec![0xdd];
dd_secret.extend_from_slice(&secret);
let dd = derive_web_capability(&dd_secret, b"proxy.example.com").unwrap();
assert_eq!(
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(dd),
"IpJrt3e7sKtzPyoXy6w-Zj6GGEvsvclN66JzQEfPYLA"
);
}
}
+71 -293
View File
@@ -7,6 +7,7 @@ const TOP_LEVEL_CONFIG_KEYS: &[&str] = &[
"logging",
"network",
"server",
"web",
"timeouts",
"censorship",
"access",
@@ -238,6 +239,7 @@ const CONNTRACK_CONTROL_CONFIG_KEYS: &[&str] = &[
const LISTENER_CONFIG_KEYS: &[&str] = &[
"ip",
"transport",
"port",
"client_mss",
"synlimit",
@@ -253,6 +255,70 @@ const LISTENER_CONFIG_KEYS: &[&str] = &[
"announce_ip",
"proxy_protocol",
"reuse_allow",
"web_client_ip_source",
"web_trusted_proxy_cidrs",
];
const WEB_CONFIG_KEYS: &[&str] = &["enabled", "limits", "timeouts", "vhosts"];
const WEB_LIMITS_CONFIG_KEYS: &[&str] = &[
"max_header_bytes",
"max_body_bytes",
"max_frame_payload_bytes",
"carrier_batch_bytes",
"max_frames_per_body",
"max_http_connections",
"max_http_handlers",
"max_body_readers",
"max_body_bytes_global",
"max_sessions_global",
"max_sessions_per_ip",
"max_streams_per_session",
"max_streams_global",
"max_stream_handshakes",
"max_tombstones_per_session",
"pending_bytes_per_session",
"pending_bytes_global",
"pending_items_per_session",
"pending_items_global",
"control_bytes_per_session",
"control_bytes_global",
"max_bootstraps_global",
"max_bootstraps_per_ip",
"max_vhosts",
"max_profiles",
"max_static_files",
"max_static_file_bytes",
"max_static_bytes",
"memory_envelope_bytes",
"new_bootstraps_per_minute",
"new_bootstraps_burst",
"new_sessions_per_minute",
"new_sessions_burst",
"new_streams_per_minute",
"new_streams_burst",
];
const WEB_TIMEOUTS_CONFIG_KEYS: &[&str] = &[
"header_secs",
"body_secs",
"stream_handshake_secs",
"long_poll_secs",
"bootstrap_lifetime_secs",
"reconnect_grace_secs",
"http_idle_secs",
"shutdown_secs",
"decoy_header_secs",
];
const WEB_VHOST_CONFIG_KEYS: &[&str] = &["host", "public_addr", "decoy", "profiles"];
const WEB_DECOY_CONFIG_KEYS: &[&str] = &["mode", "upstream", "directory", "index"];
const WEB_PROFILE_CONFIG_KEYS: &[&str] = &[
"user",
"secret_mode",
"max_sessions",
"max_streams",
"max_streams_per_session",
];
const TIMEOUTS_CONFIG_KEYS: &[&str] = &[
@@ -366,300 +432,12 @@ const LOGGING_CONFIG_KEYS: &[&str] = &[
"max_age_secs",
];
#[derive(Debug)]
struct UnknownConfigKey {
path: String,
suggestion: Option<String>,
}
fn table_at<'a>(value: &'a toml::Value, path: &[&str]) -> Option<&'a toml::Table> {
let mut current = value;
for segment in path {
current = current.get(*segment)?;
}
current.as_table()
}
fn is_strict_config(parsed_toml: &toml::Value) -> bool {
table_at(parsed_toml, &["general"])
.and_then(|table| table.get("config_strict"))
.and_then(toml::Value::as_bool)
.unwrap_or(false)
}
fn known_config_keys_for_suggestion() -> Vec<&'static str> {
let mut keys = Vec::new();
for group in [
TOP_LEVEL_CONFIG_KEYS,
GENERAL_CONFIG_KEYS,
NETWORK_CONFIG_KEYS,
SERVER_CONFIG_KEYS,
API_CONFIG_KEYS,
CONNTRACK_CONTROL_CONFIG_KEYS,
LISTENER_CONFIG_KEYS,
TIMEOUTS_CONFIG_KEYS,
CENSORSHIP_CONFIG_KEYS,
TLS_FETCH_CONFIG_KEYS,
ACCESS_CONFIG_KEYS,
RATE_LIMIT_BPS_CONFIG_KEYS,
UPSTREAM_CONFIG_KEYS,
PROXY_MODES_CONFIG_KEYS,
TELEMETRY_CONFIG_KEYS,
LINKS_CONFIG_KEYS,
LOGGING_CONFIG_KEYS,
] {
keys.extend_from_slice(group);
}
keys
}
fn levenshtein_distance(a: &str, b: &str) -> usize {
let b_chars: Vec<char> = b.chars().collect();
let mut prev: Vec<usize> = (0..=b_chars.len()).collect();
let mut curr = vec![0usize; b_chars.len() + 1];
for (i, ca) in a.chars().enumerate() {
curr[0] = i + 1;
for (j, cb) in b_chars.iter().enumerate() {
let replace = if ca == *cb { prev[j] } else { prev[j] + 1 };
curr[j + 1] = (prev[j + 1] + 1).min(curr[j] + 1).min(replace);
}
std::mem::swap(&mut prev, &mut curr);
}
prev[b_chars.len()]
}
fn unknown_key_suggestion(key: &str, known_keys: &[&'static str]) -> Option<String> {
let normalized = key.to_ascii_lowercase();
let mut best: Option<(&str, usize)> = None;
for known in known_keys {
let distance = levenshtein_distance(&normalized, known);
let is_better = match best {
Some((_, best_distance)) => distance < best_distance,
None => true,
};
if distance <= 4 && is_better {
best = Some((known, distance));
}
}
best.map(|(known, _)| known.to_string())
}
fn push_unknown_keys(
unknown: &mut Vec<UnknownConfigKey>,
known_for_suggestion: &[&'static str],
path: &str,
table: &toml::Table,
allowed: &[&str],
) {
for key in table.keys() {
if !allowed.contains(&key.as_str()) {
let full_path = if path.is_empty() {
key.clone()
} else {
format!("{path}.{key}")
};
unknown.push(UnknownConfigKey {
path: full_path,
suggestion: unknown_key_suggestion(key, known_for_suggestion),
});
}
}
}
fn check_known_table(
parsed_toml: &toml::Value,
unknown: &mut Vec<UnknownConfigKey>,
known_for_suggestion: &[&'static str],
path: &[&str],
allowed: &[&str],
) {
if let Some(table) = table_at(parsed_toml, path) {
push_unknown_keys(
unknown,
known_for_suggestion,
&path.join("."),
table,
allowed,
);
}
}
fn check_nested_table_value(
unknown: &mut Vec<UnknownConfigKey>,
known_for_suggestion: &[&'static str],
path: String,
value: &toml::Value,
allowed: &[&str],
) {
if let Some(table) = value.as_table() {
push_unknown_keys(unknown, known_for_suggestion, &path, table, allowed);
}
}
fn collect_unknown_config_keys(parsed_toml: &toml::Value) -> Vec<UnknownConfigKey> {
let known_for_suggestion = known_config_keys_for_suggestion();
let mut unknown = Vec::new();
if let Some(root) = parsed_toml.as_table() {
push_unknown_keys(
&mut unknown,
&known_for_suggestion,
"",
root,
TOP_LEVEL_CONFIG_KEYS,
);
}
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["general"],
GENERAL_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["general", "modes"],
PROXY_MODES_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["general", "telemetry"],
TELEMETRY_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["general", "links"],
LINKS_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["logging"],
LOGGING_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["network"],
NETWORK_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["server"],
SERVER_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["server", "api"],
API_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["server", "admin_api"],
API_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["server", "conntrack_control"],
CONNTRACK_CONTROL_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["timeouts"],
TIMEOUTS_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["censorship"],
CENSORSHIP_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["censorship", "tls_fetch"],
TLS_FETCH_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["access"],
ACCESS_CONFIG_KEYS,
);
if let Some(listeners) = table_at(parsed_toml, &["server"])
.and_then(|table| table.get("listeners"))
.and_then(toml::Value::as_array)
{
for (idx, listener) in listeners.iter().enumerate() {
check_nested_table_value(
&mut unknown,
&known_for_suggestion,
format!("server.listeners[{idx}]"),
listener,
LISTENER_CONFIG_KEYS,
);
}
}
if let Some(upstreams) = parsed_toml.get("upstreams").and_then(toml::Value::as_array) {
for (idx, upstream) in upstreams.iter().enumerate() {
check_nested_table_value(
&mut unknown,
&known_for_suggestion,
format!("upstreams[{idx}]"),
upstream,
UPSTREAM_CONFIG_KEYS,
);
}
}
for access_map in ["user_rate_limits", "cidr_rate_limits"] {
if let Some(table) = table_at(parsed_toml, &["access"])
.and_then(|access| access.get(access_map))
.and_then(toml::Value::as_table)
{
for (entry_name, value) in table {
check_nested_table_value(
&mut unknown,
&known_for_suggestion,
format!("access.{access_map}.{entry_name}"),
value,
RATE_LIMIT_BPS_CONFIG_KEYS,
);
}
}
}
unknown
}
// Recursive table traversal and key suggestion logic.
mod check;
/// Rejects or reports unknown configuration keys according to strict mode.
pub(super) fn handle_unknown_config_keys(parsed_toml: &toml::Value) -> Result<()> {
let unknown = collect_unknown_config_keys(parsed_toml);
let unknown = check::collect_unknown_config_keys(parsed_toml);
if unknown.is_empty() {
return Ok(());
}
@@ -676,7 +454,7 @@ pub(super) fn handle_unknown_config_keys(parsed_toml: &toml::Value) -> Result<()
}
}
if is_strict_config(parsed_toml) {
if check::is_strict_config(parsed_toml) {
let mut paths = Vec::with_capacity(unknown.len());
for item in unknown {
if let Some(suggestion) = item.suggestion {
+362
View File
@@ -0,0 +1,362 @@
use super::*;
#[derive(Debug)]
/// One rejected configuration path and its optional nearest known key.
pub(super) struct UnknownConfigKey {
/// Fully qualified configuration path.
pub(super) path: String,
/// Nearest known key when edit distance is sufficiently small.
pub(super) suggestion: Option<String>,
}
fn table_at<'a>(value: &'a toml::Value, path: &[&str]) -> Option<&'a toml::Table> {
let mut current = value;
for segment in path {
current = current.get(*segment)?;
}
current.as_table()
}
/// Reads strict-key enforcement without deserializing the full configuration.
pub(super) fn is_strict_config(parsed_toml: &toml::Value) -> bool {
table_at(parsed_toml, &["general"])
.and_then(|table| table.get("config_strict"))
.and_then(toml::Value::as_bool)
.unwrap_or(false)
}
fn known_config_keys_for_suggestion() -> Vec<&'static str> {
let mut keys = Vec::new();
for group in [
TOP_LEVEL_CONFIG_KEYS,
GENERAL_CONFIG_KEYS,
NETWORK_CONFIG_KEYS,
SERVER_CONFIG_KEYS,
API_CONFIG_KEYS,
CONNTRACK_CONTROL_CONFIG_KEYS,
LISTENER_CONFIG_KEYS,
WEB_CONFIG_KEYS,
WEB_LIMITS_CONFIG_KEYS,
WEB_TIMEOUTS_CONFIG_KEYS,
WEB_VHOST_CONFIG_KEYS,
WEB_DECOY_CONFIG_KEYS,
WEB_PROFILE_CONFIG_KEYS,
TIMEOUTS_CONFIG_KEYS,
CENSORSHIP_CONFIG_KEYS,
TLS_FETCH_CONFIG_KEYS,
ACCESS_CONFIG_KEYS,
RATE_LIMIT_BPS_CONFIG_KEYS,
UPSTREAM_CONFIG_KEYS,
PROXY_MODES_CONFIG_KEYS,
TELEMETRY_CONFIG_KEYS,
LINKS_CONFIG_KEYS,
LOGGING_CONFIG_KEYS,
] {
keys.extend_from_slice(group);
}
keys
}
fn levenshtein_distance(a: &str, b: &str) -> usize {
let b_chars: Vec<char> = b.chars().collect();
let mut prev: Vec<usize> = (0..=b_chars.len()).collect();
let mut curr = vec![0usize; b_chars.len() + 1];
for (i, ca) in a.chars().enumerate() {
curr[0] = i + 1;
for (j, cb) in b_chars.iter().enumerate() {
let replace = if ca == *cb { prev[j] } else { prev[j] + 1 };
curr[j + 1] = (prev[j + 1] + 1).min(curr[j] + 1).min(replace);
}
std::mem::swap(&mut prev, &mut curr);
}
prev[b_chars.len()]
}
fn unknown_key_suggestion(key: &str, known_keys: &[&'static str]) -> Option<String> {
let normalized = key.to_ascii_lowercase();
let mut best: Option<(&str, usize)> = None;
for known in known_keys {
let distance = levenshtein_distance(&normalized, known);
let is_better = match best {
Some((_, best_distance)) => distance < best_distance,
None => true,
};
if distance <= 4 && is_better {
best = Some((known, distance));
}
}
best.map(|(known, _)| known.to_string())
}
fn push_unknown_keys(
unknown: &mut Vec<UnknownConfigKey>,
known_for_suggestion: &[&'static str],
path: &str,
table: &toml::Table,
allowed: &[&str],
) {
for key in table.keys() {
if !allowed.contains(&key.as_str()) {
let full_path = if path.is_empty() {
key.clone()
} else {
format!("{path}.{key}")
};
unknown.push(UnknownConfigKey {
path: full_path,
suggestion: unknown_key_suggestion(key, known_for_suggestion),
});
}
}
}
fn check_known_table(
parsed_toml: &toml::Value,
unknown: &mut Vec<UnknownConfigKey>,
known_for_suggestion: &[&'static str],
path: &[&str],
allowed: &[&str],
) {
if let Some(table) = table_at(parsed_toml, path) {
push_unknown_keys(
unknown,
known_for_suggestion,
&path.join("."),
table,
allowed,
);
}
}
fn check_nested_table_value(
unknown: &mut Vec<UnknownConfigKey>,
known_for_suggestion: &[&'static str],
path: String,
value: &toml::Value,
allowed: &[&str],
) {
if let Some(table) = value.as_table() {
push_unknown_keys(unknown, known_for_suggestion, &path, table, allowed);
}
}
/// Collects unknown keys across every supported nested configuration table.
pub(super) fn collect_unknown_config_keys(parsed_toml: &toml::Value) -> Vec<UnknownConfigKey> {
let known_for_suggestion = known_config_keys_for_suggestion();
let mut unknown = Vec::new();
if let Some(root) = parsed_toml.as_table() {
push_unknown_keys(
&mut unknown,
&known_for_suggestion,
"",
root,
TOP_LEVEL_CONFIG_KEYS,
);
}
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["general"],
GENERAL_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["general", "modes"],
PROXY_MODES_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["general", "telemetry"],
TELEMETRY_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["general", "links"],
LINKS_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["logging"],
LOGGING_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["network"],
NETWORK_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["server"],
SERVER_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["server", "api"],
API_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["server", "admin_api"],
API_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["server", "conntrack_control"],
CONNTRACK_CONTROL_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["web"],
WEB_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["web", "limits"],
WEB_LIMITS_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["web", "timeouts"],
WEB_TIMEOUTS_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["timeouts"],
TIMEOUTS_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["censorship"],
CENSORSHIP_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["censorship", "tls_fetch"],
TLS_FETCH_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["access"],
ACCESS_CONFIG_KEYS,
);
if let Some(listeners) = table_at(parsed_toml, &["server"])
.and_then(|table| table.get("listeners"))
.and_then(toml::Value::as_array)
{
for (idx, listener) in listeners.iter().enumerate() {
check_nested_table_value(
&mut unknown,
&known_for_suggestion,
format!("server.listeners[{idx}]"),
listener,
LISTENER_CONFIG_KEYS,
);
}
}
if let Some(vhosts) = table_at(parsed_toml, &["web"])
.and_then(|table| table.get("vhosts"))
.and_then(toml::Value::as_array)
{
for (vhost_idx, vhost) in vhosts.iter().enumerate() {
check_nested_table_value(
&mut unknown,
&known_for_suggestion,
format!("web.vhosts[{vhost_idx}]"),
vhost,
WEB_VHOST_CONFIG_KEYS,
);
if let Some(vhost) = vhost.as_table() {
if let Some(decoy) = vhost.get("decoy") {
check_nested_table_value(
&mut unknown,
&known_for_suggestion,
format!("web.vhosts[{vhost_idx}].decoy"),
decoy,
WEB_DECOY_CONFIG_KEYS,
);
}
if let Some(profiles) = vhost.get("profiles").and_then(toml::Value::as_array) {
for (profile_idx, profile) in profiles.iter().enumerate() {
check_nested_table_value(
&mut unknown,
&known_for_suggestion,
format!("web.vhosts[{vhost_idx}].profiles[{profile_idx}]"),
profile,
WEB_PROFILE_CONFIG_KEYS,
);
}
}
}
}
}
if let Some(upstreams) = parsed_toml.get("upstreams").and_then(toml::Value::as_array) {
for (idx, upstream) in upstreams.iter().enumerate() {
check_nested_table_value(
&mut unknown,
&known_for_suggestion,
format!("upstreams[{idx}]"),
upstream,
UPSTREAM_CONFIG_KEYS,
);
}
}
for access_map in ["user_rate_limits", "cidr_rate_limits"] {
if let Some(table) = table_at(parsed_toml, &["access"])
.and_then(|access| access.get(access_map))
.and_then(toml::Value::as_table)
{
for (entry_name, value) in table {
check_nested_table_value(
&mut unknown,
&known_for_suggestion,
format!("access.{access_map}.{entry_name}"),
value,
RATE_LIMIT_BPS_CONFIG_KEYS,
);
}
}
}
unknown
}
+545
View File
@@ -0,0 +1,545 @@
use std::collections::HashSet;
use super::*;
const WEB_FRAME_HEADER_BYTES: usize = 8;
const WEB_QUEUE_ITEM_COST: usize = 256;
const WEB_CONTROL_EXTRA_ITEMS: usize = 16;
const WEB_CONTROL_ITEMS_PER_STREAM: usize = 3;
const WEB_INITIAL_STREAM_WINDOW: usize = 4 * 1024 * 1024;
const MAX_WEB_HEADER_BYTES: usize = 64 * 1024;
const MAX_WEB_BODY_BYTES: usize = 16 * 1024 * 1024;
const MAX_WEB_FRAME_BYTES: usize = 1024 * 1024;
const MAX_WEB_FRAMES_PER_BODY: usize = 4096;
const MAX_WEB_TOMBSTONES_PER_SESSION: usize = 4096;
const MAX_WEB_MEMORY_ENVELOPE_BYTES: usize = 4 * 1024 * 1024 * 1024;
/// Validates WEB policy and resource bounds before building runtime state.
pub(super) fn validate(config: &mut ProxyConfig) -> Result<()> {
let web_listener_count = config
.server
.listeners
.iter()
.filter(|listener| listener.transport == ListenerTransport::Web)
.count();
let eligible_web_listener_count = config
.server
.listeners
.iter()
.filter(|listener| listener.transport == ListenerTransport::Web)
.filter(|listener| {
(listener.ip.is_ipv4() && config.network.ipv4)
|| (listener.ip.is_ipv6() && config.network.ipv6 != Some(false))
})
.count();
for (idx, listener) in config.server.listeners.iter().enumerate() {
match listener.transport {
ListenerTransport::Mtproxy => {
if !listener.web_trusted_proxy_cidrs.is_empty() {
return Err(ProxyError::Config(format!(
"server.listeners[{idx}].web_trusted_proxy_cidrs is only valid for transport=web"
)));
}
}
ListenerTransport::Web => validate_web_listener(config, idx, listener)?,
}
}
if config.web.enabled && eligible_web_listener_count == 0 {
return Err(ProxyError::Config(
"web.enabled requires at least one network-eligible server.listeners entry with transport=web"
.to_string(),
));
}
if web_listener_count > 0 && config.web.vhosts.is_empty() {
return Err(ProxyError::Config(
"WEB listeners require at least one [[web.vhosts]] entry".to_string(),
));
}
validate_limits(&config.web.limits)?;
validate_timeouts(&config.web.timeouts)?;
validate_vhosts(config)?;
Ok(())
}
fn validate_web_listener(
config: &ProxyConfig,
idx: usize,
listener: &ListenerConfig,
) -> Result<()> {
if listener.web_trusted_proxy_cidrs.is_empty() {
return Err(ProxyError::Config(format!(
"server.listeners[{idx}].web_trusted_proxy_cidrs must be non-empty for transport=web"
)));
}
if listener
.web_trusted_proxy_cidrs
.iter()
.any(|network| network.prefix() == 0)
{
return Err(ProxyError::Config(format!(
"server.listeners[{idx}].web_trusted_proxy_cidrs must not contain a /0 network"
)));
}
let proxy_protocol = listener
.proxy_protocol
.unwrap_or(config.server.proxy_protocol);
if proxy_protocol {
return Err(ProxyError::Config(format!(
"server.listeners[{idx}].proxy_protocol must be false for transport=web; WEB identity is accepted only from the configured L7 header"
)));
}
if listener.reuse_allow {
return Err(ProxyError::Config(format!(
"server.listeners[{idx}].reuse_allow is not supported for transport=web without external session affinity"
)));
}
if listener.client_mss.is_some()
|| listener.synlimit != SynLimitMode::Off
|| listener.announce.is_some()
|| listener.announce_ip.is_some()
{
return Err(ProxyError::Config(format!(
"server.listeners[{idx}] WEB transport does not accept client_mss, synlimit, announce, or announce_ip"
)));
}
Ok(())
}
fn validate_limits(limits: &WebLimitsConfig) -> Result<()> {
if !(8192..=MAX_WEB_HEADER_BYTES).contains(&limits.max_header_bytes) {
return config_error("web.limits.max_header_bytes must be within [8192, 65536]");
}
if !(WEB_FRAME_HEADER_BYTES..=MAX_WEB_BODY_BYTES).contains(&limits.max_body_bytes) {
return config_error("web.limits.max_body_bytes must be within [8, 16777216]");
}
if !(1..=MAX_WEB_FRAME_BYTES).contains(&limits.max_frame_payload_bytes) {
return config_error("web.limits.max_frame_payload_bytes must be within [1, 1048576]");
}
if !(1..=MAX_WEB_FRAMES_PER_BODY).contains(&limits.max_frames_per_body) {
return config_error("web.limits.max_frames_per_body must be within [1, 4096]");
}
if !(1..=MAX_WEB_TOMBSTONES_PER_SESSION).contains(&limits.max_tombstones_per_session) {
return config_error("web.limits.max_tombstones_per_session must be within [1, 4096]");
}
if limits.carrier_batch_bytes > limits.max_body_bytes
|| limits.carrier_batch_bytes
< limits
.max_frame_payload_bytes
.saturating_add(WEB_FRAME_HEADER_BYTES)
{
return config_error(
"web.limits.carrier_batch_bytes must fit max_body_bytes and one maximum frame",
);
}
if limits.max_frame_payload_bytes > WEB_INITIAL_STREAM_WINDOW {
return config_error(
"web.limits.max_frame_payload_bytes must not exceed the initial stream window",
);
}
let positive = [
("max_http_connections", limits.max_http_connections),
("max_http_handlers", limits.max_http_handlers),
("max_body_readers", limits.max_body_readers),
("max_body_bytes_global", limits.max_body_bytes_global),
("max_sessions_global", limits.max_sessions_global),
("max_sessions_per_ip", limits.max_sessions_per_ip),
("max_streams_per_session", limits.max_streams_per_session),
("max_streams_global", limits.max_streams_global),
("max_stream_handshakes", limits.max_stream_handshakes),
("pending_bytes_per_session", limits.pending_bytes_per_session),
("pending_bytes_global", limits.pending_bytes_global),
("pending_items_per_session", limits.pending_items_per_session),
("pending_items_global", limits.pending_items_global),
("control_bytes_per_session", limits.control_bytes_per_session),
("control_bytes_global", limits.control_bytes_global),
("max_bootstraps_global", limits.max_bootstraps_global),
("max_bootstraps_per_ip", limits.max_bootstraps_per_ip),
("max_vhosts", limits.max_vhosts),
("max_profiles", limits.max_profiles),
("max_static_files", limits.max_static_files),
("max_static_file_bytes", limits.max_static_file_bytes),
("max_static_bytes", limits.max_static_bytes),
("memory_envelope_bytes", limits.memory_envelope_bytes),
];
if let Some((field, _)) = positive.into_iter().find(|(_, value)| *value == 0) {
return config_error(&format!("web.limits.{field} must be > 0"));
}
for (field, value) in [
("max_http_connections", limits.max_http_connections),
("max_http_handlers", limits.max_http_handlers),
("max_body_readers", limits.max_body_readers),
("max_body_bytes_global", limits.max_body_bytes_global),
("max_stream_handshakes", limits.max_stream_handshakes),
] {
if value > tokio::sync::Semaphore::MAX_PERMITS {
return config_error(&format!("web.limits.{field} exceeds Tokio semaphore capacity"));
}
}
let rates = [
("new_bootstraps_per_minute", limits.new_bootstraps_per_minute),
("new_bootstraps_burst", limits.new_bootstraps_burst),
("new_sessions_per_minute", limits.new_sessions_per_minute),
("new_sessions_burst", limits.new_sessions_burst),
("new_streams_per_minute", limits.new_streams_per_minute),
("new_streams_burst", limits.new_streams_burst),
];
if let Some((field, _)) = rates.into_iter().find(|(_, value)| *value == 0) {
return config_error(&format!("web.limits.{field} must be > 0"));
}
if limits.max_streams_per_session > u16::MAX as usize {
return config_error("web.limits.max_streams_per_session must fit synthetic source ports");
}
if limits.max_sessions_per_ip > limits.max_sessions_global
|| limits.max_streams_per_session > limits.max_streams_global
|| limits.max_stream_handshakes > limits.max_streams_global
|| limits.max_bootstraps_per_ip > limits.max_bootstraps_global
|| limits.max_http_handlers > limits.max_http_connections
|| limits.max_body_readers > limits.max_http_handlers
|| limits.pending_bytes_per_session > limits.pending_bytes_global
|| limits.pending_items_per_session > limits.pending_items_global
|| limits.control_bytes_per_session > limits.control_bytes_global
|| limits.control_bytes_per_session > limits.pending_bytes_per_session
|| limits.control_bytes_global > limits.pending_bytes_global
|| limits.max_static_file_bytes > limits.max_static_bytes
{
return config_error("web.limits per-owner ceilings must not exceed global ceilings");
}
let control_items_per_session = WEB_CONTROL_EXTRA_ITEMS
.checked_add(
limits
.max_streams_per_session
.checked_mul(WEB_CONTROL_ITEMS_PER_STREAM)
.ok_or_else(|| {
ProxyError::Config(
"web.limits control item reservation overflowed usize".to_string(),
)
})?,
)
.ok_or_else(|| {
ProxyError::Config("web.limits control item reservation overflowed usize".to_string())
})?;
let control_items_global = control_items_per_session
.checked_mul(limits.max_sessions_global)
.ok_or_else(|| {
ProxyError::Config("web.limits global control reservation overflowed usize".to_string())
})?;
let control_frame_cost = WEB_FRAME_HEADER_BYTES + 4 + WEB_QUEUE_ITEM_COST;
let required_control_bytes_per_session = control_items_per_session
.checked_mul(control_frame_cost)
.ok_or_else(|| {
ProxyError::Config("web.limits control byte reservation overflowed usize".to_string())
})?;
let required_control_bytes_global = control_items_global
.checked_mul(control_frame_cost)
.ok_or_else(|| {
ProxyError::Config("web.limits global control byte reservation overflowed usize".to_string())
})?;
if control_items_per_session >= limits.pending_items_per_session
|| control_items_global >= limits.pending_items_global
|| required_control_bytes_per_session > limits.control_bytes_per_session
|| required_control_bytes_global > limits.control_bytes_global
{
return config_error(
"web.limits control reserves must cover bounded control frames and leave data capacity",
);
}
let uplink_bytes = limits
.max_frames_per_body
.checked_mul(WEB_QUEUE_ITEM_COST)
.and_then(|value| value.checked_add(limits.max_body_bytes))
.ok_or_else(|| {
ProxyError::Config("web.limits uplink reservation overflowed usize".to_string())
})?;
let minimum_downlink_frame_bytes = WEB_FRAME_HEADER_BYTES + 1 + WEB_QUEUE_ITEM_COST;
let session_required_bytes = limits
.control_bytes_per_session
.checked_add(uplink_bytes)
.and_then(|value| value.checked_add(minimum_downlink_frame_bytes))
.ok_or_else(|| {
ProxyError::Config("web.limits session reservation overflowed usize".to_string())
})?;
let session_required_items = control_items_per_session
.checked_add(limits.max_frames_per_body)
.and_then(|value| value.checked_add(1))
.ok_or_else(|| {
ProxyError::Config("web.limits session item reservation overflowed usize".to_string())
})?;
let global_required_bytes = limits
.control_bytes_global
.checked_add(uplink_bytes)
.and_then(|value| value.checked_add(minimum_downlink_frame_bytes))
.ok_or_else(|| {
ProxyError::Config("web.limits global reservation overflowed usize".to_string())
})?;
let global_required_items = control_items_global
.checked_add(limits.max_frames_per_body)
.and_then(|value| value.checked_add(1))
.ok_or_else(|| {
ProxyError::Config("web.limits global item reservation overflowed usize".to_string())
})?;
if session_required_bytes > limits.pending_bytes_per_session
|| session_required_items > limits.pending_items_per_session
|| global_required_bytes > limits.pending_bytes_global
|| global_required_items > limits.pending_items_global
{
return config_error(
"web.limits pending ceilings must preserve one uplink batch and downlink progress",
);
}
let body_reservation = limits
.max_body_readers
.checked_mul(limits.max_body_bytes)
.ok_or_else(|| {
ProxyError::Config("web.limits body reader reservation overflowed usize".to_string())
})?;
if body_reservation > limits.max_body_bytes_global
|| limits.max_body_bytes_global > u32::MAX as usize
{
return config_error(
"web.limits max_body_readers * max_body_bytes must fit max_body_bytes_global and u32",
);
}
let http_header_reservation = limits
.max_http_connections
.checked_mul(limits.max_header_bytes)
.ok_or_else(|| {
ProxyError::Config("web.limits HTTP header reservations overflow usize".to_string())
})?;
let reserved = limits
.pending_bytes_global
.checked_add(limits.max_body_bytes_global)
.and_then(|value| value.checked_add(limits.max_static_bytes))
.and_then(|value| value.checked_add(http_header_reservation))
.ok_or_else(|| ProxyError::Config("web.limits byte ceilings overflow usize".to_string()))?;
if reserved > limits.memory_envelope_bytes
|| limits.memory_envelope_bytes > MAX_WEB_MEMORY_ENVELOPE_BYTES
{
return config_error(
"web.limits memory reservations must fit memory_envelope_bytes within 4 GiB",
);
}
Ok(())
}
fn validate_timeouts(timeouts: &WebTimeoutsConfig) -> Result<()> {
let values = [
("header_secs", timeouts.header_secs),
("body_secs", timeouts.body_secs),
("stream_handshake_secs", timeouts.stream_handshake_secs),
("long_poll_secs", timeouts.long_poll_secs),
("bootstrap_lifetime_secs", timeouts.bootstrap_lifetime_secs),
("reconnect_grace_secs", timeouts.reconnect_grace_secs),
("http_idle_secs", timeouts.http_idle_secs),
("shutdown_secs", timeouts.shutdown_secs),
("decoy_header_secs", timeouts.decoy_header_secs),
];
if let Some((field, _)) = values
.into_iter()
.find(|(_, value)| !(1..=3600).contains(value))
{
return config_error(&format!("web.timeouts.{field} must be within [1, 3600]"));
}
let request_deadline = timeouts
.header_secs
.max(timeouts.body_secs)
.max(timeouts.long_poll_secs)
.max(timeouts.decoy_header_secs);
if request_deadline >= timeouts.http_idle_secs {
return config_error("web.timeouts request deadlines must be lower than http_idle_secs");
}
Ok(())
}
fn validate_vhosts(config: &mut ProxyConfig) -> Result<()> {
let limits = &config.web.limits;
if config.web.vhosts.len() > limits.max_vhosts {
return config_error("web.vhosts exceeds web.limits.max_vhosts");
}
let mut hosts = HashSet::with_capacity(config.web.vhosts.len());
let mut profile_count = 0usize;
for (vhost_idx, vhost) in config.web.vhosts.iter_mut().enumerate() {
vhost.host = normalize_web_host(
&vhost.host,
&format!("web.vhosts[{vhost_idx}].host"),
)?;
if !hosts.insert(vhost.host.clone()) {
return config_error(&format!("duplicate WEB vhost host `{}`", vhost.host));
}
if vhost.public_addr.port() != 443 || vhost.public_addr.ip().is_unspecified() {
return config_error(&format!(
"web.vhosts[{vhost_idx}].public_addr must be a concrete socket address on port 443"
));
}
if config.web.enabled && vhost.profiles.is_empty() {
return config_error(&format!(
"web.vhosts[{vhost_idx}].profiles must be non-empty when web.enabled=true"
));
}
validate_decoy(vhost_idx, &vhost.decoy)?;
let mut profiles = HashSet::with_capacity(vhost.profiles.len());
for (profile_idx, profile) in vhost.profiles.iter().enumerate() {
if !config.access.users.contains_key(&profile.user) {
return config_error(&format!(
"web.vhosts[{vhost_idx}].profiles[{profile_idx}].user references unknown access user `{}`",
profile.user
));
}
if !profiles.insert((profile.user.as_str(), profile.secret_mode)) {
return config_error(&format!(
"duplicate WEB profile for user `{}` in vhost `{}`",
profile.user, vhost.host
));
}
let max_streams = profile.max_streams.unwrap_or(limits.max_streams_global);
let max_streams_per_session = profile
.max_streams_per_session
.unwrap_or(limits.max_streams_per_session);
if profile.max_sessions == Some(0)
|| profile.max_sessions.is_some_and(|value| value > limits.max_sessions_global)
|| profile.max_streams == Some(0)
|| profile
.max_streams
.is_some_and(|value| value > limits.max_streams_global)
|| profile.max_streams_per_session == Some(0)
|| profile
.max_streams_per_session
.is_some_and(|value| value > limits.max_streams_per_session)
|| max_streams_per_session > max_streams
{
return config_error(&format!(
"web.vhosts[{vhost_idx}].profiles[{profile_idx}] limits must be non-zero and within global WEB limits"
));
}
profile_count = profile_count.checked_add(1).ok_or_else(|| {
ProxyError::Config("WEB profile count overflowed usize".to_string())
})?;
}
}
if profile_count > limits.max_profiles {
return config_error("WEB profiles exceed web.limits.max_profiles");
}
Ok(())
}
fn normalize_web_host(value: &str, field: &str) -> Result<String> {
let input = value.trim();
if input.is_empty()
|| input.ends_with('.')
|| input
.chars()
.any(|character| matches!(character, ':' | '/' | '?' | '#' | '@'))
{
return config_error(&format!(
"{field} must be a hostname without a port, path, credentials, or trailing dot"
));
}
let host = normalize_domain_to_ascii(input, field)?;
if host.len() > 253
|| !host.contains('.')
|| host.parse::<IpAddr>().is_ok()
|| web_host_last_label_is_numeric(&host)
{
return config_error(&format!(
"{field} must be a non-IP fully-qualified hostname accepted by Telegram Desktop"
));
}
for label in host.split('.') {
if label.is_empty()
|| label.len() > 63
|| label.starts_with('-')
|| label.ends_with('-')
|| !label
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || byte == b'-')
{
return config_error(&format!(
"{field} contains a hostname label rejected by Telegram Desktop"
));
}
}
Ok(host)
}
fn web_host_last_label_is_numeric(host: &str) -> bool {
let label = host.rsplit('.').next().unwrap_or_default();
let digits = label
.strip_prefix("0x")
.or_else(|| label.strip_prefix("0X"));
if let Some(digits) = digits {
return digits.bytes().all(|byte| byte.is_ascii_hexdigit());
}
label.bytes().all(|byte| byte.is_ascii_digit())
}
fn validate_decoy(vhost_idx: usize, decoy: &WebDecoyConfig) -> Result<()> {
match decoy {
WebDecoyConfig::HttpUpstream { upstream } => {
let parsed = url::Url::parse(upstream).map_err(|error| {
ProxyError::Config(format!(
"web.vhosts[{vhost_idx}].decoy.upstream is invalid: {error}"
))
})?;
if parsed.scheme() != "http"
|| parsed.host_str().is_none()
|| !parsed.username().is_empty()
|| parsed.password().is_some()
|| parsed.query().is_some()
|| parsed.fragment().is_some()
|| parsed.path() != "/"
|| parsed.port() == Some(0)
{
return config_error(&format!(
"web.vhosts[{vhost_idx}].decoy.upstream must be an http origin without credentials, path, query, or fragment"
));
}
let ip = match parsed.host() {
Some(url::Host::Ipv4(ip)) => IpAddr::V4(ip),
Some(url::Host::Ipv6(ip)) => IpAddr::V6(ip),
_ => {
return config_error(&format!(
"web.vhosts[{vhost_idx}].decoy.upstream host must be a loopback or private IP literal"
));
}
};
let private = match ip {
IpAddr::V4(ip) => ip.is_loopback() || ip.is_private() || ip.is_link_local(),
IpAddr::V6(ip) => {
ip.is_loopback() || ip.is_unique_local() || ip.is_unicast_link_local()
}
};
if !private {
return config_error(&format!(
"web.vhosts[{vhost_idx}].decoy.upstream must remain inside loopback or a private network"
));
}
}
WebDecoyConfig::StaticDirectory { directory, index } => {
if !directory.is_absolute() {
return config_error(&format!(
"web.vhosts[{vhost_idx}].decoy.directory must be absolute"
));
}
if index.is_empty()
|| index.contains('\\')
|| std::path::Path::new(index).components().count() != 1
|| matches!(index.as_str(), "." | "..")
{
return config_error(&format!(
"web.vhosts[{vhost_idx}].decoy.index must be one safe file name"
));
}
}
}
Ok(())
}
fn config_error<T>(message: &str) -> Result<T> {
Err(ProxyError::Config(message.to_string()))
}
#[cfg(test)]
mod tests;
+26
View File
@@ -0,0 +1,26 @@
use super::*;
#[test]
fn web_host_normalization_matches_client_vectors() {
assert_eq!(
normalize_web_host(" Proxy.Example.COM ", "host").unwrap(),
"proxy.example.com"
);
assert_eq!(
normalize_web_host("bücher.example", "host").unwrap(),
"xn--bcher-kva.example"
);
for invalid in [
"localhost",
"127.0.0.1",
"127.1",
"0x7f.1",
"0177.0.0.1",
"1.2.3",
"site.example:443",
"site..example",
"site.example.",
] {
assert!(normalize_web_host(invalid, "host").is_err(), "{invalid}");
}
}
+2
View File
@@ -52,3 +52,5 @@ mod synlimit_mss_tests;
mod tls_fetch_tests;
#[path = "load_basic_tests/upstream_tests.rs"]
mod upstream_tests;
#[path = "load_basic_tests/web_tests.rs"]
mod web_tests;
@@ -0,0 +1,96 @@
use super::*;
const WEB_CONFIG: &str = r#"
[access.users]
alice = "000102030405060708090a0b0c0d0e0f"
[[server.listeners]]
ip = "127.0.0.1"
port = 18080
transport = "web"
proxy_protocol = false
web_client_ip_source = "x_forwarded_for"
web_trusted_proxy_cidrs = ["127.0.0.1/32"]
[web]
enabled = true
[[web.vhosts]]
host = "Proxy.Example.COM"
public_addr = "203.0.113.10:443"
[web.vhosts.decoy]
mode = "http_upstream"
upstream = "http://127.0.0.1:18081"
[[web.vhosts.profiles]]
user = "alice"
secret_mode = "dd"
max_sessions = 4
max_streams = 64
max_streams_per_session = 16
"#;
#[test]
fn web_config_builds_canonical_runtime_snapshot() {
let config = load_config_from_temp_toml(WEB_CONFIG);
let runtime = config.web.runtime.expect("WEB runtime snapshot");
let vhost = runtime
.vhosts
.get("proxy.example.com")
.expect("canonical WEB vhost");
assert_eq!(vhost.profiles.len(), 1);
assert_eq!(vhost.profiles[0].user, "alice");
assert_eq!(vhost.profiles[0].secret_mode, WebSecretMode::Dd);
assert_eq!(vhost.profiles[0].max_sessions, 4);
assert_eq!(vhost.profiles[0].max_streams, 64);
assert_eq!(vhost.profiles[0].max_streams_per_session, 16);
}
#[test]
fn web_listener_requires_an_explicit_trusted_proxy() {
let invalid = WEB_CONFIG.replace(
"web_trusted_proxy_cidrs = [\"127.0.0.1/32\"]",
"web_trusted_proxy_cidrs = []",
);
let error = load_config_error_from_temp_toml(&invalid);
assert!(error.contains("web_trusted_proxy_cidrs must be non-empty"));
}
#[test]
fn web_queue_limits_preserve_control_and_uplink_progress() {
let invalid = WEB_CONFIG.replace(
"[web]\nenabled = true",
"[web]\nenabled = true\n\n[web.limits]\ncontrol_bytes_per_session = 1",
);
let error = load_config_error_from_temp_toml(&invalid);
assert!(error.contains("control reserves must cover bounded control frames"));
}
#[test]
fn web_semaphore_limits_are_rejected_before_runtime_construction() {
let invalid = WEB_CONFIG.replace(
"[web]\nenabled = true",
&format!(
"[web]\nenabled = true\n\n[web.limits]\nmax_http_connections = {}",
tokio::sync::Semaphore::MAX_PERMITS + 1,
),
);
let error = load_config_error_from_temp_toml(&invalid);
assert!(error.contains("exceeds Tokio semaphore capacity"));
}
#[test]
fn web_ipv6_decoy_uses_a_valid_http_authority() {
let ipv6 = WEB_CONFIG.replace(
"http://127.0.0.1:18081",
"http://[::1]:18081",
);
let config = load_config_from_temp_toml(&ipv6);
let runtime = config.web.runtime.expect("WEB runtime snapshot");
let vhost = runtime.vhosts.get("proxy.example.com").unwrap();
let WebRuntimeDecoy::HttpUpstream { authority, .. } = &vhost.decoy else {
panic!("expected HTTP decoy");
};
assert_eq!(authority, "[::1]:18081");
}
+12 -1
View File
@@ -23,6 +23,7 @@ mod logging;
mod network;
mod policies;
mod server;
mod web;
pub use access::{AccessConfig, CidrRateLimitKey, RateLimitBps};
#[allow(unused_imports)]
@@ -43,7 +44,17 @@ pub use policies::{
pub use server::{
CLIENT_MSS_2IN8, CLIENT_MSS_EXTREME_LOW, CLIENT_MSS_MAX, CLIENT_MSS_MIN, CLIENT_MSS_TSPU,
ConntrackBackend, ConntrackControlConfig, ConntrackMode, ConntrackPressureProfile,
ListenerConfig, ServerConfig, SynLimitMode, TimeoutsConfig,
ListenerConfig, ListenerTransport, ServerConfig, SynLimitMode, TimeoutsConfig,
WebClientIpSource,
};
#[allow(unused_imports)]
pub use web::{
WebConfig, WebDecoyConfig, WebLimitsConfig, WebProfileConfig, WebSecretMode,
WebTimeoutsConfig, WebVhostConfig,
};
pub(crate) use web::{
WebRuntimeConfig, WebRuntimeDecoy, WebRuntimeProfile, WebRuntimeVhost, WebStaticAsset,
WebStaticSite,
};
fn default_quota_state_path() -> PathBuf {
+29
View File
@@ -75,6 +75,26 @@ pub enum SynLimitMode {
Pf,
}
/// Application protocol accepted by one process-owned TCP listener.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(rename_all = "lowercase")]
pub enum ListenerTransport {
/// Existing MTProxy TCP listener behavior.
#[default]
Mtproxy,
/// Plain HTTP WEB gateway behind a trusted TLS terminator.
Web,
}
/// Trusted L7 source used to recover a WEB client's identity address.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(rename_all = "snake_case")]
pub enum WebClientIpSource {
/// Require exactly one canonical IP in `X-Forwarded-For`.
#[default]
XForwardedFor,
}
impl Serialize for SynLimitMode {
fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
where
@@ -380,6 +400,9 @@ impl Default for TimeoutsConfig {
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ListenerConfig {
pub ip: IpAddr,
/// Application protocol accepted by this listener.
#[serde(default)]
pub transport: ListenerTransport,
/// Per-listener TCP port. If omitted, falls back to legacy `server.port`.
#[serde(default)]
pub port: Option<u16>,
@@ -429,6 +452,12 @@ pub struct ListenerConfig {
/// Default is false for safety.
#[serde(default)]
pub reuse_allow: bool,
/// L7 header policy used only by WEB listeners.
#[serde(default)]
pub web_client_ip_source: WebClientIpSource,
/// Immediate socket peers allowed to provide the WEB client identity header.
#[serde(default)]
pub web_trusted_proxy_cidrs: Vec<IpNetwork>,
}
/// Client-facing TCP MSS preset for extreme-low fragmentation profiles.
+434
View File
@@ -0,0 +1,434 @@
use std::collections::BTreeMap;
use std::net::SocketAddr;
use std::path::PathBuf;
use std::sync::Arc;
use bytes::Bytes;
use serde::{Deserialize, Serialize};
/// Client-facing secret representation used to derive a WEB capability.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum WebSecretMode {
/// Use the existing 16-byte access secret without a prefix.
Plain,
/// Prefix the existing access secret with `0xdd` for capability derivation.
Dd,
}
/// One access user explicitly exposed through a WEB virtual host.
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct WebProfileConfig {
/// Existing `[access.users]` key authenticated by the inner MTProxy handshake.
pub user: String,
/// Exact client-facing secret representation advertised in WEB links.
pub secret_mode: WebSecretMode,
/// Optional per-profile live session ceiling.
#[serde(default)]
pub max_sessions: Option<usize>,
/// Optional per-profile live logical-stream ceiling.
#[serde(default)]
pub max_streams: Option<usize>,
/// Optional per-profile stream ceiling for one session.
#[serde(default)]
pub max_streams_per_session: Option<usize>,
}
/// Public-site fallback used for requests that are not authenticated WEB traffic.
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "mode", rename_all = "snake_case")]
pub enum WebDecoyConfig {
/// Stream requests to one fixed private HTTP origin.
HttpUpstream {
/// Origin URL without a query or fragment.
upstream: String,
},
/// Serve an immutable, bounded snapshot of a local directory.
StaticDirectory {
/// Absolute directory containing public files.
directory: PathBuf,
/// File served for `/` and directory paths.
#[serde(default = "default_web_static_index")]
index: String,
},
}
/// One externally visible WEB hostname and its explicit access profiles.
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct WebVhostConfig {
/// Canonical lowercase ACE hostname used by Telegram Desktop.
pub host: String,
/// Stable public destination tuple used by inner relay routing and KDF metadata.
pub public_addr: SocketAddr,
/// Ordinary-site fallback for this hostname.
pub decoy: WebDecoyConfig,
/// Access users and exact secret modes enabled for this hostname.
#[serde(default)]
pub profiles: Vec<WebProfileConfig>,
}
/// Hard process and protocol limits for WEB ingress.
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct WebLimitsConfig {
/// Maximum bytes accepted while parsing one HTTP request head.
#[serde(default = "default_web_max_header_bytes")]
pub max_header_bytes: usize,
/// Maximum collected carrier request body size.
#[serde(default = "default_web_max_body_bytes")]
pub max_body_bytes: usize,
/// Maximum payload carried by one WEB frame.
#[serde(default = "default_web_max_frame_payload_bytes")]
pub max_frame_payload_bytes: usize,
/// Maximum encoded downlink batch returned by one poll.
#[serde(default = "default_web_carrier_batch_bytes")]
pub carrier_batch_bytes: usize,
/// Maximum frame count parsed or emitted in one carrier body.
#[serde(default = "default_web_max_frames_per_body")]
pub max_frames_per_body: usize,
/// Process-wide accepted WEB HTTP connection ceiling.
#[serde(default = "default_web_max_http_connections")]
pub max_http_connections: usize,
/// Process-wide concurrently executing HTTP handler ceiling.
#[serde(default = "default_web_max_http_handlers")]
pub max_http_handlers: usize,
/// Process-wide concurrently collected request body ceiling.
#[serde(default = "default_web_max_body_readers")]
pub max_body_readers: usize,
/// Process-wide byte reservation for collected request bodies.
#[serde(default = "default_web_max_body_bytes_global")]
pub max_body_bytes_global: usize,
/// Process-wide live WEB session ceiling.
#[serde(default = "default_web_max_sessions_global")]
pub max_sessions_global: usize,
/// Live WEB session ceiling for one forwarded client address.
#[serde(default = "default_web_max_sessions_per_ip")]
pub max_sessions_per_ip: usize,
/// Default live logical-stream ceiling for one WEB session.
#[serde(default = "default_web_max_streams_per_session")]
pub max_streams_per_session: usize,
/// Process-wide live logical-stream ceiling.
#[serde(default = "default_web_max_streams_global")]
pub max_streams_global: usize,
/// Process-wide concurrent inner MTProxy handshake ceiling.
#[serde(default = "default_web_max_stream_handshakes")]
pub max_stream_handshakes: usize,
/// Closed stream identifiers retained by one session.
#[serde(default = "default_web_max_tombstones")]
pub max_tombstones_per_session: usize,
/// Total queued data and control bytes allowed for one session.
#[serde(default = "default_web_pending_bytes_per_session")]
pub pending_bytes_per_session: usize,
/// Process-wide queued data and control byte ceiling.
#[serde(default = "default_web_pending_bytes_global")]
pub pending_bytes_global: usize,
/// Total queued data and control item ceiling for one session.
#[serde(default = "default_web_pending_items_per_session")]
pub pending_items_per_session: usize,
/// Process-wide queued data and control item ceiling.
#[serde(default = "default_web_pending_items_global")]
pub pending_items_global: usize,
/// Per-session byte reserve available only to control frames.
#[serde(default = "default_web_control_bytes_per_session")]
pub control_bytes_per_session: usize,
/// Process-wide byte reserve available only to control frames.
#[serde(default = "default_web_control_bytes_global")]
pub control_bytes_global: usize,
/// Process-wide live bootstrap credential ceiling.
#[serde(default = "default_web_max_bootstraps_global")]
pub max_bootstraps_global: usize,
/// Live bootstrap credential ceiling for one forwarded client address.
#[serde(default = "default_web_max_bootstraps_per_ip")]
pub max_bootstraps_per_ip: usize,
/// Maximum configured WEB virtual-host count.
#[serde(default = "default_web_max_vhosts")]
pub max_vhosts: usize,
/// Maximum configured WEB access-profile count across all virtual hosts.
#[serde(default = "default_web_max_profiles")]
pub max_profiles: usize,
/// Maximum static snapshot entry count across all virtual hosts.
#[serde(default = "default_web_max_static_files")]
pub max_static_files: usize,
/// Maximum bytes read from one static snapshot file.
#[serde(default = "default_web_max_static_file_bytes")]
pub max_static_file_bytes: usize,
/// Maximum static snapshot bytes across all virtual hosts.
#[serde(default = "default_web_max_static_bytes")]
pub max_static_bytes: usize,
/// Declared process envelope for HTTP heads, bodies, queues, and static snapshots.
#[serde(default = "default_web_memory_envelope_bytes")]
pub memory_envelope_bytes: usize,
/// Sustained process-wide bootstrap issuance rate.
#[serde(default = "default_web_new_bootstraps_per_minute")]
pub new_bootstraps_per_minute: u32,
/// Process-wide bootstrap issuance burst.
#[serde(default = "default_web_new_bootstraps_burst")]
pub new_bootstraps_burst: u32,
/// Sustained process-wide session creation rate.
#[serde(default = "default_web_new_sessions_per_minute")]
pub new_sessions_per_minute: u32,
/// Process-wide session creation burst.
#[serde(default = "default_web_new_sessions_burst")]
pub new_sessions_burst: u32,
/// Sustained process-wide logical-stream creation rate.
#[serde(default = "default_web_new_streams_per_minute")]
pub new_streams_per_minute: u32,
/// Process-wide logical-stream creation burst.
#[serde(default = "default_web_new_streams_burst")]
pub new_streams_burst: u32,
}
impl Default for WebLimitsConfig {
fn default() -> Self {
Self {
max_header_bytes: default_web_max_header_bytes(),
max_body_bytes: default_web_max_body_bytes(),
max_frame_payload_bytes: default_web_max_frame_payload_bytes(),
carrier_batch_bytes: default_web_carrier_batch_bytes(),
max_frames_per_body: default_web_max_frames_per_body(),
max_http_connections: default_web_max_http_connections(),
max_http_handlers: default_web_max_http_handlers(),
max_body_readers: default_web_max_body_readers(),
max_body_bytes_global: default_web_max_body_bytes_global(),
max_sessions_global: default_web_max_sessions_global(),
max_sessions_per_ip: default_web_max_sessions_per_ip(),
max_streams_per_session: default_web_max_streams_per_session(),
max_streams_global: default_web_max_streams_global(),
max_stream_handshakes: default_web_max_stream_handshakes(),
max_tombstones_per_session: default_web_max_tombstones(),
pending_bytes_per_session: default_web_pending_bytes_per_session(),
pending_bytes_global: default_web_pending_bytes_global(),
pending_items_per_session: default_web_pending_items_per_session(),
pending_items_global: default_web_pending_items_global(),
control_bytes_per_session: default_web_control_bytes_per_session(),
control_bytes_global: default_web_control_bytes_global(),
max_bootstraps_global: default_web_max_bootstraps_global(),
max_bootstraps_per_ip: default_web_max_bootstraps_per_ip(),
max_vhosts: default_web_max_vhosts(),
max_profiles: default_web_max_profiles(),
max_static_files: default_web_max_static_files(),
max_static_file_bytes: default_web_max_static_file_bytes(),
max_static_bytes: default_web_max_static_bytes(),
memory_envelope_bytes: default_web_memory_envelope_bytes(),
new_bootstraps_per_minute: default_web_new_bootstraps_per_minute(),
new_bootstraps_burst: default_web_new_bootstraps_burst(),
new_sessions_per_minute: default_web_new_sessions_per_minute(),
new_sessions_burst: default_web_new_sessions_burst(),
new_streams_per_minute: default_web_new_streams_per_minute(),
new_streams_burst: default_web_new_streams_burst(),
}
}
}
/// Deadlines for WEB HTTP, bootstrap, session, and shutdown lifecycle.
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct WebTimeoutsConfig {
/// Deadline for receiving one complete HTTP request head.
#[serde(default = "default_web_header_timeout_secs")]
pub header_secs: u64,
/// Deadline for collecting one authenticated carrier request body.
#[serde(default = "default_web_body_timeout_secs")]
pub body_secs: u64,
/// Deadline for the inner MTProxy handshake on one logical stream.
#[serde(default = "default_web_stream_handshake_timeout_secs")]
pub stream_handshake_secs: u64,
/// Maximum wait for one empty downlink long poll.
#[serde(default = "default_web_long_poll_timeout_secs")]
pub long_poll_secs: u64,
/// Lifetime of an unused bootstrap credential and closed-token replay marker.
#[serde(default = "default_web_bootstrap_lifetime_secs")]
pub bootstrap_lifetime_secs: u64,
/// Maximum carrier inactivity before a session is closed.
#[serde(default = "default_web_reconnect_grace_secs")]
pub reconnect_grace_secs: u64,
/// Maximum idle lifetime of a WEB HTTP keep-alive connection.
#[serde(default = "default_web_http_idle_secs")]
pub http_idle_secs: u64,
/// Maximum graceful wait for WEB connections and process-owned tasks.
#[serde(default = "default_web_shutdown_secs")]
pub shutdown_secs: u64,
/// Deadline for connecting to and receiving headers from an HTTP decoy.
#[serde(default = "default_web_decoy_header_timeout_secs")]
pub decoy_header_secs: u64,
}
impl Default for WebTimeoutsConfig {
fn default() -> Self {
Self {
header_secs: default_web_header_timeout_secs(),
body_secs: default_web_body_timeout_secs(),
stream_handshake_secs: default_web_stream_handshake_timeout_secs(),
long_poll_secs: default_web_long_poll_timeout_secs(),
bootstrap_lifetime_secs: default_web_bootstrap_lifetime_secs(),
reconnect_grace_secs: default_web_reconnect_grace_secs(),
http_idle_secs: default_web_http_idle_secs(),
shutdown_secs: default_web_shutdown_secs(),
decoy_header_secs: default_web_decoy_header_timeout_secs(),
}
}
}
/// WEB ingress, carrier, fallback, and lifecycle configuration.
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct WebConfig {
/// Enables issuance of new WEB bridge and session credentials.
#[serde(default)]
pub enabled: bool,
/// Hard process and protocol limits.
#[serde(default)]
pub limits: WebLimitsConfig,
/// WEB lifecycle deadlines.
#[serde(default)]
pub timeouts: WebTimeoutsConfig,
/// Public hostnames served by WEB listeners.
#[serde(default)]
pub vhosts: Vec<WebVhostConfig>,
/// Validated immutable runtime snapshot built during configuration loading.
#[serde(skip)]
pub(crate) runtime: Option<Arc<WebRuntimeConfig>>,
}
/// Precomputed WEB configuration consumed by listener hot paths.
#[derive(Debug)]
pub(crate) struct WebRuntimeConfig {
/// Canonical host lookup used by HTTP request routing.
pub(crate) vhosts: BTreeMap<String, Arc<WebRuntimeVhost>>,
/// Flat profile inventory used by startup link emission.
pub(crate) profiles: Vec<Arc<WebRuntimeProfile>>,
}
/// Precomputed immutable virtual-host data.
#[derive(Debug)]
pub(crate) struct WebRuntimeVhost {
/// Canonical lowercase ACE hostname.
pub(crate) host: String,
/// Immutable ordinary-site fallback snapshot.
pub(crate) decoy: WebRuntimeDecoy,
/// Upstream connect and response-head deadline.
pub(crate) decoy_header_secs: u64,
/// Exact capability profiles accepted by this host.
pub(crate) profiles: Vec<Arc<WebRuntimeProfile>>,
}
/// Precomputed exact-user capability entry.
#[derive(Debug)]
pub(crate) struct WebRuntimeProfile {
/// Canonical host that owns this profile.
pub(crate) host: String,
/// Stable public destination tuple supplied to relay routing.
pub(crate) public_addr: SocketAddr,
/// Exact access user authenticated by logical streams.
pub(crate) user: String,
/// Client secret representation and inner protocol policy.
pub(crate) secret_mode: WebSecretMode,
/// HMAC-derived bridge capability.
pub(crate) capability: [u8; 32],
/// Per-profile live session ceiling.
pub(crate) max_sessions: usize,
/// Per-profile live logical-stream ceiling.
pub(crate) max_streams: usize,
/// Per-session live relay-task ceiling.
pub(crate) max_streams_per_session: usize,
}
/// Runtime-ready ordinary-site fallback.
#[derive(Debug)]
pub(crate) enum WebRuntimeDecoy {
HttpUpstream {
addr: SocketAddr,
authority: String,
},
StaticDirectory(Arc<WebStaticSite>),
}
/// Immutable bounded static-site snapshot.
#[derive(Debug)]
pub(crate) struct WebStaticSite {
/// Canonical URL-path to immutable response asset mapping.
pub(crate) assets: BTreeMap<String, WebStaticAsset>,
/// Configured root index file name.
pub(crate) index: String,
}
/// One immutable static response body and metadata.
#[derive(Debug)]
pub(crate) struct WebStaticAsset {
/// Immutable response body retained by the runtime snapshot.
pub(crate) body: Bytes,
/// Extension-derived static content type.
pub(crate) content_type: &'static str,
/// Strong SHA-256 entity tag.
pub(crate) etag: String,
}
fn default_web_static_index() -> String {
"index.html".to_string()
}
macro_rules! usize_default {
($name:ident, $value:expr) => {
fn $name() -> usize {
$value
}
};
}
macro_rules! u32_default {
($name:ident, $value:expr) => {
fn $name() -> u32 {
$value
}
};
}
macro_rules! u64_default {
($name:ident, $value:expr) => {
fn $name() -> u64 {
$value
}
};
}
usize_default!(default_web_max_header_bytes, 16 * 1024);
usize_default!(default_web_max_body_bytes, 2 * 1024 * 1024);
usize_default!(default_web_max_frame_payload_bytes, 1024 * 1024);
usize_default!(default_web_carrier_batch_bytes, 2 * 1024 * 1024);
usize_default!(default_web_max_frames_per_body, 4096);
usize_default!(default_web_max_http_connections, 1024);
usize_default!(default_web_max_http_handlers, 512);
usize_default!(default_web_max_body_readers, 32);
usize_default!(default_web_max_body_bytes_global, 64 * 1024 * 1024);
usize_default!(default_web_max_sessions_global, 128);
usize_default!(default_web_max_sessions_per_ip, 16);
usize_default!(default_web_max_streams_per_session, 128);
usize_default!(default_web_max_streams_global, 4096);
usize_default!(default_web_max_stream_handshakes, 256);
usize_default!(default_web_max_tombstones, 4096);
usize_default!(default_web_pending_bytes_per_session, 32 * 1024 * 1024);
usize_default!(default_web_pending_bytes_global, 512 * 1024 * 1024);
usize_default!(default_web_pending_items_per_session, 16 * 1024);
usize_default!(default_web_pending_items_global, 256 * 1024);
usize_default!(default_web_control_bytes_per_session, 256 * 1024);
usize_default!(default_web_control_bytes_global, 16 * 1024 * 1024);
usize_default!(default_web_max_bootstraps_global, 512);
usize_default!(default_web_max_bootstraps_per_ip, 64);
usize_default!(default_web_max_vhosts, 8);
usize_default!(default_web_max_profiles, 32);
usize_default!(default_web_max_static_files, 4096);
usize_default!(default_web_max_static_file_bytes, 8 * 1024 * 1024);
usize_default!(default_web_max_static_bytes, 64 * 1024 * 1024);
usize_default!(default_web_memory_envelope_bytes, 768 * 1024 * 1024);
u32_default!(default_web_new_bootstraps_per_minute, 1200);
u32_default!(default_web_new_bootstraps_burst, 256);
u32_default!(default_web_new_sessions_per_minute, 600);
u32_default!(default_web_new_sessions_burst, 128);
u32_default!(default_web_new_streams_per_minute, 6000);
u32_default!(default_web_new_streams_burst, 512);
u64_default!(default_web_header_timeout_secs, 10);
u64_default!(default_web_body_timeout_secs, 30);
u64_default!(default_web_stream_handshake_timeout_secs, 10);
u64_default!(default_web_long_poll_timeout_secs, 25);
u64_default!(default_web_bootstrap_lifetime_secs, 120);
u64_default!(default_web_reconnect_grace_secs, 120);
u64_default!(default_web_http_idle_secs, 75);
u64_default!(default_web_shutdown_secs, 15);
u64_default!(default_web_decoy_header_timeout_secs, 30);
+18 -1
View File
@@ -14,6 +14,7 @@ use crate::ip_tracker::UserIpTracker;
use crate::proxy::route_mode::RelayRouteMode;
use crate::proxy::route_mode::RouteRuntimeController;
use crate::proxy::shared_state::ProxySharedState;
use crate::proxy::authenticated::ClientRuntimeDeps;
use crate::stats::beobachten::BeobachtenStore;
use crate::stats::{ReplayChecker, Stats};
use crate::stream::BufferPool;
@@ -227,6 +228,22 @@ impl RuntimeGeneration {
self.me_pool_runtime.read().await.clone()
}
/// Pins all dependencies required by a client stream without retaining the generation.
pub(crate) fn client_runtime_deps(&self) -> ClientRuntimeDeps {
ClientRuntimeDeps {
config: self.config(),
stats: Arc::clone(&self.stats),
upstream_manager: Arc::clone(&self.upstream_manager),
buffer_pool: Arc::clone(&self.buffer_pool),
rng: Arc::clone(&self.rng),
me_pool: self.me_pool.clone(),
me_pool_runtime: Some(Arc::clone(&self.me_pool_runtime)),
route_runtime: Arc::clone(&self.route_runtime),
ip_tracker: Arc::clone(&self.ip_tracker),
shared: Arc::clone(&self.proxy_shared),
}
}
/// Registers a session only while admission remains open.
pub(crate) fn spawn_session<F>(&self, future: F) -> bool
where
@@ -287,7 +304,7 @@ impl RuntimeGeneration {
#[cfg(test)]
/// Builds a lightweight runtime generation without network startup tasks.
pub(super) fn test_runtime_generation(id: u64, config: ProxyConfig) -> Arc<RuntimeGeneration> {
pub(crate) fn test_runtime_generation(id: u64, config: ProxyConfig) -> Arc<RuntimeGeneration> {
let (config_tx, config_rx) = watch::channel(Arc::new(config.clone()));
let (_admission_tx, admission_rx) = watch::channel(true);
let stats = Arc::new(Stats::new());
+40
View File
@@ -607,6 +607,46 @@ pub(crate) fn print_proxy_links(host: &str, port: u16, config: &ProxyConfig) {
}
}
/// 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()
+58 -2
View File
@@ -6,11 +6,13 @@ use tokio::net::{TcpListener, TcpStream};
use tokio::sync::OwnedSemaphorePermit;
use tokio::task::JoinHandle;
use tokio_util::sync::CancellationToken;
use tokio_util::task::TaskTracker;
use tracing::{debug, error, info, warn};
use crate::config::RstOnCloseMode;
use crate::config::{ListenerTransport, RstOnCloseMode};
use crate::proxy::ClientHandler;
use crate::transport::socket::set_linger_zero;
use crate::web::manager::WebProcessRuntime;
use super::bind::BoundTcpListener;
use super::plan::ListenerBindSpec;
@@ -19,11 +21,15 @@ use crate::maestro::helpers::{
expected_handshake_close_description, is_expected_handshake_eof, peer_close_description,
};
/// One bound listener and all connection tasks accepted through its lifecycle.
pub(super) struct ListenerSlot {
pub(super) spec: ListenerBindSpec,
listener: Arc<TcpListener>,
cancellation: CancellationToken,
task: Option<JoinHandle<()>>,
connections: TaskTracker,
web_runtime: Option<Arc<WebProcessRuntime>>,
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
}
enum PermitWait {
@@ -181,6 +187,8 @@ async fn run_accept_loop(
listener: Arc<TcpListener>,
spec: ListenerBindSpec,
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
web_runtime: Option<Arc<WebProcessRuntime>>,
connections: TaskTracker,
cancellation: CancellationToken,
) {
loop {
@@ -191,6 +199,26 @@ async fn run_accept_loop(
};
match accepted {
Ok((stream, peer_addr)) => {
if spec.transport == ListenerTransport::Web {
let Some(web_runtime) = web_runtime.as_ref() else {
error!(addr = %spec.addr, "WEB listener has no process runtime");
return;
};
let Some(connection_permit) = web_runtime.try_http_connection() else {
drop(stream);
continue;
};
connections.spawn(crate::web::http::serve_connection(
stream,
peer_addr,
spec.web_client_ip_source,
Arc::clone(&spec.web_trusted_proxy_cidrs),
Arc::clone(web_runtime),
cancellation.clone(),
connection_permit,
));
continue;
}
let runtime = active_runtime.load_full();
if !*runtime.admission_rx.borrow() {
debug!(peer = %peer_addr, "Admission gate closed, dropping connection");
@@ -232,12 +260,16 @@ impl ListenerSlot {
pub(super) fn start(
bound: BoundTcpListener,
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
web_runtime: Option<Arc<WebProcessRuntime>>,
) -> Self {
let cancellation = CancellationToken::new();
let connections = TaskTracker::new();
let task = tokio::spawn(run_accept_loop(
bound.listener.clone(),
bound.spec.clone(),
active_runtime,
active_runtime.clone(),
web_runtime.clone(),
connections.clone(),
cancellation.clone(),
));
Self {
@@ -245,6 +277,9 @@ impl ListenerSlot {
listener: bound.listener,
cancellation,
task: Some(task),
connections,
web_runtime,
active_runtime,
}
}
@@ -255,15 +290,36 @@ impl ListenerSlot {
format!("listener {} task failed: {error_value}", self.spec.addr)
})?;
}
self.connections.close();
let connection_stop_timeout = Duration::from_secs(
self.active_runtime
.load()
.config()
.web
.timeouts
.shutdown_secs,
);
tokio::time::timeout(connection_stop_timeout, self.connections.wait())
.await
.map_err(|_| {
format!(
"listener {} connection shutdown timed out",
self.spec.addr
)
})?;
Ok(())
}
pub(super) fn restart(&mut self, active_runtime: Arc<ArcSwap<RuntimeGeneration>>) {
self.active_runtime = active_runtime.clone();
self.cancellation = CancellationToken::new();
self.connections = TaskTracker::new();
self.task = Some(tokio::spawn(run_accept_loop(
self.listener.clone(),
self.spec.clone(),
active_runtime,
self.web_runtime.clone(),
self.connections.clone(),
self.cancellation.clone(),
)));
}
+15 -5
View File
@@ -8,13 +8,13 @@ use tokio::net::TcpListener;
use tokio::net::UnixListener;
use tracing::{error, info, warn};
use crate::config::ProxyConfig;
use crate::config::{ListenerTransport, ProxyConfig};
use crate::startup::{COMPONENT_LISTENERS_BIND, StartupTracker};
use crate::transport::find_listener_processes;
use crate::transport::socket::{activate_listener_socket, bind_listener_socket};
use super::plan::{ListenerBindSpec, listener_bind_plan};
use crate::maestro::helpers::print_proxy_links;
use crate::maestro::helpers::{print_proxy_links, print_web_proxy_links};
/// Owns sockets bound before process accept loops start.
pub(crate) struct BoundListeners {
@@ -57,7 +57,8 @@ fn default_link_port(config: &ProxyConfig) -> u16 {
config
.server
.listeners
.first()
.iter()
.find(|listener| listener.transport == ListenerTransport::Mtproxy)
.and_then(|listener| listener.port)
.unwrap_or(config.server.port)
}
@@ -110,7 +111,7 @@ impl PreparedTcpListener {
}
fn log_listener_profile(spec: &ListenerBindSpec) {
info!(addr = %spec.addr, "Listening on TCP endpoint");
info!(addr = %spec.addr, transport = ?spec.transport, "Listening on TCP endpoint");
if let Some(client_mss) = spec.options.client_mss {
info!(
addr = %spec.addr,
@@ -135,7 +136,11 @@ fn print_configured_links(
detected_ip_v4: Option<IpAddr>,
detected_ip_v6: Option<IpAddr>,
) {
print_web_proxy_links(config);
for listener in &config.server.listeners {
if listener.transport != ListenerTransport::Mtproxy {
continue;
}
let port = listener.port.unwrap_or(config.server.port);
let addr = SocketAddr::new(listener.ip, port);
if !plan.contains_key(&addr) || config.general.links.public_host.is_some() {
@@ -160,7 +165,12 @@ fn print_configured_links(
}
}
if config.general.links.show.is_empty() || config.general.links.public_host.is_none() {
if config.general.links.show.is_empty()
|| config.general.links.public_host.is_none()
|| !plan
.values()
.any(|spec| spec.transport == ListenerTransport::Mtproxy)
{
return;
}
let host = config
+48 -2
View File
@@ -5,11 +5,13 @@ use std::sync::Arc;
use arc_swap::ArcSwap;
use crate::config::ProxyConfig;
use crate::config::ListenerTransport;
use crate::maestro::generation::RuntimeGeneration;
use super::accept::ListenerSlot;
use super::bind::{BoundListeners, BoundTcpListener, PreparedTcpListener, prepare_listener};
use super::plan::{ListenerBindSpec, listener_bind_plan};
use crate::web::manager::WebProcessRuntime;
#[cfg(unix)]
use super::unix::UnixAcceptHandle;
@@ -17,16 +19,19 @@ use super::unix::UnixAcceptHandle;
pub(crate) struct ListenerManager {
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
slots: BTreeMap<SocketAddr, ListenerSlot>,
web_runtime: Option<Arc<WebProcessRuntime>>,
#[cfg(unix)]
unix: Option<UnixAcceptHandle>,
}
/// Socket changes prepared without activating or stopping accept loops.
pub(crate) struct PreparedListenerTransition {
target_specs: BTreeMap<SocketAddr, ListenerBindSpec>,
additions: Vec<PreparedTcpListener>,
removals: Vec<SocketAddr>,
}
/// Activated additions and stopped removals awaiting runtime publication.
pub(crate) struct PendingListenerTransition {
target_specs: BTreeMap<SocketAddr, ListenerBindSpec>,
additions: Vec<BoundTcpListener>,
@@ -39,10 +44,22 @@ impl ListenerManager {
bound: BoundListeners,
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
) -> Self {
let has_web = bound
.listeners
.iter()
.any(|listener| listener.spec.transport == ListenerTransport::Web);
let web_runtime = has_web.then(|| WebProcessRuntime::start(active_runtime.clone()));
let mut slots = BTreeMap::new();
for listener in bound.listeners {
let addr = listener.spec.addr;
slots.insert(addr, ListenerSlot::start(listener, active_runtime.clone()));
slots.insert(
addr,
ListenerSlot::start(
listener,
active_runtime.clone(),
web_runtime.clone(),
),
);
}
#[cfg(unix)]
let unix = bound
@@ -51,6 +68,7 @@ impl ListenerManager {
Self {
active_runtime,
slots,
web_runtime,
#[cfg(unix)]
unix,
}
@@ -61,6 +79,7 @@ impl ListenerManager {
Self {
active_runtime,
slots: BTreeMap::new(),
web_runtime: None,
#[cfg(unix)]
unix: None,
}
@@ -72,6 +91,20 @@ impl ListenerManager {
desired: &ProxyConfig,
) -> Result<Option<PreparedListenerTransition>, String> {
let target_specs = listener_bind_plan(desired)?;
let web_inventory_changed = self
.slots
.iter()
.filter(|(_, slot)| slot.spec.transport == ListenerTransport::Web)
.map(|(addr, slot)| (*addr, slot.spec.clone()))
.collect::<BTreeMap<_, _>>()
!= target_specs
.iter()
.filter(|(_, spec)| spec.transport == ListenerTransport::Web)
.map(|(addr, spec)| (*addr, spec.clone()))
.collect::<BTreeMap<_, _>>();
if web_inventory_changed {
return Err("WEB listener inventory is process-owned; process restart required".to_string());
}
let current_addresses: BTreeSet<_> = self.slots.keys().copied().collect();
let target_addresses: BTreeSet<_> = target_specs.keys().copied().collect();
if current_addresses == target_addresses
@@ -164,7 +197,11 @@ impl ListenerManager {
let addr = listener.spec.addr;
self.slots.insert(
addr,
ListenerSlot::start(listener, self.active_runtime.clone()),
ListenerSlot::start(
listener,
self.active_runtime.clone(),
self.web_runtime.clone(),
),
);
}
debug_assert_eq!(
@@ -191,6 +228,9 @@ impl ListenerManager {
errors.push(error_value);
}
self.slots.clear();
if let Some(web_runtime) = self.web_runtime.take() {
web_runtime.shutdown().await;
}
#[cfg(unix)]
{
self.unix = None;
@@ -214,6 +254,7 @@ mod tests {
fn listener_config(addr: SocketAddr) -> ListenerConfig {
ListenerConfig {
ip: addr.ip(),
transport: crate::config::ListenerTransport::Mtproxy,
port: Some(addr.port()),
client_mss: None,
synlimit: SynLimitMode::Off,
@@ -229,6 +270,8 @@ mod tests {
announce_ip: None,
proxy_protocol: None,
reuse_allow: false,
web_client_ip_source: crate::config::WebClientIpSource::XForwardedFor,
web_trusted_proxy_cidrs: Vec::new(),
}
}
@@ -237,12 +280,15 @@ mod tests {
let addr = listener.local_addr().unwrap();
let spec = ListenerBindSpec {
addr,
transport: crate::config::ListenerTransport::Mtproxy,
options: ListenOptions {
reuse_port: false,
..Default::default()
},
proxy_protocol: false,
tls_response_fragment_size: None,
web_client_ip_source: crate::config::WebClientIpSource::XForwardedFor,
web_trusted_proxy_cidrs: Arc::from([]),
};
(
BoundTcpListener {
+41 -5
View File
@@ -1,7 +1,10 @@
use std::collections::{BTreeMap, BTreeSet};
use std::net::SocketAddr;
use std::sync::Arc;
use crate::config::{ProxyConfig, ServerConfig, SynLimitMode};
use crate::config::{
ListenerTransport, ProxyConfig, ServerConfig, SynLimitMode, WebClientIpSource,
};
use crate::transport::ListenOptions;
use super::tcp_mss_runtime_profile;
@@ -10,9 +13,12 @@ use super::tcp_mss_runtime_profile;
#[derive(Debug, Clone, PartialEq, Eq)]
pub(super) struct ListenerBindSpec {
pub(super) addr: SocketAddr,
pub(super) transport: ListenerTransport,
pub(super) options: ListenOptions,
pub(super) proxy_protocol: bool,
pub(super) tls_response_fragment_size: Option<u16>,
pub(super) web_client_ip_source: WebClientIpSource,
pub(super) web_trusted_proxy_cidrs: Arc<[ipnetwork::IpNetwork]>,
}
fn listener_port_or_legacy(listener: &crate::config::ListenerConfig, server: &ServerConfig) -> u16 {
@@ -40,16 +46,24 @@ pub(crate) fn listener_bind_plan(
if addr.is_ipv6() && config.network.ipv6 == Some(false) {
continue;
}
let configured_client_mss = listener
.effective_client_mss(&config.server)
.map_err(|error| format!("invalid client MSS for listener {addr}: {error}"))?;
let configured_client_mss = if listener.transport == ListenerTransport::Web {
None
} else {
listener
.effective_client_mss(&config.server)
.map_err(|error| format!("invalid client MSS for listener {addr}: {error}"))?
};
let listener_bulk_mss = (listener.transport != ListenerTransport::Web)
.then_some(bulk_client_mss)
.flatten();
#[cfg(target_os = "linux")]
let (client_mss, tls_response_fragment_size) =
tcp_mss_runtime_profile(configured_client_mss, bulk_client_mss);
tcp_mss_runtime_profile(configured_client_mss, listener_bulk_mss);
#[cfg(not(target_os = "linux"))]
let (client_mss, tls_response_fragment_size) = (configured_client_mss, None);
let spec = ListenerBindSpec {
addr,
transport: listener.transport,
options: ListenOptions {
reuse_port: listener.reuse_allow,
ipv6_only: listener.ip.is_ipv6(),
@@ -61,6 +75,8 @@ pub(crate) fn listener_bind_plan(
.proxy_protocol
.unwrap_or(config.server.proxy_protocol),
tls_response_fragment_size,
web_client_ip_source: listener.web_client_ip_source,
web_trusted_proxy_cidrs: Arc::from(listener.web_trusted_proxy_cidrs.clone()),
};
if plan.insert(addr, spec).is_some() {
return Err(format!("duplicate effective listener endpoint: {addr}"));
@@ -80,6 +96,23 @@ fn any_synlimit_enabled(config: &ProxyConfig) -> bool {
/// Returns whether an endpoint-only change can use coordinated process rebind.
pub(crate) fn listener_rebind_supported(old: &ProxyConfig, desired: &ProxyConfig) -> bool {
let Ok(old_plan) = listener_bind_plan(old) else {
return false;
};
let Ok(desired_plan) = listener_bind_plan(desired) else {
return false;
};
let old_web = old_plan
.iter()
.filter(|(_, spec)| spec.transport == ListenerTransport::Web)
.collect::<BTreeMap<_, _>>();
let desired_web = desired_plan
.iter()
.filter(|(_, spec)| spec.transport == ListenerTransport::Web)
.collect::<BTreeMap<_, _>>();
if old_web != desired_web {
return false;
}
if any_synlimit_enabled(old) || any_synlimit_enabled(desired) {
return false;
}
@@ -107,6 +140,7 @@ mod tests {
fn listener(ip: &str, port: u16) -> ListenerConfig {
ListenerConfig {
ip: ip.parse().unwrap(),
transport: crate::config::ListenerTransport::Mtproxy,
port: Some(port),
client_mss: None,
synlimit: SynLimitMode::Off,
@@ -122,6 +156,8 @@ mod tests {
announce_ip: None,
proxy_protocol: None,
reuse_allow: false,
web_client_ip_source: crate::config::WebClientIpSource::XForwardedFor,
web_trusted_proxy_cidrs: Vec::new(),
}
}
+3
View File
@@ -131,6 +131,7 @@ mod tests {
fn listener_with_synlimit(synlimit: SynLimitMode) -> ListenerConfig {
ListenerConfig {
ip: "127.0.0.1".parse().unwrap(),
transport: crate::config::ListenerTransport::Mtproxy,
port: Some(443),
client_mss: None,
synlimit,
@@ -146,6 +147,8 @@ mod tests {
announce_ip: None,
proxy_protocol: None,
reuse_allow: false,
web_client_ip_source: crate::config::WebClientIpSource::XForwardedFor,
web_trusted_proxy_cidrs: Vec::new(),
}
}
+10
View File
@@ -406,6 +406,16 @@ pub(crate) fn resolve_reload_config(
fields.push("logging".to_string());
effective.logging = old.logging.clone();
}
if serde_json::to_value(&old.web.limits).ok()
!= serde_json::to_value(&desired.web.limits).ok()
{
fields.push("web.limits".to_string());
effective.web.limits = old.web.limits.clone();
if effective.rebuild_runtime_web().is_err() {
fields.push("web".to_string());
effective.web = old.web.clone();
}
}
let runtime_changed = !configs_equal(old, &effective);
ResolvedReloadConfig {
effective,
+27
View File
@@ -3,6 +3,7 @@ use super::*;
fn test_listener(port: u16) -> crate::config::ListenerConfig {
crate::config::ListenerConfig {
ip: "127.0.0.1".parse().unwrap(),
transport: crate::config::ListenerTransport::Mtproxy,
port: Some(port),
client_mss: None,
synlimit: crate::config::SynLimitMode::Off,
@@ -18,6 +19,8 @@ fn test_listener(port: u16) -> crate::config::ListenerConfig {
announce_ip: None,
proxy_protocol: None,
reuse_allow: false,
web_client_ip_source: crate::config::WebClientIpSource::XForwardedFor,
web_trusted_proxy_cidrs: Vec::new(),
}
}
@@ -80,6 +83,7 @@ fn listener_announcement_is_runtime_owned_when_bind_identity_is_stable() {
let mut old = ProxyConfig::default();
old.server.listeners.push(crate::config::ListenerConfig {
ip: "0.0.0.0".parse().unwrap(),
transport: crate::config::ListenerTransport::Mtproxy,
port: Some(443),
client_mss: None,
synlimit: crate::config::SynLimitMode::Off,
@@ -95,6 +99,8 @@ fn listener_announcement_is_runtime_owned_when_bind_identity_is_stable() {
announce_ip: None,
proxy_protocol: None,
reuse_allow: false,
web_client_ip_source: crate::config::WebClientIpSource::XForwardedFor,
web_trusted_proxy_cidrs: Vec::new(),
});
let mut desired = old.clone();
desired.server.listeners[0].announce = Some("proxy.example".to_string());
@@ -143,6 +149,27 @@ fn runtime_only_change_does_not_require_process_rebind() {
assert!(deferred_process_fields(&old, &new).is_empty());
}
#[test]
fn web_allocation_limits_are_deferred_until_restart() {
let mut old = ProxyConfig::default();
old.rebuild_runtime_user_auth().unwrap();
old.rebuild_runtime_web().unwrap();
let mut desired = old.clone();
desired.web.limits.max_sessions_global += 1;
let resolved = resolve_reload_config(&old, &desired);
assert_eq!(
resolved.deferred_process_fields,
vec!["web.limits".to_string()]
);
assert_eq!(
resolved.effective.web.limits.max_sessions_global,
old.web.limits.max_sessions_global
);
assert!(!resolved.runtime_changed);
}
#[test]
fn strict_middle_proxy_requires_a_prepared_pool() {
assert!(strict_middle_proxy_unavailable(true, false, false));
+1
View File
@@ -34,6 +34,7 @@ mod synlimit_control;
mod tls_front;
mod transport;
mod util;
mod web;
fn main() -> std::result::Result<(), Box<dyn std::error::Error>> {
// Install rustls crypto provider early
+325
View File
@@ -0,0 +1,325 @@
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use tokio::io::{AsyncRead, AsyncWrite};
use tokio::sync::RwLock;
use tracing::warn;
use crate::config::ProxyConfig;
use crate::crypto::SecureRandom;
use crate::error::{ProxyError, Result};
use crate::ip_tracker::UserIpTracker;
use crate::proxy::direct_relay::handle_via_direct_with_shared_and_conntrack;
use crate::proxy::handshake::HandshakeSuccess;
use crate::proxy::middle_relay::{
handle_via_middle_proxy, handle_via_middle_proxy_with_conntrack,
};
use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController};
use crate::proxy::shared_state::{ConntrackClosePolicy, ProxySharedState};
use crate::stats::Stats;
use crate::stream::{BufferPool, CryptoReader, CryptoWriter};
use crate::transport::UpstreamManager;
use crate::transport::middle_proxy::MePool;
/// Immutable dependency snapshot pinned by one authenticated client stream.
#[derive(Clone)]
pub(crate) struct ClientRuntimeDeps {
/// Immutable effective configuration pinned for this stream.
pub(crate) config: Arc<ProxyConfig>,
/// Process statistics registry.
pub(crate) stats: Arc<Stats>,
/// Direct Telegram upstream connector.
pub(crate) upstream_manager: Arc<UpstreamManager>,
/// Shared relay buffer pool.
pub(crate) buffer_pool: Arc<BufferPool>,
/// Process cryptographic random source.
pub(crate) rng: Arc<SecureRandom>,
/// Startup Middle-End pool, when immediately available.
pub(crate) me_pool: Option<Arc<MePool>>,
/// Hot-swappable Middle-End pool holder.
pub(crate) me_pool_runtime: Option<Arc<RwLock<Option<Arc<MePool>>>>>,
/// Route-mode controller shared by active generations.
pub(crate) route_runtime: Arc<RouteRuntimeController>,
/// Per-user source-IP admission tracker.
pub(crate) ip_tracker: Arc<UserIpTracker>,
/// Process-shared admission and relay coordination state.
pub(crate) shared: Arc<ProxySharedState>,
}
/// Runs admission and relay after a successful MTProxy handshake.
pub(crate) async fn run_authenticated<R, W>(
client_reader: CryptoReader<R>,
client_writer: CryptoWriter<W>,
success: HandshakeSuccess,
deps: ClientRuntimeDeps,
local_addr: SocketAddr,
peer_addr: SocketAddr,
conntrack_close_policy: ConntrackClosePolicy,
) -> Result<()>
where
R: AsyncRead + Unpin + Send + 'static,
W: AsyncWrite + Unpin + Send + 'static,
{
let user = success.user.clone();
if !deps.shared.is_user_enabled(&user) {
warn!(user = %user, "Disabled user rejected");
return Err(ProxyError::UserDisabled { user });
}
let user_reservation = acquire_user_connection_reservation(
&user,
&deps.config,
Arc::clone(&deps.stats),
peer_addr,
Arc::clone(&deps.ip_tracker),
)
.await
.map_err(|error| {
warn!(user = %user, error = %error, "User admission check failed");
error
})?;
let route_snapshot = deps.route_runtime.snapshot();
let session_id = deps.rng.u64();
let user_session = deps.shared.register_user_session(&user, session_id);
let session_cancel = user_session.token();
let selected_me_pool = if deps.config.general.use_middle_proxy
&& matches!(route_snapshot.mode, RelayRouteMode::Middle)
{
if let Some(pool) = &deps.me_pool {
Some(Arc::clone(pool))
} else if let Some(pool_runtime) = &deps.me_pool_runtime {
pool_runtime.read().await.clone()
} else {
None
}
} else {
None
};
let relay_result = if deps.config.general.use_middle_proxy
&& matches!(route_snapshot.mode, RelayRouteMode::Middle)
{
if let Some(pool) = selected_me_pool {
if conntrack_close_policy == ConntrackClosePolicy::Publish {
handle_via_middle_proxy(
client_reader,
client_writer,
success,
pool,
Arc::clone(&deps.stats),
Arc::clone(&deps.config),
Arc::clone(&deps.buffer_pool),
local_addr,
Arc::clone(&deps.rng),
deps.route_runtime.subscribe(),
route_snapshot,
session_id,
session_cancel.clone(),
Arc::clone(&deps.shared),
)
.await
} else {
handle_via_middle_proxy_with_conntrack(
client_reader,
client_writer,
success,
pool,
Arc::clone(&deps.stats),
Arc::clone(&deps.config),
Arc::clone(&deps.buffer_pool),
local_addr,
Arc::clone(&deps.rng),
deps.route_runtime.subscribe(),
route_snapshot,
session_id,
session_cancel.clone(),
Arc::clone(&deps.shared),
ConntrackClosePolicy::Suppress,
)
.await
}
} else {
warn!("use_middle_proxy=true but MePool not initialized, falling back to direct");
run_direct(
client_reader,
client_writer,
success,
&deps,
route_snapshot,
session_id,
local_addr,
session_cancel.clone(),
conntrack_close_policy,
)
.await
}
} else {
run_direct(
client_reader,
client_writer,
success,
&deps,
route_snapshot,
session_id,
local_addr,
session_cancel,
conntrack_close_policy,
)
.await
};
user_reservation.release().await;
relay_result
}
async fn run_direct<R, W>(
client_reader: CryptoReader<R>,
client_writer: CryptoWriter<W>,
success: HandshakeSuccess,
deps: &ClientRuntimeDeps,
route_snapshot: crate::proxy::route_mode::RouteCutoverState,
session_id: u64,
local_addr: SocketAddr,
session_cancel: tokio_util::sync::CancellationToken,
conntrack_close_policy: ConntrackClosePolicy,
) -> Result<()>
where
R: AsyncRead + Unpin + Send + 'static,
W: AsyncWrite + Unpin + Send + 'static,
{
handle_via_direct_with_shared_and_conntrack(
client_reader,
client_writer,
success,
Arc::clone(&deps.upstream_manager),
Arc::clone(&deps.stats),
Arc::clone(&deps.config),
Arc::clone(&deps.buffer_pool),
Arc::clone(&deps.rng),
deps.route_runtime.subscribe(),
route_snapshot,
session_id,
local_addr,
session_cancel,
Arc::clone(&deps.shared),
conntrack_close_policy,
)
.await
}
#[must_use = "the reservation owns user and IP admission until release or drop"]
/// Owns one authenticated user's connection and source-IP admission slots.
pub(crate) struct UserConnectionReservation {
stats: Arc<Stats>,
ip_tracker: Arc<UserIpTracker>,
user: String,
ip: IpAddr,
tracks_ip: bool,
active: bool,
}
impl UserConnectionReservation {
/// Creates an active reservation after both admission counters were acquired.
pub(crate) fn new(
stats: Arc<Stats>,
ip_tracker: Arc<UserIpTracker>,
user: String,
ip: IpAddr,
tracks_ip: bool,
) -> Self {
Self {
stats,
ip_tracker,
user,
ip,
tracks_ip,
active: true,
}
}
/// Releases both admission counters through the asynchronous cleanup path.
pub(crate) async fn release(mut self) {
if !self.active {
return;
}
self.active = false;
if self.tracks_ip {
self.ip_tracker.remove_ip(&self.user, self.ip).await;
}
self.stats.decrement_user_curr_connects(&self.user);
}
}
impl Drop for UserConnectionReservation {
fn drop(&mut self) {
if !self.active {
return;
}
self.active = false;
self.stats.increment_session_drop_fallback_total();
self.stats.decrement_user_curr_connects(&self.user);
if self.tracks_ip {
self.ip_tracker.enqueue_cleanup(self.user.clone(), self.ip);
}
}
}
/// Applies user quota, connection, and source-IP admission atomically.
pub(crate) async fn acquire_user_connection_reservation(
user: &str,
config: &ProxyConfig,
stats: Arc<Stats>,
peer_addr: SocketAddr,
ip_tracker: Arc<UserIpTracker>,
) -> Result<UserConnectionReservation> {
if let Some(expiration) = config.access.user_expirations.get(user)
&& chrono::Utc::now() > *expiration
{
return Err(ProxyError::UserExpired {
user: user.to_string(),
});
}
if let Some(quota) = config.access.user_data_quota.get(user)
&& stats.get_user_quota_used(user) >= *quota
{
return Err(ProxyError::DataQuotaExceeded {
user: user.to_string(),
});
}
let limit = config
.access
.user_max_tcp_conns
.get(user)
.copied()
.filter(|limit| *limit > 0)
.or((config.access.user_max_tcp_conns_global_each > 0)
.then_some(config.access.user_max_tcp_conns_global_each))
.map(|value| value as u64);
if !stats.try_acquire_user_curr_connects(user, limit) {
return Err(ProxyError::ConnectionLimitExceeded {
user: user.to_string(),
});
}
if let Err(reason) = ip_tracker.check_and_add(user, peer_addr.ip()).await {
stats.decrement_user_curr_connects(user);
warn!(
user = %user,
ip = %peer_addr.ip(),
reason = %reason,
"IP limit exceeded"
);
return Err(ProxyError::ConnectionLimitExceeded {
user: user.to_string(),
});
}
Ok(UserConnectionReservation::new(
stats,
ip_tracker,
user.to_string(),
peer_addr.ip(),
true,
))
}
+29 -225
View File
@@ -26,72 +26,6 @@ enum HandshakeOutcome {
NeedsMasking(PostHandshakeFuture),
}
#[must_use = "UserConnectionReservation must be kept alive to retain user/IP reservation until release or drop"]
struct UserConnectionReservation {
stats: Arc<Stats>,
ip_tracker: Arc<UserIpTracker>,
user: String,
ip: IpAddr,
tracks_ip: bool,
state: SessionReservationState,
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum SessionReservationState {
Active,
Released,
}
impl UserConnectionReservation {
fn new(
stats: Arc<Stats>,
ip_tracker: Arc<UserIpTracker>,
user: String,
ip: IpAddr,
tracks_ip: bool,
) -> Self {
Self {
stats,
ip_tracker,
user,
ip,
tracks_ip,
state: SessionReservationState::Active,
}
}
fn mark_released(&mut self) -> bool {
if self.state != SessionReservationState::Active {
return false;
}
self.state = SessionReservationState::Released;
true
}
async fn release(mut self) {
if !self.mark_released() {
return;
}
if self.tracks_ip {
self.ip_tracker.remove_ip(&self.user, self.ip).await;
}
self.stats.decrement_user_curr_connects(&self.user);
}
}
impl Drop for UserConnectionReservation {
fn drop(&mut self) {
if !self.mark_released() {
return;
}
self.stats.increment_session_drop_fallback_total();
self.stats.decrement_user_curr_connects(&self.user);
if self.tracks_ip {
self.ip_tracker.enqueue_cleanup(self.user.clone(), self.ip);
}
}
}
use crate::config::ProxyConfig;
use crate::crypto::SecureRandom;
use crate::error::{HandshakeResult, ProxyError, Result, StreamError};
@@ -107,7 +41,11 @@ use crate::transport::middle_proxy::MePool;
use crate::transport::socket::normalize_ip;
use crate::transport::{UpstreamManager, configure_client_socket, parse_proxy_protocol};
use crate::proxy::direct_relay::handle_via_direct_with_shared;
use crate::proxy::authenticated::{ClientRuntimeDeps, run_authenticated};
#[cfg(test)]
use crate::proxy::authenticated::{
UserConnectionReservation, acquire_user_connection_reservation,
};
use crate::proxy::handshake::{
HandshakeSuccess, TlsResponseWriteOptions, handle_mtproto_handshake_with_shared,
handle_tls_handshake_with_shared, handle_tls_handshake_with_shared_and_options,
@@ -115,9 +53,10 @@ use crate::proxy::handshake::{
#[cfg(test)]
use crate::proxy::handshake::{handle_mtproto_handshake, handle_tls_handshake};
use crate::proxy::masking::handle_bad_client_with_shared;
use crate::proxy::middle_relay::handle_via_middle_proxy;
use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController};
use crate::proxy::shared_state::ProxySharedState;
use crate::proxy::route_mode::RouteRuntimeController;
#[cfg(test)]
use crate::proxy::route_mode::RelayRouteMode;
use crate::proxy::shared_state::{ConntrackClosePolicy, ProxySharedState};
fn beobachten_ttl(config: &ProxyConfig) -> Duration {
const BEOBACHTEN_TTL_MAX_MINUTES: u64 = 24 * 60;
@@ -1688,112 +1627,30 @@ impl RunningClientHandler {
R: AsyncRead + Unpin + Send + 'static,
W: AsyncWrite + Unpin + Send + 'static,
{
let user = success.user.clone();
if !shared.is_user_enabled(&user) {
warn!(user = %user, "Disabled user rejected");
return Err(ProxyError::UserDisabled { user });
}
let user_limit_reservation = match Self::acquire_user_connection_reservation_static(
&user,
&config,
stats.clone(),
peer_addr,
ip_tracker,
)
.await
{
Ok(reservation) => reservation,
Err(e) => {
warn!(user = %user, error = %e, "User admission check failed");
return Err(e);
}
};
let route_snapshot = route_runtime.snapshot();
let session_id = rng.u64();
let _user_session = shared.register_user_session(&user, session_id);
let session_cancel = _user_session.token();
let selected_me_pool = if config.general.use_middle_proxy
&& matches!(route_snapshot.mode, RelayRouteMode::Middle)
{
if let Some(ref pool) = me_pool {
Some(pool.clone())
} else if let Some(pool_runtime) = me_pool_runtime.as_ref() {
pool_runtime.read().await.clone()
} else {
None
}
} else {
None
};
let relay_result = if config.general.use_middle_proxy
&& matches!(route_snapshot.mode, RelayRouteMode::Middle)
{
if let Some(pool) = selected_me_pool {
handle_via_middle_proxy(
client_reader,
client_writer,
success,
pool,
stats.clone(),
config,
buffer_pool,
local_addr,
rng,
route_runtime.subscribe(),
route_snapshot,
session_id,
session_cancel.clone(),
shared.clone(),
)
.await
} else {
warn!("use_middle_proxy=true but MePool not initialized, falling back to direct");
handle_via_direct_with_shared(
client_reader,
client_writer,
success,
upstream_manager,
stats.clone(),
config,
buffer_pool,
rng,
route_runtime.subscribe(),
route_snapshot,
session_id,
local_addr,
session_cancel.clone(),
shared.clone(),
)
.await
}
} else {
// Direct mode (original behavior)
handle_via_direct_with_shared(
client_reader,
client_writer,
success,
upstream_manager,
stats.clone(),
run_authenticated(
client_reader,
client_writer,
success,
ClientRuntimeDeps {
config,
stats,
upstream_manager,
buffer_pool,
rng,
route_runtime.subscribe(),
route_snapshot,
session_id,
local_addr,
session_cancel,
shared.clone(),
)
.await
};
user_limit_reservation.release().await;
relay_result
me_pool,
me_pool_runtime,
route_runtime,
ip_tracker,
shared,
},
local_addr,
peer_addr,
ConntrackClosePolicy::Publish,
)
.await
}
#[cfg(test)]
async fn acquire_user_connection_reservation_static(
user: &str,
config: &ProxyConfig,
@@ -1801,60 +1658,7 @@ impl RunningClientHandler {
peer_addr: SocketAddr,
ip_tracker: Arc<UserIpTracker>,
) -> Result<UserConnectionReservation> {
if let Some(expiration) = config.access.user_expirations.get(user)
&& chrono::Utc::now() > *expiration
{
return Err(ProxyError::UserExpired {
user: user.to_string(),
});
}
if let Some(quota) = config.access.user_data_quota.get(user)
&& stats.get_user_quota_used(user) >= *quota
{
return Err(ProxyError::DataQuotaExceeded {
user: user.to_string(),
});
}
let limit = config
.access
.user_max_tcp_conns
.get(user)
.copied()
.filter(|limit| *limit > 0)
.or((config.access.user_max_tcp_conns_global_each > 0)
.then_some(config.access.user_max_tcp_conns_global_each))
.map(|v| v as u64);
if !stats.try_acquire_user_curr_connects(user, limit) {
return Err(ProxyError::ConnectionLimitExceeded {
user: user.to_string(),
});
}
match ip_tracker.check_and_add(user, peer_addr.ip()).await {
Ok(()) => {}
Err(reason) => {
stats.decrement_user_curr_connects(user);
warn!(
user = %user,
ip = %peer_addr.ip(),
reason = %reason,
"IP limit exceeded"
);
return Err(ProxyError::ConnectionLimitExceeded {
user: user.to_string(),
});
}
}
Ok(UserConnectionReservation::new(
stats,
ip_tracker,
user.to_string(),
peer_addr.ip(),
true,
))
acquire_user_connection_reservation(user, config, stats, peer_addr, ip_tracker).await
}
#[cfg(test)]
+59 -12
View File
@@ -22,7 +22,8 @@ use crate::proxy::route_mode::{
RelayRouteMode, RouteCutoverState, affected_cutover_state, cutover_stagger_delay,
};
use crate::proxy::shared_state::{
ConntrackCloseEvent, ConntrackClosePublishResult, ConntrackCloseReason, ProxySharedState,
ConntrackCloseEvent, ConntrackClosePolicy, ConntrackClosePublishResult, ConntrackCloseReason,
ProxySharedState,
};
use crate::stats::Stats;
use crate::stream::{BufferPool, CryptoReader, CryptoWriter};
@@ -229,6 +230,7 @@ fn unknown_dc_test_lock() -> &'static Mutex<()> {
}
#[allow(dead_code)]
/// Runs Direct relay with standalone cancellation and shared-state defaults.
pub(crate) async fn handle_via_direct<R, W>(
client_reader: CryptoReader<R>,
client_writer: CryptoWriter<W>,
@@ -265,7 +267,49 @@ where
.await
}
/// Runs Direct relay for a kernel-backed TCP client tuple.
pub(crate) async fn handle_via_direct_with_shared<R, W>(
client_reader: CryptoReader<R>,
client_writer: CryptoWriter<W>,
success: HandshakeSuccess,
upstream_manager: Arc<UpstreamManager>,
stats: Arc<Stats>,
config: Arc<ProxyConfig>,
buffer_pool: Arc<BufferPool>,
rng: Arc<SecureRandom>,
route_rx: watch::Receiver<RouteCutoverState>,
route_snapshot: RouteCutoverState,
session_id: u64,
local_addr: SocketAddr,
session_cancel: CancellationToken,
shared: Arc<ProxySharedState>,
) -> Result<()>
where
R: AsyncRead + Unpin + Send + 'static,
W: AsyncWrite + Unpin + Send + 'static,
{
handle_via_direct_with_shared_and_conntrack(
client_reader,
client_writer,
success,
upstream_manager,
stats,
config,
buffer_pool,
rng,
route_rx,
route_snapshot,
session_id,
local_addr,
session_cancel,
shared,
ConntrackClosePolicy::Publish,
)
.await
}
/// Runs Direct relay with explicit kernel-conntrack close publication policy.
pub(crate) async fn handle_via_direct_with_shared_and_conntrack<R, W>(
client_reader: CryptoReader<R>,
client_writer: CryptoWriter<W>,
success: HandshakeSuccess,
@@ -280,6 +324,7 @@ pub(crate) async fn handle_via_direct_with_shared<R, W>(
local_addr: SocketAddr,
session_cancel: CancellationToken,
shared: Arc<ProxySharedState>,
conntrack_close_policy: ConntrackClosePolicy,
) -> Result<()>
where
R: AsyncRead + Unpin + Send + 'static,
@@ -407,17 +452,19 @@ where
pool_snapshot.allocated.saturating_sub(pool_snapshot.pooled),
);
let close_reason = classify_conntrack_close_reason(&relay_result);
let publish_result = shared.publish_conntrack_close_event(ConntrackCloseEvent {
src: success.peer,
dst: local_addr,
reason: close_reason,
});
if !matches!(
publish_result,
ConntrackClosePublishResult::Sent | ConntrackClosePublishResult::Disabled
) {
stats.increment_conntrack_close_event_drop_total();
if conntrack_close_policy == ConntrackClosePolicy::Publish {
let close_reason = classify_conntrack_close_reason(&relay_result);
let publish_result = shared.publish_conntrack_close_event(ConntrackCloseEvent {
src: success.peer,
dst: local_addr,
reason: close_reason,
});
if !matches!(
publish_result,
ConntrackClosePublishResult::Sent | ConntrackClosePublishResult::Disabled
) {
stats.increment_conntrack_close_event_drop_total();
}
}
relay_result
+2 -1
View File
@@ -20,7 +20,7 @@ use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt};
use tracing::{debug, info, trace, warn};
use zeroize::{Zeroize, Zeroizing};
use crate::config::{ProxyConfig, UnknownSniAction};
use crate::config::{ProxyConfig, UnknownSniAction, WebSecretMode};
use crate::crypto::{AesCtr, SecureRandom, sha256};
use crate::error::{HandshakeResult, ProxyError};
use crate::protocol::constants::*;
@@ -58,6 +58,7 @@ pub(crate) use self::auth_probe::{AuthProbeSaturationState, AuthProbeState};
#[cfg(test)]
pub use self::mtproto::handle_mtproto_handshake;
pub use self::mtproto::handle_mtproto_handshake_with_shared;
pub(crate) use self::mtproto::handle_mtproto_handshake_for_web_user;
#[allow(unused_imports)]
pub use self::nonce::{encrypt_tg_nonce, encrypt_tg_nonce_with_ciphers, generate_tg_nonce};
pub use self::session::HandshakeSuccess;
+70 -1
View File
@@ -11,6 +11,12 @@ pub(super) struct MtprotoCandidateValidation {
pub(super) encryptor: AesCtr,
}
#[derive(Clone, Copy)]
pub(super) enum MtprotoModePolicy {
Configured,
Web(WebSecretMode),
}
pub(super) fn sni_hint_hash(sni: &str) -> u64 {
let mut hasher = DefaultHasher::new();
for byte in sni.bytes() {
@@ -146,6 +152,7 @@ pub(super) fn validate_mtproto_secret_candidate(
secret: &[u8; ACCESS_SECRET_BYTES],
config: &ProxyConfig,
is_tls: bool,
mode_policy: MtprotoModePolicy,
) -> Option<MtprotoCandidateValidation> {
let mut dec_key_input = Zeroizing::new(Vec::with_capacity(PREKEY_LEN + secret.len()));
dec_key_input.extend_from_slice(dec_prekey);
@@ -163,7 +170,7 @@ pub(super) fn validate_mtproto_secret_candidate(
decrypted[PROTO_TAG_POS + 3],
];
let proto_tag = ProtoTag::from_bytes(tag_bytes)?;
if !mode_enabled_for_proto(config, proto_tag, is_tls) {
if !mode_enabled_for_proto_with_policy(config, proto_tag, is_tls, mode_policy) {
return None;
}
@@ -267,6 +274,28 @@ pub(super) fn mode_enabled_for_proto(
proto_tag: ProtoTag,
is_tls: bool,
) -> bool {
mode_enabled_for_proto_with_policy(
config,
proto_tag,
is_tls,
MtprotoModePolicy::Configured,
)
}
fn mode_enabled_for_proto_with_policy(
config: &ProxyConfig,
proto_tag: ProtoTag,
is_tls: bool,
policy: MtprotoModePolicy,
) -> bool {
if let MtprotoModePolicy::Web(secret_mode) = policy {
return match secret_mode {
WebSecretMode::Plain => {
matches!(proto_tag, ProtoTag::Intermediate | ProtoTag::Abridged)
}
WebSecretMode::Dd => matches!(proto_tag, ProtoTag::Secure),
};
}
match proto_tag {
ProtoTag::Secure => {
if is_tls {
@@ -279,6 +308,46 @@ pub(super) fn mode_enabled_for_proto(
}
}
#[cfg(test)]
mod web_mode_tests {
use super::*;
#[test]
fn web_secret_mode_isolates_inner_protocol_tags() {
let config = ProxyConfig::default();
assert!(mode_enabled_for_proto_with_policy(
&config,
ProtoTag::Abridged,
false,
MtprotoModePolicy::Web(WebSecretMode::Plain),
));
assert!(mode_enabled_for_proto_with_policy(
&config,
ProtoTag::Intermediate,
false,
MtprotoModePolicy::Web(WebSecretMode::Plain),
));
assert!(!mode_enabled_for_proto_with_policy(
&config,
ProtoTag::Secure,
false,
MtprotoModePolicy::Web(WebSecretMode::Plain),
));
assert!(mode_enabled_for_proto_with_policy(
&config,
ProtoTag::Secure,
false,
MtprotoModePolicy::Web(WebSecretMode::Dd),
));
assert!(!mode_enabled_for_proto_with_policy(
&config,
ProtoTag::Intermediate,
false,
MtprotoModePolicy::Web(WebSecretMode::Dd),
));
}
}
pub(super) fn decode_user_secrets_in(
shared: &ProxySharedState,
config: &ProxyConfig,
+63 -10
View File
@@ -1,6 +1,6 @@
use super::*;
/// Handle MTProto obfuscation handshake
/// Handles an MTProto obfuscation handshake with isolated test state.
#[cfg(test)]
pub async fn handle_mtproto_handshake<R, W>(
handshake: &[u8; HANDSHAKE_LEN],
@@ -26,11 +26,14 @@ where
replay_checker,
is_tls,
preferred_user,
None,
MtprotoModePolicy::Configured,
shared.as_ref(),
)
.await
}
/// Handles an MTProto obfuscation handshake with process-shared defenses.
pub async fn handle_mtproto_handshake_with_shared<R, W>(
handshake: &[u8; HANDSHAKE_LEN],
reader: R,
@@ -55,6 +58,40 @@ where
replay_checker,
is_tls,
preferred_user,
None,
MtprotoModePolicy::Configured,
shared,
)
.await
}
/// Authenticates one WEB logical stream against exactly one user and secret mode.
pub(crate) async fn handle_mtproto_handshake_for_web_user<R, W>(
handshake: &[u8; HANDSHAKE_LEN],
reader: R,
writer: W,
peer: SocketAddr,
config: &ProxyConfig,
replay_checker: &ReplayChecker,
exact_user: &str,
secret_mode: WebSecretMode,
shared: &ProxySharedState,
) -> HandshakeResult<(CryptoReader<R>, CryptoWriter<W>, HandshakeSuccess), R, W>
where
R: AsyncRead + Unpin + Send,
W: AsyncWrite + Unpin + Send,
{
handle_mtproto_handshake_impl(
handshake,
reader,
writer,
peer,
config,
replay_checker,
false,
None,
Some(exact_user),
MtprotoModePolicy::Web(secret_mode),
shared,
)
.await
@@ -69,6 +106,8 @@ async fn handle_mtproto_handshake_impl<R, W>(
replay_checker: &ReplayChecker,
is_tls: bool,
preferred_user: Option<&str>,
exact_user: Option<&str>,
mode_policy: MtprotoModePolicy,
shared: &ProxySharedState,
) -> HandshakeResult<(CryptoReader<R>, CryptoWriter<W>, HandshakeSuccess), R, W>
where
@@ -113,8 +152,11 @@ where
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 preferred_user_id = preferred_user.and_then(|user| snapshot.user_id_by_name(user));
let has_hint =
sticky_ip_hint.is_some() || sticky_prefix_hint.is_some() || preferred_user_id.is_some();
let exact_user_id = exact_user.and_then(|user| snapshot.user_id_by_name(user));
let has_hint = sticky_ip_hint.is_some()
|| sticky_prefix_hint.is_some()
|| preferred_user_id.is_some()
|| exact_user_id.is_some();
let overload = auth_probe_saturation_is_throttled_in(shared, Instant::now());
let candidate_budget = budget_for_validation(snapshot.entries().len(), overload, has_hint);
@@ -145,6 +187,7 @@ where
&entry.secret,
config,
is_tls,
mode_policy,
) {
matched_user = entry.user.clone();
matched_user_id = Some($user_id);
@@ -159,20 +202,20 @@ where
}};
}
let mut matched = false;
if let Some(user_id) = sticky_ip_hint {
let mut matched = exact_user_id.is_some_and(|user_id| try_user_id!(user_id));
if exact_user.is_none() && let Some(user_id) = sticky_ip_hint {
matched = try_user_id!(user_id);
}
if !matched && let Some(user_id) = preferred_user_id {
if exact_user.is_none() && !matched && let Some(user_id) = preferred_user_id {
matched = try_user_id!(user_id);
}
if !matched && let Some(user_id) = sticky_prefix_hint {
if exact_user.is_none() && !matched && let Some(user_id) = sticky_prefix_hint {
matched = try_user_id!(user_id);
}
if !matched && !budget_exhausted {
if exact_user.is_none() && !matched && !budget_exhausted {
let ring = &shared.handshake.recent_user_ring;
if !ring.is_empty() {
let next_seq = shared
@@ -197,7 +240,7 @@ where
}
}
if !matched && !budget_exhausted {
if exact_user.is_none() && !matched && !budget_exhausted {
for idx in 0..snapshot.entries().len() {
let Some(user_id) = u32::try_from(idx).ok() else {
break;
@@ -317,7 +360,16 @@ where
success,
));
} else {
let decoded_users = decode_user_secrets_in(shared, config, preferred_user);
let decoded_users = match exact_user {
Some(user) => config
.access
.users
.get(user)
.and_then(|secret| decode_user_secret(shared, user, secret))
.map(|secret| vec![(user.to_string(), secret)])
.unwrap_or_default(),
None => decode_user_secrets_in(shared, config, preferred_user),
};
let mut validation_checks = 0usize;
for (user, secret) in decoded_users {
@@ -337,6 +389,7 @@ where
&secret_arr,
config,
is_tls,
mode_policy,
) else {
continue;
};
+45 -2
View File
@@ -26,7 +26,8 @@ use crate::proxy::route_mode::{
RelayRouteMode, RouteCutoverState, affected_cutover_state, cutover_stagger_delay,
};
use crate::proxy::shared_state::{
ConntrackCloseEvent, ConntrackClosePublishResult, ConntrackCloseReason, ProxySharedState,
ConntrackCloseEvent, ConntrackClosePolicy, ConntrackClosePublishResult, ConntrackCloseReason,
ProxySharedState,
};
use crate::proxy::traffic_limiter::{RateDirection, TrafficLease, next_refill_delay};
use crate::stats::{
@@ -44,7 +45,7 @@ mod session;
pub(crate) use self::desync::DesyncDedupRotationState;
pub(crate) use self::idle::{RelayIdleCandidateRegistry, note_global_relay_pressure};
pub(crate) use self::session::handle_via_middle_proxy;
pub(crate) use self::session::handle_via_middle_proxy_with_conntrack;
use self::c2me::{
C2MeCommand, acquire_c2me_payload_permit, c2me_queued_permit_budget, enqueue_c2me_command_in,
@@ -91,6 +92,47 @@ pub(crate) use self::idle::{
set_relay_pressure_state_for_testing,
};
/// Runs Middle-End relay for a kernel-backed TCP client tuple.
pub(crate) async fn handle_via_middle_proxy<R, W>(
crypto_reader: CryptoReader<R>,
crypto_writer: CryptoWriter<W>,
success: HandshakeSuccess,
me_pool: Arc<MePool>,
stats: Arc<Stats>,
config: Arc<ProxyConfig>,
buffer_pool: Arc<BufferPool>,
local_addr: SocketAddr,
rng: Arc<SecureRandom>,
route_rx: watch::Receiver<RouteCutoverState>,
route_snapshot: RouteCutoverState,
session_id: u64,
session_cancel: CancellationToken,
shared: Arc<ProxySharedState>,
) -> Result<()>
where
R: AsyncRead + Unpin + Send + 'static,
W: AsyncWrite + Unpin + Send + 'static,
{
handle_via_middle_proxy_with_conntrack(
crypto_reader,
crypto_writer,
success,
me_pool,
stats,
config,
buffer_pool,
local_addr,
rng,
route_rx,
route_snapshot,
session_id,
session_cancel,
shared,
ConntrackClosePolicy::Publish,
)
.await
}
const DESYNC_DEDUP_WINDOW: Duration = Duration::from_secs(60);
const DESYNC_DEDUP_MAX_ENTRIES: usize = 65_536;
const DESYNC_FULL_CACHE_EMIT_MIN_INTERVAL: Duration = Duration::from_millis(1000);
@@ -98,6 +140,7 @@ const DESYNC_ERROR_CLASS: &str = "frame_too_large_crypto_desync";
const C2ME_CHANNEL_CAPACITY_FALLBACK: usize = 128;
const C2ME_SOFT_PRESSURE_MIN_FREE_SLOTS: usize = 64;
const C2ME_SENDER_FAIRNESS_BUDGET: usize = 32;
const C2ME_QUEUED_BYTE_PERMIT_UNIT: usize = 16 * 1024;
const C2ME_QUEUED_PERMITS_PER_SLOT: usize = 4;
const RELAY_IDLE_IO_POLL_MAX: Duration = Duration::from_secs(1);
+18 -14
View File
@@ -1,6 +1,7 @@
use super::*;
pub(crate) async fn handle_via_middle_proxy<R, W>(
/// Runs Middle-End relay with explicit kernel-conntrack close publication policy.
pub(crate) async fn handle_via_middle_proxy_with_conntrack<R, W>(
mut crypto_reader: CryptoReader<R>,
crypto_writer: CryptoWriter<W>,
success: HandshakeSuccess,
@@ -15,6 +16,7 @@ pub(crate) async fn handle_via_middle_proxy<R, W>(
session_id: u64,
session_cancel: CancellationToken,
shared: Arc<ProxySharedState>,
conntrack_close_policy: ConntrackClosePolicy,
) -> Result<()>
where
R: AsyncRead + Unpin + Send + 'static,
@@ -78,7 +80,7 @@ where
return Err(ProxyError::RouteSwitched);
}
// Per-user ad_tag from access.user_ad_tags; fallback to general.ad_tag (hot-reloadable)
// Prefer the hot-reloadable per-user ad tag over the global fallback.
let user_tag: Option<Vec<u8>> = config
.access
.user_ad_tags
@@ -785,7 +787,7 @@ where
}
};
// When client closes, but ME channel stopped as unregistered - it isnt error
// A client-initiated close can unregister the ME channel before its writer exits.
if client_closed && matches!(writer_result, Err(ProxyError::MiddleConnectionLost)) {
writer_result = Ok(());
}
@@ -808,17 +810,19 @@ where
"ME relay cleanup"
);
let close_reason = classify_conntrack_close_reason(&result);
let publish_result = shared.publish_conntrack_close_event(ConntrackCloseEvent {
src: peer,
dst: local_addr,
reason: close_reason,
});
if !matches!(
publish_result,
ConntrackClosePublishResult::Sent | ConntrackClosePublishResult::Disabled
) {
stats.increment_conntrack_close_event_drop_total();
if conntrack_close_policy == ConntrackClosePolicy::Publish {
let close_reason = classify_conntrack_close_reason(&result);
let publish_result = shared.publish_conntrack_close_event(ConntrackCloseEvent {
src: peer,
dst: local_addr,
reason: close_reason,
});
if !matches!(
publish_result,
ConntrackClosePublishResult::Sent | ConntrackClosePublishResult::Disabled
) {
stats.increment_conntrack_close_event_drop_total();
}
}
clear_relay_idle_candidate_in(shared.as_ref(), conn_id);
+2
View File
@@ -59,6 +59,8 @@
)]
pub mod adaptive_buffers;
// Shared authenticated admission and relay orchestration for TCP and WEB streams.
pub(crate) mod authenticated;
pub mod client;
// Process-wide Direct relay copy-buffer ownership and pressure policy.
pub(crate) mod direct_buffer_budget;
+9
View File
@@ -41,6 +41,15 @@ pub(crate) enum ConntrackClosePublishResult {
QueueClosed,
}
/// Controls whether a relay tuple maps to a real kernel conntrack entry.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum ConntrackClosePolicy {
/// Publish closure for a tuple backed by an accepted kernel TCP flow.
Publish,
/// Suppress closure for a virtual transport tuple with no kernel flow.
Suppress,
}
pub(crate) struct HandshakeSharedState {
pub(crate) auth_probe: DashMap<IpAddr, AuthProbeState>,
pub(crate) auth_probe_saturation: Mutex<Option<AuthProbeSaturationState>>,
+3
View File
@@ -260,6 +260,7 @@ mod tests {
fn listener(ip: IpAddr, port: Option<u16>, synlimit: SynLimitMode) -> ListenerConfig {
ListenerConfig {
ip,
transport: crate::config::ListenerTransport::Mtproxy,
port,
client_mss: None,
synlimit,
@@ -275,6 +276,8 @@ mod tests {
announce_ip: None,
proxy_protocol: None,
reuse_allow: false,
web_client_ip_source: crate::config::WebClientIpSource::XForwardedFor,
web_trusted_proxy_cidrs: Vec::new(),
}
}
+240
View File
@@ -0,0 +1,240 @@
use base64::Engine as _;
use crate::crypto::SecureRandom;
/// Browser security policy for the transient Telegram Desktop bridge page.
pub(crate) const PERMISSIONS_POLICY: &str = "accelerometer=(), autoplay=(), camera=(), clipboard-read=(), clipboard-write=(), display-capture=(), encrypted-media=(), fullscreen=(), geolocation=(), gyroscope=(), hid=(), idle-detection=(), magnetometer=(), microphone=(), midi=(), payment=(), picture-in-picture=(), publickey-credentials-create=(), publickey-credentials-get=(), screen-wake-lock=(), serial=(), usb=(), web-share=(), xr-spatial-tracking=()";
/// Fully rendered bridge response and its per-response script policy.
pub(crate) struct BridgePage {
/// Complete transient HTML document.
pub(crate) body: String,
/// Nonce-bound policy that authorizes only the embedded bridge script.
pub(crate) content_security_policy: String,
}
/// Renders the HTTPS-only WEB carrier bridge with a fresh CSP nonce.
pub(crate) fn render(
host: &str,
bootstrap: &str,
batch_limit: usize,
queue_limit: usize,
queue_items: usize,
rng: &SecureRandom,
) -> BridgePage {
let mut nonce = [0u8; 18];
rng.fill(&mut nonce);
let nonce = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(nonce);
let body = DOCUMENT
.replace("__NONCE__", &nonce)
.replace("__HOST__", host)
.replace("__BOOTSTRAP__", bootstrap)
.replace("__BATCH_LIMIT__", &batch_limit.to_string())
.replace("__QUEUE_LIMIT__", &queue_limit.to_string())
.replace("__QUEUE_ITEMS__", &queue_items.to_string());
BridgePage {
body,
content_security_policy: format!(
"default-src 'none'; base-uri 'none'; child-src 'none'; connect-src 'self' wss://{host}; font-src 'none'; form-action 'none'; frame-ancestors http://127.0.0.1:*; frame-src 'none'; img-src 'none'; manifest-src 'none'; media-src 'none'; object-src 'none'; script-src 'nonce-{nonce}'; style-src 'none'; worker-src 'none'; sandbox allow-same-origin allow-scripts"
),
}
}
const DOCUMENT: &str = r##"<!doctype html>
<html lang="en">
<head>
<meta charset="utf-8">
<meta name="viewport" content="width=device-width,initial-scale=1">
<title>Connection</title>
</head>
<body>
<script nonce="__NONCE__">
(()=>{
'use strict';
const relayOrigin='https://__HOST__',bootstrap='__BOOTSTRAP__';
const batchLimit=__BATCH_LIMIT__,queueLimit=__QUEUE_LIMIT__,queueItemLimit=__QUEUE_ITEMS__;
const fragment=location.hash,androidNonce=/^#android=([A-Za-z0-9_-]{43})$/.exec(fragment)?.[1]||'';
history.replaceState(null,'',location.pathname);
let initialized=false,closed=false,port=null,sessionToken='',createStarted=false;
let queuedBytes=0,queuedItems=0,upSequence=1,downCursor='0',upRunning=false,pollController=null;
const pending=[],upPending=[];
const status=state=>{if(port&&!closed)port.postMessage({t:'status',state})};
const pause=milliseconds=>new Promise(resolve=>setTimeout(resolve,milliseconds));
const options=(method,token,body,headers,signal,keepalive)=>({
method,body,signal,keepalive:!!keepalive,mode:'same-origin',credentials:'omit',cache:'no-store',redirect:'error',referrerPolicy:'no-referrer',
headers:Object.assign(token?{Authorization:'Bearer '+token}:{},body?{'Content-Type':'application/octet-stream'}:{},headers||{})
});
function reserve(data){
if(!data.byteLength||data.byteLength>queueLimit-queuedBytes||queuedItems>=queueItemLimit)return false;
queuedBytes+=data.byteLength;queuedItems++;return true;
}
function release(bytes,items){queuedBytes-=bytes;queuedItems-=items}
function frameBound(value,maxFrames,maxBytes){
const view=new DataView(value);let offset=0,frames=0;
while(offset<value.byteLength){
if(value.byteLength-offset<8)throw new Error('invalid frame batch');
const size=view.getUint32(offset+4),end=offset+8+size;
if(size>1048576||end>value.byteLength)throw new Error('invalid frame');
if(frames>0&&(frames>=maxFrames||end>maxBytes))break;
frames++;offset=end;
}
if(!frames)throw new Error('empty frame batch');
return {frames,bytes:offset};
}
function splitFrames(value){
const view=new DataView(value),result=[];let offset=0;
while(offset<value.byteLength){
if(value.byteLength-offset<8||result.length>=4096)throw new Error('invalid frame batch');
const size=view.getUint32(offset+4),end=offset+8+size;
if(size>1048576||end>value.byteLength)throw new Error('invalid frame');
result.push(offset===0&&end===value.byteLength?value:value.slice(offset,end));offset=end;
}
if(!result.length)throw new Error('empty frame batch');return result;
}
function joinPending(values){
let total=0,count=0,frames=0;
while(count<values.length){
const bound=frameBound(values[count],4096,batchLimit),whole=bound.bytes===values[count].byteLength;
if(count===0&&!whole){
const head=new Uint8Array(values[0],0,bound.bytes).slice();
values[0]=values[0].slice(bound.bytes);queuedItems++;
return {body:head.buffer,total:bound.bytes,count:1};
}
if(count&&(total+values[count].byteLength>batchLimit||frames+bound.frames>4096))break;
total+=values[count].byteLength;frames+=bound.frames;count++;
}
const joined=new Uint8Array(total);let offset=0;
for(const data of values.splice(0,count)){joined.set(new Uint8Array(data),offset);offset+=data.byteLength}
return {body:joined.buffer,total,count};
}
function retryAfterMs(response){
const value=Number(response.headers.get('Retry-After'));
return Number.isFinite(value)&&value>=0?Math.min(value*1000,30000):0;
}
async function request(path,makeOptions){
let delay=250,attempt=0;const deadline=Date.now()+90000;
while(true){
const requestOptions=makeOptions(),controller=new AbortController(),external=requestOptions.signal;
const abort=()=>controller.abort();if(external)external.addEventListener('abort',abort,{once:true});
requestOptions.signal=controller.signal;const timer=setTimeout(abort,90000);
let serviceUnavailable=false,wait=0;
try{
const response=await fetch(relayOrigin+path,requestOptions);
if(response.status!==503)return response;
serviceUnavailable=true;wait=retryAfterMs(response);await response.arrayBuffer();
}catch(error){
if(closed||(external&&external.aborted))throw error;
if(++attempt===9)throw new Error('carrier retry limit reached');
}finally{clearTimeout(timer);if(external)external.removeEventListener('abort',abort)}
if(serviceUnavailable&&Date.now()>=deadline)throw new Error('carrier retry limit reached');
status('reconnecting');await pause(wait||(delay+Math.floor(Math.random()*Math.max(1,delay/4))));
if(!serviceUnavailable)delay=Math.min(delay*2,5000);
}
}
function fail(){if(closed)return;status('failed');if(port)port.postMessage({t:'close'});close(true)}
async function createSession(first){
try{
status('connecting');
const response=await request('/api/v1/session',()=>options('POST',bootstrap,first));
if(response.status!==200||response.headers.get('X-Carrier-Mode')!=='https')throw new Error('session rejected');
sessionToken=response.headers.get('X-Session-Token')||'';downCursor=response.headers.get('X-Down-Cursor')||'0';
if(!/^[A-Za-z0-9_-]{43}$/.test(sessionToken)||downCursor!=='0')throw new Error('invalid session metadata');
if(closed){deleteSession();return}
const welcome=await response.arrayBuffer();
const welcomeBytes=new Uint8Array(welcome);
if(welcomeBytes.length!==8||welcomeBytes[0]!==17||welcomeBytes.slice(1).some(value=>value!==0))throw new Error('invalid welcome');
port.postMessage(welcome,[welcome]);status('connected');
for(const data of pending.splice(0)){release(data.byteLength,1);queueUp(data)}
poll();
}catch(error){fail()}
}
function queueUp(data){if(!reserve(data)){fail();return}upPending.push(data);runUp()}
async function runUp(){
if(upRunning)return;upRunning=true;
try{
while(!closed&&sessionToken&&upPending.length){
const batch=joinPending(upPending),sequence=String(upSequence);
const response=await request('/api/v1/up',()=>options('POST',sessionToken,batch.body,{'X-Up-Seq':sequence}));
if(response.status!==204||response.headers.get('X-Up-Ack')!==sequence)throw new Error('uplink rejected');
release(batch.total,batch.count);port.postMessage({t:'traffic',up:batch.total,down:0});upSequence++;
}
}catch(error){fail()}
finally{upRunning=false;if(!closed&&sessionToken&&upPending.length)runUp()}
}
async function poll(){
while(!closed&&sessionToken){
try{
pollController=new AbortController();
const response=await request('/api/v1/down',()=>options('POST',sessionToken,null,{'X-Down-Cursor':downCursor},pollController.signal));
if(response.status===204){status('connected');continue}
if(response.status!==200)throw new Error('downlink rejected');
const next=response.headers.get('X-Down-Cursor')||'',data=await response.arrayBuffer();
if(!next||!data.byteLength)throw new Error('invalid downlink response');
if(closed)return;
port.postMessage({t:'traffic',up:0,down:data.byteLength});port.postMessage(data,[data]);downCursor=next;status('connected');
}catch(error){if(!closed)fail();return}
}
}
function deleteSession(){
if(sessionToken)fetch(relayOrigin+'/api/v1/session',options('DELETE',sessionToken,null,null,undefined,true)).catch(()=>{});
}
function close(notifyServer){
if(closed)return;closed=true;if(pollController)pollController.abort();if(notifyServer)deleteSession();
pending.length=0;upPending.length=0;queuedBytes=0;queuedItems=0;if(port)port.close();
}
function activatePort(nextPort){
initialized=true;port=nextPort;
port.onmessage=message=>{
if(message.data instanceof ArrayBuffer){
if(!createStarted){createStarted=true;createSession(message.data)}
else if(!sessionToken){if(!reserve(message.data)){fail();return}pending.push(message.data)}
else queueUp(message.data);
}else if(message.data&&message.data.t==='close')close(true);
};
port.start();status('connecting');
}
addEventListener('message',event=>{
if(initialized||event.source!==parent||event.data===null||typeof event.data!=='object')return;
const keys=Object.keys(event.data).sort();
if(keys.length!==2||keys[0]!=='t'||keys[1]!=='v'||event.data.t!=='tproxy-init'||event.data.v!==1||event.ports.length!==1)return;
let source;try{source=new URL(event.origin)}catch(error){return}
if(source.protocol!=='http:'||source.hostname!=='127.0.0.1'||!source.port||source.origin!==event.origin)return;
activatePort(event.ports[0]);
});
const androidBridge=globalThis.TelegramWebProxy;
if(!initialized&&androidNonce&&androidBridge&&typeof androidBridge.postMessage==='function'){
const androidPort={onmessage:null,start(){},close(){androidBridge.onmessage=null},postMessage(value){
if(value instanceof ArrayBuffer){for(const item of splitFrames(value))androidBridge.postMessage(item)}else androidBridge.postMessage(JSON.stringify(value));
}};
androidBridge.onmessage=event=>{let data=event.data;if(typeof data==='string'){try{data=JSON.parse(data)}catch(error){return}}if(androidPort.onmessage)androidPort.onmessage({data})};
activatePort(androidPort);androidBridge.postMessage(JSON.stringify({t:'tproxy-android-init',v:1,nonce:androidNonce}));
}
addEventListener('pagehide',()=>close(true),{once:true});
})();
</script>
</body>
</html>
"##;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rendered_page_contains_no_template_markers_or_capability() {
let page = render(
"proxy.example.com",
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA",
2 * 1024 * 1024,
32 * 1024 * 1024,
16 * 1024,
&SecureRandom::new(),
);
assert!(!page.body.contains("__"));
assert!(!page.body.contains("bridge="));
assert!(page.body.contains("X-Up-Seq"));
assert!(page
.content_security_policy
.contains("frame-ancestors http://127.0.0.1:*"));
}
}
+272
View File
@@ -0,0 +1,272 @@
use bytes::{BufMut, Bytes, BytesMut};
use crate::config::WebLimitsConfig;
/// Fixed WEB frame header size.
pub(crate) const HEADER_BYTES: usize = 8;
/// Initial bidirectional stream credit.
pub(crate) const INITIAL_STREAM_WINDOW: u32 = 4 * 1024 * 1024;
/// Maximum data chunk emitted by the server.
pub(crate) const DATA_CHUNK_BYTES: usize = 64 * 1024;
/// WEB frame type codes shared with Telegram Desktop.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[repr(u8)]
pub(crate) enum FrameType {
/// Opens a logical MTProxy stream.
Open = 0x01,
/// Carries logical-stream payload bytes.
Data = 0x02,
/// Closes a logical stream.
Close = 0x03,
/// Returns consumed flow-control credit.
Window = 0x04,
/// Requests an application-level liveness response.
Ping = 0x05,
/// Answers application-level liveness traffic.
Pong = 0x06,
/// Starts one WEB carrier session.
Hello = 0x10,
/// Confirms WEB carrier session creation.
Welcome = 0x11,
/// Terminates a WEB carrier session.
Bye = 0x1f,
}
impl FrameType {
fn parse(value: u8) -> Option<Self> {
match value {
0x01 => Some(Self::Open),
0x02 => Some(Self::Data),
0x03 => Some(Self::Close),
0x04 => Some(Self::Window),
0x05 => Some(Self::Ping),
0x06 => Some(Self::Pong),
0x10 => Some(Self::Hello),
0x11 => Some(Self::Welcome),
0x1f => Some(Self::Bye),
_ => None,
}
}
}
/// One parsed frame borrowing its payload from the HTTP request body.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) struct Frame<'a> {
/// Parsed frame type.
pub(crate) frame_type: FrameType,
/// Logical 24-bit stream identifier.
pub(crate) stream_id: u32,
/// Borrowed frame payload.
pub(crate) payload: &'a [u8],
}
/// Protocol parse or shape failure.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum FrameError {
/// A carrier body contained no frame.
EmptyBatch,
/// A carrier body exceeded the configured frame count.
TooManyFrames,
/// A frame header or payload was truncated.
Incomplete,
/// A frame payload exceeded its configured ceiling.
PayloadLimit,
/// The frame type code is not defined.
UnknownType,
/// A known frame violated direction-specific grammar.
InvalidShape,
}
/// Parses and validates all frame boundaries without copying payloads.
pub(crate) fn parse_all<'a>(
input: &'a [u8],
limits: &WebLimitsConfig,
) -> std::result::Result<Vec<Frame<'a>>, FrameError> {
if input.is_empty() {
return Err(FrameError::EmptyBatch);
}
let mut remaining = input;
let mut frames = Vec::with_capacity(remaining.len().div_ceil(HEADER_BYTES).min(16));
while !remaining.is_empty() {
if frames.len() >= limits.max_frames_per_body {
return Err(FrameError::TooManyFrames);
}
if remaining.len() < HEADER_BYTES {
return Err(FrameError::Incomplete);
}
let frame_type = FrameType::parse(remaining[0]).ok_or(FrameError::UnknownType)?;
let stream_id = u32::from(remaining[1]) << 16
| u32::from(remaining[2]) << 8
| u32::from(remaining[3]);
let payload_len = u32::from_be_bytes([
remaining[4],
remaining[5],
remaining[6],
remaining[7],
]) as usize;
if payload_len > limits.max_frame_payload_bytes {
return Err(FrameError::PayloadLimit);
}
let frame_len = HEADER_BYTES
.checked_add(payload_len)
.ok_or(FrameError::PayloadLimit)?;
if frame_len > remaining.len() {
return Err(FrameError::Incomplete);
}
frames.push(Frame {
frame_type,
stream_id,
payload: &remaining[HEADER_BYTES..frame_len],
});
remaining = &remaining[frame_len..];
}
Ok(frames)
}
/// Enforces the client-to-server frame grammar.
pub(crate) fn validate_client_shape(frame: Frame<'_>) -> std::result::Result<(), FrameError> {
if frame.stream_id == 0 {
return if frame.frame_type == FrameType::Pong && frame.payload.len() <= 64 {
Ok(())
} else {
Err(FrameError::InvalidShape)
};
}
match frame.frame_type {
FrameType::Open | FrameType::Close if frame.payload.is_empty() => Ok(()),
FrameType::Data if !frame.payload.is_empty() => Ok(()),
FrameType::Window => window_amount(frame.payload).map(|_| ()),
_ => Err(FrameError::InvalidShape),
}
}
/// Validates the exact first-session HELLO body.
pub(crate) fn validate_hello(input: &[u8], limits: &WebLimitsConfig) -> bool {
let Ok(frames) = parse_all(input, limits) else {
return false;
};
frames.len() == 1
&& frames[0].frame_type == FrameType::Hello
&& frames[0].stream_id == 0
&& frames[0].payload == [1]
}
/// Encodes one complete WEB frame.
pub(crate) fn encode(frame_type: FrameType, stream_id: u32, payload: &[u8]) -> Bytes {
let mut output = BytesMut::with_capacity(HEADER_BYTES + payload.len());
output.put_u8(frame_type as u8);
output.put_u8((stream_id >> 16) as u8);
output.put_u8((stream_id >> 8) as u8);
output.put_u8(stream_id as u8);
output.put_u32(payload.len() as u32);
output.extend_from_slice(payload);
output.freeze()
}
/// Decodes a non-zero WINDOW delta.
pub(crate) fn window_amount(payload: &[u8]) -> std::result::Result<u32, FrameError> {
let bytes: [u8; 4] = payload.try_into().map_err(|_| FrameError::InvalidShape)?;
let amount = u32::from_be_bytes(bytes);
(amount != 0)
.then_some(amount)
.ok_or(FrameError::InvalidShape)
}
/// Encodes a WINDOW delta payload.
pub(crate) fn window_payload(amount: u32) -> [u8; 4] {
amount.to_be_bytes()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn hello_and_welcome_match_reference_bytes() {
let limits = WebLimitsConfig::default();
let hello = encode(FrameType::Hello, 0, &[1]);
assert_eq!(hello.as_ref(), &hex::decode("100000000000000101").unwrap());
assert!(validate_hello(&hello, &limits));
assert_eq!(
encode(FrameType::Welcome, 0, &[]).as_ref(),
&hex::decode("1100000000000000").unwrap()
);
}
#[test]
fn parser_rejects_excessive_payload_before_slicing() {
let limits = WebLimitsConfig {
max_frame_payload_bytes: 4,
..WebLimitsConfig::default()
};
let frame = encode(FrameType::Data, 1, &[0; 5]);
assert_eq!(parse_all(&frame, &limits), Err(FrameError::PayloadLimit));
}
#[test]
fn client_shape_rejects_control_types_on_stream_zero() {
let frame = Frame {
frame_type: FrameType::Ping,
stream_id: 0,
payload: &[],
};
assert_eq!(validate_client_shape(frame), Err(FrameError::InvalidShape));
}
#[test]
fn stream_frames_match_client_reference_vectors() {
assert_eq!(
encode(FrameType::Open, 17, &[]).as_ref(),
&hex::decode("0100001100000000").unwrap()
);
assert_eq!(
encode(FrameType::Data, 17, b"round trip").as_ref(),
&hex::decode("020000110000000a726f756e642074726970").unwrap()
);
assert_eq!(
encode(FrameType::Window, 17, &10u32.to_be_bytes()).as_ref(),
&hex::decode("04000011000000040000000a").unwrap()
);
assert_eq!(
encode(FrameType::Open, 0x00ff_ffff, &[]).as_ref(),
&hex::decode("01ffffff00000000").unwrap()
);
}
#[test]
fn parser_rejects_empty_truncated_and_excessive_batches() {
let mut limits = WebLimitsConfig::default();
assert_eq!(parse_all(&[], &limits), Err(FrameError::EmptyBatch));
assert_eq!(
parse_all(&hex::decode("0200000100000001").unwrap(), &limits),
Err(FrameError::Incomplete)
);
limits.max_frames_per_body = 1;
let mut body = encode(FrameType::Pong, 0, &[]).to_vec();
body.extend_from_slice(&encode(FrameType::Pong, 0, &[]));
assert_eq!(parse_all(&body, &limits), Err(FrameError::TooManyFrames));
}
#[test]
fn client_shape_rejects_empty_data_and_zero_window() {
let empty_data = Frame {
frame_type: FrameType::Data,
stream_id: 1,
payload: &[],
};
let zero_window = Frame {
frame_type: FrameType::Window,
stream_id: 1,
payload: &[0; 4],
};
assert_eq!(
validate_client_shape(empty_data),
Err(FrameError::InvalidShape)
);
assert_eq!(
validate_client_shape(zero_window),
Err(FrameError::InvalidShape)
);
}
}
+511
View File
@@ -0,0 +1,511 @@
use std::convert::Infallible;
use std::error::Error;
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use std::time::{Duration, Instant};
use bytes::Bytes;
use http_body_util::combinators::UnsyncBoxBody;
use http_body_util::{BodyExt, Full};
use hyper::body::Incoming;
use hyper::header::{self, HeaderName, HeaderValue};
use hyper::server::conn::http1;
use hyper::service::service_fn;
use hyper::{Method, Request, Response, StatusCode};
use hyper_util::rt::{TokioIo, TokioTimer};
use ipnetwork::IpNetwork;
use parking_lot::Mutex;
use tokio::net::TcpStream;
use tokio_util::sync::CancellationToken;
use crate::config::{WebClientIpSource, WebRuntimeVhost};
use crate::web::bridge;
use crate::web::frame::{self, FrameType};
use crate::web::manager::{ManagerError, WebProcessRuntime};
// Response-body activity keeps connection idle accounting lifecycle-correct.
mod activity;
// Body collection retains allocation permits through request processing.
mod body;
// Decoy routing and upstream proxying are isolated from carrier authentication.
mod decoy;
// Canonical request parsing rejects ambiguous credentials before routing.
mod request;
#[cfg(test)]
mod tests;
use decoy::serve_decoy;
use activity::{ActivityBody, RequestActivity};
use body::{CollectBodyError, CollectedBody, collect_body};
use request::{
bearer_token_hash, binary_content_type, bridge_candidate, canonical_request_host,
canonical_u64_header, client_ip, match_profile,
};
type BoxError = Box<dyn Error + Send + Sync>;
type HttpBody = UnsyncBoxBody<Bytes, BoxError>;
type HttpResponse = Response<HttpBody>;
const CREATE_BODY_LIMIT: usize = 64;
const TRANSPORT_PATHS: [&str; 3] = ["/api/v1/session", "/api/v1/up", "/api/v1/down"];
/// Serves one bounded HTTP/1.1 connection accepted from an external TLS terminator.
pub(crate) async fn serve_connection(
stream: TcpStream,
peer: SocketAddr,
client_ip_source: WebClientIpSource,
trusted_proxy_cidrs: Arc<[IpNetwork]>,
runtime: Arc<WebProcessRuntime>,
cancellation: CancellationToken,
connection_permit: tokio::sync::OwnedSemaphorePermit,
) {
let config = runtime.active_generation().config();
let max_header_bytes = config.web.limits.max_header_bytes;
let header_timeout = Duration::from_secs(config.web.timeouts.header_secs);
let idle_timeout = Duration::from_secs(config.web.timeouts.http_idle_secs);
let last_activity = Arc::new(Mutex::new(Instant::now()));
let service_last_activity = Arc::clone(&last_activity);
let service = service_fn(move |request| {
let runtime = Arc::clone(&runtime);
let trusted_proxy_cidrs = Arc::clone(&trusted_proxy_cidrs);
let last_activity = Arc::clone(&service_last_activity);
let client_ip_source = client_ip_source;
async move {
let activity = RequestActivity::begin(last_activity);
let response = if let Some(_handler_permit) = runtime.try_http_handler() {
handle_request(
request,
peer,
client_ip_source,
&trusted_proxy_cidrs,
runtime,
)
.await
} else {
service_unavailable()
};
let response = response.map(|body| {
ActivityBody::new(body, activity)
.boxed_unsync()
});
Ok::<_, Infallible>(response)
}
});
let connection = http1::Builder::new()
.timer(TokioTimer::new())
.header_read_timeout(header_timeout)
.max_buf_size(max_header_bytes)
.keep_alive(true)
.serve_connection(TokioIo::new(stream), service);
tokio::pin!(connection);
let mut idle_check = tokio::time::interval((idle_timeout / 2).max(Duration::from_secs(1)));
idle_check.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
loop {
tokio::select! {
biased;
_ = cancellation.cancelled() => break,
_ = &mut connection => break,
_ = idle_check.tick() => {
if Instant::now().saturating_duration_since(*last_activity.lock())
>= idle_timeout
{
break;
}
}
}
}
drop(connection_permit);
}
async fn handle_request(
request: Request<Incoming>,
peer: SocketAddr,
client_ip_source: WebClientIpSource,
trusted_proxy_cidrs: &[IpNetwork],
runtime: Arc<WebProcessRuntime>,
) -> HttpResponse {
let generation = runtime.active_generation();
let config = generation.config();
let Some(web_runtime) = config.web.runtime.as_ref() else {
return generic_not_found();
};
let Some(host) = canonical_request_host(&request) else {
return generic_not_found();
};
let Some(vhost) = web_runtime.vhosts.get(host).cloned() else {
return generic_not_found();
};
let path = request.uri().path();
if TRANSPORT_PATHS.contains(&path) {
return handle_api(
request,
peer,
client_ip_source,
trusted_proxy_cidrs,
runtime,
vhost,
)
.await;
}
if path == "/" && matches!(*request.method(), Method::GET | Method::HEAD) {
return handle_root(
request,
peer,
client_ip_source,
trusted_proxy_cidrs,
runtime,
vhost,
)
.await;
}
serve_decoy(request, vhost, false, &runtime).await
}
async fn handle_root(
mut request: Request<Incoming>,
peer: SocketAddr,
client_ip_source: WebClientIpSource,
trusted_proxy_cidrs: &[IpNetwork],
runtime: Arc<WebProcessRuntime>,
vhost: Arc<WebRuntimeVhost>,
) -> HttpResponse {
let (candidate, canonical) = bridge_candidate(request.uri().query());
let profile = match_profile(&vhost, &candidate);
let Some(profile) = profile.filter(|_| canonical && request.method() == Method::GET) else {
return serve_decoy(request, vhost, false, &runtime).await;
};
let Some(client_ip) = client_ip(
&request,
peer,
client_ip_source,
trusted_proxy_cidrs,
) else {
strip_query(&mut request);
return serve_decoy(request, vhost, true, &runtime).await;
};
let Ok(bootstrap) = runtime.issue_bootstrap(profile, client_ip) else {
strip_query(&mut request);
return serve_decoy(request, vhost, true, &runtime).await;
};
let generation = runtime.active_generation();
let config = generation.config();
let page = bridge::render(
&vhost.host,
&bootstrap,
config.web.limits.carrier_batch_bytes,
config.web.limits.pending_bytes_per_session,
config.web.limits.pending_items_per_session,
&generation.rng,
);
let mut response = full_response(StatusCode::OK, Bytes::from(page.body));
response.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static("text/html; charset=utf-8"),
);
insert_header(
&mut response,
header::CONTENT_SECURITY_POLICY,
&page.content_security_policy,
);
response
.headers_mut()
.insert(header::CACHE_CONTROL, HeaderValue::from_static("no-store"));
response.headers_mut().insert(
header::REFERRER_POLICY,
HeaderValue::from_static("no-referrer"),
);
response.headers_mut().insert(
header::X_CONTENT_TYPE_OPTIONS,
HeaderValue::from_static("nosniff"),
);
response.headers_mut().insert(
HeaderName::from_static("x-dns-prefetch-control"),
HeaderValue::from_static("off"),
);
insert_header(
&mut response,
HeaderName::from_static("permissions-policy"),
bridge::PERMISSIONS_POLICY,
);
response
}
async fn handle_api(
request: Request<Incoming>,
peer: SocketAddr,
client_ip_source: WebClientIpSource,
trusted_proxy_cidrs: &[IpNetwork],
runtime: Arc<WebProcessRuntime>,
vhost: Arc<WebRuntimeVhost>,
) -> HttpResponse {
if request.uri().query().is_some()
|| request.headers().contains_key(header::COOKIE)
|| request.headers().contains_key("x-lane-id")
{
return serve_decoy(request, vhost, true, &runtime).await;
}
let Some(client_ip) = client_ip(
&request,
peer,
client_ip_source,
trusted_proxy_cidrs,
) else {
return serve_decoy(request, vhost, true, &runtime).await;
};
let Some(token_hash) = bearer_token_hash(&request) else {
return serve_decoy(request, vhost, true, &runtime).await;
};
match request.uri().path() {
"/api/v1/session" => {
handle_session(request, runtime, vhost, token_hash, client_ip).await
}
"/api/v1/up" => handle_up(request, runtime, vhost, token_hash).await,
"/api/v1/down" => handle_down(request, runtime, vhost, token_hash).await,
_ => serve_decoy(request, vhost, true, &runtime).await,
}
}
async fn handle_session(
request: Request<Incoming>,
runtime: Arc<WebProcessRuntime>,
vhost: Arc<WebRuntimeVhost>,
token_hash: crate::web::manager::TokenHash,
client_ip: IpAddr,
) -> HttpResponse {
if request.method() == Method::DELETE {
if request.headers().contains_key(header::CONTENT_TYPE) {
return serve_decoy(request, vhost, true, &runtime).await;
}
let CollectedBody {
request,
body,
_body_budget,
} = match collect_body(request, &runtime, 1, true).await {
Ok(result) => result,
Err(CollectBodyError::Limit) => return service_unavailable(),
Err(CollectBodyError::Invalid(request)) => {
return serve_decoy(request, vhost, true, &runtime).await;
}
};
if !body.is_empty() || runtime.close_token(token_hash, &vhost.host).is_err() {
return serve_decoy(request, vhost, true, &runtime).await;
}
return carrier_empty(StatusCode::NO_CONTENT);
}
if request.method() != Method::POST || !binary_content_type(&request) {
return serve_decoy(request, vhost, true, &runtime).await;
}
if !runtime.has_bootstrap(token_hash, &vhost.host) {
return serve_decoy(request, vhost, true, &runtime).await;
}
let CollectedBody {
request,
body,
_body_budget,
} = match collect_body(request, &runtime, CREATE_BODY_LIMIT, false).await {
Ok(result) => result,
Err(CollectBodyError::Limit) => return service_unavailable(),
Err(CollectBodyError::Invalid(request)) => {
return serve_decoy(request, vhost, true, &runtime).await;
}
};
match runtime.create_session(token_hash, &vhost.host, client_ip, &body) {
Ok(result) => {
let welcome = frame::encode(FrameType::Welcome, 0, &[]);
let mut response = full_response(StatusCode::OK, welcome);
carrier_headers(&mut response);
insert_header(
&mut response,
HeaderName::from_static("x-session-token"),
&result.token,
);
response.headers_mut().insert(
HeaderName::from_static("x-carrier-mode"),
HeaderValue::from_static("https"),
);
response.headers_mut().insert(
HeaderName::from_static("x-down-cursor"),
HeaderValue::from_static("0"),
);
response
}
Err(ManagerError::Limit | ManagerError::Backpressure | ManagerError::Concurrent) => {
service_unavailable()
}
Err(_) => serve_decoy(request, vhost, true, &runtime).await,
}
}
async fn handle_up(
request: Request<Incoming>,
runtime: Arc<WebProcessRuntime>,
vhost: Arc<WebRuntimeVhost>,
token_hash: crate::web::manager::TokenHash,
) -> HttpResponse {
if request.method() != Method::POST || !binary_content_type(&request) {
return serve_decoy(request, vhost, true, &runtime).await;
}
let Some(sequence) = canonical_u64_header(&request, "x-up-seq").filter(|value| *value != 0)
else {
return serve_decoy(request, vhost, true, &runtime).await;
};
let Ok(session) = runtime.get_session(token_hash, &vhost.host) else {
return serve_decoy(request, vhost, true, &runtime).await;
};
let limit = runtime.active_generation().config().web.limits.max_body_bytes;
let CollectedBody {
request,
body,
_body_budget,
} = match collect_body(request, &runtime, limit, false).await {
Ok(result) => result,
Err(CollectBodyError::Limit) => return service_unavailable(),
Err(CollectBodyError::Invalid(request)) => {
return serve_decoy(request, vhost, true, &runtime).await;
}
};
match session.process_up(sequence, &body) {
Ok(ack) => {
let mut response = carrier_empty(StatusCode::NO_CONTENT);
insert_header(
&mut response,
HeaderName::from_static("x-up-ack"),
&ack.to_string(),
);
response
}
Err(ManagerError::Backpressure | ManagerError::Concurrent | ManagerError::Limit) => {
service_unavailable()
}
Err(_) => serve_decoy(request, vhost, true, &runtime).await,
}
}
async fn handle_down(
request: Request<Incoming>,
runtime: Arc<WebProcessRuntime>,
vhost: Arc<WebRuntimeVhost>,
token_hash: crate::web::manager::TokenHash,
) -> HttpResponse {
if request.method() != Method::POST || request.headers().contains_key(header::CONTENT_TYPE) {
return serve_decoy(request, vhost, true, &runtime).await;
}
let Some(cursor) = canonical_u64_header(&request, "x-down-cursor") else {
return serve_decoy(request, vhost, true, &runtime).await;
};
let Ok(session) = runtime.get_session(token_hash, &vhost.host) else {
return serve_decoy(request, vhost, true, &runtime).await;
};
let CollectedBody {
request,
body,
_body_budget,
} = match collect_body(request, &runtime, 1, true).await {
Ok(result) => result,
Err(CollectBodyError::Limit) => return service_unavailable(),
Err(CollectBodyError::Invalid(request)) => {
return serve_decoy(request, vhost, true, &runtime).await;
}
};
if !body.is_empty() {
return serve_decoy(request, vhost, true, &runtime).await;
}
match session.poll_down(cursor).await {
Ok(result) if result.body.is_empty() => {
let mut response = carrier_empty(StatusCode::NO_CONTENT);
insert_header(
&mut response,
HeaderName::from_static("x-down-cursor"),
&result.next_cursor.to_string(),
);
response
}
Ok(result) => {
let mut response = full_response(StatusCode::OK, result.body);
carrier_headers(&mut response);
insert_header(
&mut response,
HeaderName::from_static("x-down-cursor"),
&result.next_cursor.to_string(),
);
response
}
Err(ManagerError::Concurrent | ManagerError::Backpressure | ManagerError::Limit) => {
service_unavailable()
}
Err(_) => serve_decoy(request, vhost, true, &runtime).await,
}
}
fn carrier_headers(response: &mut HttpResponse) {
response.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static("application/octet-stream"),
);
response
.headers_mut()
.insert(header::CACHE_CONTROL, HeaderValue::from_static("no-store"));
}
fn carrier_empty(status: StatusCode) -> HttpResponse {
let mut response = empty_response(status);
response
.headers_mut()
.insert(header::CACHE_CONTROL, HeaderValue::from_static("no-store"));
response
}
fn service_unavailable() -> HttpResponse {
let mut response = carrier_empty(StatusCode::SERVICE_UNAVAILABLE);
response
.headers_mut()
.insert(header::RETRY_AFTER, HeaderValue::from_static("1"));
response
}
fn bad_gateway() -> HttpResponse {
full_response(
StatusCode::BAD_GATEWAY,
Bytes::from_static(b"site unavailable\n"),
)
}
fn generic_not_found() -> HttpResponse {
full_response(
StatusCode::NOT_FOUND,
Bytes::from_static(b"not found\n"),
)
}
fn full_response(status: StatusCode, body: Bytes) -> HttpResponse {
let length = body.len();
let body = Full::new(body)
.map_err(|never| -> BoxError { match never {} })
.boxed_unsync();
let mut response = Response::new(body);
*response.status_mut() = status;
insert_header(
&mut response,
header::CONTENT_LENGTH,
&length.to_string(),
);
response
}
fn empty_response(status: StatusCode) -> HttpResponse {
full_response(status, Bytes::new())
}
fn insert_header(response: &mut HttpResponse, name: HeaderName, value: &str) {
if let Ok(value) = HeaderValue::from_str(value) {
response.headers_mut().insert(name, value);
}
}
fn strip_query<B>(request: &mut Request<B>) {
if request.uri().query().is_some()
&& let Ok(uri) = request.uri().path().parse()
{
*request.uri_mut() = uri;
}
}
+69
View File
@@ -0,0 +1,69 @@
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::Instant;
use bytes::Bytes;
use hyper::body::{Body, Frame, SizeHint};
use parking_lot::Mutex;
use super::{BoxError, HttpBody};
/// Request lifecycle guard that refreshes HTTP connection activity on completion.
pub(super) struct RequestActivity {
last_activity: Arc<Mutex<Instant>>,
}
impl RequestActivity {
/// Starts activity accounting for one HTTP request.
pub(super) fn begin(last_activity: Arc<Mutex<Instant>>) -> Self {
*last_activity.lock() = Instant::now();
Self { last_activity }
}
}
impl Drop for RequestActivity {
fn drop(&mut self) {
*self.last_activity.lock() = Instant::now();
}
}
/// Response body wrapper that refreshes activity while downstream data progresses.
pub(super) struct ActivityBody {
inner: HttpBody,
activity: RequestActivity,
}
impl ActivityBody {
/// Binds one response body to its request activity guard.
pub(super) fn new(inner: HttpBody, activity: RequestActivity) -> Self {
Self {
inner,
activity,
}
}
}
impl Body for ActivityBody {
type Data = Bytes;
type Error = BoxError;
fn poll_frame(
mut self: Pin<&mut Self>,
context: &mut Context<'_>,
) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
let result = Pin::new(&mut self.inner).poll_frame(context);
if result.is_ready() {
*self.activity.last_activity.lock() = Instant::now();
}
result
}
fn is_end_stream(&self) -> bool {
self.inner.is_end_stream()
}
fn size_hint(&self) -> SizeHint {
self.inner.size_hint()
}
}
+80
View File
@@ -0,0 +1,80 @@
use std::time::Duration;
use bytes::Bytes;
use http_body_util::{BodyExt, Empty, Limited};
use hyper::body::{Body as _, Incoming};
use hyper::Request;
use crate::web::manager::WebProcessRuntime;
/// Collected carrier request retaining its process-wide body reservation.
pub(super) struct CollectedBody {
/// Request head reconstructed without the consumed network body.
pub(super) request: Request<Empty<Bytes>>,
/// Fully collected bounded carrier payload.
pub(super) body: Bytes,
/// Byte-budget reservation held through request processing.
pub(super) _body_budget: tokio::sync::OwnedSemaphorePermit,
}
// Keep rejected requests inline to avoid attacker-controlled allocations on invalid bodies.
#[allow(clippy::large_enum_variant)]
/// Body collection failure with sanitized request context when decoy routing is safe.
pub(super) enum CollectBodyError {
/// The body shape, size, or deadline failed after retaining the request head.
Invalid(Request<Empty<Bytes>>),
/// Process-wide body reader or byte capacity is temporarily exhausted.
Limit,
}
/// Collects one bounded carrier body under reader, byte, and deadline ownership.
pub(super) async fn collect_body(
request: Request<Incoming>,
runtime: &WebProcessRuntime,
limit: usize,
allow_empty: bool,
) -> Result<CollectedBody, CollectBodyError> {
let exceeds_limit = request.body().size_hint().lower() > limit as u64
|| request
.body()
.size_hint()
.upper()
.is_some_and(|upper| upper > limit as u64);
let (parts, body) = request.into_parts();
if exceeds_limit {
return Err(CollectBodyError::Invalid(Request::from_parts(
parts,
Empty::new(),
)));
}
let Some((reader_budget, body_budget)) = runtime.try_body_budget(limit) else {
return Err(CollectBodyError::Limit);
};
let body_timeout = Duration::from_secs(
runtime
.active_generation()
.config()
.web
.timeouts
.body_secs,
);
let body = match tokio::time::timeout(body_timeout, Limited::new(body, limit).collect()).await {
Ok(Ok(body)) => body.to_bytes(),
_ => {
return Err(CollectBodyError::Invalid(Request::from_parts(
parts,
Empty::new(),
)));
}
};
drop(reader_budget);
let request = Request::from_parts(parts, Empty::new());
if !allow_empty && body.is_empty() {
return Err(CollectBodyError::Invalid(request));
}
Ok(CollectedBody {
request,
body,
_body_budget: body_budget,
})
}
+284
View File
@@ -0,0 +1,284 @@
use std::error::Error;
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Duration;
use bytes::Bytes;
use http_body_util::{BodyExt, Empty};
use hyper::header::{self, HeaderName, HeaderValue};
use hyper::{Method, Request, StatusCode, Uri};
use hyper_util::rt::TokioIo;
use tokio::net::TcpStream;
use super::{
BoxError, HttpBody, HttpResponse, bad_gateway, full_response, generic_not_found,
insert_header,
};
use crate::config::{WebRuntimeDecoy, WebRuntimeVhost};
use crate::web::manager::WebProcessRuntime;
/// Serves the configured ordinary site after optionally removing carrier material.
pub(super) async fn serve_decoy<B>(
mut request: Request<B>,
vhost: Arc<WebRuntimeVhost>,
sanitize_transport: bool,
runtime: &WebProcessRuntime,
) -> HttpResponse
where
B: hyper::body::Body<Data = Bytes> + Send + 'static,
B::Error: Error + Send + Sync + 'static,
{
if sanitize_transport {
sanitize_transport_request(&mut request);
}
let (parts, body) = request.into_parts();
let body = if sanitize_transport {
Empty::<Bytes>::new()
.map_err(|never| -> BoxError { match never {} })
.boxed_unsync()
} else {
body.map_err(|error| -> BoxError { Box::new(error) })
.boxed_unsync()
};
let request = Request::from_parts(parts, body);
match &vhost.decoy {
WebRuntimeDecoy::StaticDirectory(site) => serve_static(request, site),
WebRuntimeDecoy::HttpUpstream { addr, authority } => {
proxy_to_upstream(
request,
*addr,
authority,
Duration::from_secs(vhost.decoy_header_secs),
runtime,
)
.await
}
}
}
fn serve_static<B>(request: Request<B>, site: &crate::config::WebStaticSite) -> HttpResponse {
if !matches!(*request.method(), Method::GET | Method::HEAD) {
return static_entry(request, site, None, StatusCode::NOT_FOUND);
}
let path = request.uri().path();
let resolved = resolve_static_path(path, site);
let status = if resolved.is_some() {
StatusCode::OK
} else {
StatusCode::NOT_FOUND
};
static_entry(request, site, resolved, status)
}
fn static_entry<B>(
request: Request<B>,
site: &crate::config::WebStaticSite,
route: Option<&str>,
status: StatusCode,
) -> HttpResponse {
let fallback = format!("/{}", site.index);
let not_found = site.assets.contains_key("/404.html").then_some("/404.html");
let route = route.or(not_found).unwrap_or(&fallback);
let Some(asset) = site.assets.get(route) else {
return generic_not_found();
};
let not_modified = status == StatusCode::OK
&& request
.headers()
.get(header::IF_NONE_MATCH)
.and_then(|value| value.to_str().ok())
== Some(asset.etag.as_str());
let response_body = if request.method() == Method::HEAD || not_modified {
Bytes::new()
} else {
asset.body.clone()
};
let mut response = full_response(
if not_modified {
StatusCode::NOT_MODIFIED
} else {
status
},
response_body,
);
insert_header(&mut response, header::CONTENT_TYPE, asset.content_type);
insert_header(&mut response, header::ETAG, &asset.etag);
insert_header(
&mut response,
header::CONTENT_LENGTH,
&asset.body.len().to_string(),
);
response.headers_mut().insert(
header::CACHE_CONTROL,
if status.is_client_error() || request.uri().query().is_some() {
HeaderValue::from_static("no-store")
} else {
HeaderValue::from_static("public, max-age=300")
},
);
response.headers_mut().insert(
header::CONTENT_SECURITY_POLICY,
HeaderValue::from_static("default-src 'self'; style-src 'self'; img-src 'self'; worker-src 'none'; frame-ancestors 'none'; base-uri 'none'; form-action 'none'"),
);
response.headers_mut().insert(
header::REFERRER_POLICY,
HeaderValue::from_static("strict-origin-when-cross-origin"),
);
response.headers_mut().insert(
header::X_CONTENT_TYPE_OPTIONS,
HeaderValue::from_static("nosniff"),
);
response.headers_mut().insert(header::X_FRAME_OPTIONS, HeaderValue::from_static("DENY"));
response
}
fn resolve_static_path<'a>(
path: &str,
site: &'a crate::config::WebStaticSite,
) -> Option<&'a str> {
if !path.starts_with('/')
|| path.contains('\\')
|| path.contains("//")
|| path.split('/').any(|part| matches!(part, "." | ".."))
{
return None;
}
let root;
let route = if path == "/" {
root = format!("/{}", site.index);
root.as_str()
} else {
path
};
if site.assets.contains_key(route) {
return site.assets.get_key_value(route).map(|(key, _)| key.as_str());
}
if route == "/favicon.ico" && site.assets.contains_key("/favicon.svg") {
return Some("/favicon.svg");
}
if !route.rsplit('/').next().unwrap_or_default().contains('.') {
let html = format!("{route}.html");
return site
.assets
.get_key_value(&html)
.map(|(key, _)| key.as_str());
}
None
}
async fn proxy_to_upstream(
mut request: Request<HttpBody>,
addr: SocketAddr,
authority: &str,
header_timeout: Duration,
runtime: &WebProcessRuntime,
) -> HttpResponse {
remove_hop_by_hop(request.headers_mut());
if let Ok(host) = HeaderValue::from_str(authority) {
request.headers_mut().insert(header::HOST, host);
}
let path_and_query = request
.uri()
.path_and_query()
.map(|value| value.as_str())
.unwrap_or("/");
let Ok(uri) = path_and_query.parse::<Uri>() else {
return bad_gateway();
};
*request.uri_mut() = uri;
let stream = match tokio::time::timeout(header_timeout, TcpStream::connect(addr)).await {
Ok(Ok(stream)) => stream,
_ => return bad_gateway(),
};
let max_header_bytes = runtime
.active_generation()
.config()
.web
.limits
.max_header_bytes;
let mut builder = hyper::client::conn::http1::Builder::new();
builder.max_buf_size(max_header_bytes);
let (mut sender, connection) = match tokio::time::timeout(
header_timeout,
builder.handshake(TokioIo::new(stream)),
)
.await
{
Ok(Ok(parts)) => parts,
_ => return bad_gateway(),
};
runtime.spawn_auxiliary(async move {
let _ = connection.await;
});
let mut response = match tokio::time::timeout(header_timeout, sender.send_request(request)).await
{
Ok(Ok(response)) => response,
_ => return bad_gateway(),
};
remove_hop_by_hop(response.headers_mut());
response.map(|body| {
body.map_err(|error| -> BoxError { Box::new(error) })
.boxed_unsync()
})
}
fn sanitize_transport_request<B>(request: &mut Request<B>) {
for name in [
header::AUTHORIZATION,
header::CONTENT_LENGTH,
header::CONTENT_TYPE,
header::UPGRADE,
HeaderName::from_static("sec-websocket-key"),
HeaderName::from_static("sec-websocket-protocol"),
HeaderName::from_static("sec-websocket-version"),
HeaderName::from_static("x-down-cursor"),
HeaderName::from_static("x-lane-id"),
HeaderName::from_static("x-up-seq"),
] {
request.headers_mut().remove(name);
}
request
.headers_mut()
.insert(header::CONNECTION, HeaderValue::from_static("close"));
}
fn remove_hop_by_hop(headers: &mut hyper::HeaderMap) {
let nominated = headers
.get_all(header::CONNECTION)
.iter()
.filter_map(|value| value.to_str().ok())
.flat_map(|value| value.split(','))
.filter_map(|value| HeaderName::from_bytes(value.trim().as_bytes()).ok())
.collect::<Vec<_>>();
for name in nominated {
headers.remove(name);
}
for name in [
header::CONNECTION,
header::PROXY_AUTHENTICATE,
header::PROXY_AUTHORIZATION,
header::TE,
header::TRAILER,
header::TRANSFER_ENCODING,
header::UPGRADE,
HeaderName::from_static("keep-alive"),
HeaderName::from_static("proxy-connection"),
] {
headers.remove(name);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn static_resolver_rejects_rewritten_paths() {
let site = crate::config::WebStaticSite {
assets: std::collections::BTreeMap::new(),
index: "index.html".to_string(),
};
assert!(resolve_static_path("/../index.html", &site).is_none());
assert!(resolve_static_path("//index.html", &site).is_none());
}
}
+240
View File
@@ -0,0 +1,240 @@
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use base64::Engine as _;
use hyper::header;
use hyper::Request;
use ipnetwork::IpNetwork;
use sha2::{Digest, Sha256};
use subtle::ConstantTimeEq;
use crate::config::{
WebClientIpSource, WebRuntimeProfile, WebRuntimeVhost,
};
use crate::web::manager::TokenHash;
/// Parses one lowercase canonical Host value restricted to the public HTTPS port.
pub(super) fn canonical_request_host<B>(request: &Request<B>) -> Option<&str> {
let values = request.headers().get_all(header::HOST);
let mut values = values.iter();
let value = values.next()?.to_str().ok()?;
if values.next().is_some() {
return None;
}
let authority = value.parse::<hyper::http::uri::Authority>().ok()?;
if authority.port_u16().is_some_and(|port| port != 443) {
return None;
}
let host = value.strip_suffix(":443").unwrap_or(value);
if authority.host() != host || host.bytes().any(|byte| byte.is_ascii_uppercase())
{
return None;
}
Some(host)
}
/// Accepts one canonical forwarded client address from an explicitly trusted peer.
pub(super) fn client_ip<B>(
request: &Request<B>,
peer: SocketAddr,
source: WebClientIpSource,
trusted_proxy_cidrs: &[IpNetwork],
) -> Option<IpAddr> {
if !trusted_proxy_cidrs
.iter()
.any(|network| network.contains(peer.ip()))
{
return None;
}
let header_name = match source {
WebClientIpSource::XForwardedFor => "x-forwarded-for",
};
let values = request.headers().get_all(header_name);
let mut values = values.iter();
let value = values.next()?.to_str().ok()?;
if values.next().is_some()
|| value.is_empty()
|| value.trim() != value
|| value.contains(',')
{
return None;
}
let ip = value.parse::<IpAddr>().ok()?;
(ip.to_string() == value).then_some(ip)
}
/// Decodes an exact canonical bridge query without allocating credential strings.
pub(super) fn bridge_candidate(query: Option<&str>) -> ([u8; 32], bool) {
let mut candidate = [0u8; 32];
let Some(value) = query.and_then(|query| query.strip_prefix("bridge=")) else {
return (candidate, false);
};
if value.len() != 43 {
return (candidate, false);
}
let mut decoded = [0u8; 32];
let Ok(decoded_len) = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode_slice(value, &mut decoded)
else {
return (candidate, false);
};
let mut canonical = [0u8; 43];
let Ok(encoded_len) = base64::engine::general_purpose::URL_SAFE_NO_PAD
.encode_slice(decoded, &mut canonical)
else {
return (candidate, false);
};
if decoded_len != decoded.len()
|| encoded_len != canonical.len()
|| !bool::from(canonical.ct_eq(value.as_bytes()))
{
return (candidate, false);
}
candidate = decoded;
(candidate, true)
}
/// Matches a capability in constant time across every profile of one virtual host.
pub(super) fn match_profile(
vhost: &WebRuntimeVhost,
candidate: &[u8; 32],
) -> Option<Arc<WebRuntimeProfile>> {
let mut matched = None;
for profile in &vhost.profiles {
if bool::from(profile.capability.ct_eq(candidate)) {
matched = Some(Arc::clone(profile));
}
}
matched
}
/// Validates and hashes one canonical bearer credential for map lookup.
pub(super) fn bearer_token_hash<B>(request: &Request<B>) -> Option<TokenHash> {
let values = request.headers().get_all(header::AUTHORIZATION);
let mut values = values.iter();
let value = values.next()?.to_str().ok()?;
if values.next().is_some() || !value.starts_with("Bearer ") || value.matches(' ').count() != 1
{
return None;
}
let token = value.strip_prefix("Bearer ")?;
if token.len() != 43 {
return None;
}
let mut decoded = [0u8; 32];
let decoded_len = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode_slice(token, &mut decoded)
.ok()?;
let mut canonical = [0u8; 43];
let encoded_len = base64::engine::general_purpose::URL_SAFE_NO_PAD
.encode_slice(decoded, &mut canonical)
.ok()?;
(decoded_len == decoded.len()
&& encoded_len == canonical.len()
&& bool::from(canonical.ct_eq(token.as_bytes())))
.then(|| Sha256::digest(decoded).into())
}
/// Checks the exact carrier media type without accepting duplicate headers.
pub(super) fn binary_content_type<B>(request: &Request<B>) -> bool {
let values = request.headers().get_all(header::CONTENT_TYPE);
let mut values = values.iter();
let value = values.next().and_then(|value| value.to_str().ok());
values.next().is_none()
&& value.is_some_and(|value| value.eq_ignore_ascii_case("application/octet-stream"))
}
/// Parses one canonical unsigned decimal carrier sequence header.
pub(super) fn canonical_u64_header<B>(
request: &Request<B>,
name: &'static str,
) -> Option<u64> {
let values = request.headers().get_all(name);
let mut values = values.iter();
let value = values.next()?.to_str().ok()?;
if values.next().is_some()
|| value.is_empty()
|| value.starts_with('+')
|| (value.len() > 1 && value.starts_with('0'))
{
return None;
}
let parsed = value.parse::<u64>().ok()?;
(parsed.to_string() == value).then_some(parsed)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn canonical_bridge_query_rejects_aliases() {
let token = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode([7u8; 32]);
assert!(bridge_candidate(Some(&format!("bridge={token}"))).1);
assert!(!bridge_candidate(Some(&format!("x=1&bridge={token}"))).1);
assert!(!bridge_candidate(Some(&format!("bridge={token}="))).1);
}
#[test]
fn host_and_forwarded_identity_require_canonical_single_values() {
let request = Request::builder()
.header(header::HOST, "proxy.example.com:443")
.header("x-forwarded-for", "192.0.2.10")
.body(())
.unwrap();
assert_eq!(
canonical_request_host(&request),
Some("proxy.example.com")
);
let trusted: [IpNetwork; 1] = ["127.0.0.1/32".parse().unwrap()];
assert_eq!(
client_ip(
&request,
"127.0.0.1:40000".parse().unwrap(),
WebClientIpSource::XForwardedFor,
&trusted,
),
Some("192.0.2.10".parse().unwrap())
);
let uppercase = Request::builder()
.header(header::HOST, "Proxy.Example.com")
.body(())
.unwrap();
assert!(canonical_request_host(&uppercase).is_none());
let appended = Request::builder()
.header("x-forwarded-for", "192.0.2.10, 198.51.100.4")
.body(())
.unwrap();
assert!(
client_ip(
&appended,
"127.0.0.1:40000".parse().unwrap(),
WebClientIpSource::XForwardedFor,
&trusted,
)
.is_none()
);
}
#[test]
fn bearer_and_sequence_headers_reject_noncanonical_aliases() {
let token = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode([1u8; 32]);
let request = Request::builder()
.header(header::AUTHORIZATION, format!("Bearer {token}"))
.header("x-up-seq", "17")
.body(())
.unwrap();
assert_eq!(
bearer_token_hash(&request),
Some(Sha256::digest([1u8; 32]).into())
);
assert_eq!(canonical_u64_header(&request, "x-up-seq"), Some(17));
let leading_zero = Request::builder()
.header("x-up-seq", "017")
.body(())
.unwrap();
assert!(canonical_u64_header(&leading_zero, "x-up-seq").is_none());
}
}
+206
View File
@@ -0,0 +1,206 @@
use std::collections::BTreeMap;
use std::sync::Arc;
use arc_swap::ArcSwap;
use base64::Engine as _;
use bytes::Bytes;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use tokio_util::sync::CancellationToken;
use super::serve_connection;
use crate::config::{
ProxyConfig, WebClientIpSource, WebRuntimeConfig, WebRuntimeDecoy,
WebRuntimeProfile, WebRuntimeVhost, WebSecretMode, WebStaticAsset, WebStaticSite,
};
use crate::maestro::generation::test_runtime_generation;
use crate::web::frame::{self, FrameType};
use crate::web::manager::WebProcessRuntime;
fn runtime_config(capability: [u8; 32]) -> ProxyConfig {
let profile = Arc::new(WebRuntimeProfile {
host: "proxy.example.com".to_string(),
public_addr: "203.0.113.10:443".parse().unwrap(),
user: "alice".to_string(),
secret_mode: WebSecretMode::Plain,
capability,
max_sessions: 4,
max_streams: 16,
max_streams_per_session: 4,
});
let mut assets = BTreeMap::new();
assets.insert(
"/index.html".to_string(),
WebStaticAsset {
body: Bytes::from_static(b"<!doctype html><title>decoy</title>"),
content_type: "text/html; charset=utf-8",
etag: "\"test\"".to_string(),
},
);
let site = Arc::new(WebStaticSite {
assets,
index: "index.html".to_string(),
});
let vhost = Arc::new(WebRuntimeVhost {
host: "proxy.example.com".to_string(),
decoy: WebRuntimeDecoy::StaticDirectory(Arc::clone(&site)),
decoy_header_secs: 1,
profiles: vec![Arc::clone(&profile)],
});
let mut vhosts = BTreeMap::new();
vhosts.insert("proxy.example.com".to_string(), vhost);
vhosts.insert(
"other.example.com".to_string(),
Arc::new(WebRuntimeVhost {
host: "other.example.com".to_string(),
decoy: WebRuntimeDecoy::StaticDirectory(site),
decoy_header_secs: 1,
profiles: Vec::new(),
}),
);
let mut config = ProxyConfig::default();
config.web.enabled = true;
config.web.limits.max_bootstraps_per_ip = 1;
config.web.timeouts.shutdown_secs = 1;
config.web.runtime = Some(Arc::new(WebRuntimeConfig {
vhosts,
profiles: vec![profile],
}));
config
}
async fn request(
listener: &TcpListener,
runtime: &Arc<WebProcessRuntime>,
request: Vec<u8>,
) -> Vec<u8> {
let addr = listener.local_addr().unwrap();
let (accepted, client) = tokio::join!(listener.accept(), TcpStream::connect(addr));
let (server, peer) = accepted.unwrap();
let mut client = client.unwrap();
let permit = runtime.try_http_connection().unwrap();
let task = tokio::spawn(serve_connection(
server,
peer,
WebClientIpSource::XForwardedFor,
Arc::from(["127.0.0.1/32".parse().unwrap()]),
Arc::clone(runtime),
CancellationToken::new(),
permit,
));
client.write_all(&request).await.unwrap();
let mut response = Vec::new();
client.read_to_end(&mut response).await.unwrap();
task.await.unwrap();
response
}
fn split_response(response: &[u8]) -> (&[u8], &[u8]) {
let separator = response
.windows(4)
.position(|window| window == b"\r\n\r\n")
.unwrap();
(&response[..separator], &response[separator + 4..])
}
fn response_header<'a>(headers: &'a [u8], name: &str) -> &'a str {
std::str::from_utf8(headers)
.unwrap()
.lines()
.filter_map(|line| line.split_once(':'))
.find_map(|(header, value)| header.eq_ignore_ascii_case(name).then_some(value.trim()))
.unwrap()
}
#[tokio::test]
async fn https_carrier_bootstraps_and_closes_one_session() {
let capability = [7u8; 32];
let generation = test_runtime_generation(1, runtime_config(capability));
let active_runtime = Arc::new(ArcSwap::from(Arc::clone(&generation)));
let runtime = WebProcessRuntime::start(Arc::clone(&active_runtime));
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let encoded = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(capability);
let wrong_family = format!(
"GET /?bridge={encoded} HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 2001:db8::10\r\nConnection: close\r\n\r\n"
)
.into_bytes();
let wrong_family_response = request(&listener, &runtime, wrong_family).await;
let (_, wrong_family_body) = split_response(&wrong_family_response);
assert!(!wrong_family_body
.windows(11)
.any(|value| value == b"bootstrap='"));
let root = format!(
"GET /?bridge={encoded} HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nConnection: close\r\n\r\n"
)
.into_bytes();
let root_response = request(&listener, &runtime, root).await;
let (root_headers, root_body) = split_response(&root_response);
assert!(root_headers.starts_with(b"HTTP/1.1 200"));
let root_body = std::str::from_utf8(root_body).unwrap();
let bootstrap = root_body
.split_once("bootstrap='")
.and_then(|(_, suffix)| suffix.split_once('\''))
.map(|(token, _)| token)
.unwrap();
assert_eq!(bootstrap.len(), 43);
let hello = frame::encode(FrameType::Hello, 0, &[1]);
let mut wrong_host = format!(
"POST /api/v1/session HTTP/1.1\r\nHost: other.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nAuthorization: Bearer {bootstrap}\r\nContent-Type: application/octet-stream\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
hello.len()
)
.into_bytes();
wrong_host.extend_from_slice(&hello);
let wrong_host_response = request(&listener, &runtime, wrong_host).await;
assert!(wrong_host_response.starts_with(b"HTTP/1.1 404"));
let mut create = format!(
"POST /api/v1/session HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nAuthorization: Bearer {bootstrap}\r\nContent-Type: application/octet-stream\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
hello.len()
)
.into_bytes();
let create_retry = create.clone();
create.extend_from_slice(&hello);
let mut create_retry = create_retry;
create_retry.extend_from_slice(&hello);
let create_response = request(&listener, &runtime, create).await;
let (create_headers, create_body) = split_response(&create_response);
assert!(create_headers.starts_with(b"HTTP/1.1 200"));
assert_eq!(response_header(create_headers, "x-carrier-mode"), "https");
assert_eq!(create_body, frame::encode(FrameType::Welcome, 0, &[]));
let session = response_header(create_headers, "x-session-token");
assert_eq!(session.len(), 43);
let replacement = test_runtime_generation(2, runtime_config(capability));
active_runtime.store(Arc::clone(&replacement));
tokio::time::sleep(std::time::Duration::from_millis(1100)).await;
let retry_response = request(&listener, &runtime, create_retry).await;
let (retry_headers, retry_body) = split_response(&retry_response);
assert!(retry_headers.starts_with(b"HTTP/1.1 200"));
assert_eq!(response_header(retry_headers, "x-session-token"), session);
assert_eq!(retry_body, frame::encode(FrameType::Welcome, 0, &[]));
let next_root = format!(
"GET /?bridge={encoded} HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nConnection: close\r\n\r\n"
)
.into_bytes();
let next_root_response = request(&listener, &runtime, next_root).await;
let (_, next_root_body) = split_response(&next_root_response);
assert!(next_root_body.windows(11).any(|value| value == b"bootstrap='"));
let close = format!(
"DELETE /api/v1/session HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nAuthorization: Bearer {session}\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
)
.into_bytes();
let close_retry = close.clone();
let close_response = request(&listener, &runtime, close).await;
assert!(close_response.starts_with(b"HTTP/1.1 204"));
let close_retry_response = request(&listener, &runtime, close_retry).await;
assert!(close_retry_response.starts_with(b"HTTP/1.1 204"));
runtime.shutdown().await;
generation.stop_sessions().await;
generation.stop_background_tasks().await;
replacement.stop_sessions().await;
replacement.stop_background_tasks().await;
}
+532
View File
@@ -0,0 +1,532 @@
use std::future::Future;
use std::net::IpAddr;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::time::{Duration, Instant};
use arc_swap::ArcSwap;
use parking_lot::Mutex;
use sha2::{Digest, Sha256};
use subtle::ConstantTimeEq;
use tokio::sync::{Notify, OwnedSemaphorePermit, Semaphore};
use tokio_util::sync::CancellationToken;
use tokio_util::task::TaskTracker;
use zeroize::Zeroizing;
use crate::config::{WebLimitsConfig, WebRuntimeProfile};
use crate::maestro::generation::RuntimeGeneration;
use crate::web::frame;
use crate::web::session::WebSession;
// Credential maps, quotas, and token-bucket helpers remain private to the manager.
mod state;
// Stream admission and synthetic tuple ownership are process-scoped.
mod admission;
// Shutdown and expiry work remain outside request-path coordination.
mod lifecycle;
use state::{
Bootstrap, ManagerState, allow_rate, control_item_reserve, decrement_map,
evict_oldest_unused_bootstrap, matching_profile, new_unique_token, profile_key,
remove_expired_locked,
};
const TOKEN_BYTES: usize = 32;
const CLEANUP_INTERVAL: Duration = Duration::from_secs(1);
/// Stable hash key used for bootstrap and session credentials.
pub(crate) type TokenHash = [u8; TOKEN_BYTES];
/// Stable non-allocating key used for per-profile quotas.
pub(crate) type ProfileKey = [u8; TOKEN_BYTES];
/// WEB manager operation failure category.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum ManagerError {
/// Credential, hostname, or ownership validation failed.
Authentication,
/// Bounded queue capacity is temporarily unavailable.
Backpressure,
/// A configured admission or rate ceiling was reached.
Limit,
/// Carrier framing or sequencing violated the protocol.
Protocol,
/// The operation conflicts with another in-flight operation.
Concurrent,
/// The process or session has stopped accepting work.
Closed,
}
/// Successful idempotent session creation result.
pub(crate) struct CreateResult {
/// Opaque bearer token for the created or replayed session.
pub(crate) token: String,
}
/// Process-owned bounded WEB credential, session, and memory coordinator.
pub(crate) struct WebProcessRuntime {
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
limits: WebLimitsConfig,
state: Mutex<ManagerState>,
http_connections: Arc<Semaphore>,
http_handlers: Arc<Semaphore>,
body_readers: Arc<Semaphore>,
body_bytes: Arc<Semaphore>,
stream_handshakes: Arc<Semaphore>,
budget_notify: Arc<Notify>,
budget_saturated: AtomicBool,
shutdown: CancellationToken,
tasks: TaskTracker,
sessions_created: AtomicU64,
sessions_closed: AtomicU64,
streams_opened: AtomicU64,
streams_rejected: AtomicU64,
bytes_up: AtomicU64,
bytes_down: AtomicU64,
limit_hits: AtomicU64,
}
impl WebProcessRuntime {
/// Starts one process-scoped manager using immutable allocation ceilings.
pub(crate) fn start(
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
) -> Arc<Self> {
let limits = active_runtime.load().config().web.limits.clone();
let runtime = Arc::new(Self {
active_runtime,
http_connections: Arc::new(Semaphore::new(limits.max_http_connections)),
http_handlers: Arc::new(Semaphore::new(limits.max_http_handlers)),
body_readers: Arc::new(Semaphore::new(limits.max_body_readers)),
body_bytes: Arc::new(Semaphore::new(limits.max_body_bytes_global)),
stream_handshakes: Arc::new(Semaphore::new(limits.max_stream_handshakes)),
limits,
state: Mutex::new(ManagerState::default()),
budget_notify: Arc::new(Notify::new()),
budget_saturated: AtomicBool::new(false),
shutdown: CancellationToken::new(),
tasks: TaskTracker::new(),
sessions_created: AtomicU64::new(0),
sessions_closed: AtomicU64::new(0),
streams_opened: AtomicU64::new(0),
streams_rejected: AtomicU64::new(0),
bytes_up: AtomicU64::new(0),
bytes_down: AtomicU64::new(0),
limit_hits: AtomicU64::new(0),
});
let weak = Arc::downgrade(&runtime);
let shutdown = runtime.shutdown.clone();
runtime.tasks.spawn(async move {
let mut interval = tokio::time::interval(CLEANUP_INTERVAL);
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
loop {
tokio::select! {
_ = shutdown.cancelled() => break,
_ = interval.tick() => {
let Some(runtime) = weak.upgrade() else {
break;
};
runtime.cleanup();
}
}
}
});
runtime
}
/// Loads the currently active generation without retaining older generations.
pub(crate) fn active_generation(&self) -> Arc<RuntimeGeneration> {
self.active_runtime.load_full()
}
/// Reserves one accepted HTTP connection.
pub(crate) fn try_http_connection(&self) -> Option<OwnedSemaphorePermit> {
let permit = Arc::clone(&self.http_connections).try_acquire_owned().ok();
if permit.is_none() {
self.record_limit_hit();
}
permit
}
/// Reserves one concurrently executing HTTP request handler.
pub(crate) fn try_http_handler(&self) -> Option<OwnedSemaphorePermit> {
let permit = Arc::clone(&self.http_handlers).try_acquire_owned().ok();
if permit.is_none() {
self.record_limit_hit();
}
permit
}
/// Reserves one logical stream in the inner MTProxy handshake phase.
pub(crate) fn try_stream_handshake(&self) -> Option<OwnedSemaphorePermit> {
let permit = Arc::clone(&self.stream_handshakes)
.try_acquire_owned()
.ok();
if permit.is_none() {
self.record_stream_rejected();
}
permit
}
/// Spawns one process-owned auxiliary task with shutdown cancellation.
pub(crate) fn spawn_auxiliary<F>(&self, future: F)
where
F: Future<Output = ()> + Send + 'static,
{
let shutdown = self.shutdown.clone();
self.tasks.spawn(async move {
tokio::select! {
_ = shutdown.cancelled() => {}
_ = future => {}
}
});
}
/// Reserves one body reader and its declared bounded body allocation.
pub(crate) fn try_body_budget(
&self,
bytes: usize,
) -> Option<(OwnedSemaphorePermit, OwnedSemaphorePermit)> {
let Some(bytes) = u32::try_from(bytes).ok() else {
self.record_limit_hit();
return None;
};
let Some(reader) = Arc::clone(&self.body_readers).try_acquire_owned().ok() else {
self.record_limit_hit();
return None;
};
let Some(body) = Arc::clone(&self.body_bytes)
.try_acquire_many_owned(bytes)
.ok()
else {
self.record_limit_hit();
return None;
};
Some((reader, body))
}
/// Issues a one-use bootstrap credential for the active generation.
pub(crate) fn issue_bootstrap(
&self,
profile: Arc<WebRuntimeProfile>,
client_ip: IpAddr,
) -> std::result::Result<String, ManagerError> {
let generation = self.active_generation();
let config = generation.config();
let profile = config
.web
.runtime
.as_ref()
.and_then(|runtime| matching_profile(runtime, &profile))
.ok_or(ManagerError::Authentication)?;
if !config.web.enabled
|| profile.public_addr.is_ipv4() != client_ip.is_ipv4()
|| !generation.proxy_shared.is_user_enabled(&profile.user)
{
return Err(ManagerError::Closed);
}
let now = Instant::now();
let mut state = self.state.lock();
remove_expired_locked(&mut state, now);
if state.closed
|| state.bootstraps_per_ip.get(&client_ip).copied().unwrap_or(0)
>= self.limits.max_bootstraps_per_ip
|| !allow_rate(
&mut state.bootstrap_rate,
now,
self.limits.new_bootstraps_per_minute,
self.limits.new_bootstraps_burst,
)
{
self.limit_hits.fetch_add(1, Ordering::Relaxed);
return Err(ManagerError::Limit);
}
if state.bootstraps.len() >= self.limits.max_bootstraps_global
&& !evict_oldest_unused_bootstrap(&mut state)
{
self.limit_hits.fetch_add(1, Ordering::Relaxed);
return Err(ManagerError::Limit);
}
let Some((token, hash)) = new_unique_token(&generation, &state) else {
self.limit_hits.fetch_add(1, Ordering::Relaxed);
return Err(ManagerError::Limit);
};
state.bootstraps.insert(
hash,
Bootstrap {
generation_id: generation.id,
expires_at: now + Duration::from_secs(config.web.timeouts.bootstrap_lifetime_secs),
issued_at: now,
issuance_ip: client_ip,
profile,
body_digest: [0; TOKEN_BYTES],
session_token: Zeroizing::new(String::new()),
session: None,
used: false,
},
);
*state.bootstraps_per_ip.entry(client_ip).or_insert(0) += 1;
Ok(token)
}
/// Checks whether a bootstrap token is live before reading a request body.
pub(crate) fn has_bootstrap(&self, hash: TokenHash, host: &str) -> bool {
let generation_id = self.active_runtime.load().id;
let now = Instant::now();
let state = self.state.lock();
state.bootstraps.get(&hash).is_some_and(|entry| {
entry.profile.host == host
&& now <= entry.expires_at
&& (entry.generation_id == generation_id
|| entry.used && entry.session.is_some())
})
}
/// Creates a session exactly once or replays the original successful result.
pub(crate) fn create_session(
self: &Arc<Self>,
bootstrap_hash: TokenHash,
host: &str,
client_ip: IpAddr,
body: &[u8],
) -> std::result::Result<CreateResult, ManagerError> {
if !frame::validate_hello(body, &self.limits) {
return Err(ManagerError::Protocol);
}
let body_digest: TokenHash = Sha256::digest(body).into();
let generation = self.active_generation();
let config = generation.config();
let now = Instant::now();
let mut state = self.state.lock();
remove_expired_locked(&mut state, now);
let Some(entry) = state.bootstraps.get(&bootstrap_hash) else {
return Err(ManagerError::Authentication);
};
if entry.profile.host != host || now > entry.expires_at {
return Err(ManagerError::Authentication);
}
if entry.used {
let digest_matches = bool::from(entry.body_digest.ct_eq(&body_digest));
if !digest_matches {
return Err(ManagerError::Authentication);
}
if entry.session.is_none() {
return Err(ManagerError::Authentication);
}
return Ok(CreateResult {
token: entry.session_token.as_str().to_owned(),
});
}
if entry.generation_id != generation.id {
return Err(ManagerError::Authentication);
}
if state.closed || !config.web.enabled {
return Err(ManagerError::Closed);
}
let profile = config
.web
.runtime
.as_ref()
.and_then(|runtime| matching_profile(runtime, &entry.profile))
.filter(|profile| {
profile.public_addr.is_ipv4() == client_ip.is_ipv4()
&& generation.proxy_shared.is_user_enabled(&profile.user)
})
.ok_or(ManagerError::Authentication)?;
let profile_key = profile_key(&profile);
if state.sessions.len() >= self.limits.max_sessions_global
|| state.sessions_per_ip.get(&client_ip).copied().unwrap_or(0)
>= self.limits.max_sessions_per_ip
|| state
.sessions_per_profile
.get(&profile_key)
.copied()
.unwrap_or(0)
>= profile.max_sessions
|| !allow_rate(
&mut state.session_rate,
now,
self.limits.new_sessions_per_minute,
self.limits.new_sessions_burst,
)
{
self.limit_hits.fetch_add(1, Ordering::Relaxed);
return Err(ManagerError::Limit);
}
let Some((session_token, session_hash)) = new_unique_token(&generation, &state) else {
self.limit_hits.fetch_add(1, Ordering::Relaxed);
return Err(ManagerError::Limit);
};
let session = WebSession::new(
Arc::downgrade(self),
session_hash,
client_ip,
profile,
profile_key,
self.limits.clone(),
config.web.timeouts.clone(),
);
state.sessions.insert(session_hash, Arc::clone(&session));
*state.sessions_per_ip.entry(client_ip).or_insert(0) += 1;
*state.sessions_per_profile.entry(profile_key).or_insert(0) += 1;
let entry = state
.bootstraps
.get_mut(&bootstrap_hash)
.ok_or(ManagerError::Authentication)?;
entry.used = true;
entry.body_digest = body_digest;
entry.session_token = Zeroizing::new(session_token.clone());
entry.session = Some(Arc::clone(&session));
let issuance_ip = entry.issuance_ip;
decrement_map(&mut state.bootstraps_per_ip, &issuance_ip);
self.sessions_created.fetch_add(1, Ordering::Relaxed);
Ok(CreateResult {
token: session_token,
})
}
/// Resolves an authenticated session token.
pub(crate) fn get_session(
&self,
hash: TokenHash,
host: &str,
) -> std::result::Result<Arc<WebSession>, ManagerError> {
self.state
.lock()
.sessions
.get(&hash)
.cloned()
.filter(|session| session.matches_host(host))
.ok_or(ManagerError::Authentication)
}
/// Closes a live token and accepts bounded tombstone retries.
pub(crate) fn close_token(
&self,
hash: TokenHash,
host: &str,
) -> std::result::Result<(), ManagerError> {
let state = self.state.lock();
let session = state
.sessions
.get(&hash)
.filter(|session| session.matches_host(host))
.cloned();
let closed = state
.closed_tokens
.get(&hash)
.is_some_and(|closed| closed.host == host);
drop(state);
if let Some(session) = session {
session.close();
return Ok(());
}
closed.then_some(()).ok_or(ManagerError::Authentication)
}
/// Reserves bounded process-wide queue capacity for data or control traffic.
pub(crate) fn try_reserve_pending(
&self,
bytes: usize,
items: usize,
control: bool,
downlink: bool,
) -> bool {
let mut state = self.state.lock();
let data_byte_limit = self
.limits
.pending_bytes_global
.saturating_sub(self.limits.control_bytes_global);
let control_item_reserve = control_item_reserve(&self.limits);
let data_item_limit = self
.limits
.pending_items_global
.saturating_sub(control_item_reserve);
if state.closed {
return false;
}
let fits = if control {
bytes <= self.limits.control_bytes_global
&& items <= control_item_reserve
&& state.pending_bytes
<= self.limits.pending_bytes_global.saturating_sub(bytes)
&& state.pending_items
<= self.limits.pending_items_global.saturating_sub(items)
&& state.pending_control_bytes
<= self.limits.control_bytes_global.saturating_sub(bytes)
&& state.pending_control_items
<= control_item_reserve.saturating_sub(items)
} else {
let data_bytes = state
.pending_bytes
.saturating_sub(state.pending_control_bytes);
let data_items = state
.pending_items
.saturating_sub(state.pending_control_items);
let (byte_limit, item_limit) = if downlink {
let uplink_bytes = self
.limits
.max_body_bytes
.saturating_add(
self.limits
.max_frames_per_body
.saturating_mul(crate::web::session::QUEUE_ITEM_COST),
);
(
data_byte_limit.saturating_sub(uplink_bytes),
data_item_limit.saturating_sub(self.limits.max_frames_per_body),
)
} else {
(data_byte_limit, data_item_limit)
};
bytes <= byte_limit
&& items <= item_limit
&& data_bytes <= byte_limit - bytes
&& data_items <= item_limit - items
};
if !fits {
self.budget_saturated.store(true, Ordering::Release);
self.record_limit_hit();
return false;
}
state.pending_bytes += bytes;
state.pending_items += items;
if control {
state.pending_control_bytes += bytes;
state.pending_control_items += items;
}
true
}
/// Releases process-wide queue capacity and wakes blocked relay writers.
pub(crate) fn release_pending(&self, bytes: usize, items: usize, control: bool) {
let mut state = self.state.lock();
state.pending_bytes = state.pending_bytes.saturating_sub(bytes);
state.pending_items = state.pending_items.saturating_sub(items);
if control {
state.pending_control_bytes = state.pending_control_bytes.saturating_sub(bytes);
state.pending_control_items = state.pending_control_items.saturating_sub(items);
}
drop(state);
if self.budget_saturated.swap(false, Ordering::AcqRel) {
self.budget_notify.notify_waiters();
}
}
/// Returns the shared notification source for global queue capacity changes.
pub(crate) fn budget_notify(&self) -> Arc<Notify> {
Arc::clone(&self.budget_notify)
}
/// Accounts one successfully committed carrier uplink body.
pub(crate) fn record_up(&self, bytes: usize) {
self.bytes_up.fetch_add(bytes as u64, Ordering::Relaxed);
}
/// Accounts one emitted carrier downlink body.
pub(crate) fn record_down(&self, bytes: usize) {
self.bytes_down.fetch_add(bytes as u64, Ordering::Relaxed);
}
fn record_limit_hit(&self) {
self.limit_hits.fetch_add(1, Ordering::Relaxed);
}
}
+130
View File
@@ -0,0 +1,130 @@
use std::net::{IpAddr, SocketAddr};
use std::sync::atomic::Ordering;
use std::time::Instant;
use super::state::{
allocate_stream_port, allow_rate, decrement_map, release_stream_port,
};
use super::{ProfileKey, WebProcessRuntime};
impl WebProcessRuntime {
/// Reserves one process-wide and per-profile live logical-stream slot.
pub(crate) fn try_acquire_stream(
&self,
profile_key: ProfileKey,
max_streams: usize,
client_ip: IpAddr,
public_addr: SocketAddr,
) -> Option<u16> {
let now = Instant::now();
let mut state = self.state.lock();
if state.closed
|| state.streams_live >= self.limits.max_streams_global
|| state
.streams_per_profile
.get(&profile_key)
.copied()
.unwrap_or(0)
>= max_streams
|| !allow_rate(
&mut state.stream_rate,
now,
self.limits.new_streams_per_minute,
self.limits.new_streams_burst,
)
{
self.streams_rejected.fetch_add(1, Ordering::Relaxed);
self.limit_hits.fetch_add(1, Ordering::Relaxed);
return None;
}
let Some(peer_port) = allocate_stream_port(&mut state, client_ip, public_addr) else {
self.streams_rejected.fetch_add(1, Ordering::Relaxed);
self.limit_hits.fetch_add(1, Ordering::Relaxed);
return None;
};
state.streams_live += 1;
*state
.streams_per_profile
.entry(profile_key)
.or_insert(0) += 1;
self.streams_opened.fetch_add(1, Ordering::Relaxed);
Some(peer_port)
}
/// Releases one live logical-stream slot after its relay task exits.
pub(crate) fn release_stream(
&self,
profile_key: ProfileKey,
client_ip: IpAddr,
public_addr: SocketAddr,
peer_port: u16,
) {
let mut state = self.state.lock();
if !release_stream_port(&mut state, client_ip, public_addr, peer_port) {
return;
}
state.streams_live = state.streams_live.saturating_sub(1);
decrement_map(&mut state.streams_per_profile, &profile_key);
}
/// Records a logical stream rejected outside manager quota acquisition.
pub(crate) fn record_stream_rejected(&self) {
self.streams_rejected.fetch_add(1, Ordering::Relaxed);
self.record_limit_hit();
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use arc_swap::ArcSwap;
use super::*;
use crate::config::ProxyConfig;
use crate::maestro::generation::test_runtime_generation;
use crate::web::session::QUEUE_ITEM_COST;
#[tokio::test]
async fn global_downlink_budget_preserves_one_maximum_uplink_batch() {
let generation = test_runtime_generation(1, ProxyConfig::default());
let runtime = WebProcessRuntime::start(Arc::new(ArcSwap::from(generation)));
let control_items = super::super::state::control_item_reserve(&runtime.limits);
let data_bytes = runtime
.limits
.pending_bytes_global
.saturating_sub(runtime.limits.control_bytes_global);
let data_items = runtime
.limits
.pending_items_global
.saturating_sub(control_items);
let uplink_bytes = runtime
.limits
.max_body_bytes
.saturating_add(runtime.limits.max_frames_per_body * QUEUE_ITEM_COST);
let downlink_bytes = data_bytes - uplink_bytes;
let downlink_items = data_items - runtime.limits.max_frames_per_body;
assert!(runtime.try_reserve_pending(
downlink_bytes,
downlink_items,
false,
true,
));
assert!(runtime.try_reserve_pending(
uplink_bytes,
runtime.limits.max_frames_per_body,
false,
false,
));
assert!(!runtime.try_reserve_pending(1, 1, false, true));
runtime.release_pending(downlink_bytes, downlink_items, false);
runtime.release_pending(
uplink_bytes,
runtime.limits.max_frames_per_body,
false,
);
runtime.shutdown().await;
}
}
+149
View File
@@ -0,0 +1,149 @@
use std::net::IpAddr;
use std::sync::atomic::Ordering;
use std::time::{Duration, Instant};
use tracing::info;
use super::{ProfileKey, TokenHash, WebProcessRuntime};
use super::state::{
ClosedToken, decrement_map, remove_bootstrap_locked, remove_expired_locked,
};
impl WebProcessRuntime {
/// Removes one closed session and retains a bounded host-bound replay marker.
pub(crate) fn session_finished(
&self,
hash: TokenHash,
client_ip: IpAddr,
profile_key: ProfileKey,
profile_host: &str,
) {
let mut state = self.state.lock();
if state.sessions.remove(&hash).is_none() {
return;
}
decrement_map(&mut state.sessions_per_ip, &client_ip);
decrement_map(&mut state.sessions_per_profile, &profile_key);
let expiry = Instant::now()
+ Duration::from_secs(
self.active_runtime
.load()
.config()
.web
.timeouts
.bootstrap_lifetime_secs,
);
state.closed_tokens.insert(
hash,
ClosedToken {
expires_at: expiry,
host: profile_host.to_string(),
},
);
while state.closed_tokens.len() > self.limits.max_sessions_global.saturating_mul(16) {
let Some(oldest) = state
.closed_tokens
.iter()
.min_by_key(|(_, closed)| closed.expires_at)
.map(|(hash, _)| *hash)
else {
break;
};
state.closed_tokens.remove(&oldest);
}
let bootstrap_hashes = state
.bootstraps
.iter()
.filter_map(|(bootstrap_hash, bootstrap)| {
bootstrap
.session
.as_ref()
.is_some_and(|session| session.token_hash() == hash)
.then_some(*bootstrap_hash)
})
.collect::<Vec<_>>();
for bootstrap_hash in bootstrap_hashes {
remove_bootstrap_locked(&mut state, bootstrap_hash);
}
self.sessions_closed.fetch_add(1, Ordering::Relaxed);
}
/// Stops issuance, closes all sessions, and joins bounded child work.
pub(crate) async fn shutdown(&self) {
self.shutdown.cancel();
let sessions = {
let mut state = self.state.lock();
state.closed = true;
state.bootstraps.clear();
state.bootstraps_per_ip.clear();
state.sessions.values().cloned().collect::<Vec<_>>()
};
for session in &sessions {
session.close();
}
let timeout_secs = self
.active_runtime
.load()
.config()
.web
.timeouts
.shutdown_secs;
let waits = async {
for session in sessions {
session.wait().await;
}
};
let _ = tokio::time::timeout(Duration::from_secs(timeout_secs), waits).await;
self.tasks.close();
let _ = tokio::time::timeout(Duration::from_secs(timeout_secs), self.tasks.wait()).await;
let (sessions_live, streams_live, pending_bytes, pending_items) = {
let state = self.state.lock();
(
state.sessions.len(),
state.streams_live,
state.pending_bytes,
state.pending_items,
)
};
info!(
target: "telemt::web",
sessions_created = self.sessions_created.load(Ordering::Relaxed),
sessions_closed = self.sessions_closed.load(Ordering::Relaxed),
sessions_live,
streams_opened = self.streams_opened.load(Ordering::Relaxed),
streams_rejected = self.streams_rejected.load(Ordering::Relaxed),
streams_live,
pending_bytes,
pending_items,
bytes_up = self.bytes_up.load(Ordering::Relaxed),
bytes_down = self.bytes_down.load(Ordering::Relaxed),
limit_hits = self.limit_hits.load(Ordering::Relaxed),
"WEB runtime stopped"
);
}
/// Expires credentials and closes idle sessions without holding locks across callbacks.
pub(super) fn cleanup(&self) {
let generation_id = self.active_runtime.load().id;
let now = Instant::now();
let sessions = {
let mut state = self.state.lock();
remove_expired_locked(&mut state, now);
let stale_bootstraps = state
.bootstraps
.iter()
.filter_map(|(hash, bootstrap)| {
(bootstrap.generation_id != generation_id && !bootstrap.used)
.then_some(*hash)
})
.collect::<Vec<_>>();
for hash in stale_bootstraps {
remove_bootstrap_locked(&mut state, hash);
}
state.sessions.values().cloned().collect::<Vec<_>>()
};
for session in sessions.into_iter().filter(|session| session.is_idle(now)) {
session.close();
}
}
}
+293
View File
@@ -0,0 +1,293 @@
use std::collections::{HashMap, HashSet};
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use std::time::Instant;
use base64::Engine as _;
use sha2::{Digest, Sha256};
use zeroize::Zeroizing;
use super::{ProfileKey, TOKEN_BYTES, TokenHash};
use crate::config::{WebLimitsConfig, WebRuntimeConfig, WebRuntimeProfile};
use crate::maestro::generation::RuntimeGeneration;
use crate::web::session::WebSession;
/// One issued bootstrap and optional idempotent session-creation replay state.
pub(super) struct Bootstrap {
/// Generation that issued the bootstrap.
pub(super) generation_id: u64,
/// Credential and replay-state expiry deadline.
pub(super) expires_at: Instant,
/// Stable ordering point used for bounded eviction.
pub(super) issued_at: Instant,
/// Forwarded client address that owns this credential.
pub(super) issuance_ip: IpAddr,
/// Immutable profile selected during capability validation.
pub(super) profile: Arc<WebRuntimeProfile>,
/// Digest of the accepted HELLO body for idempotent retry matching.
pub(super) body_digest: TokenHash,
/// Zeroizing copy returned only for an exact session-creation retry.
pub(super) session_token: Zeroizing<String>,
/// Created session retained while retry replay remains valid.
pub(super) session: Option<Arc<WebSession>>,
/// Distinguishes unused issuance quota from completed creation replay state.
pub(super) used: bool,
}
/// Bounded replay marker for one explicitly or naturally closed session token.
pub(super) struct ClosedToken {
/// Deadline after which the token hash may be forgotten.
pub(super) expires_at: Instant,
/// Canonical host that owned the session.
pub(super) host: String,
}
/// Token-bucket state for one process-wide creation class.
#[derive(Default)]
pub(super) struct RateState {
tokens: f64,
last: Option<Instant>,
}
struct StreamPortState {
active: HashSet<u16>,
next: u16,
}
/// Process-wide WEB registries and quota accounting protected by one short lock.
#[derive(Default)]
pub(super) struct ManagerState {
/// Bootstrap credentials indexed by their SHA-256 token hash.
pub(super) bootstraps: HashMap<TokenHash, Bootstrap>,
/// Unused bootstrap ownership counts by forwarded client address.
pub(super) bootstraps_per_ip: HashMap<IpAddr, usize>,
/// Live sessions indexed by bearer-token hash.
pub(super) sessions: HashMap<TokenHash, Arc<WebSession>>,
/// Recently closed token hashes retained for idempotent DELETE semantics.
pub(super) closed_tokens: HashMap<TokenHash, ClosedToken>,
/// Live session counts by forwarded client address.
pub(super) sessions_per_ip: HashMap<IpAddr, usize>,
/// Live session counts by stable profile key.
pub(super) sessions_per_profile: HashMap<ProfileKey, usize>,
/// Live relay-task counts by stable profile key.
pub(super) streams_per_profile: HashMap<ProfileKey, usize>,
/// Process-wide live relay-task count.
pub(super) streams_live: usize,
stream_ports: HashMap<(IpAddr, SocketAddr), StreamPortState>,
/// Total process-wide queued byte reservation.
pub(super) pending_bytes: usize,
/// Total process-wide queued item reservation.
pub(super) pending_items: usize,
/// Portion of queued bytes charged to the control reserve.
pub(super) pending_control_bytes: usize,
/// Portion of queued items charged to the control reserve.
pub(super) pending_control_items: usize,
/// Bootstrap issuance rate limiter.
pub(super) bootstrap_rate: RateState,
/// Session creation rate limiter.
pub(super) session_rate: RateState,
/// Logical-stream creation rate limiter.
pub(super) stream_rate: RateState,
/// Process shutdown admission latch.
pub(super) closed: bool,
}
/// Generates one collision-checked credential and its stable hash key.
pub(super) fn new_unique_token(
generation: &RuntimeGeneration,
state: &ManagerState,
) -> Option<(String, TokenHash)> {
for _ in 0..8 {
let mut raw = [0u8; TOKEN_BYTES];
generation.rng.fill(&mut raw);
let token = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(raw);
let hash = Sha256::digest(raw).into();
if !state.bootstraps.contains_key(&hash)
&& !state.sessions.contains_key(&hash)
&& !state.closed_tokens.contains_key(&hash)
{
return Some((token, hash));
}
}
None
}
/// Returns the precomputed capability as the stable process profile key.
pub(super) fn profile_key(profile: &WebRuntimeProfile) -> ProfileKey {
profile.capability
}
/// Re-resolves an issued profile against the active generation without weakening identity.
pub(super) fn matching_profile(
runtime: &WebRuntimeConfig,
expected: &WebRuntimeProfile,
) -> Option<Arc<WebRuntimeProfile>> {
runtime
.profiles
.iter()
.find(|profile| {
profile.host == expected.host
&& profile.public_addr == expected.public_addr
&& profile.user == expected.user
&& profile.secret_mode == expected.secret_mode
&& profile.capability == expected.capability
})
.cloned()
}
/// Applies one token-bucket admission decision at a caller-supplied monotonic time.
pub(super) fn allow_rate(
state: &mut RateState,
now: Instant,
per_minute: u32,
burst: u32,
) -> bool {
let burst = f64::from(burst);
if let Some(last) = state.last {
let elapsed = now.saturating_duration_since(last).as_secs_f64();
state.tokens =
(state.tokens + elapsed * f64::from(per_minute) / 60.0).min(burst);
} else {
state.tokens = burst;
}
state.last = Some(now);
if state.tokens < 1.0 {
return false;
}
state.tokens -= 1.0;
true
}
/// Evicts the oldest unused bootstrap while preserving used retry state.
pub(super) fn evict_oldest_unused_bootstrap(state: &mut ManagerState) -> bool {
let Some(hash) = state
.bootstraps
.iter()
.filter(|(_, bootstrap)| !bootstrap.used)
.min_by_key(|(_, bootstrap)| bootstrap.issued_at)
.map(|(hash, _)| *hash)
else {
return false;
};
remove_bootstrap_locked(state, hash);
true
}
/// Removes expired bootstrap and closed-token entries while the manager lock is held.
pub(super) fn remove_expired_locked(state: &mut ManagerState, now: Instant) {
let expired = state
.bootstraps
.iter()
.filter_map(|(hash, bootstrap)| (now > bootstrap.expires_at).then_some(*hash))
.collect::<Vec<_>>();
for hash in expired {
remove_bootstrap_locked(state, hash);
}
state
.closed_tokens
.retain(|_, closed| now <= closed.expires_at);
}
/// Removes one bootstrap and releases its per-address issuance quota when unused.
pub(super) fn remove_bootstrap_locked(state: &mut ManagerState, hash: TokenHash) {
let Some(bootstrap) = state.bootstraps.remove(&hash) else {
return;
};
if !bootstrap.used {
decrement_map(&mut state.bootstraps_per_ip, &bootstrap.issuance_ip);
}
}
/// Decrements one counted owner and removes its map entry at zero.
pub(super) fn decrement_map<K, Q>(values: &mut HashMap<K, usize>, key: &Q)
where
K: std::borrow::Borrow<Q> + std::hash::Hash + Eq,
Q: std::hash::Hash + Eq + ?Sized,
{
let remove = if let Some(value) = values.get_mut(key) {
*value = value.saturating_sub(1);
*value == 0
} else {
false
};
if remove {
values.remove(key);
}
}
/// Computes the process-wide item reserve required for session control progress.
pub(super) fn control_item_reserve(limits: &WebLimitsConfig) -> usize {
limits.max_sessions_global.saturating_mul(
16usize.saturating_add(limits.max_streams_per_session.saturating_mul(3)),
)
}
/// Allocates a non-zero source port unique among live streams for one KDF route.
pub(super) fn allocate_stream_port(
state: &mut ManagerState,
client_ip: IpAddr,
public_addr: SocketAddr,
) -> Option<u16> {
let ports = state
.stream_ports
.entry((client_ip, public_addr))
.or_insert_with(|| StreamPortState {
active: HashSet::new(),
next: 1,
});
for _ in 0..u16::MAX {
let candidate = ports.next;
ports.next = ports.next.checked_add(1).unwrap_or(1);
if ports.active.insert(candidate) {
return Some(candidate);
}
}
None
}
/// Releases one source port and reclaims empty per-route allocator state.
pub(super) fn release_stream_port(
state: &mut ManagerState,
client_ip: IpAddr,
public_addr: SocketAddr,
peer_port: u16,
) -> bool {
let key = (client_ip, public_addr);
let Some(ports) = state.stream_ports.get_mut(&key) else {
return false;
};
let removed = ports.active.remove(&peer_port);
if ports.active.is_empty() {
state.stream_ports.remove(&key);
}
removed
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn synthetic_ports_are_unique_per_live_route_and_state_is_reclaimed() {
let mut state = ManagerState::default();
let client_ip = "192.0.2.10".parse().unwrap();
let public_addr = "203.0.113.10:443".parse().unwrap();
let first = allocate_stream_port(&mut state, client_ip, public_addr).unwrap();
let second = allocate_stream_port(&mut state, client_ip, public_addr).unwrap();
assert_ne!(first, second);
assert!(release_stream_port(
&mut state,
client_ip,
public_addr,
first,
));
assert!(release_stream_port(
&mut state,
client_ip,
public_addr,
second,
));
assert!(state.stream_ports.is_empty());
}
}
+14
View File
@@ -0,0 +1,14 @@
//! Bounded WEB carrier ingress behind a trusted external TLS terminator.
/// Browser bridge generation for the serialized HTTPS carrier.
pub(crate) mod bridge;
/// Shared binary frame codec and protocol constants.
pub(crate) mod frame;
/// Plain HTTP ingress and decoy routing behind external TLS termination.
pub(crate) mod http;
/// Process-wide credentials, quotas, memory budgets, and shutdown ownership.
pub(crate) mod manager;
/// Resumable carrier sessions and logical-stream state machines.
pub(crate) mod session;
/// AsyncRead and AsyncWrite adapter for one logical MTProxy stream.
pub(crate) mod stream;
+355
View File
@@ -0,0 +1,355 @@
use std::collections::{HashMap, HashSet, VecDeque};
use std::io;
use std::net::IpAddr;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::task::{Context, Poll, Waker};
use std::time::{Duration, Instant};
use bytes::{Bytes, BytesMut};
use parking_lot::Mutex;
use tokio::io::ReadBuf;
use tokio::sync::Notify;
use tokio_util::sync::CancellationToken;
use crate::config::{WebLimitsConfig, WebRuntimeProfile, WebTimeoutsConfig};
use crate::web::frame::{self, FrameType};
use crate::web::manager::{ProfileKey, TokenHash, WebProcessRuntime};
// Backend tasks own generation admission and authenticated MTProxy relay lifetimes.
mod backend;
// Downlink queues own cursor replay, flow control, and memory reservations.
mod downlink;
// Uplink batches own exactly-once sequencing and client-frame validation.
mod uplink;
/// Conservative allocator and container overhead charged to every queued item.
pub(crate) const QUEUE_ITEM_COST: usize = 256;
#[derive(Clone, Copy, PartialEq, Eq)]
enum PendingClass {
Uplink,
Downlink,
Control,
}
struct InboundChunk {
bytes: Bytes,
offset: usize,
}
struct StreamState {
inbound: VecDeque<InboundChunk>,
receive_window: u32,
send_credit: u64,
read_waker: Option<Waker>,
write_waker: Option<Waker>,
}
struct QueuedFrame {
encoded: BytesMut,
frame_type: FrameType,
stream_id: u32,
control: bool,
cost: usize,
}
struct DownBatch {
body: Bytes,
base_cursor: u64,
next_cursor: u64,
data_bytes: usize,
data_items: usize,
control_bytes: usize,
control_items: usize,
}
struct SessionState {
streams: HashMap<u32, StreamState>,
active_peer_ports: HashSet<u16>,
closed_streams: HashSet<u32>,
closed_order: VecDeque<u32>,
pending_frames: VecDeque<QueuedFrame>,
pending_windows: HashMap<u32, usize>,
unacked: Option<DownBatch>,
down_cursor: u64,
down_epoch: u64,
last_up_sequence: u64,
last_up_digest: TokenHash,
pending_bytes: usize,
pending_items: usize,
pending_control_bytes: usize,
pending_control_items: usize,
last_activity: Instant,
closed: bool,
}
/// One bounded WEB carrier session containing logical MTProxy streams.
pub(crate) struct WebSession {
manager: std::sync::Weak<WebProcessRuntime>,
token_hash: TokenHash,
client_ip: IpAddr,
profile: Arc<WebRuntimeProfile>,
profile_key: ProfileKey,
limits: WebLimitsConfig,
timeouts: WebTimeoutsConfig,
state: Mutex<SessionState>,
down_notify: Arc<Notify>,
cancel: CancellationToken,
tasks_live: AtomicUsize,
tasks_done: Arc<Notify>,
finished: AtomicBool,
up_active: AtomicBool,
}
/// One successful downlink poll result.
pub(crate) struct PollResult {
/// Encoded downlink frame batch, or an empty long-poll result.
pub(crate) body: Bytes,
/// Cursor the client must present on its next downlink request.
pub(crate) next_cursor: u64,
}
impl WebSession {
#[allow(clippy::too_many_arguments)]
/// Creates one carrier session with immutable ownership and allocation policy.
pub(crate) fn new(
manager: std::sync::Weak<WebProcessRuntime>,
token_hash: TokenHash,
client_ip: IpAddr,
profile: Arc<WebRuntimeProfile>,
profile_key: ProfileKey,
limits: WebLimitsConfig,
timeouts: WebTimeoutsConfig,
) -> Arc<Self> {
Arc::new(Self {
manager,
token_hash,
client_ip,
profile,
profile_key,
limits,
timeouts,
state: Mutex::new(SessionState {
streams: HashMap::new(),
active_peer_ports: HashSet::new(),
closed_streams: HashSet::new(),
closed_order: VecDeque::new(),
pending_frames: VecDeque::new(),
pending_windows: HashMap::new(),
unacked: None,
down_cursor: 0,
down_epoch: 0,
last_up_sequence: 0,
last_up_digest: [0; 32],
pending_bytes: 0,
pending_items: 0,
pending_control_bytes: 0,
pending_control_items: 0,
last_activity: Instant::now(),
closed: false,
}),
down_notify: Arc::new(Notify::new()),
cancel: CancellationToken::new(),
tasks_live: AtomicUsize::new(0),
tasks_done: Arc::new(Notify::new()),
finished: AtomicBool::new(false),
up_active: AtomicBool::new(false),
})
}
/// Returns the stable hashed token identity without exposing the credential.
pub(crate) fn token_hash(&self) -> TokenHash {
self.token_hash
}
/// Checks the canonical virtual host that owns this bearer session.
pub(crate) fn matches_host(&self, host: &str) -> bool {
self.profile.host == host
}
/// Closes carrier state while relay tasks retain their admission until exit.
pub(crate) fn close(&self) {
let (data_bytes, data_items, control_bytes, control_items) = {
let mut state = self.state.lock();
if state.closed {
return;
}
state.closed = true;
for stream in state.streams.values_mut() {
if let Some(waker) = stream.read_waker.take() {
waker.wake();
}
if let Some(waker) = stream.write_waker.take() {
waker.wake();
}
}
state.streams.clear();
state.pending_frames.clear();
state.pending_windows.clear();
state.unacked = None;
let control_bytes = state.pending_control_bytes;
let control_items = state.pending_control_items;
let data_bytes = state.pending_bytes.saturating_sub(control_bytes);
let data_items = state.pending_items.saturating_sub(control_items);
state.pending_bytes = 0;
state.pending_items = 0;
state.pending_control_bytes = 0;
state.pending_control_items = 0;
(data_bytes, data_items, control_bytes, control_items)
};
self.cancel.cancel();
self.down_notify.notify_waiters();
if let Some(manager) = self.manager.upgrade() {
manager.release_pending(data_bytes, data_items, false);
manager.release_pending(control_bytes, control_items, true);
if !self.finished.swap(true, Ordering::AcqRel) {
manager.session_finished(
self.token_hash,
self.client_ip,
self.profile_key,
&self.profile.host,
);
}
}
}
/// Waits for all logical-stream tasks after admission has closed.
pub(crate) async fn wait(&self) {
loop {
let notified = self.tasks_done.notified();
if self.tasks_live.load(Ordering::Acquire) == 0 {
return;
}
notified.await;
}
}
/// Returns whether reconnect grace elapsed without activity.
pub(crate) fn is_idle(&self, now: Instant) -> bool {
let state = self.state.lock();
!state.closed
&& now.saturating_duration_since(state.last_activity)
>= Duration::from_secs(self.timeouts.reconnect_grace_secs)
}
/// Polls client-to-server bytes and returns consumed flow-control credit.
pub(super) fn poll_read(
&self,
stream_id: u32,
cx: &mut Context<'_>,
output: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
let mut state = self.state.lock();
let (count, finished) = {
let Some(stream) = state.streams.get_mut(&stream_id) else {
return Poll::Ready(Ok(()));
};
let Some(chunk) = stream.inbound.front_mut() else {
stream.read_waker = Some(cx.waker().clone());
return Poll::Pending;
};
let available = &chunk.bytes[chunk.offset..];
let count = available.len().min(output.remaining());
output.put_slice(&available[..count]);
chunk.offset += count;
let finished = chunk.offset == chunk.bytes.len();
if finished {
stream.inbound.pop_front();
}
stream.receive_window = stream.receive_window.saturating_add(count as u32);
(count, finished)
};
let overhead = if finished { QUEUE_ITEM_COST } else { 0 };
self.release_locked(&mut state, count + overhead, usize::from(finished), false);
if !self.queue_window_locked(&mut state, stream_id, count as u32) {
drop(state);
self.close();
return Poll::Ready(Err(io::Error::other("WEB session control budget exhausted")));
}
Poll::Ready(Ok(()))
}
/// Polls server-to-client writes against stream credit and bounded queues.
pub(super) fn poll_write(
&self,
stream_id: u32,
cx: &mut Context<'_>,
input: &[u8],
) -> Poll<io::Result<usize>> {
if input.is_empty() {
return Poll::Ready(Ok(0));
}
let mut state = self.state.lock();
let Some(stream) = state.streams.get_mut(&stream_id) else {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"WEB logical stream is closed",
)));
};
let count = input
.len()
.min(frame::DATA_CHUNK_BYTES)
.min(self.limits.max_frame_payload_bytes)
.min(stream.send_credit as usize);
if count == 0 {
stream.write_waker = Some(cx.waker().clone());
return Poll::Pending;
}
if !self.queue_data_locked(&mut state, stream_id, &input[..count]) {
if let Some(stream) = state.streams.get_mut(&stream_id) {
stream.write_waker = Some(cx.waker().clone());
}
return Poll::Pending;
}
let Some(stream) = state.streams.get_mut(&stream_id) else {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"WEB logical stream is closed",
)));
};
stream.send_credit -= count as u64;
state.last_activity = Instant::now();
drop(state);
self.down_notify.notify_waiters();
Poll::Ready(Ok(count))
}
/// Returns the process queue-capacity notification source while the manager lives.
pub(super) fn budget_notify(&self) -> Option<Arc<Notify>> {
self.manager.upgrade().map(|manager| manager.budget_notify())
}
fn release_stream_reservation(&self, peer_port: u16) {
let removed = self.state.lock().active_peer_ports.remove(&peer_port);
if removed
&& let Some(manager) = self.manager.upgrade()
{
manager.release_stream(
self.profile_key,
self.client_ip,
self.profile.public_addr,
peer_port,
);
}
}
}
fn inbound_queue_cost(queue: &VecDeque<InboundChunk>) -> (usize, usize) {
let bytes = queue.iter().fold(0usize, |total, chunk| {
total.saturating_add(chunk.bytes.len().saturating_sub(chunk.offset) + QUEUE_ITEM_COST)
});
(bytes, queue.len())
}
fn remember_closed(state: &mut SessionState, stream_id: u32, limit: usize) {
if !state.closed_streams.insert(stream_id) {
return;
}
state.closed_order.push_back(stream_id);
while state.closed_order.len() > limit {
if let Some(oldest) = state.closed_order.pop_front() {
state.closed_streams.remove(&oldest);
}
}
}
+168
View File
@@ -0,0 +1,168 @@
use std::io;
use std::sync::Arc;
use std::sync::atomic::Ordering;
use std::time::Duration;
use crate::web::frame::FrameType;
use crate::web::stream::WebLogicalStream;
use crate::proxy::shared_state::ConntrackClosePolicy;
use super::{WebSession, inbound_queue_cost, remember_closed};
impl WebSession {
/// Starts one owned inner handshake and relay task for an admitted stream.
pub(super) fn spawn_stream(self: &Arc<Self>, stream_id: u32, peer_port: u16) {
let Some(manager) = self.manager.upgrade() else {
self.stream_finished(stream_id, peer_port);
return;
};
let generation = manager.active_generation();
let Ok(connection_permit) = generation.max_connections.clone().try_acquire_owned() else {
manager.record_stream_rejected();
self.stream_finished(stream_id, peer_port);
return;
};
let Some(handshake_permit) = manager.try_stream_handshake() else {
self.stream_finished(stream_id, peer_port);
return;
};
let deps = generation.client_runtime_deps();
let replay_checker = Arc::clone(&generation.replay_checker);
let session = Arc::clone(self);
let cancel = self.cancel.clone();
self.tasks_live.fetch_add(1, Ordering::AcqRel);
let spawned = generation.spawn_session(async move {
let _connection_permit = connection_permit;
let _completion = StreamCompletion {
session: Arc::clone(&session),
stream_id,
peer_port,
};
let stream = WebLogicalStream::new(Arc::clone(&session), stream_id);
tokio::select! {
_ = cancel.cancelled() => {}
_ = run_stream(
Arc::clone(&session),
stream,
deps,
replay_checker,
handshake_permit,
peer_port,
) => {}
}
});
if !spawned {
self.tasks_live.fetch_sub(1, Ordering::AcqRel);
self.stream_finished(stream_id, peer_port);
self.tasks_done.notify_waiters();
}
}
fn stream_finished(&self, stream_id: u32, peer_port: u16) {
let (queued, reserved) = {
let mut state = self.state.lock();
let reserved = state.active_peer_ports.remove(&peer_port);
let queued = state.streams.remove(&stream_id).map(|stream| {
let (bytes, items) = inbound_queue_cost(&stream.inbound);
self.release_locked(&mut state, bytes, items, false);
remember_closed(
&mut state,
stream_id,
self.limits.max_tombstones_per_session,
);
self.queue_control_locked(&mut state, FrameType::Close, stream_id, &[])
});
(queued, reserved)
};
if reserved
&& let Some(manager) = self.manager.upgrade()
{
manager.release_stream(
self.profile_key,
self.client_ip,
self.profile.public_addr,
peer_port,
);
}
if let Some(queued) = queued {
if !queued {
self.close();
}
self.down_notify.notify_waiters();
}
}
}
struct StreamCompletion {
session: Arc<WebSession>,
stream_id: u32,
peer_port: u16,
}
impl Drop for StreamCompletion {
fn drop(&mut self) {
self.session.stream_finished(self.stream_id, self.peer_port);
if self.session.tasks_live.fetch_sub(1, Ordering::AcqRel) == 1 {
self.session.tasks_done.notify_waiters();
}
}
}
async fn run_stream(
session: Arc<WebSession>,
stream: WebLogicalStream,
deps: crate::proxy::authenticated::ClientRuntimeDeps,
replay_checker: Arc<crate::stats::ReplayChecker>,
handshake_permit: tokio::sync::OwnedSemaphorePermit,
peer_port: u16,
) {
use tokio::io::AsyncReadExt;
use crate::protocol::constants::HANDSHAKE_LEN;
use crate::proxy::authenticated::run_authenticated;
use crate::proxy::handshake::handle_mtproto_handshake_for_web_user;
let (mut reader, writer) = tokio::io::split(stream);
let mut handshake = [0u8; HANDSHAKE_LEN];
let peer = std::net::SocketAddr::new(session.client_ip, peer_port);
deps.stats.increment_connects_all();
let handshake_result = tokio::time::timeout(
Duration::from_secs(session.timeouts.stream_handshake_secs),
async {
reader.read_exact(&mut handshake).await?;
Ok::<_, io::Error>(
handle_mtproto_handshake_for_web_user(
&handshake,
reader,
writer,
peer,
&deps.config,
&replay_checker,
&session.profile.user,
session.profile.secret_mode,
&deps.shared,
)
.await,
)
},
)
.await;
drop(handshake_permit);
let Ok(Ok(crate::error::HandshakeResult::Success((reader, writer, success)))) =
handshake_result
else {
deps.stats
.increment_connects_bad_with_class("web_mtproto_bad_client");
return;
};
let _ = run_authenticated(
reader,
writer,
success,
deps,
session.profile.public_addr,
peer,
ConntrackClosePolicy::Suppress,
)
.await;
}
+494
View File
@@ -0,0 +1,494 @@
use std::time::{Duration, Instant};
use bytes::{BufMut, Bytes, BytesMut};
use super::{
DownBatch, PendingClass, PollResult, QUEUE_ITEM_COST, QueuedFrame, SessionState,
WebSession,
};
use crate::web::frame::{self, FrameType};
use crate::web::manager::ManagerError;
impl WebSession {
/// Polls pending downlink frames with cursor replay and newest-poll-wins semantics.
pub(crate) async fn poll_down(&self, cursor: u64) -> Result<PollResult, ManagerError> {
let epoch = {
let mut state = self.state.lock();
if state.closed {
return Err(ManagerError::Closed);
}
state.last_activity = Instant::now();
if let Some(unacked) = &state.unacked {
if cursor == unacked.base_cursor {
return Ok(PollResult {
body: unacked.body.clone(),
next_cursor: unacked.next_cursor,
});
}
if cursor != unacked.next_cursor {
drop(state);
self.close();
return Err(ManagerError::Protocol);
}
self.release_unacked_locked(&mut state);
} else if cursor != state.down_cursor {
drop(state);
self.close();
return Err(ManagerError::Protocol);
}
state.down_epoch = state.down_epoch.wrapping_add(1).max(1);
state.down_epoch
};
self.down_notify.notify_waiters();
let deadline = Duration::from_secs(self.timeouts.long_poll_secs);
let poll = async {
loop {
let notified = self.down_notify.notified();
{
let mut state = self.state.lock();
if state.down_epoch != epoch {
return Ok(PollResult {
body: Bytes::new(),
next_cursor: cursor,
});
}
if !state.pending_frames.is_empty() {
let batch = match self.take_down_batch_locked(&mut state, cursor) {
Ok(batch) => batch,
Err(error) => {
drop(state);
self.close();
return Err(error);
}
};
let result = PollResult {
body: batch.body.clone(),
next_cursor: batch.next_cursor,
};
if let Some(manager) = self.manager.upgrade() {
manager.record_down(result.body.len());
}
state.unacked = Some(batch);
return Ok(result);
}
if state.closed {
return Err(ManagerError::Closed);
}
}
notified.await;
}
};
match tokio::time::timeout(deadline, poll).await {
Ok(result) => result,
Err(_) => {
let mut state = self.state.lock();
if state.down_epoch == epoch {
state.last_activity = Instant::now();
}
Ok(PollResult {
body: Bytes::new(),
next_cursor: cursor,
})
}
}
}
/// Reserves session and process queue capacity while the session lock is held.
pub(super) fn reserve_locked(
&self,
state: &mut SessionState,
bytes: usize,
items: usize,
class: PendingClass,
) -> bool {
if bytes == 0 && items == 0 {
return true;
}
let data_byte_limit = self
.limits
.pending_bytes_per_session
.saturating_sub(self.limits.control_bytes_per_session);
let item_reserve = 16usize.saturating_add(
self.limits.max_streams_per_session.saturating_mul(3),
);
let data_item_limit = self
.limits
.pending_items_per_session
.saturating_sub(item_reserve);
if state.closed {
return false;
}
let control = class == PendingClass::Control;
let fits = if control {
bytes <= self.limits.control_bytes_per_session
&& items <= item_reserve
&& state.pending_bytes
<= self.limits.pending_bytes_per_session.saturating_sub(bytes)
&& state.pending_items
<= self.limits.pending_items_per_session.saturating_sub(items)
&& state.pending_control_bytes
<= self.limits.control_bytes_per_session.saturating_sub(bytes)
&& state.pending_control_items <= item_reserve.saturating_sub(items)
} else {
let data_bytes = state
.pending_bytes
.saturating_sub(state.pending_control_bytes);
let data_items = state
.pending_items
.saturating_sub(state.pending_control_items);
let (byte_limit, item_limit) = if class == PendingClass::Downlink {
let uplink_bytes = self
.limits
.max_body_bytes
.saturating_add(
self.limits
.max_frames_per_body
.saturating_mul(QUEUE_ITEM_COST),
);
(
data_byte_limit.saturating_sub(uplink_bytes),
data_item_limit.saturating_sub(self.limits.max_frames_per_body),
)
} else {
(data_byte_limit, data_item_limit)
};
bytes <= byte_limit
&& items <= item_limit
&& data_bytes <= byte_limit - bytes
&& data_items <= item_limit - items
};
if !fits {
return false;
}
let Some(manager) = self.manager.upgrade() else {
return false;
};
if !manager.try_reserve_pending(
bytes,
items,
control,
class == PendingClass::Downlink,
) {
return false;
}
state.pending_bytes += bytes;
state.pending_items += items;
if control {
state.pending_control_bytes += bytes;
state.pending_control_items += items;
}
true
}
/// Releases session and process queue capacity while the session lock is held.
pub(super) fn release_locked(
&self,
state: &mut SessionState,
bytes: usize,
items: usize,
control: bool,
) {
state.pending_bytes = state.pending_bytes.saturating_sub(bytes);
state.pending_items = state.pending_items.saturating_sub(items);
if control {
state.pending_control_bytes = state.pending_control_bytes.saturating_sub(bytes);
state.pending_control_items = state.pending_control_items.saturating_sub(items);
}
if let Some(manager) = self.manager.upgrade() {
manager.release_pending(bytes, items, control);
}
}
/// Coalesces one flow-control update into the bounded control queue.
pub(super) fn queue_window_locked(
&self,
state: &mut SessionState,
stream_id: u32,
amount: u32,
) -> bool {
if amount == 0 {
return true;
}
if let Some(index) = state.pending_windows.get(&stream_id).copied()
&& let Some(queued) = state.pending_frames.get_mut(index)
{
let previous = u32::from_be_bytes(
queued.encoded[frame::HEADER_BYTES..frame::HEADER_BYTES + 4]
.try_into()
.unwrap_or([0; 4]),
);
if let Some(total) = previous.checked_add(amount) {
queued.encoded[frame::HEADER_BYTES..frame::HEADER_BYTES + 4]
.copy_from_slice(&total.to_be_bytes());
self.down_notify.notify_waiters();
return true;
}
}
self.queue_control_locked(
state,
FrameType::Window,
stream_id,
&frame::window_payload(amount),
)
}
/// Appends one control frame under both reserved queue budgets.
pub(super) fn queue_control_locked(
&self,
state: &mut SessionState,
frame_type: FrameType,
stream_id: u32,
payload: &[u8],
) -> bool {
self.queue_frame_locked(state, frame_type, stream_id, payload, true)
}
/// Appends one server-to-client DATA frame under downlink data budgets.
pub(super) fn queue_data_locked(
&self,
state: &mut SessionState,
stream_id: u32,
payload: &[u8],
) -> bool {
let can_coalesce = state.pending_frames.back().is_some_and(|last| {
last.frame_type == FrameType::Data
&& last.stream_id == stream_id
&& last.encoded.len() - frame::HEADER_BYTES + payload.len()
<= self.limits.max_frame_payload_bytes
});
if can_coalesce {
if !self.reserve_locked(state, payload.len(), 0, PendingClass::Downlink) {
return false;
}
let Some(last) = state.pending_frames.back_mut() else {
return false;
};
last.encoded.extend_from_slice(payload);
last.cost += payload.len();
let payload_len = (last.encoded.len() - frame::HEADER_BYTES) as u32;
last.encoded[4..8].copy_from_slice(&payload_len.to_be_bytes());
return true;
}
self.queue_frame_locked(state, FrameType::Data, stream_id, payload, false)
}
fn queue_frame_locked(
&self,
state: &mut SessionState,
frame_type: FrameType,
stream_id: u32,
payload: &[u8],
control: bool,
) -> bool {
let cost = frame::HEADER_BYTES + payload.len() + QUEUE_ITEM_COST;
let class = if control {
PendingClass::Control
} else {
PendingClass::Downlink
};
if !self.reserve_locked(state, cost, 1, class) {
return false;
}
let mut encoded = BytesMut::with_capacity(frame::HEADER_BYTES + payload.len());
encoded.put_u8(frame_type as u8);
encoded.put_u8((stream_id >> 16) as u8);
encoded.put_u8((stream_id >> 8) as u8);
encoded.put_u8(stream_id as u8);
encoded.put_u32(payload.len() as u32);
encoded.extend_from_slice(payload);
let index = state.pending_frames.len();
state.pending_frames.push_back(QueuedFrame {
encoded,
frame_type,
stream_id,
control,
cost,
});
if frame_type == FrameType::Window {
state.pending_windows.insert(stream_id, index);
}
self.down_notify.notify_waiters();
true
}
fn take_down_batch_locked(
&self,
state: &mut SessionState,
cursor: u64,
) -> Result<DownBatch, ManagerError> {
let next_cursor = state
.down_cursor
.checked_add(1)
.ok_or(ManagerError::Protocol)?;
let mut count = 0usize;
let mut body_len = 0usize;
for queued in &state.pending_frames {
if count >= self.limits.max_frames_per_body
|| (count != 0
&& body_len.saturating_add(queued.encoded.len())
> self.limits.carrier_batch_bytes)
{
break;
}
body_len += queued.encoded.len();
count += 1;
}
let mut body = BytesMut::with_capacity(body_len);
let mut data_bytes = 0usize;
let mut data_items = 0usize;
let mut control_bytes = 0usize;
let mut control_items = 0usize;
for index in 0..count {
let Some(queued) = state.pending_frames.get(index) else {
break;
};
if queued.frame_type == FrameType::Window
&& state.pending_windows.get(&queued.stream_id) == Some(&index)
{
state.pending_windows.remove(&queued.stream_id);
}
}
for _ in 0..count {
let Some(queued) = state.pending_frames.pop_front() else {
break;
};
body.extend_from_slice(&queued.encoded);
if queued.control {
control_bytes += queued.cost;
control_items += 1;
} else {
data_bytes += queued.cost;
data_items += 1;
}
}
for index in state.pending_windows.values_mut() {
*index = index.saturating_sub(count);
}
state.down_cursor = next_cursor;
Ok(DownBatch {
body: body.freeze(),
base_cursor: cursor,
next_cursor,
data_bytes,
data_items,
control_bytes,
control_items,
})
}
fn release_unacked_locked(&self, state: &mut SessionState) {
let Some(batch) = state.unacked.take() else {
return;
};
self.release_locked(state, batch.data_bytes, batch.data_items, false);
self.release_locked(state, batch.control_bytes, batch.control_items, true);
for stream in state.streams.values_mut() {
if let Some(waker) = stream.write_waker.take() {
waker.wake();
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::net::SocketAddr;
use std::sync::Arc;
use crate::config::{
WebLimitsConfig, WebRuntimeProfile, WebSecretMode, WebTimeoutsConfig,
};
use crate::web::manager::WebProcessRuntime;
fn session() -> Arc<WebSession> {
let profile = Arc::new(WebRuntimeProfile {
host: "proxy.example.com".to_string(),
public_addr: SocketAddr::from(([203, 0, 113, 10], 443)),
user: "alice".to_string(),
secret_mode: WebSecretMode::Plain,
capability: [0; 32],
max_sessions: 1,
max_streams: 1,
max_streams_per_session: 1,
});
WebSession::new(
std::sync::Weak::<WebProcessRuntime>::new(),
[1; 32],
"192.0.2.10".parse().unwrap(),
profile,
[2; 32],
WebLimitsConfig::default(),
WebTimeoutsConfig::default(),
)
}
fn queue_close(session: &WebSession) {
let encoded = frame::encode(FrameType::Close, 1, &[]);
session.state.lock().pending_frames.push_back(QueuedFrame {
encoded: BytesMut::from(encoded.as_ref()),
frame_type: FrameType::Close,
stream_id: 1,
control: true,
cost: frame::HEADER_BYTES + QUEUE_ITEM_COST,
});
}
#[tokio::test]
async fn downlink_replays_unacknowledged_batch_byte_for_byte() {
let session = session();
queue_close(&session);
let first = session.poll_down(0).await.unwrap();
let replay = session.poll_down(0).await.unwrap();
assert_eq!(first.next_cursor, 1);
assert_eq!(replay.next_cursor, 1);
assert_eq!(first.body, replay.body);
}
#[tokio::test]
async fn invalid_or_overflowing_cursor_closes_session() {
let invalid = session();
assert!(matches!(
invalid.poll_down(1).await,
Err(ManagerError::Protocol)
));
assert!(invalid.state.lock().closed);
let overflow = session();
{
let mut state = overflow.state.lock();
state.down_cursor = u64::MAX;
}
queue_close(&overflow);
assert!(matches!(
overflow.poll_down(u64::MAX).await,
Err(ManagerError::Protocol)
));
assert!(overflow.state.lock().closed);
}
#[tokio::test]
async fn newer_poll_supersedes_older_poll_without_closing_session() {
let session = session();
let first_session = Arc::clone(&session);
let first = tokio::spawn(async move { first_session.poll_down(0).await });
while session.state.lock().down_epoch < 1 {
tokio::task::yield_now().await;
}
let second_session = Arc::clone(&session);
let second = tokio::spawn(async move { second_session.poll_down(0).await });
while session.state.lock().down_epoch < 2 {
tokio::task::yield_now().await;
}
let superseded = tokio::time::timeout(Duration::from_secs(1), first)
.await
.unwrap()
.unwrap()
.unwrap();
assert!(superseded.body.is_empty());
assert_eq!(superseded.next_cursor, 0);
assert!(!session.state.lock().closed);
second.abort();
}
}
+431
View File
@@ -0,0 +1,431 @@
use std::collections::{HashMap, HashSet, VecDeque};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Instant;
use bytes::Bytes;
use sha2::{Digest, Sha256};
use subtle::ConstantTimeEq;
use super::{
InboundChunk, PendingClass, QUEUE_ITEM_COST, SessionState, StreamState, WebSession,
inbound_queue_cost, remember_closed,
};
use crate::web::frame::{self, Frame, FrameType};
use crate::web::manager::{ManagerError, TokenHash};
impl WebSession {
/// Applies one exactly-once uplink batch.
pub(crate) fn process_up(
self: &Arc<Self>,
sequence: u64,
body: &[u8],
) -> Result<u64, ManagerError> {
if self
.up_active
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
.is_err()
{
return Err(ManagerError::Concurrent);
}
let _uplink = UplinkGuard(&self.up_active);
let frames = match frame::parse_all(body, &self.limits) {
Ok(frames) => frames,
Err(_) => {
self.close();
return Err(ManagerError::Protocol);
}
};
if frames
.iter()
.copied()
.any(|value| frame::validate_client_shape(value).is_err())
{
self.close();
return Err(ManagerError::Protocol);
}
let digest: TokenHash = Sha256::digest(body).into();
let mut opened = Vec::new();
let result = {
let mut state = self.state.lock();
if state.closed {
return Err(ManagerError::Closed);
}
state.last_activity = Instant::now();
if sequence == state.last_up_sequence && sequence != 0 {
return if bool::from(state.last_up_digest.ct_eq(&digest)) {
Ok(sequence)
} else {
drop(state);
self.close();
Err(ManagerError::Protocol)
};
}
if sequence == 0 || sequence != state.last_up_sequence.saturating_add(1) {
drop(state);
self.close();
return Err(ManagerError::Protocol);
}
if !validate_batch(&state, &frames) {
drop(state);
self.close();
return Err(ManagerError::Protocol);
}
let (reserve_bytes, reserve_items) = inbound_reservation(&state, &frames);
if !self.reserve_locked(
&mut state,
reserve_bytes,
reserve_items,
PendingClass::Uplink,
) {
return Err(ManagerError::Backpressure);
}
let mut unused_bytes = reserve_bytes;
let mut unused_items = reserve_items;
let applied = self.apply_batch_locked(
&mut state,
&frames,
&mut opened,
&mut unused_bytes,
&mut unused_items,
);
self.release_locked(&mut state, unused_bytes, unused_items, false);
if !applied {
Err(ManagerError::Closed)
} else {
state.last_up_sequence = sequence;
state.last_up_digest = digest;
Ok(sequence)
}
};
if matches!(result, Err(ManagerError::Backpressure)) {
return result;
}
if result.is_err() {
self.close();
for (_, peer_port) in opened {
self.release_stream_reservation(peer_port);
}
return result;
}
for (stream_id, peer_port) in opened {
self.spawn_stream(stream_id, peer_port);
}
if let Some(manager) = self.manager.upgrade() {
manager.record_up(body.len());
}
result
}
fn apply_batch_locked(
&self,
state: &mut SessionState,
frames: &[Frame<'_>],
opened: &mut Vec<(u32, u16)>,
unused_bytes: &mut usize,
unused_items: &mut usize,
) -> bool {
for value in frames {
if value.stream_id == 0 {
continue;
}
let was_closed = state.closed_streams.contains(&value.stream_id);
match value.frame_type {
FrameType::Open => {
let Some(peer_port) = self.reserve_stream_locked(state) else {
remember_closed(
state,
value.stream_id,
self.limits.max_tombstones_per_session,
);
if !self.queue_control_locked(
state,
FrameType::Close,
value.stream_id,
&[],
) {
return false;
}
continue;
};
state.streams.insert(
value.stream_id,
StreamState {
inbound: VecDeque::new(),
receive_window: frame::INITIAL_STREAM_WINDOW,
send_credit: u64::from(frame::INITIAL_STREAM_WINDOW),
read_waker: None,
write_waker: None,
},
);
opened.push((value.stream_id, peer_port));
}
FrameType::Data if !was_closed => {
let Some(stream) = state.streams.get_mut(&value.stream_id) else {
return false;
};
stream.receive_window -= value.payload.len() as u32;
stream.inbound.push_back(InboundChunk {
bytes: Bytes::copy_from_slice(value.payload),
offset: 0,
});
*unused_bytes = unused_bytes
.saturating_sub(value.payload.len() + QUEUE_ITEM_COST);
*unused_items = unused_items.saturating_sub(1);
if let Some(waker) = stream.read_waker.take() {
waker.wake();
}
}
FrameType::Window if !was_closed => {
let Some(stream) = state.streams.get_mut(&value.stream_id) else {
return false;
};
let amount = frame::window_amount(value.payload).unwrap_or(0);
stream.send_credit = stream
.send_credit
.saturating_add(u64::from(amount))
.min(u64::from(u32::MAX));
if let Some(waker) = stream.write_waker.take() {
waker.wake();
}
}
FrameType::Close if !was_closed => {
let Some(stream) = state.streams.remove(&value.stream_id) else {
return false;
};
let (bytes, items) = inbound_queue_cost(&stream.inbound);
self.release_locked(state, bytes, items, false);
remember_closed(
state,
value.stream_id,
self.limits.max_tombstones_per_session,
);
if let Some(waker) = stream.read_waker {
waker.wake();
}
if let Some(waker) = stream.write_waker {
waker.wake();
}
}
FrameType::Data | FrameType::Window | FrameType::Close => {}
_ => return false,
}
}
true
}
fn reserve_stream_locked(&self, state: &mut SessionState) -> Option<u16> {
if state.active_peer_ports.len() >= self.profile.max_streams_per_session {
return None;
}
let manager = self.manager.upgrade()?;
let peer_port = manager.try_acquire_stream(
self.profile_key,
self.profile.max_streams,
self.client_ip,
self.profile.public_addr,
)?;
if state.active_peer_ports.insert(peer_port) {
return Some(peer_port);
}
manager.release_stream(
self.profile_key,
self.client_ip,
self.profile.public_addr,
peer_port,
);
None
}
}
struct UplinkGuard<'a>(&'a AtomicBool);
impl Drop for UplinkGuard<'_> {
fn drop(&mut self) {
self.0.store(false, Ordering::Release);
}
}
fn validate_batch(state: &SessionState, frames: &[Frame<'_>]) -> bool {
let mut live = state
.streams
.iter()
.map(|(id, stream)| (*id, (stream.receive_window, stream.send_credit)))
.collect::<HashMap<_, _>>();
let mut closed = HashSet::new();
for value in frames {
if value.stream_id == 0 {
if value.frame_type != FrameType::Pong {
return false;
}
continue;
}
let was_closed = state.closed_streams.contains(&value.stream_id)
|| closed.contains(&value.stream_id);
match value.frame_type {
FrameType::Open => {
if live.contains_key(&value.stream_id) || was_closed {
return false;
}
live.insert(
value.stream_id,
(
frame::INITIAL_STREAM_WINDOW,
u64::from(frame::INITIAL_STREAM_WINDOW),
),
);
}
FrameType::Data if !was_closed => {
let Some((receive_window, send_credit)) = live.get_mut(&value.stream_id) else {
return false;
};
let Ok(payload_len) = u32::try_from(value.payload.len()) else {
return false;
};
if payload_len > *receive_window {
return false;
}
*receive_window -= payload_len;
let _ = send_credit;
}
FrameType::Window if !was_closed => {
let Some((_, send_credit)) = live.get_mut(&value.stream_id) else {
return false;
};
let Ok(amount) = frame::window_amount(value.payload) else {
return false;
};
*send_credit = send_credit
.saturating_add(u64::from(amount))
.min(u64::from(u32::MAX));
}
FrameType::Close if !was_closed => {
if live.remove(&value.stream_id).is_none() {
return false;
}
closed.insert(value.stream_id);
}
FrameType::Data | FrameType::Window | FrameType::Close => {}
_ => return false,
}
}
true
}
fn inbound_reservation(state: &SessionState, frames: &[Frame<'_>]) -> (usize, usize) {
let mut live = state.streams.keys().copied().collect::<HashSet<_>>();
let mut bytes = 0usize;
let mut items = 0usize;
for value in frames {
match value.frame_type {
FrameType::Open => {
live.insert(value.stream_id);
}
FrameType::Data if live.contains(&value.stream_id) => {
bytes = bytes.saturating_add(value.payload.len() + QUEUE_ITEM_COST);
items = items.saturating_add(1);
}
FrameType::Close => {
live.remove(&value.stream_id);
}
_ => {}
}
}
(bytes, items)
}
#[cfg(test)]
mod tests {
use super::*;
use std::net::SocketAddr;
use crate::config::{
WebLimitsConfig, WebRuntimeProfile, WebSecretMode, WebTimeoutsConfig,
};
use crate::web::manager::WebProcessRuntime;
fn session() -> Arc<WebSession> {
let profile = Arc::new(WebRuntimeProfile {
host: "proxy.example.com".to_string(),
public_addr: SocketAddr::from(([203, 0, 113, 10], 443)),
user: "alice".to_string(),
secret_mode: WebSecretMode::Plain,
capability: [0; 32],
max_sessions: 1,
max_streams: 1,
max_streams_per_session: 1,
});
WebSession::new(
std::sync::Weak::<WebProcessRuntime>::new(),
[1; 32],
"192.0.2.10".parse().unwrap(),
profile,
[2; 32],
WebLimitsConfig::default(),
WebTimeoutsConfig::default(),
)
}
#[test]
fn uplink_retry_commits_only_one_exact_body() {
let session = session();
let first = frame::encode(FrameType::Pong, 0, &[1, 2, 3]);
assert_eq!(session.process_up(1, &first), Ok(1));
assert_eq!(session.process_up(1, &first), Ok(1));
let changed = frame::encode(FrameType::Pong, 0, &[1, 2, 4]);
assert_eq!(session.process_up(1, &changed), Err(ManagerError::Protocol));
assert!(session.state.lock().closed);
}
#[test]
fn concurrent_uplink_does_not_commit_sequence() {
let session = session();
let body = frame::encode(FrameType::Pong, 0, &[]);
session.up_active.store(true, Ordering::Release);
assert_eq!(
session.process_up(1, &body),
Err(ManagerError::Concurrent)
);
assert_eq!(session.state.lock().last_up_sequence, 0);
session.up_active.store(false, Ordering::Release);
assert_eq!(session.process_up(1, &body), Ok(1));
}
#[test]
fn backpressured_uplink_does_not_commit_or_close() {
let session = session();
{
let mut state = session.state.lock();
state.streams.insert(
1,
StreamState {
inbound: VecDeque::new(),
receive_window: frame::INITIAL_STREAM_WINDOW,
send_credit: u64::from(frame::INITIAL_STREAM_WINDOW),
read_waker: None,
write_waker: None,
},
);
state.pending_bytes = session.limits.pending_bytes_per_session;
}
let body = frame::encode(FrameType::Data, 1, &[1]);
assert_eq!(
session.process_up(1, &body),
Err(ManagerError::Backpressure)
);
let state = session.state.lock();
assert!(!state.closed);
assert_eq!(state.last_up_sequence, 0);
assert!(state.streams.get(&1).unwrap().inbound.is_empty());
}
#[test]
fn uplink_gap_is_fatal() {
let session = session();
let body = frame::encode(FrameType::Pong, 0, &[]);
assert_eq!(session.process_up(2, &body), Err(ManagerError::Protocol));
assert!(session.state.lock().closed);
}
}
+83
View File
@@ -0,0 +1,83 @@
use std::future::Future;
use std::io;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use tokio::sync::futures::OwnedNotified;
use crate::web::session::WebSession;
/// Async byte stream that maps one WEB stream identifier onto carrier frames.
pub(crate) struct WebLogicalStream {
session: Arc<WebSession>,
stream_id: u32,
budget_wait: Option<Pin<Box<OwnedNotified>>>,
}
impl WebLogicalStream {
/// Binds a virtual byte stream to one live carrier stream identifier.
pub(crate) fn new(session: Arc<WebSession>, stream_id: u32) -> Self {
Self {
session,
stream_id,
budget_wait: None,
}
}
}
impl AsyncRead for WebLogicalStream {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
output: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
self.session.poll_read(self.stream_id, cx, output)
}
}
impl AsyncWrite for WebLogicalStream {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
input: &[u8],
) -> Poll<io::Result<usize>> {
let result = self.session.poll_write(self.stream_id, cx, input);
if !result.is_pending() {
self.budget_wait = None;
return result;
}
// Register before retrying so a concurrent global-capacity release cannot be lost.
loop {
if self.budget_wait.is_none()
&& let Some(notify) = self.session.budget_notify()
{
self.budget_wait = Some(Box::pin(notify.notified_owned()));
}
let Some(wait) = self.budget_wait.as_mut() else {
break;
};
if wait.as_mut().poll(cx).is_pending() {
break;
}
self.budget_wait = None;
}
match self.session.poll_write(self.stream_id, cx, input) {
Poll::Ready(result) => {
self.budget_wait = None;
Poll::Ready(result)
}
Poll::Pending => Poll::Pending,
}
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
}