diff --git a/src/api/handler.rs b/src/api/handler.rs new file mode 100644 index 0000000..ed9180a --- /dev/null +++ b/src/api/handler.rs @@ -0,0 +1,187 @@ +use super::*; + +// Read-only fixed API endpoints. +mod read_routes; +// Fixed configuration and lifecycle mutations. +mod fixed_routes; +// Dynamic reload and user-resource routes. +mod user_routes; + +pub(super) async fn handle( + req: Request, + peer: SocketAddr, + shared: Arc, +) -> Result>, IoError> { + let runtime = shared.active_runtime.load_full(); + let previous_cache_generation = shared.cache_generation.swap(runtime.id, Ordering::AcqRel); + if previous_cache_generation != runtime.id { + *shared.minimal_cache.lock().await = None; + *shared.runtime_edge_connections_cache.lock().await = None; + } + let shared = Arc::new(shared.for_runtime(runtime.as_ref())); + let config_rx = runtime.config_rx.clone(); + shared + .runtime_state + .admission_open + .store(*runtime.admission_rx.borrow(), Ordering::Relaxed); + let request_id = shared.next_request_id(); + let cfg = config_rx.borrow().clone(); + let api_cfg = &cfg.server.api; + + if !api_cfg.enabled { + return Ok(error_response( + request_id, + ApiFailure::new( + StatusCode::SERVICE_UNAVAILABLE, + "api_disabled", + "API is disabled", + ), + )); + } + + if !api_cfg.whitelist.is_empty() && !api_cfg.whitelist.iter().any(|net| net.contains(peer.ip())) + { + return match api_cfg.gray_action { + ApiGrayAction::Api => Ok(error_response( + request_id, + ApiFailure::new( + StatusCode::FORBIDDEN, + "forbidden", + "Source IP is not allowed", + ), + )), + ApiGrayAction::Ok200 => Ok(Response::builder() + .status(StatusCode::OK) + .header("content-type", "text/html; charset=utf-8") + .body(Full::new(Bytes::new())) + .unwrap()), + ApiGrayAction::Drop => Err(IoError::new( + ErrorKind::ConnectionAborted, + "api request dropped by gray_action=drop", + )), + }; + } + + if !api_cfg.auth_header.is_empty() { + let auth_ok = req + .headers() + .get(AUTHORIZATION) + .and_then(|v| v.to_str().ok()) + .map(|v| auth_header_matches(v, &api_cfg.auth_header)) + .unwrap_or(false); + if !auth_ok { + return Ok(error_response( + request_id, + ApiFailure::new( + StatusCode::UNAUTHORIZED, + "unauthorized", + "Missing or invalid Authorization header", + ), + )); + } + } + + let method = req.method().clone(); + let path = req.uri().path().to_string(); + let normalized_path = if path.len() > 1 { + path.trim_end_matches('/') + } else { + path.as_str() + }; + let query = req.uri().query().map(str::to_string); + let body_limit = api_cfg.request_body_limit_bytes; + + let result = dispatch( + req, + method, + &path, + normalized_path, + query.as_deref(), + body_limit, + &shared, + cfg.as_ref(), + &config_rx, + request_id, + ) + .await; + match result { + Ok(resp) => Ok(resp), + Err(error) => Ok(error_response(request_id, error)), + } +} + +async fn dispatch( + req: Request, + method: Method, + path: &str, + normalized_path: &str, + query: Option<&str>, + body_limit: usize, + shared: &Arc, + cfg: &ProxyConfig, + config_rx: &watch::Receiver>, + request_id: u64, +) -> Result>, ApiFailure> { + if web_runtime::is_route(normalized_path) { + let web_mutation = method == Method::POST; + let result = web_runtime::handle( + method, + normalized_path, + query, + req, + shared.as_ref(), + cfg, + request_id, + body_limit, + ) + .await; + if web_mutation && let Err(error) = &result { + shared.runtime_events.record( + "api.web.control.failed", + format!("path={} code={}", normalized_path, error.code), + ); + } + return result; + } + + if let Some(response) = read_routes::handle( + &method, + normalized_path, + query, + shared.as_ref(), + cfg, + config_rx, + ) + .await? + { + return Ok(response); + } + + match (method.as_str(), normalized_path) { + ("POST", "/v1/users") => { + fixed_routes::create_user_route(req, shared, cfg, config_rx, request_id, body_limit) + .await + } + ("GET", "/v1/config") => fixed_routes::get_config_route(shared).await, + ("POST", "/v1/system/reload") => { + fixed_routes::reload_route(req, shared, cfg, request_id, body_limit).await + } + ("PATCH", "/v1/config") => { + fixed_routes::patch_config_route(req, shared, cfg, query, request_id, body_limit).await + } + _ => { + user_routes::handle( + req, + &method, + path, + normalized_path, + shared, + cfg, + config_rx, + request_id, + body_limit, + ) + .await + } + } +} diff --git a/src/api/handler/fixed_routes.rs b/src/api/handler/fixed_routes.rs new file mode 100644 index 0000000..57cdf99 --- /dev/null +++ b/src/api/handler/fixed_routes.rs @@ -0,0 +1,149 @@ +use super::*; + +pub(super) async fn create_user_route( + req: Request, + shared: &Arc, + cfg: &ProxyConfig, + config_rx: &watch::Receiver>, + request_id: u64, + body_limit: usize, +) -> Result>, ApiFailure> { + let api_cfg = &cfg.server.api; + if api_cfg.read_only { + return Ok(error_response( + request_id, + ApiFailure::new( + StatusCode::FORBIDDEN, + "read_only", + "API runs in read-only mode", + ), + )); + } + let expected_revision = parse_if_match(req.headers()); + let body = read_json::(req.into_body(), body_limit).await?; + let requested_enabled = body.enabled; + let result = create_user(body, expected_revision, shared).await; + let (mut data, revision) = match result { + Ok(ok) => ok, + Err(error) => { + shared + .runtime_events + .record("api.user.create.failed", error.code); + return Err(error); + } + }; + let runtime_cfg = config_rx.borrow().clone(); + data.user.in_runtime = runtime_cfg.access.users.contains_key(&data.user.username); + if let Some(enabled) = requested_enabled { + shared + .proxy_shared + .set_user_enabled(&data.user.username, enabled); + if !enabled { + let cancelled = shared + .proxy_shared + .cancel_user_sessions(&data.user.username); + if cancelled > 0 { + shared.runtime_events.record( + "api.user.disable.runtime", + format!( + "username={} cancelled_sessions={}", + data.user.username, cancelled + ), + ); + } + } + } + shared.runtime_events.record( + "api.user.create.ok", + format!("username={}", data.user.username), + ); + let status = if data.user.in_runtime { + StatusCode::CREATED + } else { + StatusCode::ACCEPTED + }; + Ok(success_response(status, data, revision)) +} + +pub(super) async fn get_config_route( + shared: &Arc, +) -> Result>, ApiFailure> { + let (value, revision) = config_edit::read_managed_config(&shared.config_path).await?; + Ok(success_response(StatusCode::OK, value, revision)) +} + +pub(super) async fn reload_route( + req: Request, + shared: &Arc, + cfg: &ProxyConfig, + request_id: u64, + body_limit: usize, +) -> Result>, ApiFailure> { + let api_cfg = &cfg.server.api; + if api_cfg.read_only { + return Ok(error_response( + request_id, + ApiFailure::new( + StatusCode::FORBIDDEN, + "read_only", + "API runs in read-only mode", + ), + )); + } + let expected_revision = parse_if_match(req.headers()); + let request = read_optional_json::(req.into_body(), body_limit) + .await? + .unwrap_or_default(); + request.validate().map_err(ApiFailure::bad_request)?; + + let (accepted, revision) = submit_reload_from_disk( + &shared.config_path, + shared.mutation_lock.as_ref(), + &shared.reload_control, + expected_revision.as_deref(), + request, + ) + .await?; + Ok(success_response(StatusCode::ACCEPTED, accepted, revision)) +} + +pub(super) async fn patch_config_route( + req: Request, + shared: &Arc, + cfg: &ProxyConfig, + query: Option<&str>, + request_id: u64, + body_limit: usize, +) -> Result>, ApiFailure> { + let api_cfg = &cfg.server.api; + if api_cfg.read_only { + return Ok(error_response( + request_id, + ApiFailure::new( + StatusCode::FORBIDDEN, + "read_only", + "API runs in read-only mode", + ), + )); + } + let expected_revision = parse_if_match(req.headers()); + let reload_request = ReloadRequest::from_query(query).map_err(ApiFailure::bad_request)?; + let body = read_json::(req.into_body(), body_limit).await?; + match config_edit::patch_config(body, expected_revision, reload_request, shared).await { + Ok(resp) => { + let revision = resp.revision.clone(); + let status = if resp.reload.is_some() { + StatusCode::ACCEPTED + } else { + StatusCode::OK + }; + Ok(success_response(status, resp, revision)) + } + Err(error) => { + shared + .runtime_events + .record("api.config.patch.failed", error.code); + Err(error) + } + } +} diff --git a/src/api/handler/read_routes.rs b/src/api/handler/read_routes.rs new file mode 100644 index 0000000..8c0ef4a --- /dev/null +++ b/src/api/handler/read_routes.rs @@ -0,0 +1,210 @@ +use super::*; + +pub(super) async fn handle( + method: &Method, + normalized_path: &str, + query: Option<&str>, + shared: &ApiShared, + cfg: &ProxyConfig, + config_rx: &watch::Receiver>, +) -> Result>>, ApiFailure> { + let api_cfg = &cfg.server.api; + match (method.as_str(), normalized_path) { + ("GET", "/web-status") => Ok(web_status::render(query, &shared.web_trace).await), + ("GET", "/v1/health") => { + let revision = current_revision(&shared.config_path).await?; + let data = HealthData { + status: "ok", + read_only: api_cfg.read_only, + }; + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/health/ready") => { + let revision = current_revision(&shared.config_path).await?; + let admission_open = shared.runtime_state.admission_open.load(Ordering::Relaxed); + let upstream_health = shared.upstream_manager.api_health_summary().await; + let ready = admission_open && upstream_health.healthy_total > 0; + let reason = if ready { + None + } else if !admission_open { + Some("admission_closed") + } else { + Some("no_healthy_upstreams") + }; + let data = HealthReadyData { + ready, + status: if ready { "ready" } else { "not_ready" }, + reason, + admission_open, + healthy_upstreams: upstream_health.healthy_total, + total_upstreams: upstream_health.configured_total, + }; + let status_code = if ready { + StatusCode::OK + } else { + StatusCode::SERVICE_UNAVAILABLE + }; + Ok(success_response(status_code, data, revision)) + } + ("GET", "/v1/system/info") => { + let revision = current_revision(&shared.config_path).await?; + let data = build_system_info_data(shared, cfg, &revision); + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/runtime/gates") => { + let revision = current_revision(&shared.config_path).await?; + let data = build_runtime_gates_data(shared, cfg).await; + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/runtime/initialization") => { + let revision = current_revision(&shared.config_path).await?; + let data = build_runtime_initialization_data(shared).await; + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/limits/effective") => { + let revision = current_revision(&shared.config_path).await?; + let data = build_limits_effective_data(cfg); + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/security/posture") => { + let revision = current_revision(&shared.config_path).await?; + let data = build_security_posture_data(cfg); + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/security/whitelist") => { + let revision = current_revision(&shared.config_path).await?; + let data = build_security_whitelist_data(cfg); + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/stats/summary") => { + let revision = current_revision(&shared.config_path).await?; + let connections_bad_by_class = shared + .stats + .get_connects_bad_class_counts() + .into_iter() + .map(|(class, total)| ClassCount { class, total }) + .collect(); + let handshake_failures_by_class = shared + .stats + .get_handshake_failure_class_counts() + .into_iter() + .map(|(class, total)| ClassCount { class, total }) + .collect(); + let data = SummaryData { + uptime_seconds: shared.stats.uptime_secs(), + connections_total: shared.stats.get_connects_all(), + connections_bad_total: shared.stats.get_connects_bad(), + connections_bad_by_class, + handshake_failures_by_class, + handshake_timeouts_total: shared.stats.get_handshake_timeouts(), + configured_users: cfg.access.users.len(), + }; + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/stats/zero/all") => { + let revision = current_revision(&shared.config_path).await?; + let data = build_zero_all_data(&shared.stats, cfg.access.users.len()); + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/stats/upstreams") => { + let revision = current_revision(&shared.config_path).await?; + let data = build_upstreams_data(shared, api_cfg); + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/stats/minimal/all") => { + let revision = current_revision(&shared.config_path).await?; + let data = build_minimal_all_data(shared, api_cfg).await; + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/stats/me-writers") => { + let revision = current_revision(&shared.config_path).await?; + let data = build_me_writers_data(shared, api_cfg).await; + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/stats/dcs") => { + let revision = current_revision(&shared.config_path).await?; + let data = build_dcs_data(shared, api_cfg).await; + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/runtime/me-pool-state") | ("GET", "/v1/runtime/me_pool_state") => { + let revision = current_revision(&shared.config_path).await?; + let data = build_runtime_me_pool_state_data(shared).await; + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/runtime/me-quality") | ("GET", "/v1/runtime/me_quality") => { + let revision = current_revision(&shared.config_path).await?; + let data = build_runtime_me_quality_data(shared).await; + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/runtime/upstream-quality") | ("GET", "/v1/runtime/upstream_quality") => { + let revision = current_revision(&shared.config_path).await?; + let data = build_runtime_upstream_quality_data(shared).await; + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/runtime/nat-stun") | ("GET", "/v1/runtime/nat_stun") => { + let revision = current_revision(&shared.config_path).await?; + let data = build_runtime_nat_stun_data(shared).await; + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/runtime/me-selftest") => { + let revision = current_revision(&shared.config_path).await?; + let data = build_runtime_me_selftest_data(shared, cfg).await; + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/runtime/connections/summary") => { + let revision = current_revision(&shared.config_path).await?; + let data = build_runtime_connections_summary_data(shared, cfg).await; + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/runtime/events/recent") => { + let revision = current_revision(&shared.config_path).await?; + let data = build_runtime_events_recent_data(shared, cfg, query); + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/runtime/tls-fingerprints") => { + let revision = current_revision(&shared.config_path).await?; + let data = build_runtime_tls_fingerprints_data(shared, cfg, query); + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/stats/users/active-ips") => { + let revision = current_revision(&shared.config_path).await?; + let usernames: Vec<_> = cfg.access.users.keys().cloned().collect(); + let active_ips_map = shared.ip_tracker.get_active_ips_for_users(&usernames).await; + let mut data: Vec = active_ips_map + .into_iter() + .filter(|(_, ips)| !ips.is_empty()) + .map(|(username, active_ips)| UserActiveIps { + username, + active_ips, + }) + .collect(); + data.sort_by(|a, b| a.username.cmp(&b.username)); + Ok(success_response(StatusCode::OK, data, revision)) + } + ("GET", "/v1/stats/users") | ("GET", "/v1/users") => { + let revision = current_revision(&shared.config_path).await?; + let disk_cfg = load_config_from_disk(&shared.config_path).await?; + let runtime_cfg = config_rx.borrow().clone(); + let (detected_ip_v4, detected_ip_v6) = shared.detected_link_ips(); + let users = users_from_config( + &disk_cfg, + &shared.stats, + &shared.ip_tracker, + detected_ip_v4, + detected_ip_v6, + Some(runtime_cfg.as_ref()), + ) + .await; + Ok(success_response(StatusCode::OK, users, revision)) + } + ("GET", "/v1/stats/users/quota") => { + let revision = current_revision(&shared.config_path).await?; + let disk_cfg = load_config_from_disk(&shared.config_path).await?; + let data = build_user_quota_list(&disk_cfg, shared.stats.as_ref()); + Ok(success_response(StatusCode::OK, data, revision)) + } + + _ => return Ok(None), + } + .map(Some) +} diff --git a/src/api/handler/user_routes.rs b/src/api/handler/user_routes.rs new file mode 100644 index 0000000..004f01d --- /dev/null +++ b/src/api/handler/user_routes.rs @@ -0,0 +1,386 @@ +use super::*; + +pub(super) async fn handle( + req: Request, + method: &Method, + path: &str, + normalized_path: &str, + shared: &Arc, + cfg: &ProxyConfig, + config_rx: &watch::Receiver>, + request_id: u64, + body_limit: usize, +) -> Result>, ApiFailure> { + let api_cfg = &cfg.server.api; + if method == Method::GET + && let Some(reload_id) = reload_status_route_id(normalized_path) + { + let revision = current_revision(&shared.config_path).await?; + let status = shared + .reload_control + .status(reload_id) + .await + .ok_or_else(|| { + ApiFailure::new( + StatusCode::NOT_FOUND, + "reload_not_found", + format!("Reload {} was not found", reload_id), + ) + })?; + return Ok(success_response(StatusCode::OK, status, revision)); + } + if method == Method::POST + && let Some(base_user) = normalized_path + .strip_prefix("/v1/users/") + .and_then(|path| path.strip_suffix("/enable")) + && !base_user.is_empty() + && !base_user.contains('/') + { + let base_user = parse_route_username(base_user)?; + if api_cfg.read_only { + return Ok(error_response( + request_id, + ApiFailure::new( + StatusCode::FORBIDDEN, + "read_only", + "API runs in read-only mode", + ), + )); + } + let expected_revision = parse_if_match(req.headers()); + let result = set_user_enabled(base_user, true, expected_revision, shared).await; + let (mut data, revision) = match result { + Ok(ok) => ok, + Err(error) => { + shared.runtime_events.record( + "api.user.enable.failed", + format!("username={} code={}", base_user, error.code), + ); + return Err(error); + } + }; + let runtime_cfg = config_rx.borrow().clone(); + data.in_runtime = runtime_cfg.access.users.contains_key(&data.username); + shared.proxy_shared.set_user_enabled(base_user, true); + shared + .runtime_events + .record("api.user.enable.ok", format!("username={}", base_user)); + let status = if data.in_runtime { + StatusCode::OK + } else { + StatusCode::ACCEPTED + }; + return Ok(success_response(status, data, revision)); + } + if method == Method::POST + && let Some(base_user) = normalized_path + .strip_prefix("/v1/users/") + .and_then(|path| path.strip_suffix("/disable")) + && !base_user.is_empty() + && !base_user.contains('/') + { + let base_user = parse_route_username(base_user)?; + if api_cfg.read_only { + return Ok(error_response( + request_id, + ApiFailure::new( + StatusCode::FORBIDDEN, + "read_only", + "API runs in read-only mode", + ), + )); + } + let expected_revision = parse_if_match(req.headers()); + let result = set_user_enabled(base_user, false, expected_revision, shared).await; + let (mut data, revision) = match result { + Ok(ok) => ok, + Err(error) => { + shared.runtime_events.record( + "api.user.disable.failed", + format!("username={} code={}", base_user, error.code), + ); + return Err(error); + } + }; + let runtime_cfg = config_rx.borrow().clone(); + data.in_runtime = runtime_cfg.access.users.contains_key(&data.username); + let newly_disabled = shared.proxy_shared.set_user_enabled(base_user, false); + let cancelled = shared.proxy_shared.cancel_user_sessions(base_user); + shared.runtime_events.record( + "api.user.disable.ok", + format!( + "username={} newly_disabled={} cancelled_sessions={}", + base_user, newly_disabled, cancelled + ), + ); + let status = if data.in_runtime { + StatusCode::OK + } else { + StatusCode::ACCEPTED + }; + return Ok(success_response(status, data, revision)); + } + if method == Method::POST + && let Some(user) = normalized_path + .strip_prefix("/v1/users/") + .and_then(|path| path.strip_suffix("/reset-quota")) + && !user.is_empty() + && !user.contains('/') + { + let user = parse_route_username(user)?; + if api_cfg.read_only { + return Ok(error_response( + request_id, + ApiFailure::new( + StatusCode::FORBIDDEN, + "read_only", + "API runs in read-only mode", + ), + )); + } + let expected_revision = parse_if_match(req.headers()); + let _mutation_guard = shared.mutation_lock.lock().await; + let disk_cfg = load_config_from_disk(&shared.config_path).await?; + ensure_expected_revision(&shared.config_path, expected_revision.as_deref()).await?; + if !disk_cfg.access.users.contains_key(user) { + return Ok(error_response( + request_id, + ApiFailure::new(StatusCode::NOT_FOUND, "not_found", "User not found"), + )); + } + let configured_users = disk_cfg + .access + .users + .keys() + .cloned() + .collect::>(); + let snapshot = match shared.quota_state.reset_user(&configured_users, user).await { + Ok(snapshot) => snapshot, + Err(error) => { + shared.runtime_events.record( + "api.user.reset_quota.failed", + format!("username={} error={}", user, error), + ); + return Err(ApiFailure::internal(format!( + "Failed to reset user quota: {}", + error + ))); + } + }; + shared + .runtime_events + .record("api.user.reset_quota.ok", format!("username={}", user)); + let revision = current_revision(&shared.config_path).await?; + return Ok(success_response( + StatusCode::OK, + ResetUserQuotaResponse { + username: user.to_string(), + used_bytes: snapshot.used_bytes, + last_reset_epoch_secs: snapshot.last_reset_epoch_secs, + }, + revision, + )); + } + if method == Method::POST + && let Some(base_user) = normalized_path + .strip_prefix("/v1/users/") + .and_then(|path| path.strip_suffix("/rotate-secret")) + && !base_user.is_empty() + && !base_user.contains('/') + { + let base_user = parse_route_username(base_user)?; + if api_cfg.read_only { + return Ok(error_response( + request_id, + ApiFailure::new( + StatusCode::FORBIDDEN, + "read_only", + "API runs in read-only mode", + ), + )); + } + let expected_revision = parse_if_match(req.headers()); + let body = read_optional_json::(req.into_body(), body_limit).await?; + let result = rotate_secret( + base_user, + body.unwrap_or_default(), + expected_revision, + shared, + ) + .await; + let (mut data, revision) = match result { + Ok(ok) => ok, + Err(error) => { + shared.runtime_events.record( + "api.user.rotate_secret.failed", + format!("username={} code={}", base_user, error.code), + ); + return Err(error); + } + }; + let runtime_cfg = config_rx.borrow().clone(); + data.user.in_runtime = runtime_cfg.access.users.contains_key(&data.user.username); + shared.runtime_events.record( + "api.user.rotate_secret.ok", + format!("username={}", base_user), + ); + let status = if data.user.in_runtime { + StatusCode::OK + } else { + StatusCode::ACCEPTED + }; + return Ok(success_response(status, data, revision)); + } + if let Some(user) = normalized_path.strip_prefix("/v1/users/") + && !user.is_empty() + && !user.contains('/') + { + let user = parse_route_username(user)?; + if method == Method::GET { + let revision = current_revision(&shared.config_path).await?; + let disk_cfg = load_config_from_disk(&shared.config_path).await?; + let runtime_cfg = config_rx.borrow().clone(); + let (detected_ip_v4, detected_ip_v6) = shared.detected_link_ips(); + let users = users_from_config( + &disk_cfg, + &shared.stats, + &shared.ip_tracker, + detected_ip_v4, + detected_ip_v6, + Some(runtime_cfg.as_ref()), + ) + .await; + if let Some(user_info) = users.into_iter().find(|entry| entry.username == user) { + return Ok(success_response(StatusCode::OK, user_info, revision)); + } + return Ok(error_response( + request_id, + ApiFailure::new(StatusCode::NOT_FOUND, "not_found", "User not found"), + )); + } + if method == Method::PATCH { + if api_cfg.read_only { + return Ok(error_response( + request_id, + ApiFailure::new( + StatusCode::FORBIDDEN, + "read_only", + "API runs in read-only mode", + ), + )); + } + let expected_revision = parse_if_match(req.headers()); + let body = read_json::(req.into_body(), body_limit).await?; + let enabled_update = match &body.enabled { + Patch::Unchanged => None, + Patch::Remove => Some(true), + Patch::Set(enabled) => Some(*enabled), + }; + let result = patch_user(user, body, expected_revision, shared).await; + let (mut data, revision) = match result { + Ok(ok) => ok, + Err(error) => { + shared.runtime_events.record( + "api.user.patch.failed", + format!("username={} code={}", user, error.code), + ); + return Err(error); + } + }; + let runtime_cfg = config_rx.borrow().clone(); + data.in_runtime = runtime_cfg.access.users.contains_key(&data.username); + if let Some(enabled) = enabled_update { + shared + .proxy_shared + .set_user_enabled(&data.username, enabled); + if !enabled { + let cancelled = shared.proxy_shared.cancel_user_sessions(&data.username); + shared.runtime_events.record( + "api.user.disable.runtime", + format!( + "username={} cancelled_sessions={}", + data.username, cancelled + ), + ); + } + } + shared + .runtime_events + .record("api.user.patch.ok", format!("username={}", data.username)); + let status = if data.in_runtime { + StatusCode::OK + } else { + StatusCode::ACCEPTED + }; + return Ok(success_response(status, data, revision)); + } + if method == Method::DELETE { + if api_cfg.read_only { + return Ok(error_response( + request_id, + ApiFailure::new( + StatusCode::FORBIDDEN, + "read_only", + "API runs in read-only mode", + ), + )); + } + let expected_revision = parse_if_match(req.headers()); + let result = delete_user(user, expected_revision, shared).await; + let (deleted_user, revision) = match result { + Ok(ok) => ok, + Err(error) => { + shared.runtime_events.record( + "api.user.delete.failed", + format!("username={} code={}", user, error.code), + ); + return Err(error); + } + }; + shared.proxy_shared.set_user_enabled(&deleted_user, true); + let cancelled = shared.proxy_shared.cancel_user_sessions(&deleted_user); + shared.runtime_events.record( + "api.user.delete.ok", + format!("username={} cancelled_sessions={}", deleted_user, cancelled), + ); + let runtime_cfg = config_rx.borrow().clone(); + let in_runtime = runtime_cfg.access.users.contains_key(&deleted_user); + let response = DeleteUserResponse { + username: deleted_user, + in_runtime, + }; + let status = if response.in_runtime { + StatusCode::ACCEPTED + } else { + StatusCode::OK + }; + return Ok(success_response(status, response, revision)); + } + if method == Method::POST { + return Ok(error_response( + request_id, + ApiFailure::method_not_allowed(ALLOW_GET_PATCH_DELETE), + )); + } + return Ok(error_response( + request_id, + ApiFailure::method_not_allowed(ALLOW_GET_PATCH_DELETE), + )); + } + if let Some(allow) = allowed_methods_for_path(normalized_path) { + return Ok(error_response( + request_id, + ApiFailure::method_not_allowed(allow), + )); + } + debug!( + method = method.as_str(), + path = %path, + normalized_path = %normalized_path, + "API route not found" + ); + Ok(error_response( + request_id, + ApiFailure::new(StatusCode::NOT_FOUND, "not_found", "Route not found"), + )) +} diff --git a/src/api/mod.rs b/src/api/mod.rs index ecda81e..e4aad93 100644 --- a/src/api/mod.rs +++ b/src/api/mod.rs @@ -21,7 +21,7 @@ use tokio::sync::{Mutex, RwLock, Semaphore, watch}; use tokio::time::timeout; use tracing::{debug, info, warn}; -use crate::config::ApiGrayAction; +use crate::config::{ApiGrayAction, ProxyConfig}; use crate::ip_tracker::UserIpTracker; use crate::maestro::control_plane::ProcessControlPlane; use crate::maestro::generation::{RuntimeGeneration, RuntimeWatchState}; @@ -40,6 +40,8 @@ mod config_edit; pub(crate) mod config_store; mod events; mod http_utils; +// Authenticated request admission and route dispatch. +mod handler; mod model; mod patch; #[cfg(test)] @@ -61,6 +63,7 @@ use config_store::{ parse_if_match, }; use events::ApiEventStore; +use handler::handle; use http_utils::{error_response, read_json, read_optional_json, success_response}; use model::{ ApiFailure, ClassCount, CreateUserRequest, DeleteUserResponse, HealthData, HealthReadyData, @@ -424,835 +427,3 @@ pub(crate) async fn serve( }); } } - -async fn handle( - req: Request, - peer: SocketAddr, - shared: Arc, -) -> Result>, IoError> { - let runtime = shared.active_runtime.load_full(); - let previous_cache_generation = shared.cache_generation.swap(runtime.id, Ordering::AcqRel); - if previous_cache_generation != runtime.id { - *shared.minimal_cache.lock().await = None; - *shared.runtime_edge_connections_cache.lock().await = None; - } - let shared = Arc::new(shared.for_runtime(runtime.as_ref())); - let config_rx = runtime.config_rx.clone(); - shared - .runtime_state - .admission_open - .store(*runtime.admission_rx.borrow(), Ordering::Relaxed); - let request_id = shared.next_request_id(); - let cfg = config_rx.borrow().clone(); - let api_cfg = &cfg.server.api; - - if !api_cfg.enabled { - return Ok(error_response( - request_id, - ApiFailure::new( - StatusCode::SERVICE_UNAVAILABLE, - "api_disabled", - "API is disabled", - ), - )); - } - - if !api_cfg.whitelist.is_empty() && !api_cfg.whitelist.iter().any(|net| net.contains(peer.ip())) - { - return match api_cfg.gray_action { - ApiGrayAction::Api => Ok(error_response( - request_id, - ApiFailure::new( - StatusCode::FORBIDDEN, - "forbidden", - "Source IP is not allowed", - ), - )), - ApiGrayAction::Ok200 => Ok(Response::builder() - .status(StatusCode::OK) - .header("content-type", "text/html; charset=utf-8") - .body(Full::new(Bytes::new())) - .unwrap()), - ApiGrayAction::Drop => Err(IoError::new( - ErrorKind::ConnectionAborted, - "api request dropped by gray_action=drop", - )), - }; - } - - if !api_cfg.auth_header.is_empty() { - let auth_ok = req - .headers() - .get(AUTHORIZATION) - .and_then(|v| v.to_str().ok()) - .map(|v| auth_header_matches(v, &api_cfg.auth_header)) - .unwrap_or(false); - if !auth_ok { - return Ok(error_response( - request_id, - ApiFailure::new( - StatusCode::UNAUTHORIZED, - "unauthorized", - "Missing or invalid Authorization header", - ), - )); - } - } - - let method = req.method().clone(); - let path = req.uri().path().to_string(); - let normalized_path = if path.len() > 1 { - path.trim_end_matches('/') - } else { - path.as_str() - }; - let query = req.uri().query().map(str::to_string); - let body_limit = api_cfg.request_body_limit_bytes; - - let result: Result>, ApiFailure> = async { - if web_runtime::is_route(normalized_path) { - let web_mutation = method == Method::POST; - let result = web_runtime::handle( - method, - normalized_path, - query.as_deref(), - req, - shared.as_ref(), - cfg.as_ref(), - request_id, - body_limit, - ) - .await; - if web_mutation && let Err(error) = &result { - shared.runtime_events.record( - "api.web.control.failed", - format!("path={} code={}", normalized_path, error.code), - ); - } - return result; - } - match (method.as_str(), normalized_path) { - ("GET", "/web-status") => { - Ok(web_status::render(query.as_deref(), &shared.web_trace).await) - } - ("GET", "/v1/health") => { - let revision = current_revision(&shared.config_path).await?; - let data = HealthData { - status: "ok", - read_only: api_cfg.read_only, - }; - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/health/ready") => { - let revision = current_revision(&shared.config_path).await?; - let admission_open = shared.runtime_state.admission_open.load(Ordering::Relaxed); - let upstream_health = shared.upstream_manager.api_health_summary().await; - let ready = admission_open && upstream_health.healthy_total > 0; - let reason = if ready { - None - } else if !admission_open { - Some("admission_closed") - } else { - Some("no_healthy_upstreams") - }; - let data = HealthReadyData { - ready, - status: if ready { "ready" } else { "not_ready" }, - reason, - admission_open, - healthy_upstreams: upstream_health.healthy_total, - total_upstreams: upstream_health.configured_total, - }; - let status_code = if ready { - StatusCode::OK - } else { - StatusCode::SERVICE_UNAVAILABLE - }; - Ok(success_response(status_code, data, revision)) - } - ("GET", "/v1/system/info") => { - let revision = current_revision(&shared.config_path).await?; - let data = build_system_info_data(shared.as_ref(), cfg.as_ref(), &revision); - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/runtime/gates") => { - let revision = current_revision(&shared.config_path).await?; - let data = build_runtime_gates_data(shared.as_ref(), cfg.as_ref()).await; - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/runtime/initialization") => { - let revision = current_revision(&shared.config_path).await?; - let data = build_runtime_initialization_data(shared.as_ref()).await; - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/limits/effective") => { - let revision = current_revision(&shared.config_path).await?; - let data = build_limits_effective_data(cfg.as_ref()); - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/security/posture") => { - let revision = current_revision(&shared.config_path).await?; - let data = build_security_posture_data(cfg.as_ref()); - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/security/whitelist") => { - let revision = current_revision(&shared.config_path).await?; - let data = build_security_whitelist_data(cfg.as_ref()); - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/stats/summary") => { - let revision = current_revision(&shared.config_path).await?; - let connections_bad_by_class = shared - .stats - .get_connects_bad_class_counts() - .into_iter() - .map(|(class, total)| ClassCount { class, total }) - .collect(); - let handshake_failures_by_class = shared - .stats - .get_handshake_failure_class_counts() - .into_iter() - .map(|(class, total)| ClassCount { class, total }) - .collect(); - let data = SummaryData { - uptime_seconds: shared.stats.uptime_secs(), - connections_total: shared.stats.get_connects_all(), - connections_bad_total: shared.stats.get_connects_bad(), - connections_bad_by_class, - handshake_failures_by_class, - handshake_timeouts_total: shared.stats.get_handshake_timeouts(), - configured_users: cfg.access.users.len(), - }; - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/stats/zero/all") => { - let revision = current_revision(&shared.config_path).await?; - let data = build_zero_all_data(&shared.stats, cfg.access.users.len()); - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/stats/upstreams") => { - let revision = current_revision(&shared.config_path).await?; - let data = build_upstreams_data(shared.as_ref(), api_cfg); - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/stats/minimal/all") => { - let revision = current_revision(&shared.config_path).await?; - let data = build_minimal_all_data(shared.as_ref(), api_cfg).await; - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/stats/me-writers") => { - let revision = current_revision(&shared.config_path).await?; - let data = build_me_writers_data(shared.as_ref(), api_cfg).await; - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/stats/dcs") => { - let revision = current_revision(&shared.config_path).await?; - let data = build_dcs_data(shared.as_ref(), api_cfg).await; - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/runtime/me-pool-state") | ("GET", "/v1/runtime/me_pool_state") => { - let revision = current_revision(&shared.config_path).await?; - let data = build_runtime_me_pool_state_data(shared.as_ref()).await; - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/runtime/me-quality") | ("GET", "/v1/runtime/me_quality") => { - let revision = current_revision(&shared.config_path).await?; - let data = build_runtime_me_quality_data(shared.as_ref()).await; - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/runtime/upstream-quality") | ("GET", "/v1/runtime/upstream_quality") => { - let revision = current_revision(&shared.config_path).await?; - let data = build_runtime_upstream_quality_data(shared.as_ref()).await; - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/runtime/nat-stun") | ("GET", "/v1/runtime/nat_stun") => { - let revision = current_revision(&shared.config_path).await?; - let data = build_runtime_nat_stun_data(shared.as_ref()).await; - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/runtime/me-selftest") => { - let revision = current_revision(&shared.config_path).await?; - let data = build_runtime_me_selftest_data(shared.as_ref(), cfg.as_ref()).await; - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/runtime/connections/summary") => { - let revision = current_revision(&shared.config_path).await?; - let data = - build_runtime_connections_summary_data(shared.as_ref(), cfg.as_ref()).await; - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/runtime/events/recent") => { - let revision = current_revision(&shared.config_path).await?; - let data = build_runtime_events_recent_data( - shared.as_ref(), - cfg.as_ref(), - query.as_deref(), - ); - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/runtime/tls-fingerprints") => { - let revision = current_revision(&shared.config_path).await?; - let data = build_runtime_tls_fingerprints_data( - shared.as_ref(), - cfg.as_ref(), - query.as_deref(), - ); - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/stats/users/active-ips") => { - let revision = current_revision(&shared.config_path).await?; - let usernames: Vec<_> = cfg.access.users.keys().cloned().collect(); - let active_ips_map = shared.ip_tracker.get_active_ips_for_users(&usernames).await; - let mut data: Vec = active_ips_map - .into_iter() - .filter(|(_, ips)| !ips.is_empty()) - .map(|(username, active_ips)| UserActiveIps { - username, - active_ips, - }) - .collect(); - data.sort_by(|a, b| a.username.cmp(&b.username)); - Ok(success_response(StatusCode::OK, data, revision)) - } - ("GET", "/v1/stats/users") | ("GET", "/v1/users") => { - let revision = current_revision(&shared.config_path).await?; - let disk_cfg = load_config_from_disk(&shared.config_path).await?; - let runtime_cfg = config_rx.borrow().clone(); - let (detected_ip_v4, detected_ip_v6) = shared.detected_link_ips(); - let users = users_from_config( - &disk_cfg, - &shared.stats, - &shared.ip_tracker, - detected_ip_v4, - detected_ip_v6, - Some(runtime_cfg.as_ref()), - ) - .await; - Ok(success_response(StatusCode::OK, users, revision)) - } - ("GET", "/v1/stats/users/quota") => { - let revision = current_revision(&shared.config_path).await?; - let disk_cfg = load_config_from_disk(&shared.config_path).await?; - let data = build_user_quota_list(&disk_cfg, shared.stats.as_ref()); - Ok(success_response(StatusCode::OK, data, revision)) - } - ("POST", "/v1/users") => { - if api_cfg.read_only { - return Ok(error_response( - request_id, - ApiFailure::new( - StatusCode::FORBIDDEN, - "read_only", - "API runs in read-only mode", - ), - )); - } - let expected_revision = parse_if_match(req.headers()); - let body = read_json::(req.into_body(), body_limit).await?; - let requested_enabled = body.enabled; - let result = create_user(body, expected_revision, &shared).await; - let (mut data, revision) = match result { - Ok(ok) => ok, - Err(error) => { - shared - .runtime_events - .record("api.user.create.failed", error.code); - return Err(error); - } - }; - let runtime_cfg = config_rx.borrow().clone(); - data.user.in_runtime = runtime_cfg.access.users.contains_key(&data.user.username); - if let Some(enabled) = requested_enabled { - shared - .proxy_shared - .set_user_enabled(&data.user.username, enabled); - if !enabled { - let cancelled = shared - .proxy_shared - .cancel_user_sessions(&data.user.username); - if cancelled > 0 { - shared.runtime_events.record( - "api.user.disable.runtime", - format!( - "username={} cancelled_sessions={}", - data.user.username, cancelled - ), - ); - } - } - } - shared.runtime_events.record( - "api.user.create.ok", - format!("username={}", data.user.username), - ); - let status = if data.user.in_runtime { - StatusCode::CREATED - } else { - StatusCode::ACCEPTED - }; - Ok(success_response(status, data, revision)) - } - ("GET", "/v1/config") => { - let (value, revision) = - config_edit::read_managed_config(&shared.config_path).await?; - Ok(success_response(StatusCode::OK, value, revision)) - } - ("POST", "/v1/system/reload") => { - if api_cfg.read_only { - return Ok(error_response( - request_id, - ApiFailure::new( - StatusCode::FORBIDDEN, - "read_only", - "API runs in read-only mode", - ), - )); - } - let expected_revision = parse_if_match(req.headers()); - let request = read_optional_json::(req.into_body(), body_limit) - .await? - .unwrap_or_default(); - request.validate().map_err(ApiFailure::bad_request)?; - - let (accepted, revision) = submit_reload_from_disk( - &shared.config_path, - shared.mutation_lock.as_ref(), - &shared.reload_control, - expected_revision.as_deref(), - request, - ) - .await?; - Ok(success_response(StatusCode::ACCEPTED, accepted, revision)) - } - ("PATCH", "/v1/config") => { - if api_cfg.read_only { - return Ok(error_response( - request_id, - ApiFailure::new( - StatusCode::FORBIDDEN, - "read_only", - "API runs in read-only mode", - ), - )); - } - let expected_revision = parse_if_match(req.headers()); - let reload_request = - ReloadRequest::from_query(query.as_deref()).map_err(ApiFailure::bad_request)?; - let body = read_json::(req.into_body(), body_limit).await?; - match config_edit::patch_config(body, expected_revision, reload_request, &shared) - .await - { - Ok(resp) => { - let revision = resp.revision.clone(); - let status = if resp.reload.is_some() { - StatusCode::ACCEPTED - } else { - StatusCode::OK - }; - Ok(success_response(status, resp, revision)) - } - Err(error) => { - shared - .runtime_events - .record("api.config.patch.failed", error.code); - Err(error) - } - } - } - _ => { - if method == Method::GET - && let Some(reload_id) = reload_status_route_id(normalized_path) - { - let revision = current_revision(&shared.config_path).await?; - let status = - shared - .reload_control - .status(reload_id) - .await - .ok_or_else(|| { - ApiFailure::new( - StatusCode::NOT_FOUND, - "reload_not_found", - format!("Reload {} was not found", reload_id), - ) - })?; - return Ok(success_response(StatusCode::OK, status, revision)); - } - if method == Method::POST - && let Some(base_user) = normalized_path - .strip_prefix("/v1/users/") - .and_then(|path| path.strip_suffix("/enable")) - && !base_user.is_empty() - && !base_user.contains('/') - { - let base_user = parse_route_username(base_user)?; - if api_cfg.read_only { - return Ok(error_response( - request_id, - ApiFailure::new( - StatusCode::FORBIDDEN, - "read_only", - "API runs in read-only mode", - ), - )); - } - let expected_revision = parse_if_match(req.headers()); - let result = - set_user_enabled(base_user, true, expected_revision, &shared).await; - let (mut data, revision) = match result { - Ok(ok) => ok, - Err(error) => { - shared.runtime_events.record( - "api.user.enable.failed", - format!("username={} code={}", base_user, error.code), - ); - return Err(error); - } - }; - let runtime_cfg = config_rx.borrow().clone(); - data.in_runtime = runtime_cfg.access.users.contains_key(&data.username); - shared.proxy_shared.set_user_enabled(base_user, true); - shared - .runtime_events - .record("api.user.enable.ok", format!("username={}", base_user)); - let status = if data.in_runtime { - StatusCode::OK - } else { - StatusCode::ACCEPTED - }; - return Ok(success_response(status, data, revision)); - } - if method == Method::POST - && let Some(base_user) = normalized_path - .strip_prefix("/v1/users/") - .and_then(|path| path.strip_suffix("/disable")) - && !base_user.is_empty() - && !base_user.contains('/') - { - let base_user = parse_route_username(base_user)?; - if api_cfg.read_only { - return Ok(error_response( - request_id, - ApiFailure::new( - StatusCode::FORBIDDEN, - "read_only", - "API runs in read-only mode", - ), - )); - } - let expected_revision = parse_if_match(req.headers()); - let result = - set_user_enabled(base_user, false, expected_revision, &shared).await; - let (mut data, revision) = match result { - Ok(ok) => ok, - Err(error) => { - shared.runtime_events.record( - "api.user.disable.failed", - format!("username={} code={}", base_user, error.code), - ); - return Err(error); - } - }; - let runtime_cfg = config_rx.borrow().clone(); - data.in_runtime = runtime_cfg.access.users.contains_key(&data.username); - let newly_disabled = shared.proxy_shared.set_user_enabled(base_user, false); - let cancelled = shared.proxy_shared.cancel_user_sessions(base_user); - shared.runtime_events.record( - "api.user.disable.ok", - format!( - "username={} newly_disabled={} cancelled_sessions={}", - base_user, newly_disabled, cancelled - ), - ); - let status = if data.in_runtime { - StatusCode::OK - } else { - StatusCode::ACCEPTED - }; - return Ok(success_response(status, data, revision)); - } - if method == Method::POST - && let Some(user) = normalized_path - .strip_prefix("/v1/users/") - .and_then(|path| path.strip_suffix("/reset-quota")) - && !user.is_empty() - && !user.contains('/') - { - let user = parse_route_username(user)?; - if api_cfg.read_only { - return Ok(error_response( - request_id, - ApiFailure::new( - StatusCode::FORBIDDEN, - "read_only", - "API runs in read-only mode", - ), - )); - } - let expected_revision = parse_if_match(req.headers()); - let _mutation_guard = shared.mutation_lock.lock().await; - let disk_cfg = load_config_from_disk(&shared.config_path).await?; - ensure_expected_revision(&shared.config_path, expected_revision.as_deref()) - .await?; - if !disk_cfg.access.users.contains_key(user) { - return Ok(error_response( - request_id, - ApiFailure::new(StatusCode::NOT_FOUND, "not_found", "User not found"), - )); - } - let configured_users = disk_cfg - .access - .users - .keys() - .cloned() - .collect::>(); - let snapshot = match shared - .quota_state - .reset_user(&configured_users, user) - .await - { - Ok(snapshot) => snapshot, - Err(error) => { - shared.runtime_events.record( - "api.user.reset_quota.failed", - format!("username={} error={}", user, error), - ); - return Err(ApiFailure::internal(format!( - "Failed to reset user quota: {}", - error - ))); - } - }; - shared - .runtime_events - .record("api.user.reset_quota.ok", format!("username={}", user)); - let revision = current_revision(&shared.config_path).await?; - return Ok(success_response( - StatusCode::OK, - ResetUserQuotaResponse { - username: user.to_string(), - used_bytes: snapshot.used_bytes, - last_reset_epoch_secs: snapshot.last_reset_epoch_secs, - }, - revision, - )); - } - if method == Method::POST - && let Some(base_user) = normalized_path - .strip_prefix("/v1/users/") - .and_then(|path| path.strip_suffix("/rotate-secret")) - && !base_user.is_empty() - && !base_user.contains('/') - { - let base_user = parse_route_username(base_user)?; - if api_cfg.read_only { - return Ok(error_response( - request_id, - ApiFailure::new( - StatusCode::FORBIDDEN, - "read_only", - "API runs in read-only mode", - ), - )); - } - let expected_revision = parse_if_match(req.headers()); - let body = - read_optional_json::(req.into_body(), body_limit) - .await?; - let result = rotate_secret( - base_user, - body.unwrap_or_default(), - expected_revision, - &shared, - ) - .await; - let (mut data, revision) = match result { - Ok(ok) => ok, - Err(error) => { - shared.runtime_events.record( - "api.user.rotate_secret.failed", - format!("username={} code={}", base_user, error.code), - ); - return Err(error); - } - }; - let runtime_cfg = config_rx.borrow().clone(); - data.user.in_runtime = - runtime_cfg.access.users.contains_key(&data.user.username); - shared.runtime_events.record( - "api.user.rotate_secret.ok", - format!("username={}", base_user), - ); - let status = if data.user.in_runtime { - StatusCode::OK - } else { - StatusCode::ACCEPTED - }; - return Ok(success_response(status, data, revision)); - } - if let Some(user) = normalized_path.strip_prefix("/v1/users/") - && !user.is_empty() - && !user.contains('/') - { - let user = parse_route_username(user)?; - if method == Method::GET { - let revision = current_revision(&shared.config_path).await?; - let disk_cfg = load_config_from_disk(&shared.config_path).await?; - let runtime_cfg = config_rx.borrow().clone(); - let (detected_ip_v4, detected_ip_v6) = shared.detected_link_ips(); - let users = users_from_config( - &disk_cfg, - &shared.stats, - &shared.ip_tracker, - detected_ip_v4, - detected_ip_v6, - Some(runtime_cfg.as_ref()), - ) - .await; - if let Some(user_info) = - users.into_iter().find(|entry| entry.username == user) - { - return Ok(success_response(StatusCode::OK, user_info, revision)); - } - return Ok(error_response( - request_id, - ApiFailure::new(StatusCode::NOT_FOUND, "not_found", "User not found"), - )); - } - if method == Method::PATCH { - if api_cfg.read_only { - return Ok(error_response( - request_id, - ApiFailure::new( - StatusCode::FORBIDDEN, - "read_only", - "API runs in read-only mode", - ), - )); - } - let expected_revision = parse_if_match(req.headers()); - let body = - read_json::(req.into_body(), body_limit).await?; - let enabled_update = match &body.enabled { - Patch::Unchanged => None, - Patch::Remove => Some(true), - Patch::Set(enabled) => Some(*enabled), - }; - let result = patch_user(user, body, expected_revision, &shared).await; - let (mut data, revision) = match result { - Ok(ok) => ok, - Err(error) => { - shared.runtime_events.record( - "api.user.patch.failed", - format!("username={} code={}", user, error.code), - ); - return Err(error); - } - }; - let runtime_cfg = config_rx.borrow().clone(); - data.in_runtime = runtime_cfg.access.users.contains_key(&data.username); - if let Some(enabled) = enabled_update { - shared - .proxy_shared - .set_user_enabled(&data.username, enabled); - if !enabled { - let cancelled = - shared.proxy_shared.cancel_user_sessions(&data.username); - shared.runtime_events.record( - "api.user.disable.runtime", - format!( - "username={} cancelled_sessions={}", - data.username, cancelled - ), - ); - } - } - shared - .runtime_events - .record("api.user.patch.ok", format!("username={}", data.username)); - let status = if data.in_runtime { - StatusCode::OK - } else { - StatusCode::ACCEPTED - }; - return Ok(success_response(status, data, revision)); - } - if method == Method::DELETE { - if api_cfg.read_only { - return Ok(error_response( - request_id, - ApiFailure::new( - StatusCode::FORBIDDEN, - "read_only", - "API runs in read-only mode", - ), - )); - } - let expected_revision = parse_if_match(req.headers()); - let result = delete_user(user, expected_revision, &shared).await; - let (deleted_user, revision) = match result { - Ok(ok) => ok, - Err(error) => { - shared.runtime_events.record( - "api.user.delete.failed", - format!("username={} code={}", user, error.code), - ); - return Err(error); - } - }; - shared.proxy_shared.set_user_enabled(&deleted_user, true); - let cancelled = shared.proxy_shared.cancel_user_sessions(&deleted_user); - shared.runtime_events.record( - "api.user.delete.ok", - format!("username={} cancelled_sessions={}", deleted_user, cancelled), - ); - let runtime_cfg = config_rx.borrow().clone(); - let in_runtime = runtime_cfg.access.users.contains_key(&deleted_user); - let response = DeleteUserResponse { - username: deleted_user, - in_runtime, - }; - let status = if response.in_runtime { - StatusCode::ACCEPTED - } else { - StatusCode::OK - }; - return Ok(success_response(status, response, revision)); - } - if method == Method::POST { - return Ok(error_response( - request_id, - ApiFailure::method_not_allowed(ALLOW_GET_PATCH_DELETE), - )); - } - return Ok(error_response( - request_id, - ApiFailure::method_not_allowed(ALLOW_GET_PATCH_DELETE), - )); - } - if let Some(allow) = allowed_methods_for_path(normalized_path) { - return Ok(error_response( - request_id, - ApiFailure::method_not_allowed(allow), - )); - } - debug!( - method = method.as_str(), - path = %path, - normalized_path = %normalized_path, - "API route not found" - ); - Ok(error_response( - request_id, - ApiFailure::new(StatusCode::NOT_FOUND, "not_found", "Route not found"), - )) - } - } - } - .await; - - match result { - Ok(resp) => Ok(resp), - Err(error) => Ok(error_response(request_id, error)), - } -} diff --git a/src/api/model.rs b/src/api/model.rs index 2bae5fe..3c7497d 100644 --- a/src/api/model.rs +++ b/src/api/model.rs @@ -466,164 +466,6 @@ pub(super) struct MinimalAllData { pub(super) data: Option, } -#[derive(Serialize)] -pub(super) struct UserLinks { - pub(super) classic: Vec, - pub(super) secure: Vec, - pub(super) tls: Vec, - pub(super) tls_domains: Vec, -} - -#[derive(Serialize)] -pub(super) struct TlsDomainLink { - pub(super) domain: String, - pub(super) link: String, -} - -#[derive(Serialize)] -pub(super) struct UserInfo { - pub(super) username: String, - pub(super) enabled: bool, - pub(super) in_runtime: bool, - pub(super) user_ad_tag: Option, - pub(super) max_tcp_conns: Option, - pub(super) expiration_rfc3339: Option, - pub(super) data_quota_bytes: Option, - pub(super) rate_limit_up_bps: Option, - pub(super) rate_limit_down_bps: Option, - pub(super) max_unique_ips: Option, - pub(super) current_connections: u64, - pub(super) active_unique_ips: usize, - pub(super) active_unique_ips_list: Vec, - pub(super) recent_unique_ips: usize, - pub(super) recent_unique_ips_list: Vec, - pub(super) total_octets: u64, - pub(super) links: UserLinks, -} - -#[derive(Serialize)] -pub(super) struct UserActiveIps { - pub(super) username: String, - pub(super) active_ips: Vec, -} - -#[derive(Serialize)] -pub(super) struct CreateUserResponse { - pub(super) user: UserInfo, - pub(super) secret: String, -} - -#[derive(Serialize)] -pub(super) struct DeleteUserResponse { - pub(super) username: String, - pub(super) in_runtime: bool, -} - -#[derive(Serialize)] -pub(super) struct ResetUserQuotaResponse { - pub(super) username: String, - pub(super) used_bytes: u64, - pub(super) last_reset_epoch_secs: u64, -} - -#[derive(Serialize)] -pub(super) struct UserQuotaListData { - pub(super) users: Vec, -} - -#[derive(Serialize)] -pub(super) struct UserQuotaEntry { - pub(super) username: String, - pub(super) data_quota_bytes: u64, - pub(super) used_bytes: u64, - pub(super) last_reset_epoch_secs: u64, -} - -#[derive(Deserialize)] -pub(super) struct CreateUserRequest { - pub(super) username: String, - pub(super) secret: Option, - pub(super) user_ad_tag: Option, - pub(super) max_tcp_conns: Option, - pub(super) expiration_rfc3339: Option, - pub(super) data_quota_bytes: Option, - pub(super) rate_limit_up_bps: Option, - pub(super) rate_limit_down_bps: Option, - pub(super) max_unique_ips: Option, - pub(super) enabled: Option, -} - -#[derive(Deserialize)] -pub(super) struct PatchUserRequest { - pub(super) secret: Option, - #[serde(default, deserialize_with = "patch_field")] - pub(super) user_ad_tag: Patch, - #[serde(default, deserialize_with = "patch_field")] - pub(super) max_tcp_conns: Patch, - #[serde(default, deserialize_with = "patch_field")] - pub(super) expiration_rfc3339: Patch, - #[serde(default, deserialize_with = "patch_field")] - pub(super) data_quota_bytes: Patch, - #[serde(default, deserialize_with = "patch_field")] - pub(super) rate_limit_up_bps: Patch, - #[serde(default, deserialize_with = "patch_field")] - pub(super) rate_limit_down_bps: Patch, - #[serde(default, deserialize_with = "patch_field")] - pub(super) max_unique_ips: Patch, - #[serde(default, deserialize_with = "patch_field")] - pub(super) enabled: Patch, -} - -#[derive(Default, Deserialize)] -pub(super) struct RotateSecretRequest { - pub(super) secret: Option, -} - -pub(super) fn parse_optional_expiration( - value: Option<&str>, -) -> Result>, ApiFailure> { - let Some(raw) = value else { - return Ok(None); - }; - let parsed = DateTime::parse_from_rfc3339(raw) - .map_err(|_| ApiFailure::bad_request("expiration_rfc3339 must be valid RFC3339"))?; - Ok(Some(parsed.with_timezone(&Utc))) -} - -pub(super) fn parse_patch_expiration( - value: &Patch, -) -> Result>, ApiFailure> { - match value { - Patch::Unchanged => Ok(Patch::Unchanged), - Patch::Remove => Ok(Patch::Remove), - Patch::Set(raw) => { - let parsed = DateTime::parse_from_rfc3339(raw) - .map_err(|_| ApiFailure::bad_request("expiration_rfc3339 must be valid RFC3339"))?; - Ok(Patch::Set(parsed.with_timezone(&Utc))) - } - } -} - -pub(super) fn is_valid_user_secret(secret: &str) -> bool { - secret.len() == 32 && secret.chars().all(|c| c.is_ascii_hexdigit()) -} - -pub(super) fn is_valid_ad_tag(tag: &str) -> bool { - tag.len() == 32 && tag.chars().all(|c| c.is_ascii_hexdigit()) -} - -pub(super) fn is_valid_username(user: &str) -> bool { - !user.is_empty() - && user.len() <= MAX_USERNAME_LEN - && user - .chars() - .all(|ch| ch.is_ascii_alphanumeric() || matches!(ch, '_' | '-' | '.')) -} - -pub(super) fn random_user_secret() -> String { - static API_SECRET_RNG: OnceLock = OnceLock::new(); - let rng = API_SECRET_RNG.get_or_init(SecureRandom::new); - let mut bytes = [0u8; 16]; - rng.fill(&mut bytes); - hex::encode(bytes) -} +// User-management request, response, and validation models. +mod users; +pub(super) use users::*; diff --git a/src/api/model/users.rs b/src/api/model/users.rs new file mode 100644 index 0000000..8d6390b --- /dev/null +++ b/src/api/model/users.rs @@ -0,0 +1,163 @@ +use super::*; + +#[derive(Serialize)] +pub(in crate::api) struct UserLinks { + pub(in crate::api) classic: Vec, + pub(in crate::api) secure: Vec, + pub(in crate::api) tls: Vec, + pub(in crate::api) tls_domains: Vec, +} + +#[derive(Serialize)] +pub(in crate::api) struct TlsDomainLink { + pub(in crate::api) domain: String, + pub(in crate::api) link: String, +} + +#[derive(Serialize)] +pub(in crate::api) struct UserInfo { + pub(in crate::api) username: String, + pub(in crate::api) enabled: bool, + pub(in crate::api) in_runtime: bool, + pub(in crate::api) user_ad_tag: Option, + pub(in crate::api) max_tcp_conns: Option, + pub(in crate::api) expiration_rfc3339: Option, + pub(in crate::api) data_quota_bytes: Option, + pub(in crate::api) rate_limit_up_bps: Option, + pub(in crate::api) rate_limit_down_bps: Option, + pub(in crate::api) max_unique_ips: Option, + pub(in crate::api) current_connections: u64, + pub(in crate::api) active_unique_ips: usize, + pub(in crate::api) active_unique_ips_list: Vec, + pub(in crate::api) recent_unique_ips: usize, + pub(in crate::api) recent_unique_ips_list: Vec, + pub(in crate::api) total_octets: u64, + pub(in crate::api) links: UserLinks, +} + +#[derive(Serialize)] +pub(in crate::api) struct UserActiveIps { + pub(in crate::api) username: String, + pub(in crate::api) active_ips: Vec, +} + +#[derive(Serialize)] +pub(in crate::api) struct CreateUserResponse { + pub(in crate::api) user: UserInfo, + pub(in crate::api) secret: String, +} + +#[derive(Serialize)] +pub(in crate::api) struct DeleteUserResponse { + pub(in crate::api) username: String, + pub(in crate::api) in_runtime: bool, +} + +#[derive(Serialize)] +pub(in crate::api) struct ResetUserQuotaResponse { + pub(in crate::api) username: String, + pub(in crate::api) used_bytes: u64, + pub(in crate::api) last_reset_epoch_secs: u64, +} + +#[derive(Serialize)] +pub(in crate::api) struct UserQuotaListData { + pub(in crate::api) users: Vec, +} + +#[derive(Serialize)] +pub(in crate::api) struct UserQuotaEntry { + pub(in crate::api) username: String, + pub(in crate::api) data_quota_bytes: u64, + pub(in crate::api) used_bytes: u64, + pub(in crate::api) last_reset_epoch_secs: u64, +} + +#[derive(Deserialize)] +pub(in crate::api) struct CreateUserRequest { + pub(in crate::api) username: String, + pub(in crate::api) secret: Option, + pub(in crate::api) user_ad_tag: Option, + pub(in crate::api) max_tcp_conns: Option, + pub(in crate::api) expiration_rfc3339: Option, + pub(in crate::api) data_quota_bytes: Option, + pub(in crate::api) rate_limit_up_bps: Option, + pub(in crate::api) rate_limit_down_bps: Option, + pub(in crate::api) max_unique_ips: Option, + pub(in crate::api) enabled: Option, +} + +#[derive(Deserialize)] +pub(in crate::api) struct PatchUserRequest { + pub(in crate::api) secret: Option, + #[serde(default, deserialize_with = "patch_field")] + pub(in crate::api) user_ad_tag: Patch, + #[serde(default, deserialize_with = "patch_field")] + pub(in crate::api) max_tcp_conns: Patch, + #[serde(default, deserialize_with = "patch_field")] + pub(in crate::api) expiration_rfc3339: Patch, + #[serde(default, deserialize_with = "patch_field")] + pub(in crate::api) data_quota_bytes: Patch, + #[serde(default, deserialize_with = "patch_field")] + pub(in crate::api) rate_limit_up_bps: Patch, + #[serde(default, deserialize_with = "patch_field")] + pub(in crate::api) rate_limit_down_bps: Patch, + #[serde(default, deserialize_with = "patch_field")] + pub(in crate::api) max_unique_ips: Patch, + #[serde(default, deserialize_with = "patch_field")] + pub(in crate::api) enabled: Patch, +} + +#[derive(Default, Deserialize)] +pub(in crate::api) struct RotateSecretRequest { + pub(in crate::api) secret: Option, +} + +pub(in crate::api) fn parse_optional_expiration( + value: Option<&str>, +) -> Result>, ApiFailure> { + let Some(raw) = value else { + return Ok(None); + }; + let parsed = DateTime::parse_from_rfc3339(raw) + .map_err(|_| ApiFailure::bad_request("expiration_rfc3339 must be valid RFC3339"))?; + Ok(Some(parsed.with_timezone(&Utc))) +} + +pub(in crate::api) fn parse_patch_expiration( + value: &Patch, +) -> Result>, ApiFailure> { + match value { + Patch::Unchanged => Ok(Patch::Unchanged), + Patch::Remove => Ok(Patch::Remove), + Patch::Set(raw) => { + let parsed = DateTime::parse_from_rfc3339(raw) + .map_err(|_| ApiFailure::bad_request("expiration_rfc3339 must be valid RFC3339"))?; + Ok(Patch::Set(parsed.with_timezone(&Utc))) + } + } +} + +pub(in crate::api) fn is_valid_user_secret(secret: &str) -> bool { + secret.len() == 32 && secret.chars().all(|c| c.is_ascii_hexdigit()) +} + +pub(in crate::api) fn is_valid_ad_tag(tag: &str) -> bool { + tag.len() == 32 && tag.chars().all(|c| c.is_ascii_hexdigit()) +} + +pub(in crate::api) fn is_valid_username(user: &str) -> bool { + !user.is_empty() + && user.len() <= MAX_USERNAME_LEN + && user + .chars() + .all(|ch| ch.is_ascii_alphanumeric() || matches!(ch, '_' | '-' | '.')) +} + +pub(in crate::api) fn random_user_secret() -> String { + static API_SECRET_RNG: OnceLock = OnceLock::new(); + let rng = API_SECRET_RNG.get_or_init(SecureRandom::new); + let mut bytes = [0u8; 16]; + rng.fill(&mut bytes); + hex::encode(bytes) +} diff --git a/src/api/runtime_min.rs b/src/api/runtime_min.rs index 794fd42..5338baf 100644 --- a/src/api/runtime_min.rs +++ b/src/api/runtime_min.rs @@ -541,55 +541,7 @@ pub(super) async fn build_runtime_upstream_quality_data( } } -pub(super) async fn build_runtime_nat_stun_data(shared: &ApiShared) -> RuntimeNatStunData { - let now_epoch_secs = now_epoch_secs(); - let Some(pool) = shared.me_pool.read().await.clone() else { - return RuntimeNatStunData { - enabled: false, - reason: Some(SOURCE_UNAVAILABLE_REASON), - generated_at_epoch_secs: now_epoch_secs, - data: None, - }; - }; - - let snapshot = pool.api_nat_stun_snapshot().await; - RuntimeNatStunData { - enabled: true, - reason: None, - generated_at_epoch_secs: now_epoch_secs, - data: Some(RuntimeNatStunPayload { - flags: RuntimeNatStunFlagsData { - nat_probe_enabled: snapshot.nat_probe_enabled, - nat_probe_disabled_runtime: snapshot.nat_probe_disabled_runtime, - nat_probe_attempts: snapshot.nat_probe_attempts, - }, - servers: RuntimeNatStunServersData { - configured: snapshot.configured_servers, - live: snapshot.live_servers.clone(), - live_total: snapshot.live_servers.len(), - }, - reflection: RuntimeNatStunReflectionBlockData { - v4: snapshot - .reflection_v4 - .map(|entry| RuntimeNatStunReflectionData { - addr: entry.addr.to_string(), - age_secs: entry.age_secs, - }), - v6: snapshot - .reflection_v6 - .map(|entry| RuntimeNatStunReflectionData { - addr: entry.addr.to_string(), - age_secs: entry.age_secs, - }), - }, - stun_backoff_remaining_ms: snapshot.stun_backoff_remaining_ms, - }), - } -} - -fn now_epoch_secs() -> u64 { - SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap_or_default() - .as_secs() -} +// NAT/STUN runtime projection and timestamping. +mod nat; +pub(super) use nat::build_runtime_nat_stun_data; +use nat::now_epoch_secs; diff --git a/src/api/runtime_min/nat.rs b/src/api/runtime_min/nat.rs new file mode 100644 index 0000000..f22b63b --- /dev/null +++ b/src/api/runtime_min/nat.rs @@ -0,0 +1,54 @@ +use super::*; + +pub(in crate::api) async fn build_runtime_nat_stun_data(shared: &ApiShared) -> RuntimeNatStunData { + let now_epoch_secs = now_epoch_secs(); + let Some(pool) = shared.me_pool.read().await.clone() else { + return RuntimeNatStunData { + enabled: false, + reason: Some(SOURCE_UNAVAILABLE_REASON), + generated_at_epoch_secs: now_epoch_secs, + data: None, + }; + }; + + let snapshot = pool.api_nat_stun_snapshot().await; + RuntimeNatStunData { + enabled: true, + reason: None, + generated_at_epoch_secs: now_epoch_secs, + data: Some(RuntimeNatStunPayload { + flags: RuntimeNatStunFlagsData { + nat_probe_enabled: snapshot.nat_probe_enabled, + nat_probe_disabled_runtime: snapshot.nat_probe_disabled_runtime, + nat_probe_attempts: snapshot.nat_probe_attempts, + }, + servers: RuntimeNatStunServersData { + configured: snapshot.configured_servers, + live: snapshot.live_servers.clone(), + live_total: snapshot.live_servers.len(), + }, + reflection: RuntimeNatStunReflectionBlockData { + v4: snapshot + .reflection_v4 + .map(|entry| RuntimeNatStunReflectionData { + addr: entry.addr.to_string(), + age_secs: entry.age_secs, + }), + v6: snapshot + .reflection_v6 + .map(|entry| RuntimeNatStunReflectionData { + addr: entry.addr.to_string(), + age_secs: entry.age_secs, + }), + }, + stun_backoff_remaining_ms: snapshot.stun_backoff_remaining_ms, + }), + } +} + +pub(super) fn now_epoch_secs() -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs() +} diff --git a/src/api/runtime_stats.rs b/src/api/runtime_stats.rs index 298e5a0..d9536f1 100644 --- a/src/api/runtime_stats.rs +++ b/src/api/runtime_stats.rs @@ -84,8 +84,7 @@ pub(super) fn build_zero_all_data(stats: &Stats, configured_users: usize) -> Zer reconnect_success_total: stats.get_me_reconnect_success(), handshake_reject_total: stats.get_me_handshake_reject_total(), handshake_error_codes, - handshake_error_code_overflow_total: stats - .get_me_handshake_error_code_overflow_total(), + handshake_error_code_overflow_total: stats.get_me_handshake_error_code_overflow_total(), reader_eof_total: stats.get_me_reader_eof_total(), idle_close_by_peer_total: stats.get_me_idle_close_by_peer_total(), route_drop_no_conn_total: stats.get_me_route_drop_no_conn(), @@ -527,57 +526,6 @@ async fn get_minimal_payload_cached( Some((generated_at_epoch_secs, payload)) } -fn disabled_me_writers(now_epoch_secs: u64, reason: &'static str) -> MeWritersData { - MeWritersData { - middle_proxy_enabled: false, - reason: Some(reason), - generated_at_epoch_secs: now_epoch_secs, - summary: MeWritersSummary { - configured_dc_groups: 0, - configured_endpoints: 0, - available_endpoints: 0, - available_pct: 0.0, - required_writers: 0, - alive_writers: 0, - coverage_pct: 0.0, - fresh_alive_writers: 0, - fresh_coverage_pct: 0.0, - }, - writers: Vec::new(), - } -} - -fn disabled_dcs(now_epoch_secs: u64, reason: &'static str) -> DcStatusData { - DcStatusData { - middle_proxy_enabled: false, - reason: Some(reason), - generated_at_epoch_secs: now_epoch_secs, - dcs: Vec::new(), - } -} - -fn map_route_kind(value: UpstreamRouteKind) -> &'static str { - match value { - UpstreamRouteKind::Direct => "direct", - UpstreamRouteKind::Socks4 => "socks4", - UpstreamRouteKind::Socks5 => "socks5", - UpstreamRouteKind::Shadowsocks => "shadowsocks", - } -} - -fn map_ip_preference(value: IpPreference) -> &'static str { - match value { - IpPreference::Unknown => "unknown", - IpPreference::PreferV6 => "prefer_v6", - IpPreference::PreferV4 => "prefer_v4", - IpPreference::BothWork => "both_work", - IpPreference::Unavailable => "unavailable", - } -} - -fn now_epoch_secs() -> u64 { - SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap_or_default() - .as_secs() -} +// Disabled-state builders and stable upstream enum mappings. +mod helpers; +use helpers::*; diff --git a/src/api/runtime_stats/helpers.rs b/src/api/runtime_stats/helpers.rs new file mode 100644 index 0000000..078be1b --- /dev/null +++ b/src/api/runtime_stats/helpers.rs @@ -0,0 +1,56 @@ +use super::*; + +pub(super) fn disabled_me_writers(now_epoch_secs: u64, reason: &'static str) -> MeWritersData { + MeWritersData { + middle_proxy_enabled: false, + reason: Some(reason), + generated_at_epoch_secs: now_epoch_secs, + summary: MeWritersSummary { + configured_dc_groups: 0, + configured_endpoints: 0, + available_endpoints: 0, + available_pct: 0.0, + required_writers: 0, + alive_writers: 0, + coverage_pct: 0.0, + fresh_alive_writers: 0, + fresh_coverage_pct: 0.0, + }, + writers: Vec::new(), + } +} + +pub(super) fn disabled_dcs(now_epoch_secs: u64, reason: &'static str) -> DcStatusData { + DcStatusData { + middle_proxy_enabled: false, + reason: Some(reason), + generated_at_epoch_secs: now_epoch_secs, + dcs: Vec::new(), + } +} + +pub(super) fn map_route_kind(value: UpstreamRouteKind) -> &'static str { + match value { + UpstreamRouteKind::Direct => "direct", + UpstreamRouteKind::Socks4 => "socks4", + UpstreamRouteKind::Socks5 => "socks5", + UpstreamRouteKind::Shadowsocks => "shadowsocks", + } +} + +pub(super) fn map_ip_preference(value: IpPreference) -> &'static str { + match value { + IpPreference::Unknown => "unknown", + IpPreference::PreferV6 => "prefer_v6", + IpPreference::PreferV4 => "prefer_v4", + IpPreference::BothWork => "both_work", + IpPreference::Unavailable => "unavailable", + } +} + +pub(super) fn now_epoch_secs() -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs() +} diff --git a/src/api/runtime_watch.rs b/src/api/runtime_watch.rs index 4788a96..307021c 100644 --- a/src/api/runtime_watch.rs +++ b/src/api/runtime_watch.rs @@ -4,8 +4,8 @@ use std::time::{SystemTime, UNIX_EPOCH}; use tokio::sync::watch; -use crate::maestro::generation::RuntimeWatchState; use crate::maestro::control_plane::ProcessControlPlane; +use crate::maestro::generation::RuntimeWatchState; use super::ApiRuntimeState; use super::events::ApiEventStore; diff --git a/src/config/defaults.rs b/src/config/defaults.rs index 3236b74..b2ad135 100644 --- a/src/config/defaults.rs +++ b/src/config/defaults.rs @@ -1,6 +1,9 @@ use ipnetwork::IpNetwork; -use serde::Deserialize; -use std::collections::HashMap; + +// Extended transport, masking, and ME default values. +mod extended; + +pub(crate) use extended::*; // Helper defaults kept private to the config module. const DEFAULT_NETWORK_IPV6: Option = Some(false); @@ -520,451 +523,3 @@ pub(crate) fn default_direct_relay_copy_buf_s2c_bytes() -> usize { pub(crate) fn default_direct_relay_buffer_budget_max_bytes() -> usize { DEFAULT_DIRECT_RELAY_BUFFER_BUDGET_MAX_BYTES } - -pub(crate) fn default_me_writer_pick_sample_size() -> u8 { - DEFAULT_ME_WRITER_PICK_SAMPLE_SIZE -} - -pub(crate) fn default_me_health_interval_ms_unhealthy() -> u64 { - DEFAULT_ME_HEALTH_INTERVAL_MS_UNHEALTHY -} - -pub(crate) fn default_me_health_interval_ms_healthy() -> u64 { - DEFAULT_ME_HEALTH_INTERVAL_MS_HEALTHY -} - -pub(crate) fn default_me_admission_poll_ms() -> u64 { - DEFAULT_ME_ADMISSION_POLL_MS -} - -pub(crate) fn default_me_warn_rate_limit_ms() -> u64 { - DEFAULT_ME_WARN_RATE_LIMIT_MS -} - -pub(crate) fn default_me_route_hybrid_max_wait_ms() -> u64 { - DEFAULT_ME_ROUTE_HYBRID_MAX_WAIT_MS -} - -pub(crate) fn default_me_route_blocking_send_timeout_ms() -> u64 { - DEFAULT_ME_ROUTE_BLOCKING_SEND_TIMEOUT_MS -} - -pub(crate) fn default_me_c2me_send_timeout_ms() -> u64 { - DEFAULT_ME_C2ME_SEND_TIMEOUT_MS -} - -pub(crate) fn default_upstream_connect_retry_attempts() -> u32 { - DEFAULT_UPSTREAM_CONNECT_RETRY_ATTEMPTS -} - -pub(crate) fn default_upstream_connect_retry_backoff_ms() -> u64 { - 100 -} - -pub(crate) fn default_upstream_unhealthy_fail_threshold() -> u32 { - DEFAULT_UPSTREAM_UNHEALTHY_FAIL_THRESHOLD -} - -pub(crate) fn default_upstream_connect_budget_ms() -> u64 { - DEFAULT_UPSTREAM_CONNECT_BUDGET_MS -} - -pub(crate) fn default_upstream_connect_failfast_hard_errors() -> bool { - false -} - -pub(crate) fn default_rpc_proxy_req_every() -> u64 { - 0 -} - -pub(crate) fn default_crypto_pending_buffer() -> usize { - 256 * 1024 -} - -pub(crate) fn default_max_client_frame() -> usize { - 16 * 1024 * 1024 -} - -pub(crate) fn default_desync_all_full() -> bool { - false -} - -pub(crate) fn default_me_route_backpressure_base_timeout_ms() -> u64 { - 25 -} - -pub(crate) fn default_me_route_backpressure_enabled() -> bool { - DEFAULT_ME_ROUTE_BACKPRESSURE_ENABLED -} - -pub(crate) fn default_me_route_fairshare_enabled() -> bool { - DEFAULT_ME_ROUTE_FAIRSHARE_ENABLED -} - -pub(crate) fn default_me_route_backpressure_high_timeout_ms() -> u64 { - 120 -} - -pub(crate) fn default_me_route_backpressure_high_watermark_pct() -> u8 { - 80 -} - -pub(crate) fn default_me_route_no_writer_wait_ms() -> u64 { - 250 -} - -pub(crate) fn default_me_route_inline_recovery_attempts() -> u32 { - 3 -} - -pub(crate) fn default_me_route_inline_recovery_wait_ms() -> u64 { - 3000 -} - -pub(crate) fn default_beobachten_minutes() -> u64 { - 10 -} - -pub(crate) fn default_beobachten_flush_secs() -> u64 { - 15 -} - -pub(crate) fn default_beobachten_file() -> String { - "beobachten.txt".to_string() -} - -pub(crate) fn default_tls_new_session_tickets() -> u8 { - 0 -} - -pub(crate) fn default_serverhello_compact() -> bool { - false -} - -pub(crate) fn default_tls_full_cert_ttl_secs() -> u64 { - 90 -} - -pub(crate) fn default_server_hello_delay_min_ms() -> u64 { - 8 -} - -pub(crate) fn default_server_hello_delay_max_ms() -> u64 { - 24 -} - -pub(crate) fn default_alpn_enforce() -> bool { - true -} - -pub(crate) fn default_mask_shape_hardening() -> bool { - true -} - -pub(crate) fn default_mask_shape_hardening_aggressive_mode() -> bool { - false -} - -pub(crate) fn default_mask_shape_bucket_floor_bytes() -> usize { - 512 -} - -pub(crate) fn default_mask_shape_bucket_cap_bytes() -> usize { - 4096 -} - -pub(crate) fn default_mask_shape_above_cap_blur() -> bool { - false -} - -pub(crate) fn default_mask_shape_above_cap_blur_max_bytes() -> usize { - 512 -} - -#[cfg(not(test))] -pub(crate) fn default_mask_relay_max_bytes() -> usize { - 5 * 1024 * 1024 -} - -#[cfg(test)] -pub(crate) fn default_mask_relay_max_bytes() -> usize { - 32 * 1024 -} - -#[cfg(not(test))] -pub(crate) fn default_mask_relay_timeout_ms() -> u64 { - 60_000 -} - -#[cfg(test)] -pub(crate) fn default_mask_relay_timeout_ms() -> u64 { - 200 -} - -#[cfg(not(test))] -pub(crate) fn default_mask_relay_idle_timeout_ms() -> u64 { - 5_000 -} - -#[cfg(test)] -pub(crate) fn default_mask_relay_idle_timeout_ms() -> u64 { - 100 -} - -pub(crate) fn default_mask_classifier_prefetch_timeout_ms() -> u64 { - 5 -} - -pub(crate) fn default_mask_timing_normalization_enabled() -> bool { - false -} - -pub(crate) fn default_mask_timing_normalization_floor_ms() -> u64 { - 0 -} - -pub(crate) fn default_mask_timing_normalization_ceiling_ms() -> u64 { - 0 -} - -pub(crate) fn default_stun_servers() -> Vec { - vec![ - "stun.l.google.com:5349".to_string(), - "stun1.l.google.com:3478".to_string(), - "stun.gmx.net:3478".to_string(), - "stun.l.google.com:19302".to_string(), - "stun.1und1.de:3478".to_string(), - "stun1.l.google.com:19302".to_string(), - "stun2.l.google.com:19302".to_string(), - "stun3.l.google.com:19302".to_string(), - "stun4.l.google.com:19302".to_string(), - "stun.services.mozilla.com:3478".to_string(), - "stun.stunprotocol.org:3478".to_string(), - "stun.nextcloud.com:3478".to_string(), - "stun.voip.eutelia.it:3478".to_string(), - ] -} - -pub(crate) fn default_http_ip_detect_urls() -> Vec { - vec![ - "https://ifconfig.me/ip".to_string(), - "https://api.ipify.org".to_string(), - ] -} - -pub(crate) fn default_cache_public_ip_path() -> String { - "cache/public_ip.txt".to_string() -} - -pub(crate) fn default_proxy_secret_reload_secs() -> u64 { - 60 * 60 -} - -pub(crate) fn default_proxy_config_reload_secs() -> u64 { - 60 * 60 -} - -pub(crate) fn default_update_every_secs() -> u64 { - 5 * 60 -} - -pub(crate) fn default_update_every() -> Option { - Some(default_update_every_secs()) -} - -pub(crate) fn default_me_reinit_every_secs() -> u64 { - 15 * 60 -} - -pub(crate) fn default_me_reinit_singleflight() -> bool { - true -} - -pub(crate) fn default_me_reinit_max_concurrency() -> usize { - 2 -} - -pub(crate) fn default_me_reinit_trigger_channel() -> usize { - 64 -} - -pub(crate) fn default_me_reinit_coalesce_window_ms() -> u64 { - 200 -} - -pub(crate) fn default_me_hardswap_warmup_delay_min_ms() -> u64 { - 1000 -} - -pub(crate) fn default_me_hardswap_warmup_delay_max_ms() -> u64 { - 2000 -} - -pub(crate) fn default_me_hardswap_warmup_extra_passes() -> u8 { - 3 -} - -pub(crate) fn default_me_hardswap_warmup_pass_backoff_base_ms() -> u64 { - 500 -} - -pub(crate) fn default_me_config_stable_snapshots() -> u8 { - 2 -} - -pub(crate) fn default_me_config_apply_cooldown_secs() -> u64 { - 300 -} - -pub(crate) fn default_me_snapshot_require_http_2xx() -> bool { - true -} - -pub(crate) fn default_me_snapshot_reject_empty_map() -> bool { - true -} - -pub(crate) fn default_me_snapshot_min_proxy_for_lines() -> u32 { - 1 -} - -pub(crate) fn default_proxy_secret_stable_snapshots() -> u8 { - 2 -} - -pub(crate) fn default_proxy_secret_rotate_runtime() -> bool { - true -} - -pub(crate) fn default_me_secret_atomic_snapshot() -> bool { - true -} - -pub(crate) fn default_proxy_secret_len_max() -> usize { - 256 -} - -pub(crate) fn default_me_reinit_drain_timeout_secs() -> u64 { - 90 -} - -pub(crate) fn default_me_pool_drain_ttl_secs() -> u64 { - 90 -} - -pub(crate) fn default_me_instadrain() -> bool { - false -} - -pub(crate) fn default_me_pool_drain_threshold() -> u64 { - 32 -} - -pub(crate) fn default_me_pool_drain_soft_evict_enabled() -> bool { - DEFAULT_ME_POOL_DRAIN_SOFT_EVICT_ENABLED -} - -pub(crate) fn default_me_pool_drain_soft_evict_grace_secs() -> u64 { - DEFAULT_ME_POOL_DRAIN_SOFT_EVICT_GRACE_SECS -} - -pub(crate) fn default_me_pool_drain_soft_evict_per_writer() -> u8 { - DEFAULT_ME_POOL_DRAIN_SOFT_EVICT_PER_WRITER -} - -pub(crate) fn default_me_pool_drain_soft_evict_budget_per_core() -> u16 { - DEFAULT_ME_POOL_DRAIN_SOFT_EVICT_BUDGET_PER_CORE -} - -pub(crate) fn default_me_pool_drain_soft_evict_cooldown_ms() -> u64 { - DEFAULT_ME_POOL_DRAIN_SOFT_EVICT_COOLDOWN_MS -} - -pub(crate) fn default_me_bind_stale_ttl_secs() -> u64 { - default_me_pool_drain_ttl_secs() -} - -pub(crate) fn default_me_pool_min_fresh_ratio() -> f32 { - 0.8 -} - -pub(crate) fn default_me_deterministic_writer_sort() -> bool { - true -} - -pub(crate) fn default_hardswap() -> bool { - true -} - -pub(crate) fn default_ntp_check() -> bool { - true -} - -pub(crate) fn default_ntp_servers() -> Vec { - vec!["pool.ntp.org".to_string()] -} - -pub(crate) fn default_fast_mode_min_tls_record() -> usize { - 0 -} - -pub(crate) fn default_degradation_min_unavailable_dc_groups() -> u8 { - 2 -} - -pub(crate) fn default_listen_addr_ipv6() -> String { - DEFAULT_LISTEN_ADDR_IPV6.to_string() -} - -pub(crate) fn default_listen_addr_ipv6_opt() -> Option { - Some(default_listen_addr_ipv6()) -} - -pub(crate) fn default_access_users() -> HashMap { - HashMap::from([( - DEFAULT_ACCESS_USER.to_string(), - DEFAULT_ACCESS_SECRET.to_string(), - )]) -} - -pub(crate) fn default_user_max_unique_ips_window_secs() -> u64 { - DEFAULT_USER_MAX_UNIQUE_IPS_WINDOW_SECS -} - -pub(crate) fn default_user_max_tcp_conns_global_each() -> usize { - 0 -} - -pub(crate) fn default_user_max_unique_ips_global_each() -> usize { - 0 -} - -// Custom deserializer helpers - -#[derive(Deserialize)] -#[serde(untagged)] -pub(crate) enum OneOrMany { - One(String), - Many(Vec), -} - -pub(crate) fn deserialize_dc_overrides<'de, D>( - deserializer: D, -) -> std::result::Result>, D::Error> -where - D: serde::de::Deserializer<'de>, -{ - let raw: HashMap = HashMap::deserialize(deserializer)?; - let mut out = HashMap::new(); - for (dc, val) in raw { - let mut addrs = match val { - OneOrMany::One(s) => vec![s], - OneOrMany::Many(v) => v, - }; - addrs.retain(|s| !s.trim().is_empty()); - if !addrs.is_empty() { - out.insert(dc, addrs); - } - } - Ok(out) -} diff --git a/src/config/defaults/extended.rs b/src/config/defaults/extended.rs new file mode 100644 index 0000000..9246605 --- /dev/null +++ b/src/config/defaults/extended.rs @@ -0,0 +1,453 @@ +use std::collections::HashMap; + +use serde::Deserialize; + +use super::*; + +pub(crate) fn default_me_writer_pick_sample_size() -> u8 { + DEFAULT_ME_WRITER_PICK_SAMPLE_SIZE +} + +pub(crate) fn default_me_health_interval_ms_unhealthy() -> u64 { + DEFAULT_ME_HEALTH_INTERVAL_MS_UNHEALTHY +} + +pub(crate) fn default_me_health_interval_ms_healthy() -> u64 { + DEFAULT_ME_HEALTH_INTERVAL_MS_HEALTHY +} + +pub(crate) fn default_me_admission_poll_ms() -> u64 { + DEFAULT_ME_ADMISSION_POLL_MS +} + +pub(crate) fn default_me_warn_rate_limit_ms() -> u64 { + DEFAULT_ME_WARN_RATE_LIMIT_MS +} + +pub(crate) fn default_me_route_hybrid_max_wait_ms() -> u64 { + DEFAULT_ME_ROUTE_HYBRID_MAX_WAIT_MS +} + +pub(crate) fn default_me_route_blocking_send_timeout_ms() -> u64 { + DEFAULT_ME_ROUTE_BLOCKING_SEND_TIMEOUT_MS +} + +pub(crate) fn default_me_c2me_send_timeout_ms() -> u64 { + DEFAULT_ME_C2ME_SEND_TIMEOUT_MS +} + +pub(crate) fn default_upstream_connect_retry_attempts() -> u32 { + DEFAULT_UPSTREAM_CONNECT_RETRY_ATTEMPTS +} + +pub(crate) fn default_upstream_connect_retry_backoff_ms() -> u64 { + 100 +} + +pub(crate) fn default_upstream_unhealthy_fail_threshold() -> u32 { + DEFAULT_UPSTREAM_UNHEALTHY_FAIL_THRESHOLD +} + +pub(crate) fn default_upstream_connect_budget_ms() -> u64 { + DEFAULT_UPSTREAM_CONNECT_BUDGET_MS +} + +pub(crate) fn default_upstream_connect_failfast_hard_errors() -> bool { + false +} + +pub(crate) fn default_rpc_proxy_req_every() -> u64 { + 0 +} + +pub(crate) fn default_crypto_pending_buffer() -> usize { + 256 * 1024 +} + +pub(crate) fn default_max_client_frame() -> usize { + 16 * 1024 * 1024 +} + +pub(crate) fn default_desync_all_full() -> bool { + false +} + +pub(crate) fn default_me_route_backpressure_base_timeout_ms() -> u64 { + 25 +} + +pub(crate) fn default_me_route_backpressure_enabled() -> bool { + DEFAULT_ME_ROUTE_BACKPRESSURE_ENABLED +} + +pub(crate) fn default_me_route_fairshare_enabled() -> bool { + DEFAULT_ME_ROUTE_FAIRSHARE_ENABLED +} + +pub(crate) fn default_me_route_backpressure_high_timeout_ms() -> u64 { + 120 +} + +pub(crate) fn default_me_route_backpressure_high_watermark_pct() -> u8 { + 80 +} + +pub(crate) fn default_me_route_no_writer_wait_ms() -> u64 { + 250 +} + +pub(crate) fn default_me_route_inline_recovery_attempts() -> u32 { + 3 +} + +pub(crate) fn default_me_route_inline_recovery_wait_ms() -> u64 { + 3000 +} + +pub(crate) fn default_beobachten_minutes() -> u64 { + 10 +} + +pub(crate) fn default_beobachten_flush_secs() -> u64 { + 15 +} + +pub(crate) fn default_beobachten_file() -> String { + "beobachten.txt".to_string() +} + +pub(crate) fn default_tls_new_session_tickets() -> u8 { + 0 +} + +pub(crate) fn default_serverhello_compact() -> bool { + false +} + +pub(crate) fn default_tls_full_cert_ttl_secs() -> u64 { + 90 +} + +pub(crate) fn default_server_hello_delay_min_ms() -> u64 { + 8 +} + +pub(crate) fn default_server_hello_delay_max_ms() -> u64 { + 24 +} + +pub(crate) fn default_alpn_enforce() -> bool { + true +} + +pub(crate) fn default_mask_shape_hardening() -> bool { + true +} + +pub(crate) fn default_mask_shape_hardening_aggressive_mode() -> bool { + false +} + +pub(crate) fn default_mask_shape_bucket_floor_bytes() -> usize { + 512 +} + +pub(crate) fn default_mask_shape_bucket_cap_bytes() -> usize { + 4096 +} + +pub(crate) fn default_mask_shape_above_cap_blur() -> bool { + false +} + +pub(crate) fn default_mask_shape_above_cap_blur_max_bytes() -> usize { + 512 +} + +#[cfg(not(test))] +pub(crate) fn default_mask_relay_max_bytes() -> usize { + 5 * 1024 * 1024 +} + +#[cfg(test)] +pub(crate) fn default_mask_relay_max_bytes() -> usize { + 32 * 1024 +} + +#[cfg(not(test))] +pub(crate) fn default_mask_relay_timeout_ms() -> u64 { + 60_000 +} + +#[cfg(test)] +pub(crate) fn default_mask_relay_timeout_ms() -> u64 { + 200 +} + +#[cfg(not(test))] +pub(crate) fn default_mask_relay_idle_timeout_ms() -> u64 { + 5_000 +} + +#[cfg(test)] +pub(crate) fn default_mask_relay_idle_timeout_ms() -> u64 { + 100 +} + +pub(crate) fn default_mask_classifier_prefetch_timeout_ms() -> u64 { + 5 +} + +pub(crate) fn default_mask_timing_normalization_enabled() -> bool { + false +} + +pub(crate) fn default_mask_timing_normalization_floor_ms() -> u64 { + 0 +} + +pub(crate) fn default_mask_timing_normalization_ceiling_ms() -> u64 { + 0 +} + +pub(crate) fn default_stun_servers() -> Vec { + vec![ + "stun.l.google.com:5349".to_string(), + "stun1.l.google.com:3478".to_string(), + "stun.gmx.net:3478".to_string(), + "stun.l.google.com:19302".to_string(), + "stun.1und1.de:3478".to_string(), + "stun1.l.google.com:19302".to_string(), + "stun2.l.google.com:19302".to_string(), + "stun3.l.google.com:19302".to_string(), + "stun4.l.google.com:19302".to_string(), + "stun.services.mozilla.com:3478".to_string(), + "stun.stunprotocol.org:3478".to_string(), + "stun.nextcloud.com:3478".to_string(), + "stun.voip.eutelia.it:3478".to_string(), + ] +} + +pub(crate) fn default_http_ip_detect_urls() -> Vec { + vec![ + "https://ifconfig.me/ip".to_string(), + "https://api.ipify.org".to_string(), + ] +} + +pub(crate) fn default_cache_public_ip_path() -> String { + "cache/public_ip.txt".to_string() +} + +pub(crate) fn default_proxy_secret_reload_secs() -> u64 { + 60 * 60 +} + +pub(crate) fn default_proxy_config_reload_secs() -> u64 { + 60 * 60 +} + +pub(crate) fn default_update_every_secs() -> u64 { + 5 * 60 +} + +pub(crate) fn default_update_every() -> Option { + Some(default_update_every_secs()) +} + +pub(crate) fn default_me_reinit_every_secs() -> u64 { + 15 * 60 +} + +pub(crate) fn default_me_reinit_singleflight() -> bool { + true +} + +pub(crate) fn default_me_reinit_max_concurrency() -> usize { + 2 +} + +pub(crate) fn default_me_reinit_trigger_channel() -> usize { + 64 +} + +pub(crate) fn default_me_reinit_coalesce_window_ms() -> u64 { + 200 +} + +pub(crate) fn default_me_hardswap_warmup_delay_min_ms() -> u64 { + 1000 +} + +pub(crate) fn default_me_hardswap_warmup_delay_max_ms() -> u64 { + 2000 +} + +pub(crate) fn default_me_hardswap_warmup_extra_passes() -> u8 { + 3 +} + +pub(crate) fn default_me_hardswap_warmup_pass_backoff_base_ms() -> u64 { + 500 +} + +pub(crate) fn default_me_config_stable_snapshots() -> u8 { + 2 +} + +pub(crate) fn default_me_config_apply_cooldown_secs() -> u64 { + 300 +} + +pub(crate) fn default_me_snapshot_require_http_2xx() -> bool { + true +} + +pub(crate) fn default_me_snapshot_reject_empty_map() -> bool { + true +} + +pub(crate) fn default_me_snapshot_min_proxy_for_lines() -> u32 { + 1 +} + +pub(crate) fn default_proxy_secret_stable_snapshots() -> u8 { + 2 +} + +pub(crate) fn default_proxy_secret_rotate_runtime() -> bool { + true +} + +pub(crate) fn default_me_secret_atomic_snapshot() -> bool { + true +} + +pub(crate) fn default_proxy_secret_len_max() -> usize { + 256 +} + +pub(crate) fn default_me_reinit_drain_timeout_secs() -> u64 { + 90 +} + +pub(crate) fn default_me_pool_drain_ttl_secs() -> u64 { + 90 +} + +pub(crate) fn default_me_instadrain() -> bool { + false +} + +pub(crate) fn default_me_pool_drain_threshold() -> u64 { + 32 +} + +pub(crate) fn default_me_pool_drain_soft_evict_enabled() -> bool { + DEFAULT_ME_POOL_DRAIN_SOFT_EVICT_ENABLED +} + +pub(crate) fn default_me_pool_drain_soft_evict_grace_secs() -> u64 { + DEFAULT_ME_POOL_DRAIN_SOFT_EVICT_GRACE_SECS +} + +pub(crate) fn default_me_pool_drain_soft_evict_per_writer() -> u8 { + DEFAULT_ME_POOL_DRAIN_SOFT_EVICT_PER_WRITER +} + +pub(crate) fn default_me_pool_drain_soft_evict_budget_per_core() -> u16 { + DEFAULT_ME_POOL_DRAIN_SOFT_EVICT_BUDGET_PER_CORE +} + +pub(crate) fn default_me_pool_drain_soft_evict_cooldown_ms() -> u64 { + DEFAULT_ME_POOL_DRAIN_SOFT_EVICT_COOLDOWN_MS +} + +pub(crate) fn default_me_bind_stale_ttl_secs() -> u64 { + default_me_pool_drain_ttl_secs() +} + +pub(crate) fn default_me_pool_min_fresh_ratio() -> f32 { + 0.8 +} + +pub(crate) fn default_me_deterministic_writer_sort() -> bool { + true +} + +pub(crate) fn default_hardswap() -> bool { + true +} + +pub(crate) fn default_ntp_check() -> bool { + true +} + +pub(crate) fn default_ntp_servers() -> Vec { + vec!["pool.ntp.org".to_string()] +} + +pub(crate) fn default_fast_mode_min_tls_record() -> usize { + 0 +} + +pub(crate) fn default_degradation_min_unavailable_dc_groups() -> u8 { + 2 +} + +pub(crate) fn default_listen_addr_ipv6() -> String { + DEFAULT_LISTEN_ADDR_IPV6.to_string() +} + +pub(crate) fn default_listen_addr_ipv6_opt() -> Option { + Some(default_listen_addr_ipv6()) +} + +pub(crate) fn default_access_users() -> HashMap { + HashMap::from([( + DEFAULT_ACCESS_USER.to_string(), + DEFAULT_ACCESS_SECRET.to_string(), + )]) +} + +pub(crate) fn default_user_max_unique_ips_window_secs() -> u64 { + DEFAULT_USER_MAX_UNIQUE_IPS_WINDOW_SECS +} + +pub(crate) fn default_user_max_tcp_conns_global_each() -> usize { + 0 +} + +pub(crate) fn default_user_max_unique_ips_global_each() -> usize { + 0 +} + +// Custom deserializer helpers + +#[derive(Deserialize)] +#[serde(untagged)] +pub(crate) enum OneOrMany { + One(String), + Many(Vec), +} + +pub(crate) fn deserialize_dc_overrides<'de, D>( + deserializer: D, +) -> std::result::Result>, D::Error> +where + D: serde::de::Deserializer<'de>, +{ + let raw: HashMap = HashMap::deserialize(deserializer)?; + let mut out = HashMap::new(); + for (dc, val) in raw { + let mut addrs = match val { + OneOrMany::One(s) => vec![s], + OneOrMany::Many(v) => v, + }; + addrs.retain(|s| !s.trim().is_empty()); + if !addrs.is_empty() { + out.insert(dc, addrs); + } + } + Ok(out) +} diff --git a/src/config/load/validate_web.rs b/src/config/load/validate_web.rs index 51b3433..76a51ac 100644 --- a/src/config/load/validate_web.rs +++ b/src/config/load/validate_web.rs @@ -414,196 +414,9 @@ fn validate_limits(limits: &WebLimitsConfig) -> Result<()> { 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 profile.user.is_empty() || profile.user.len() > 64 { - return config_error(&format!( - "web.vhosts[{vhost_idx}].profiles[{profile_idx}].user must contain 1..64 bytes" - )); - } - 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 { - 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::().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(message: &str) -> Result { - Err(ProxyError::Config(message.to_string())) -} +// Virtual-host, hostname, and decoy validation. +mod vhosts; +use vhosts::*; #[cfg(test)] mod tests; diff --git a/src/config/load/validate_web/vhosts.rs b/src/config/load/validate_web/vhosts.rs new file mode 100644 index 0000000..ec82852 --- /dev/null +++ b/src/config/load/validate_web/vhosts.rs @@ -0,0 +1,192 @@ +use super::*; + +pub(super) 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 profile.user.is_empty() || profile.user.len() > 64 { + return config_error(&format!( + "web.vhosts[{vhost_idx}].profiles[{profile_idx}].user must contain 1..64 bytes" + )); + } + 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(()) +} + +pub(super) fn normalize_web_host(value: &str, field: &str) -> Result { + 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::().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) +} + +pub(super) 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()) +} + +pub(super) 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(()) +} + +pub(super) fn config_error(message: &str) -> Result { + Err(ProxyError::Config(message.to_string())) +} diff --git a/src/config/types/web.rs b/src/config/types/web.rs index 857d579..ade2523 100644 --- a/src/config/types/web.rs +++ b/src/config/types/web.rs @@ -471,84 +471,9 @@ impl Default for WebConfig { } } -/// 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>, - /// Flat profile inventory used by startup link emission. - pub(crate) profiles: Vec>, -} - -/// 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>, -} - -/// 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, - /// Sole carrier or final fallback frozen into the issued bridge policy. - pub(crate) carrier: WebCarrier, - /// Whether an explicit carrier list enabled automatic negotiation. - pub(crate) carrier_negotiation_enabled: bool, - /// Whether automatic outcomes consult and update process-local evidence. - pub(crate) carrier_learning: bool, - /// Ordered negotiation candidates including the fallback carrier exactly once. - pub(crate) carriers: Arc<[WebCarrier]>, - /// Cumulative carrier-attempt deadlines frozen when the bridge is issued. - pub(crate) carrier_negotiation_deadlines_secs: [u64; 4], - /// HMAC-derived bridge capability. - pub(crate) capability: [u8; 32], - /// Non-secret domain-separated client-secret fingerprint for debugging. - pub(crate) key_fingerprint: String, - /// 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), -} - -/// Immutable bounded static-site snapshot. -#[derive(Debug)] -pub(crate) struct WebStaticSite { - /// Canonical URL-path to immutable response asset mapping. - pub(crate) assets: BTreeMap, - /// 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, -} +// Immutable runtime WEB configuration consumed by hot paths. +mod runtime; +pub(crate) use runtime::{ + WebRuntimeConfig, WebRuntimeDecoy, WebRuntimeProfile, WebRuntimeVhost, WebStaticAsset, + WebStaticSite, +}; diff --git a/src/config/types/web/runtime.rs b/src/config/types/web/runtime.rs new file mode 100644 index 0000000..c59e8b7 --- /dev/null +++ b/src/config/types/web/runtime.rs @@ -0,0 +1,83 @@ +use super::*; + +/// 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>, + /// Flat profile inventory used by startup link emission. + pub(crate) profiles: Vec>, +} + +/// 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>, +} + +/// 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, + /// Sole carrier or final fallback frozen into the issued bridge policy. + pub(crate) carrier: WebCarrier, + /// Whether an explicit carrier list enabled automatic negotiation. + pub(crate) carrier_negotiation_enabled: bool, + /// Whether automatic outcomes consult and update process-local evidence. + pub(crate) carrier_learning: bool, + /// Ordered negotiation candidates including the fallback carrier exactly once. + pub(crate) carriers: Arc<[WebCarrier]>, + /// Cumulative carrier-attempt deadlines frozen when the bridge is issued. + pub(crate) carrier_negotiation_deadlines_secs: [u64; 4], + /// HMAC-derived bridge capability. + pub(crate) capability: [u8; 32], + /// Non-secret domain-separated client-secret fingerprint for debugging. + pub(crate) key_fingerprint: String, + /// 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), +} + +/// Immutable bounded static-site snapshot. +#[derive(Debug)] +pub(crate) struct WebStaticSite { + /// Canonical URL-path to immutable response asset mapping. + pub(crate) assets: BTreeMap, + /// 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, +} diff --git a/src/healthcheck.rs b/src/healthcheck.rs index ddbced1..811fc73 100644 --- a/src/healthcheck.rs +++ b/src/healthcheck.rs @@ -192,8 +192,8 @@ fn validate_payload(mode: HealthcheckMode, body: &str) -> Result<(), String> { #[cfg(test)] mod tests { use super::{ - HEALTHCHECK_RESPONSE_MAX_BYTES, HealthcheckMode, parse_status_code, - read_response_bounded, split_response, validate_payload, + HEALTHCHECK_RESPONSE_MAX_BYTES, HealthcheckMode, parse_status_code, read_response_bounded, + split_response, validate_payload, }; #[test] diff --git a/src/maestro/control_plane.rs b/src/maestro/control_plane.rs index a3100fb..89cda3a 100644 --- a/src/maestro/control_plane.rs +++ b/src/maestro/control_plane.rs @@ -129,12 +129,10 @@ impl ProcessControlPlane { if self.inner.shutdown_completed.load(Ordering::Acquire) { return true; } - let registrations_stopped = tokio::time::timeout_at( - deadline, - self.inner.admission.wait_for_registrations(), - ) - .await - .is_ok(); + let registrations_stopped = + tokio::time::timeout_at(deadline, self.inner.admission.wait_for_registrations()) + .await + .is_ok(); let tasks_stopped = tokio::time::timeout_at(deadline, self.inner.tasks.wait()) .await .is_ok(); @@ -188,7 +186,8 @@ mod tests { let first = tokio::spawn(async move { first_scope.shutdown(Duration::from_secs(1)).await }); tokio::task::yield_now().await; let second_scope = scope.clone(); - let second = tokio::spawn(async move { second_scope.shutdown(Duration::from_secs(1)).await }); + let second = + tokio::spawn(async move { second_scope.shutdown(Duration::from_secs(1)).await }); tokio::task::yield_now().await; assert!(!first.is_finished()); @@ -204,7 +203,8 @@ mod tests { let scope = ProcessControlPlane::new(); let registration = scope.inner.admission.try_register().unwrap(); let first_scope = scope.clone(); - let first = tokio::spawn(async move { first_scope.shutdown(Duration::from_secs(30)).await }); + let first = + tokio::spawn(async move { first_scope.shutdown(Duration::from_secs(30)).await }); tokio::task::yield_now().await; first.abort(); diff --git a/src/maestro/listeners/accept.rs b/src/maestro/listeners/accept.rs index c847c77..395233f 100644 --- a/src/maestro/listeners/accept.rs +++ b/src/maestro/listeners/accept.rs @@ -213,11 +213,9 @@ async fn run_accept_loop( continue; } if web_runtime.is_shutdown() { - web_runtime - .telemetry() - .record_rejection( - crate::web::telemetry::WebRejectionReason::RuntimeClosed, - ); + web_runtime.telemetry().record_rejection( + crate::web::telemetry::WebRejectionReason::RuntimeClosed, + ); drop(stream); continue; } @@ -233,9 +231,8 @@ async fn run_accept_loop( Err(HttpConnectionAdmissionError::AtCapacity) => { let config = web_runtime.active_generation().config(); let action = config.web.http_connection_capacity_action; - let phase_timeout = Duration::from_millis( - config.web.timeouts.http_overload_timeout_ms, - ); + let phase_timeout = + Duration::from_millis(config.web.timeouts.http_overload_timeout_ms); drop(config); if action == crate::config::WebHttpConnectionCapacityAction::Drop { web_runtime.telemetry().record_rejection( @@ -247,30 +244,29 @@ async fn run_accept_loop( drop(stream); continue; } - let overload_permit = - match web_runtime.try_http_overload_connection() { - Ok(permit) => permit, - Err(HttpConnectionAdmissionError::Closed) => { - web_runtime.telemetry().record_rejection( - crate::web::telemetry::WebRejectionReason::RuntimeClosed, - ); - web_runtime.telemetry().record_overload( - WebHttpConnectionOverloadOutcome::ShutdownDrop, - ); - drop(stream); - continue; - } - Err(HttpConnectionAdmissionError::AtCapacity) => { - web_runtime.telemetry().record_rejection( + let overload_permit = match web_runtime.try_http_overload_connection() { + Ok(permit) => permit, + Err(HttpConnectionAdmissionError::Closed) => { + web_runtime.telemetry().record_rejection( + crate::web::telemetry::WebRejectionReason::RuntimeClosed, + ); + web_runtime.telemetry().record_overload( + WebHttpConnectionOverloadOutcome::ShutdownDrop, + ); + drop(stream); + continue; + } + Err(HttpConnectionAdmissionError::AtCapacity) => { + web_runtime.telemetry().record_rejection( crate::web::telemetry::WebRejectionReason::HttpConnectionCapacity, ); - web_runtime.telemetry().record_overload( - WebHttpConnectionOverloadOutcome::OverflowCapacityDrop, - ); - drop(stream); - continue; - } - }; + web_runtime.telemetry().record_overload( + WebHttpConnectionOverloadOutcome::OverflowCapacityDrop, + ); + drop(stream); + continue; + } + }; connections.spawn(web_overload::serve( stream, peer_addr, diff --git a/src/maestro/me_startup.rs b/src/maestro/me_startup.rs index e96f8ec..4fe5273 100644 --- a/src/maestro/me_startup.rs +++ b/src/maestro/me_startup.rs @@ -22,57 +22,11 @@ use crate::transport::middle_proxy::MePool; use super::generation::RuntimeTaskScope; use super::helpers::load_startup_proxy_config_snapshot; -async fn supervise_me_task(task_name: &'static str, mut task: F) -where - F: FnMut() -> Fut, - Fut: Future + Send + 'static, -{ - loop { - let result = AbortOnDropHandle::new(tokio::spawn(task())).await; - match result { - Ok(()) => warn!( - task = task_name, - "Middle-End supervisor task exited unexpectedly, restarting" - ), - Err(error) => { - error!(task = task_name, error = %error, "Middle-End supervisor task panicked, restarting in 1s"); - tokio::time::sleep(Duration::from_secs(1)).await; - } - } - } -} - -fn spawn_me_supervisors( - task_scope: RuntimeTaskScope, - pool: Arc, - rng: Arc, - min_connections: usize, -) { - let health_pool = pool.clone(); - let health_rng = rng; - task_scope.spawn(supervise_me_task("health_monitor", move || { - let pool = health_pool.clone(); - let rng = health_rng.clone(); - async move { - crate::transport::middle_proxy::me_health_monitor(pool, rng, min_connections).await; - } - })); - - let drain_pool = pool.clone(); - task_scope.spawn(supervise_me_task("drain_timeout_enforcer", move || { - let pool = drain_pool.clone(); - async move { - crate::transport::middle_proxy::me_drain_timeout_enforcer(pool).await; - } - })); - - task_scope.spawn(supervise_me_task("zombie_writer_watchdog", move || { - let pool = pool.clone(); - async move { - crate::transport::middle_proxy::me_zombie_writer_watchdog(pool).await; - } - })); -} +// Restarting supervisors for long-lived ME maintenance tasks. +mod supervisor; +use supervisor::spawn_me_supervisors; +#[cfg(test)] +use supervisor::supervise_me_task; pub(crate) async fn initialize_me_pool( use_middle_proxy: bool, @@ -587,67 +541,4 @@ pub(crate) async fn initialize_me_pool( } #[cfg(test)] -mod tests { - use super::*; - use std::sync::atomic::{AtomicUsize, Ordering}; - use tokio::sync::Notify; - - struct DropSignal(Arc); - - impl Drop for DropSignal { - fn drop(&mut self) { - self.0.notify_one(); - } - } - - #[tokio::test] - async fn scoped_supervisor_aborts_its_current_child() { - let scope = RuntimeTaskScope::new(); - let dropped = Arc::new(Notify::new()); - let dropped_for_task = dropped.clone(); - scope.spawn(supervise_me_task("test", move || { - let dropped = dropped_for_task.clone(); - async move { - let _signal = DropSignal(dropped); - std::future::pending::<()>().await; - } - })); - tokio::task::yield_now().await; - - scope.stop().await; - - tokio::time::timeout(Duration::from_secs(1), dropped.notified()) - .await - .unwrap(); - } - - #[tokio::test] - async fn supervisor_restarts_exited_child_and_stops_with_runtime_scope() { - let scope = RuntimeTaskScope::new(); - let starts = Arc::new(AtomicUsize::new(0)); - let restarted = Arc::new(Notify::new()); - let starts_task = starts.clone(); - let restarted_task = restarted.clone(); - scope.spawn(supervise_me_task("restart_test", move || { - let starts = starts_task.clone(); - let restarted = restarted_task.clone(); - async move { - if starts.fetch_add(1, Ordering::AcqRel) + 1 >= 3 { - restarted.notify_one(); - } - } - })); - - tokio::time::timeout(Duration::from_secs(1), restarted.notified()) - .await - .unwrap(); - scope.stop().await; - let stopped_at = starts.load(Ordering::Acquire); - for _ in 0..100 { - tokio::task::yield_now().await; - } - - assert!(stopped_at >= 3); - assert_eq!(starts.load(Ordering::Acquire), stopped_at); - } -} +mod tests; diff --git a/src/maestro/me_startup/supervisor.rs b/src/maestro/me_startup/supervisor.rs new file mode 100644 index 0000000..49d84e8 --- /dev/null +++ b/src/maestro/me_startup/supervisor.rs @@ -0,0 +1,53 @@ +use super::*; + +pub(super) async fn supervise_me_task(task_name: &'static str, mut task: F) +where + F: FnMut() -> Fut, + Fut: Future + Send + 'static, +{ + loop { + let result = AbortOnDropHandle::new(tokio::spawn(task())).await; + match result { + Ok(()) => warn!( + task = task_name, + "Middle-End supervisor task exited unexpectedly, restarting" + ), + Err(error) => { + error!(task = task_name, error = %error, "Middle-End supervisor task panicked, restarting in 1s"); + tokio::time::sleep(Duration::from_secs(1)).await; + } + } + } +} + +pub(super) fn spawn_me_supervisors( + task_scope: RuntimeTaskScope, + pool: Arc, + rng: Arc, + min_connections: usize, +) { + let health_pool = pool.clone(); + let health_rng = rng; + task_scope.spawn(supervise_me_task("health_monitor", move || { + let pool = health_pool.clone(); + let rng = health_rng.clone(); + async move { + crate::transport::middle_proxy::me_health_monitor(pool, rng, min_connections).await; + } + })); + + let drain_pool = pool.clone(); + task_scope.spawn(supervise_me_task("drain_timeout_enforcer", move || { + let pool = drain_pool.clone(); + async move { + crate::transport::middle_proxy::me_drain_timeout_enforcer(pool).await; + } + })); + + task_scope.spawn(supervise_me_task("zombie_writer_watchdog", move || { + let pool = pool.clone(); + async move { + crate::transport::middle_proxy::me_zombie_writer_watchdog(pool).await; + } + })); +} diff --git a/src/maestro/me_startup/tests.rs b/src/maestro/me_startup/tests.rs new file mode 100644 index 0000000..a1b16a7 --- /dev/null +++ b/src/maestro/me_startup/tests.rs @@ -0,0 +1,62 @@ +use super::*; +use std::sync::atomic::{AtomicUsize, Ordering}; +use tokio::sync::Notify; + +struct DropSignal(Arc); + +impl Drop for DropSignal { + fn drop(&mut self) { + self.0.notify_one(); + } +} + +#[tokio::test] +async fn scoped_supervisor_aborts_its_current_child() { + let scope = RuntimeTaskScope::new(); + let dropped = Arc::new(Notify::new()); + let dropped_for_task = dropped.clone(); + scope.spawn(supervise_me_task("test", move || { + let dropped = dropped_for_task.clone(); + async move { + let _signal = DropSignal(dropped); + std::future::pending::<()>().await; + } + })); + tokio::task::yield_now().await; + + scope.stop().await; + + tokio::time::timeout(Duration::from_secs(1), dropped.notified()) + .await + .unwrap(); +} + +#[tokio::test] +async fn supervisor_restarts_exited_child_and_stops_with_runtime_scope() { + let scope = RuntimeTaskScope::new(); + let starts = Arc::new(AtomicUsize::new(0)); + let restarted = Arc::new(Notify::new()); + let starts_task = starts.clone(); + let restarted_task = restarted.clone(); + scope.spawn(supervise_me_task("restart_test", move || { + let starts = starts_task.clone(); + let restarted = restarted_task.clone(); + async move { + if starts.fetch_add(1, Ordering::AcqRel) + 1 >= 3 { + restarted.notify_one(); + } + } + })); + + tokio::time::timeout(Duration::from_secs(1), restarted.notified()) + .await + .unwrap(); + scope.stop().await; + let stopped_at = starts.load(Ordering::Acquire); + for _ in 0..100 { + tokio::task::yield_now().await; + } + + assert!(stopped_at >= 3); + assert_eq!(starts.load(Ordering::Acquire), stopped_at); +} diff --git a/src/maestro/orchestrator.rs b/src/maestro/orchestrator.rs index c49c9fc..cf0fc61 100644 --- a/src/maestro/orchestrator.rs +++ b/src/maestro/orchestrator.rs @@ -52,16 +52,9 @@ pub(super) async fn run_telemt_core( let runtime_task_scope = generation::RuntimeTaskScope::new(); stats.apply_telemetry_policy(TelemetryPolicy::from_config(&config.general.telemetry)); let quota_state_path = config.general.quota_state_path.clone(); - let quota_state = crate::quota_state::QuotaStateOwner::new( - quota_state_path, - quota_store.clone(), - ); - let configured_quota_users = config - .access - .users - .keys() - .cloned() - .collect::>(); + let quota_state = + crate::quota_state::QuotaStateOwner::new(quota_state_path, quota_store.clone()); + let configured_quota_users = config.access.users.keys().cloned().collect::>(); quota_state.load(&configured_quota_users).await; let upstream_manager = Arc::new( diff --git a/src/maestro/shutdown.rs b/src/maestro/shutdown.rs index 5ef3cee..226060e 100644 --- a/src/maestro/shutdown.rs +++ b/src/maestro/shutdown.rs @@ -23,9 +23,9 @@ use super::control_plane::ProcessControlPlane; use super::generation::RuntimeGeneration; use super::helpers::{format_uptime, unit_label}; use super::reload_supervisor::ReloadSupervisorHandle; +use crate::quota_state::QuotaStateOwner; use crate::stats::Stats; use crate::synlimit_control; -use crate::quota_state::QuotaStateOwner; /// Signal that triggered shutdown. #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -130,10 +130,7 @@ async fn perform_shutdown( warn!(error = %error, "Failed to clear SYN limiter rules during shutdown"); } - if !process_control_plane - .shutdown(Duration::from_secs(5)) - .await - { + if !process_control_plane.shutdown(Duration::from_secs(5)).await { warn!("Process control-plane task shutdown deadline expired"); } diff --git a/src/metrics.rs b/src/metrics.rs index 87b18b1..dfa69ba 100644 --- a/src/metrics.rs +++ b/src/metrics.rs @@ -451,3943 +451,9 @@ async fn render_tls_front_profile_health( } } -async fn render_metrics( - stats: &Stats, - shared_state: &ProxySharedState, - config: &ProxyConfig, - ip_tracker: &UserIpTracker, - tls_cache: Option<&TlsFrontCache>, - tls_full_cert_budget: &TlsFullCertBudget, - web_publication: &crate::web::control::WebRuntimePublication, -) -> String { - use std::fmt::Write; - let mut out = String::with_capacity(4096); - let telemetry = stats.telemetry_policy(); - let core_enabled = telemetry.core_enabled; - let user_enabled = telemetry.user_enabled; - let me_allows_normal = telemetry.me_level.allows_normal(); - let me_allows_debug = telemetry.me_level.allows_debug(); - - let _ = writeln!( - out, - "# HELP telemt_build_info Build information for the running telemt binary" - ); - let _ = writeln!(out, "# TYPE telemt_build_info gauge"); - let _ = writeln!( - out, - "telemt_build_info{{version=\"{}\"}} 1", - env!("CARGO_PKG_VERSION") - ); - - let _ = writeln!(out, "# HELP telemt_uptime_seconds Proxy uptime"); - let _ = writeln!(out, "# TYPE telemt_uptime_seconds gauge"); - let _ = writeln!(out, "telemt_uptime_seconds {:.1}", stats.uptime_secs()); - - let _ = writeln!( - out, - "# HELP telemt_telemetry_core_enabled Runtime core telemetry switch" - ); - let _ = writeln!(out, "# TYPE telemt_telemetry_core_enabled gauge"); - let _ = writeln!( - out, - "telemt_telemetry_core_enabled {}", - if core_enabled { 1 } else { 0 } - ); - - let _ = writeln!( - out, - "# HELP telemt_telemetry_user_enabled Runtime per-user telemetry switch" - ); - let _ = writeln!(out, "# TYPE telemt_telemetry_user_enabled gauge"); - let _ = writeln!( - out, - "telemt_telemetry_user_enabled {}", - if user_enabled { 1 } else { 0 } - ); - let _ = writeln!( - out, - "# HELP telemt_stats_user_entries Retained per-user stats entries" - ); - let _ = writeln!(out, "# TYPE telemt_stats_user_entries gauge"); - let _ = writeln!(out, "telemt_stats_user_entries {}", stats.user_stats_len()); - - let _ = writeln!( - out, - "# HELP telemt_telemetry_me_level Runtime ME telemetry level flag" - ); - let _ = writeln!(out, "# TYPE telemt_telemetry_me_level gauge"); - let _ = writeln!( - out, - "telemt_telemetry_me_level{{level=\"silent\"}} {}", - if matches!(telemetry.me_level, crate::config::MeTelemetryLevel::Silent) { - 1 - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_telemetry_me_level{{level=\"normal\"}} {}", - if matches!(telemetry.me_level, crate::config::MeTelemetryLevel::Normal) { - 1 - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_telemetry_me_level{{level=\"debug\"}} {}", - if matches!(telemetry.me_level, crate::config::MeTelemetryLevel::Debug) { - 1 - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_buffer_pool_buffers_total Snapshot of pooled and allocated buffers" - ); - let _ = writeln!(out, "# TYPE telemt_buffer_pool_buffers_total gauge"); - let _ = writeln!( - out, - "telemt_buffer_pool_buffers_total{{kind=\"pooled\"}} {}", - stats.get_buffer_pool_pooled_gauge() - ); - let _ = writeln!( - out, - "telemt_buffer_pool_buffers_total{{kind=\"allocated\"}} {}", - stats.get_buffer_pool_allocated_gauge() - ); - let _ = writeln!( - out, - "telemt_buffer_pool_buffers_total{{kind=\"in_use\"}} {}", - stats.get_buffer_pool_in_use_gauge() - ); - let _ = writeln!( - out, - "# HELP telemt_buffer_pool_events_total Buffer-pool allocation lifecycle events" - ); - let _ = writeln!(out, "# TYPE telemt_buffer_pool_events_total counter"); - let _ = writeln!( - out, - "telemt_buffer_pool_events_total{{event=\"replaced_nonstandard\"}} {}", - stats.get_buffer_pool_replaced_nonstandard_total() - ); - - let direct_budget = shared_state.direct_buffer_budget.snapshot(); - let _ = writeln!( - out, - "# HELP telemt_direct_relay_buffer_budget_bytes Direct relay copy-buffer budget and memory inputs" - ); - let _ = writeln!(out, "# TYPE telemt_direct_relay_buffer_budget_bytes gauge"); - for (kind, value) in [ - ("hard_limit", direct_budget.hard_limit_bytes), - ("target", direct_budget.target_bytes), - ("reserved", direct_budget.reserved_bytes), - ("memory_total", direct_budget.memory_total_bytes), - ("memory_available", direct_budget.memory_available_bytes), - ("process_rss", direct_budget.process_rss_bytes), - ] { - let _ = writeln!( - out, - "telemt_direct_relay_buffer_budget_bytes{{kind=\"{}\"}} {}", - kind, value - ); - } - let _ = writeln!( - out, - "# HELP telemt_direct_relay_buffer_budget_events_total Direct relay buffer-budget lifecycle events" - ); - let _ = writeln!( - out, - "# TYPE telemt_direct_relay_buffer_budget_events_total counter" - ); - for (result, value) in [ - ("promotion", direct_budget.promotion_total), - ("promotion_denied", direct_budget.promotion_denied_total), - ("minimum_fallback", direct_budget.minimum_fallback_total), - ("admission_rejected", direct_budget.admission_rejected_total), - ("quiet_demotion", direct_budget.quiet_demotion_total), - ( - "write_pressure_demotion", - direct_budget.write_pressure_demotion_total, - ), - ( - "global_pressure_demotion", - direct_budget.global_pressure_demotion_total, - ), - ] { - let _ = writeln!( - out, - "telemt_direct_relay_buffer_budget_events_total{{result=\"{}\"}} {}", - result, value - ); - } - let _ = writeln!( - out, - "# HELP telemt_direct_relay_buffer_sessions Current Direct relay sessions by adaptive tier" - ); - let _ = writeln!(out, "# TYPE telemt_direct_relay_buffer_sessions gauge"); - for (tier, value) in ["base", "tier1", "tier2", "tier3"] - .into_iter() - .zip(direct_budget.tier_sessions) - { - let _ = writeln!( - out, - "telemt_direct_relay_buffer_sessions{{tier=\"{}\"}} {}", - tier, value - ); - } - - let _ = writeln!( - out, - "# HELP telemt_tls_fetch_profile_cache_entries Current adaptive TLS fetch profile-cache entries" - ); - let _ = writeln!(out, "# TYPE telemt_tls_fetch_profile_cache_entries gauge"); - let _ = writeln!( - out, - "telemt_tls_fetch_profile_cache_entries {}", - fetcher::profile_cache_entries_for_metrics() - ); - let _ = writeln!( - out, - "# HELP telemt_tls_fetch_profile_cache_cap_drops_total Profile-cache winner inserts skipped because the cache cap was reached" - ); - let _ = writeln!( - out, - "# TYPE telemt_tls_fetch_profile_cache_cap_drops_total counter" - ); - let _ = writeln!( - out, - "telemt_tls_fetch_profile_cache_cap_drops_total {}", - fetcher::profile_cache_cap_drops_for_metrics() - ); - let _ = writeln!( - out, - "# HELP telemt_tls_front_full_cert_budget_entries Current domain and IP entries tracked by the process-owned TLS full-cert budget" - ); - let _ = writeln!(out, "# TYPE telemt_tls_front_full_cert_budget_entries gauge"); - let _ = writeln!( - out, - "telemt_tls_front_full_cert_budget_entries {}", - tls_full_cert_budget.entries_for_metrics() - ); - let _ = writeln!( - out, - "# HELP telemt_tls_front_full_cert_budget_cap_drops_total New domain and IP entries denied full-cert budget tracking because a bound was reached" - ); - let _ = writeln!( - out, - "# TYPE telemt_tls_front_full_cert_budget_cap_drops_total counter" - ); - let _ = writeln!( - out, - "telemt_tls_front_full_cert_budget_cap_drops_total {}", - tls_full_cert_budget.cap_drops_for_metrics() - ); - render_tls_front_profile_health(&mut out, config, tls_cache).await; - - let _ = writeln!( - out, - "# HELP telemt_connections_total Total accepted connections" - ); - let _ = writeln!(out, "# TYPE telemt_connections_total counter"); - let _ = writeln!( - out, - "telemt_connections_total {}", - if core_enabled { - stats.get_connects_all() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_connections_bad_total Bad/rejected connections" - ); - let _ = writeln!(out, "# TYPE telemt_connections_bad_total counter"); - let _ = writeln!( - out, - "telemt_connections_bad_total {}", - if core_enabled { - stats.get_connects_bad() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_connections_bad_by_class_total Bad/rejected connections by class" - ); - let _ = writeln!(out, "# TYPE telemt_connections_bad_by_class_total counter"); - if core_enabled { - for (class, total) in stats.get_connects_bad_class_counts() { - let _ = writeln!( - out, - "telemt_connections_bad_by_class_total{{class=\"{}\"}} {}", - class, total - ); - } - } - - let _ = writeln!( - out, - "# HELP telemt_handshake_timeouts_total Handshake timeouts" - ); - let _ = writeln!(out, "# TYPE telemt_handshake_timeouts_total counter"); - let _ = writeln!( - out, - "telemt_handshake_timeouts_total {}", - if core_enabled { - stats.get_handshake_timeouts() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_handshake_failures_by_class_total Handshake failures by class" - ); - let _ = writeln!( - out, - "# TYPE telemt_handshake_failures_by_class_total counter" - ); - if core_enabled { - for (class, total) in stats.get_handshake_failure_class_counts() { - let _ = writeln!( - out, - "telemt_handshake_failures_by_class_total{{class=\"{}\"}} {}", - class, total - ); - } - } - - let _ = writeln!( - out, - "# HELP telemt_auth_expensive_checks_total Expensive authentication candidate checks executed during handshake validation" - ); - let _ = writeln!(out, "# TYPE telemt_auth_expensive_checks_total counter"); - let _ = writeln!( - out, - "telemt_auth_expensive_checks_total {}", - if core_enabled { - shared_state - .handshake - .auth_expensive_checks_total - .load(std::sync::atomic::Ordering::Relaxed) - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_auth_budget_exhausted_total Handshake validations that hit authentication candidate budget limits" - ); - let _ = writeln!(out, "# TYPE telemt_auth_budget_exhausted_total counter"); - let _ = writeln!( - out, - "telemt_auth_budget_exhausted_total {}", - if core_enabled { - shared_state - .handshake - .auth_budget_exhausted_total - .load(std::sync::atomic::Ordering::Relaxed) - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_accept_permit_timeout_total Accepted connections dropped due to permit wait timeout" - ); - let _ = writeln!(out, "# TYPE telemt_accept_permit_timeout_total counter"); - let _ = writeln!( - out, - "telemt_accept_permit_timeout_total {}", - if core_enabled { - stats.get_accept_permit_timeout_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_route_cutover_parked_current Sessions currently parked in route cutover stagger delay" - ); - let _ = writeln!(out, "# TYPE telemt_route_cutover_parked_current gauge"); - let _ = writeln!( - out, - "telemt_route_cutover_parked_current{{route=\"direct\"}} {}", - stats.get_route_cutover_parked_direct_current() - ); - let _ = writeln!( - out, - "telemt_route_cutover_parked_current{{route=\"middle\"}} {}", - stats.get_route_cutover_parked_middle_current() - ); - let _ = writeln!( - out, - "# HELP telemt_route_cutover_parked_total Sessions parked in route cutover stagger delay" - ); - let _ = writeln!(out, "# TYPE telemt_route_cutover_parked_total counter"); - let _ = writeln!( - out, - "telemt_route_cutover_parked_total{{route=\"direct\"}} {}", - stats.get_route_cutover_parked_direct_total() - ); - let _ = writeln!( - out, - "telemt_route_cutover_parked_total{{route=\"middle\"}} {}", - stats.get_route_cutover_parked_middle_total() - ); - - let _ = writeln!( - out, - "# HELP telemt_quota_refund_bytes_total Reserved quota bytes returned before commit" - ); - let _ = writeln!(out, "# TYPE telemt_quota_refund_bytes_total counter"); - let _ = writeln!( - out, - "telemt_quota_refund_bytes_total {}", - if core_enabled { - stats.get_quota_refund_bytes_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_quota_contention_total Quota reservation CAS contention events" - ); - let _ = writeln!(out, "# TYPE telemt_quota_contention_total counter"); - let _ = writeln!( - out, - "telemt_quota_contention_total {}", - if core_enabled { - stats.get_quota_contention_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_quota_contention_timeout_total Quota reservations that hit the bounded contention budget" - ); - let _ = writeln!(out, "# TYPE telemt_quota_contention_timeout_total counter"); - let _ = writeln!( - out, - "telemt_quota_contention_timeout_total {}", - if core_enabled { - stats.get_quota_contention_timeout_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_quota_acquire_cancelled_total Quota acquisitions cancelled before reservation completed" - ); - let _ = writeln!(out, "# TYPE telemt_quota_acquire_cancelled_total counter"); - let _ = writeln!( - out, - "telemt_quota_acquire_cancelled_total {}", - if core_enabled { - stats.get_quota_acquire_cancelled_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_conntrack_control_state Runtime conntrack control state flags" - ); - let _ = writeln!(out, "# TYPE telemt_conntrack_control_state gauge"); - let _ = writeln!( - out, - "telemt_conntrack_control_state{{flag=\"enabled\"}} {}", - if stats.get_conntrack_control_enabled() { - 1 - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_conntrack_control_state{{flag=\"available\"}} {}", - if stats.get_conntrack_control_available() { - 1 - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_conntrack_control_state{{flag=\"pressure_active\"}} {}", - if stats.get_conntrack_pressure_active() { - 1 - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_conntrack_control_state{{flag=\"rule_apply_ok\"}} {}", - if stats.get_conntrack_rule_apply_ok() { - 1 - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_conntrack_event_queue_depth Pending close events in conntrack control queue" - ); - let _ = writeln!(out, "# TYPE telemt_conntrack_event_queue_depth gauge"); - let _ = writeln!( - out, - "telemt_conntrack_event_queue_depth {}", - stats.get_conntrack_event_queue_depth() - ); - - let _ = writeln!( - out, - "# HELP telemt_conntrack_delete_total Conntrack delete attempts by outcome" - ); - let _ = writeln!(out, "# TYPE telemt_conntrack_delete_total counter"); - let _ = writeln!( - out, - "telemt_conntrack_delete_total{{result=\"attempt\"}} {}", - if core_enabled { - stats.get_conntrack_delete_attempt_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_conntrack_delete_total{{result=\"success\"}} {}", - if core_enabled { - stats.get_conntrack_delete_success_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_conntrack_delete_total{{result=\"not_found\"}} {}", - if core_enabled { - stats.get_conntrack_delete_not_found_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_conntrack_delete_total{{result=\"error\"}} {}", - if core_enabled { - stats.get_conntrack_delete_error_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_conntrack_close_event_drop_total Dropped conntrack close events due to queue pressure or unavailable sender" - ); - let _ = writeln!( - out, - "# TYPE telemt_conntrack_close_event_drop_total counter" - ); - let _ = writeln!( - out, - "telemt_conntrack_close_event_drop_total {}", - if core_enabled { - stats.get_conntrack_close_event_drop_total() - } else { - 0 - } - ); - - let limiter_metrics = shared_state.traffic_limiter.metrics_snapshot(); - let _ = writeln!( - out, - "# HELP telemt_rate_limiter_burst_bound_bytes Configured upper bound for one direct relay rate-limit burst" - ); - let _ = writeln!(out, "# TYPE telemt_rate_limiter_burst_bound_bytes gauge"); - let _ = writeln!( - out, - "telemt_rate_limiter_burst_bound_bytes{{direction=\"up\"}} {}", - if core_enabled { - config.general.direct_relay_copy_buf_c2s_bytes - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_rate_limiter_burst_bound_bytes{{direction=\"down\"}} {}", - if core_enabled { - config.general.direct_relay_copy_buf_s2c_bytes - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_rate_limiter_throttle_total Traffic limiter throttle events by scope and direction" - ); - let _ = writeln!(out, "# TYPE telemt_rate_limiter_throttle_total counter"); - let _ = writeln!( - out, - "telemt_rate_limiter_throttle_total{{scope=\"user\",direction=\"up\"}} {}", - if core_enabled { - limiter_metrics.user_throttle_up_total - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_rate_limiter_throttle_total{{scope=\"user\",direction=\"down\"}} {}", - if core_enabled { - limiter_metrics.user_throttle_down_total - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_rate_limiter_throttle_total{{scope=\"cidr\",direction=\"up\"}} {}", - if core_enabled { - limiter_metrics.cidr_throttle_up_total - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_rate_limiter_throttle_total{{scope=\"cidr\",direction=\"down\"}} {}", - if core_enabled { - limiter_metrics.cidr_throttle_down_total - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_rate_limiter_wait_ms_total Traffic limiter accumulated wait time in milliseconds by scope and direction" - ); - let _ = writeln!(out, "# TYPE telemt_rate_limiter_wait_ms_total counter"); - let _ = writeln!( - out, - "telemt_rate_limiter_wait_ms_total{{scope=\"user\",direction=\"up\"}} {}", - if core_enabled { - limiter_metrics.user_wait_up_ms_total - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_rate_limiter_wait_ms_total{{scope=\"user\",direction=\"down\"}} {}", - if core_enabled { - limiter_metrics.user_wait_down_ms_total - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_rate_limiter_wait_ms_total{{scope=\"cidr\",direction=\"up\"}} {}", - if core_enabled { - limiter_metrics.cidr_wait_up_ms_total - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_rate_limiter_wait_ms_total{{scope=\"cidr\",direction=\"down\"}} {}", - if core_enabled { - limiter_metrics.cidr_wait_down_ms_total - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_rate_limiter_active_leases Active relay leases under rate limiting by scope" - ); - let _ = writeln!(out, "# TYPE telemt_rate_limiter_active_leases gauge"); - let _ = writeln!( - out, - "telemt_rate_limiter_active_leases{{scope=\"user\"}} {}", - if core_enabled { - limiter_metrics.user_active_leases - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_rate_limiter_active_leases{{scope=\"cidr\"}} {}", - if core_enabled { - limiter_metrics.cidr_active_leases - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_rate_limiter_policy_entries Active rate-limit policy entries by scope" - ); - let _ = writeln!(out, "# TYPE telemt_rate_limiter_policy_entries gauge"); - let _ = writeln!( - out, - "telemt_rate_limiter_policy_entries{{scope=\"user\"}} {}", - if core_enabled { - limiter_metrics.user_policy_entries - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_rate_limiter_policy_entries{{scope=\"cidr\"}} {}", - if core_enabled { - limiter_metrics.cidr_policy_entries - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_upstream_connect_attempt_total Upstream connect attempts across all requests" - ); - let _ = writeln!(out, "# TYPE telemt_upstream_connect_attempt_total counter"); - let _ = writeln!( - out, - "telemt_upstream_connect_attempt_total {}", - if core_enabled { - stats.get_upstream_connect_attempt_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_upstream_connect_success_total Successful upstream connect request cycles" - ); - let _ = writeln!(out, "# TYPE telemt_upstream_connect_success_total counter"); - let _ = writeln!( - out, - "telemt_upstream_connect_success_total {}", - if core_enabled { - stats.get_upstream_connect_success_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_upstream_connect_fail_total Failed upstream connect request cycles" - ); - let _ = writeln!(out, "# TYPE telemt_upstream_connect_fail_total counter"); - let _ = writeln!( - out, - "telemt_upstream_connect_fail_total {}", - if core_enabled { - stats.get_upstream_connect_fail_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_upstream_connect_failfast_hard_error_total Hard errors that triggered upstream connect failfast" - ); - let _ = writeln!( - out, - "# TYPE telemt_upstream_connect_failfast_hard_error_total counter" - ); - let _ = writeln!( - out, - "telemt_upstream_connect_failfast_hard_error_total {}", - if core_enabled { - stats.get_upstream_connect_failfast_hard_error_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_upstream_connect_attempts_per_request Histogram-like buckets for attempts per upstream connect request cycle" - ); - let _ = writeln!( - out, - "# TYPE telemt_upstream_connect_attempts_per_request counter" - ); - let _ = writeln!( - out, - "telemt_upstream_connect_attempts_per_request{{bucket=\"1\"}} {}", - if core_enabled { - stats.get_upstream_connect_attempts_bucket_1() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_upstream_connect_attempts_per_request{{bucket=\"2\"}} {}", - if core_enabled { - stats.get_upstream_connect_attempts_bucket_2() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_upstream_connect_attempts_per_request{{bucket=\"3_4\"}} {}", - if core_enabled { - stats.get_upstream_connect_attempts_bucket_3_4() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_upstream_connect_attempts_per_request{{bucket=\"gt_4\"}} {}", - if core_enabled { - stats.get_upstream_connect_attempts_bucket_gt_4() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_upstream_connect_duration_success_total Histogram-like buckets of successful upstream connect cycle duration" - ); - let _ = writeln!( - out, - "# TYPE telemt_upstream_connect_duration_success_total counter" - ); - let _ = writeln!( - out, - "telemt_upstream_connect_duration_success_total{{bucket=\"le_100ms\"}} {}", - if core_enabled { - stats.get_upstream_connect_duration_success_bucket_le_100ms() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_upstream_connect_duration_success_total{{bucket=\"101_500ms\"}} {}", - if core_enabled { - stats.get_upstream_connect_duration_success_bucket_101_500ms() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_upstream_connect_duration_success_total{{bucket=\"501_1000ms\"}} {}", - if core_enabled { - stats.get_upstream_connect_duration_success_bucket_501_1000ms() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_upstream_connect_duration_success_total{{bucket=\"gt_1000ms\"}} {}", - if core_enabled { - stats.get_upstream_connect_duration_success_bucket_gt_1000ms() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_upstream_connect_duration_fail_total Histogram-like buckets of failed upstream connect cycle duration" - ); - let _ = writeln!( - out, - "# TYPE telemt_upstream_connect_duration_fail_total counter" - ); - let _ = writeln!( - out, - "telemt_upstream_connect_duration_fail_total{{bucket=\"le_100ms\"}} {}", - if core_enabled { - stats.get_upstream_connect_duration_fail_bucket_le_100ms() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_upstream_connect_duration_fail_total{{bucket=\"101_500ms\"}} {}", - if core_enabled { - stats.get_upstream_connect_duration_fail_bucket_101_500ms() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_upstream_connect_duration_fail_total{{bucket=\"501_1000ms\"}} {}", - if core_enabled { - stats.get_upstream_connect_duration_fail_bucket_501_1000ms() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_upstream_connect_duration_fail_total{{bucket=\"gt_1000ms\"}} {}", - if core_enabled { - stats.get_upstream_connect_duration_fail_bucket_gt_1000ms() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_keepalive_sent_total ME keepalive frames sent" - ); - let _ = writeln!(out, "# TYPE telemt_me_keepalive_sent_total counter"); - let _ = writeln!( - out, - "telemt_me_keepalive_sent_total {}", - if me_allows_debug { - stats.get_me_keepalive_sent() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_keepalive_failed_total ME keepalive send failures" - ); - let _ = writeln!(out, "# TYPE telemt_me_keepalive_failed_total counter"); - let _ = writeln!( - out, - "telemt_me_keepalive_failed_total {}", - if me_allows_normal { - stats.get_me_keepalive_failed() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_keepalive_pong_total ME keepalive pong replies" - ); - let _ = writeln!(out, "# TYPE telemt_me_keepalive_pong_total counter"); - let _ = writeln!( - out, - "telemt_me_keepalive_pong_total {}", - if me_allows_debug { - stats.get_me_keepalive_pong() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_keepalive_timeout_total ME keepalive ping timeouts" - ); - let _ = writeln!(out, "# TYPE telemt_me_keepalive_timeout_total counter"); - let _ = writeln!( - out, - "telemt_me_keepalive_timeout_total {}", - if me_allows_normal { - stats.get_me_keepalive_timeout() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_rpc_proxy_req_signal_sent_total Service RPC_PROXY_REQ activity signals sent" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_rpc_proxy_req_signal_sent_total counter" - ); - let _ = writeln!( - out, - "telemt_me_rpc_proxy_req_signal_sent_total {}", - if me_allows_normal { - stats.get_me_rpc_proxy_req_signal_sent_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_rpc_proxy_req_signal_failed_total Service RPC_PROXY_REQ activity signal failures" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_rpc_proxy_req_signal_failed_total counter" - ); - let _ = writeln!( - out, - "telemt_me_rpc_proxy_req_signal_failed_total {}", - if me_allows_normal { - stats.get_me_rpc_proxy_req_signal_failed_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_rpc_proxy_req_signal_skipped_no_meta_total Service RPC_PROXY_REQ skipped due to missing writer metadata" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_rpc_proxy_req_signal_skipped_no_meta_total counter" - ); - let _ = writeln!( - out, - "telemt_me_rpc_proxy_req_signal_skipped_no_meta_total {}", - if me_allows_normal { - stats.get_me_rpc_proxy_req_signal_skipped_no_meta_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_rpc_proxy_req_signal_response_total Service RPC_PROXY_REQ responses observed" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_rpc_proxy_req_signal_response_total counter" - ); - let _ = writeln!( - out, - "telemt_me_rpc_proxy_req_signal_response_total {}", - if me_allows_normal { - stats.get_me_rpc_proxy_req_signal_response_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_rpc_proxy_req_signal_close_sent_total Service RPC_CLOSE_EXT sent after activity signals" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_rpc_proxy_req_signal_close_sent_total counter" - ); - let _ = writeln!( - out, - "telemt_me_rpc_proxy_req_signal_close_sent_total {}", - if me_allows_normal { - stats.get_me_rpc_proxy_req_signal_close_sent_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_reconnect_attempts_total ME reconnect attempts" - ); - let _ = writeln!(out, "# TYPE telemt_me_reconnect_attempts_total counter"); - let _ = writeln!( - out, - "telemt_me_reconnect_attempts_total {}", - if me_allows_normal { - stats.get_me_reconnect_attempts() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_reconnect_success_total ME reconnect successes" - ); - let _ = writeln!(out, "# TYPE telemt_me_reconnect_success_total counter"); - let _ = writeln!( - out, - "telemt_me_reconnect_success_total {}", - if me_allows_normal { - stats.get_me_reconnect_success() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_handshake_reject_total ME handshake rejects from upstream" - ); - let _ = writeln!(out, "# TYPE telemt_me_handshake_reject_total counter"); - let _ = writeln!( - out, - "telemt_me_handshake_reject_total {}", - if me_allows_normal { - stats.get_me_handshake_reject_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_handshake_error_code_total ME handshake reject errors by code" - ); - let _ = writeln!(out, "# TYPE telemt_me_handshake_error_code_total counter"); - if me_allows_normal { - for (error_code, count) in stats.get_me_handshake_error_code_counts() { - let _ = writeln!( - out, - "telemt_me_handshake_error_code_total{{error_code=\"{}\"}} {}", - error_code, count - ); - } - let _ = writeln!( - out, - "telemt_me_handshake_error_code_total{{error_code=\"overflow\"}} {}", - stats.get_me_handshake_error_code_overflow_total() - ); - } - - let _ = writeln!( - out, - "# HELP telemt_me_reader_eof_total ME reader EOF terminations" - ); - let _ = writeln!(out, "# TYPE telemt_me_reader_eof_total counter"); - let _ = writeln!( - out, - "telemt_me_reader_eof_total {}", - if me_allows_normal { - stats.get_me_reader_eof_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_idle_close_by_peer_total ME idle writers closed by peer" - ); - let _ = writeln!(out, "# TYPE telemt_me_idle_close_by_peer_total counter"); - let _ = writeln!( - out, - "telemt_me_idle_close_by_peer_total {}", - if me_allows_normal { - stats.get_me_idle_close_by_peer_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_relay_idle_soft_mark_total Middle-relay sessions marked as soft-idle candidates" - ); - let _ = writeln!(out, "# TYPE telemt_relay_idle_soft_mark_total counter"); - let _ = writeln!( - out, - "telemt_relay_idle_soft_mark_total {}", - if me_allows_normal { - stats.get_relay_idle_soft_mark_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_relay_idle_hard_close_total Middle-relay sessions closed by hard-idle policy" - ); - let _ = writeln!(out, "# TYPE telemt_relay_idle_hard_close_total counter"); - let _ = writeln!( - out, - "telemt_relay_idle_hard_close_total {}", - if me_allows_normal { - stats.get_relay_idle_hard_close_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_relay_pressure_evict_total Middle-relay sessions evicted under resource pressure" - ); - let _ = writeln!(out, "# TYPE telemt_relay_pressure_evict_total counter"); - let _ = writeln!( - out, - "telemt_relay_pressure_evict_total {}", - if me_allows_normal { - stats.get_relay_pressure_evict_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_relay_protocol_desync_close_total Middle-relay sessions closed due to protocol desync" - ); - let _ = writeln!( - out, - "# TYPE telemt_relay_protocol_desync_close_total counter" - ); - let _ = writeln!( - out, - "telemt_relay_protocol_desync_close_total {}", - if me_allows_normal { - stats.get_relay_protocol_desync_close_total() - } else { - 0 - } - ); - - let _ = writeln!(out, "# HELP telemt_me_crc_mismatch_total ME CRC mismatches"); - let _ = writeln!(out, "# TYPE telemt_me_crc_mismatch_total counter"); - let _ = writeln!( - out, - "telemt_me_crc_mismatch_total {}", - if me_allows_normal { - stats.get_me_crc_mismatch() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_seq_mismatch_total ME sequence mismatches" - ); - let _ = writeln!(out, "# TYPE telemt_me_seq_mismatch_total counter"); - let _ = writeln!( - out, - "telemt_me_seq_mismatch_total {}", - if me_allows_normal { - stats.get_me_seq_mismatch() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_route_drop_no_conn_total ME route drops: no conn" - ); - let _ = writeln!(out, "# TYPE telemt_me_route_drop_no_conn_total counter"); - let _ = writeln!( - out, - "telemt_me_route_drop_no_conn_total {}", - if me_allows_normal { - stats.get_me_route_drop_no_conn() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_route_drop_channel_closed_total ME route drops: channel closed" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_route_drop_channel_closed_total counter" - ); - let _ = writeln!( - out, - "telemt_me_route_drop_channel_closed_total {}", - if me_allows_normal { - stats.get_me_route_drop_channel_closed() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_route_drop_queue_full_total ME route drops: queue full" - ); - let _ = writeln!(out, "# TYPE telemt_me_route_drop_queue_full_total counter"); - let _ = writeln!( - out, - "telemt_me_route_drop_queue_full_total {}", - if me_allows_normal { - stats.get_me_route_drop_queue_full() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_route_drop_queue_full_profile_total ME route drops: queue full by adaptive profile" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_route_drop_queue_full_profile_total counter" - ); - let _ = writeln!( - out, - "telemt_me_route_drop_queue_full_profile_total{{profile=\"base\"}} {}", - if me_allows_normal { - stats.get_me_route_drop_queue_full_base() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_route_drop_queue_full_profile_total{{profile=\"high\"}} {}", - if me_allows_normal { - stats.get_me_route_drop_queue_full_high() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_fair_pressure_state Worker-local fairness pressure state" - ); - let _ = writeln!(out, "# TYPE telemt_me_fair_pressure_state gauge"); - let _ = writeln!( - out, - "telemt_me_fair_pressure_state {}", - if me_allows_normal { - stats.get_me_fair_pressure_state_gauge() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_fair_active_flows Fair-scheduler active flow count" - ); - let _ = writeln!(out, "# TYPE telemt_me_fair_active_flows gauge"); - let _ = writeln!( - out, - "telemt_me_fair_active_flows {}", - if me_allows_normal { - stats.get_me_fair_active_flows_gauge() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_fair_queued_bytes Fair-scheduler queued bytes" - ); - let _ = writeln!(out, "# TYPE telemt_me_fair_queued_bytes gauge"); - let _ = writeln!( - out, - "telemt_me_fair_queued_bytes {}", - if me_allows_normal { - stats.get_me_fair_queued_bytes_gauge() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_fair_flow_state_gauge Fair-scheduler flow health classes" - ); - let _ = writeln!(out, "# TYPE telemt_me_fair_flow_state_gauge gauge"); - let _ = writeln!( - out, - "telemt_me_fair_flow_state_gauge{{class=\"standing\"}} {}", - if me_allows_normal { - stats.get_me_fair_standing_flows_gauge() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_fair_flow_state_gauge{{class=\"backpressured\"}} {}", - if me_allows_normal { - stats.get_me_fair_backpressured_flows_gauge() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_fair_events_total Fair-scheduler event counters" - ); - let _ = writeln!(out, "# TYPE telemt_me_fair_events_total counter"); - let _ = writeln!( - out, - "telemt_me_fair_events_total{{event=\"scheduler_round\"}} {}", - if me_allows_normal { - stats.get_me_fair_scheduler_rounds_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_fair_events_total{{event=\"deficit_grant\"}} {}", - if me_allows_normal { - stats.get_me_fair_deficit_grants_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_fair_events_total{{event=\"deficit_skip\"}} {}", - if me_allows_normal { - stats.get_me_fair_deficit_skips_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_fair_events_total{{event=\"enqueue_reject\"}} {}", - if me_allows_normal { - stats.get_me_fair_enqueue_rejects_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_fair_events_total{{event=\"shed_drop\"}} {}", - if me_allows_normal { - stats.get_me_fair_shed_drops_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_fair_events_total{{event=\"penalty\"}} {}", - if me_allows_normal { - stats.get_me_fair_penalties_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_fair_events_total{{event=\"downstream_stall\"}} {}", - if me_allows_normal { - stats.get_me_fair_downstream_stalls_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_c2me_enqueue_events_total ME client->ME enqueue outcomes" - ); - let _ = writeln!(out, "# TYPE telemt_me_c2me_enqueue_events_total counter"); - let _ = writeln!( - out, - "telemt_me_c2me_enqueue_events_total{{event=\"full\"}} {}", - if me_allows_normal { - stats.get_me_c2me_send_full_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_c2me_enqueue_events_total{{event=\"high_water\"}} {}", - if me_allows_normal { - stats.get_me_c2me_send_high_water_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_c2me_enqueue_events_total{{event=\"timeout\"}} {}", - if me_allows_normal { - stats.get_me_c2me_send_timeout_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_d2c_batches_total Total DC->Client flush batches" - ); - let _ = writeln!(out, "# TYPE telemt_me_d2c_batches_total counter"); - let _ = writeln!( - out, - "telemt_me_d2c_batches_total {}", - if me_allows_normal { - stats.get_me_d2c_batches_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_d2c_batch_frames_total Total DC->Client frames flushed in batches" - ); - let _ = writeln!(out, "# TYPE telemt_me_d2c_batch_frames_total counter"); - let _ = writeln!( - out, - "telemt_me_d2c_batch_frames_total {}", - if me_allows_normal { - stats.get_me_d2c_batch_frames_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_d2c_batch_bytes_total Total DC->Client bytes flushed in batches" - ); - let _ = writeln!(out, "# TYPE telemt_me_d2c_batch_bytes_total counter"); - let _ = writeln!( - out, - "telemt_me_d2c_batch_bytes_total {}", - if me_allows_normal { - stats.get_me_d2c_batch_bytes_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_d2c_flush_reason_total DC->Client flush reasons" - ); - let _ = writeln!(out, "# TYPE telemt_me_d2c_flush_reason_total counter"); - let _ = writeln!( - out, - "telemt_me_d2c_flush_reason_total{{reason=\"queue_drain\"}} {}", - if me_allows_normal { - stats.get_me_d2c_flush_reason_queue_drain_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_flush_reason_total{{reason=\"batch_frames\"}} {}", - if me_allows_normal { - stats.get_me_d2c_flush_reason_batch_frames_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_flush_reason_total{{reason=\"batch_bytes\"}} {}", - if me_allows_normal { - stats.get_me_d2c_flush_reason_batch_bytes_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_flush_reason_total{{reason=\"max_delay\"}} {}", - if me_allows_normal { - stats.get_me_d2c_flush_reason_max_delay_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_flush_reason_total{{reason=\"ack_immediate\"}} {}", - if me_allows_normal { - stats.get_me_d2c_flush_reason_ack_immediate_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_flush_reason_total{{reason=\"close\"}} {}", - if me_allows_normal { - stats.get_me_d2c_flush_reason_close_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_d2c_data_frames_total DC->Client data frames" - ); - let _ = writeln!(out, "# TYPE telemt_me_d2c_data_frames_total counter"); - let _ = writeln!( - out, - "telemt_me_d2c_data_frames_total {}", - if me_allows_normal { - stats.get_me_d2c_data_frames_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_d2c_ack_frames_total DC->Client quick-ack frames" - ); - let _ = writeln!(out, "# TYPE telemt_me_d2c_ack_frames_total counter"); - let _ = writeln!( - out, - "telemt_me_d2c_ack_frames_total {}", - if me_allows_normal { - stats.get_me_d2c_ack_frames_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_d2c_payload_bytes_total DC->Client payload bytes before transport framing" - ); - let _ = writeln!(out, "# TYPE telemt_me_d2c_payload_bytes_total counter"); - let _ = writeln!( - out, - "telemt_me_d2c_payload_bytes_total {}", - if me_allows_normal { - stats.get_me_d2c_payload_bytes_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_d2c_write_mode_total DC->Client writer mode selection" - ); - let _ = writeln!(out, "# TYPE telemt_me_d2c_write_mode_total counter"); - let _ = writeln!( - out, - "telemt_me_d2c_write_mode_total{{mode=\"coalesced\"}} {}", - if me_allows_normal { - stats.get_me_d2c_write_mode_coalesced_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_write_mode_total{{mode=\"split\"}} {}", - if me_allows_normal { - stats.get_me_d2c_write_mode_split_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_d2c_quota_reject_total DC->Client quota rejects" - ); - let _ = writeln!(out, "# TYPE telemt_me_d2c_quota_reject_total counter"); - let _ = writeln!( - out, - "telemt_me_d2c_quota_reject_total{{stage=\"pre_write\"}} {}", - if me_allows_normal { - stats.get_me_d2c_quota_reject_pre_write_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_quota_reject_total{{stage=\"post_write\"}} {}", - if me_allows_normal { - stats.get_me_d2c_quota_reject_post_write_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_child_join_timeout_total Middle relay child tasks that did not join before cleanup deadline" - ); - let _ = writeln!(out, "# TYPE telemt_me_child_join_timeout_total counter"); - let _ = writeln!( - out, - "telemt_me_child_join_timeout_total {}", - if core_enabled { - stats.get_me_child_join_timeout_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_child_abort_total Middle relay child tasks aborted after bounded cleanup timeout" - ); - let _ = writeln!(out, "# TYPE telemt_me_child_abort_total counter"); - let _ = writeln!( - out, - "telemt_me_child_abort_total {}", - if core_enabled { - stats.get_me_child_abort_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_flow_wait_events_total Flow wait events by reason, direction, and outcome" - ); - let _ = writeln!(out, "# TYPE telemt_flow_wait_events_total counter"); - let _ = writeln!( - out, - "telemt_flow_wait_events_total{{reason=\"middle_rate_limit\",direction=\"down\",outcome=\"waited\"}} {}", - if core_enabled { - stats.get_flow_wait_middle_rate_limit_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_flow_wait_events_total{{reason=\"middle_rate_limit\",direction=\"down\",outcome=\"cancelled\"}} {}", - if core_enabled { - stats.get_flow_wait_middle_rate_limit_cancelled_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_flow_wait_ms_total Flow wait time in milliseconds by reason and direction" - ); - let _ = writeln!(out, "# TYPE telemt_flow_wait_ms_total counter"); - let _ = writeln!( - out, - "telemt_flow_wait_ms_total{{reason=\"middle_rate_limit\",direction=\"down\"}} {}", - if core_enabled { - stats.get_flow_wait_middle_rate_limit_ms_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_session_drop_fallback_total Session reservations cleaned by Drop instead of explicit async release" - ); - let _ = writeln!(out, "# TYPE telemt_session_drop_fallback_total counter"); - let _ = writeln!( - out, - "telemt_session_drop_fallback_total {}", - if core_enabled { - stats.get_session_drop_fallback_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_d2c_frame_buf_shrink_total DC->Client reusable frame buffer shrink events" - ); - let _ = writeln!(out, "# TYPE telemt_me_d2c_frame_buf_shrink_total counter"); - let _ = writeln!( - out, - "telemt_me_d2c_frame_buf_shrink_total {}", - if me_allows_normal { - stats.get_me_d2c_frame_buf_shrink_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_d2c_frame_buf_shrink_bytes_total DC->Client reusable frame buffer bytes released" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_d2c_frame_buf_shrink_bytes_total counter" - ); - let _ = writeln!( - out, - "telemt_me_d2c_frame_buf_shrink_bytes_total {}", - if me_allows_normal { - stats.get_me_d2c_frame_buf_shrink_bytes_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_d2c_batch_frames_bucket_total DC->Client batch frame count buckets" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_d2c_batch_frames_bucket_total counter" - ); - let _ = writeln!( - out, - "telemt_me_d2c_batch_frames_bucket_total{{bucket=\"1\"}} {}", - if me_allows_debug { - stats.get_me_d2c_batch_frames_bucket_1() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_batch_frames_bucket_total{{bucket=\"2_4\"}} {}", - if me_allows_debug { - stats.get_me_d2c_batch_frames_bucket_2_4() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_batch_frames_bucket_total{{bucket=\"5_8\"}} {}", - if me_allows_debug { - stats.get_me_d2c_batch_frames_bucket_5_8() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_batch_frames_bucket_total{{bucket=\"9_16\"}} {}", - if me_allows_debug { - stats.get_me_d2c_batch_frames_bucket_9_16() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_batch_frames_bucket_total{{bucket=\"17_32\"}} {}", - if me_allows_debug { - stats.get_me_d2c_batch_frames_bucket_17_32() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_batch_frames_bucket_total{{bucket=\"gt_32\"}} {}", - if me_allows_debug { - stats.get_me_d2c_batch_frames_bucket_gt_32() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_d2c_batch_bytes_bucket_total DC->Client batch byte size buckets" - ); - let _ = writeln!(out, "# TYPE telemt_me_d2c_batch_bytes_bucket_total counter"); - let _ = writeln!( - out, - "telemt_me_d2c_batch_bytes_bucket_total{{bucket=\"0_1k\"}} {}", - if me_allows_debug { - stats.get_me_d2c_batch_bytes_bucket_0_1k() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_batch_bytes_bucket_total{{bucket=\"1k_4k\"}} {}", - if me_allows_debug { - stats.get_me_d2c_batch_bytes_bucket_1k_4k() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_batch_bytes_bucket_total{{bucket=\"4k_16k\"}} {}", - if me_allows_debug { - stats.get_me_d2c_batch_bytes_bucket_4k_16k() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_batch_bytes_bucket_total{{bucket=\"16k_64k\"}} {}", - if me_allows_debug { - stats.get_me_d2c_batch_bytes_bucket_16k_64k() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_batch_bytes_bucket_total{{bucket=\"64k_128k\"}} {}", - if me_allows_debug { - stats.get_me_d2c_batch_bytes_bucket_64k_128k() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_batch_bytes_bucket_total{{bucket=\"gt_128k\"}} {}", - if me_allows_debug { - stats.get_me_d2c_batch_bytes_bucket_gt_128k() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_d2c_flush_duration_us_bucket_total DC->Client flush duration buckets" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_d2c_flush_duration_us_bucket_total counter" - ); - let _ = writeln!( - out, - "telemt_me_d2c_flush_duration_us_bucket_total{{bucket=\"0_50\"}} {}", - if me_allows_debug { - stats.get_me_d2c_flush_duration_us_bucket_0_50() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_flush_duration_us_bucket_total{{bucket=\"51_200\"}} {}", - if me_allows_debug { - stats.get_me_d2c_flush_duration_us_bucket_51_200() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_flush_duration_us_bucket_total{{bucket=\"201_1000\"}} {}", - if me_allows_debug { - stats.get_me_d2c_flush_duration_us_bucket_201_1000() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_flush_duration_us_bucket_total{{bucket=\"1001_5000\"}} {}", - if me_allows_debug { - stats.get_me_d2c_flush_duration_us_bucket_1001_5000() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_flush_duration_us_bucket_total{{bucket=\"5001_20000\"}} {}", - if me_allows_debug { - stats.get_me_d2c_flush_duration_us_bucket_5001_20000() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_d2c_flush_duration_us_bucket_total{{bucket=\"gt_20000\"}} {}", - if me_allows_debug { - stats.get_me_d2c_flush_duration_us_bucket_gt_20000() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_d2c_batch_timeout_armed_total DC->Client max-delay timer armed events" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_d2c_batch_timeout_armed_total counter" - ); - let _ = writeln!( - out, - "telemt_me_d2c_batch_timeout_armed_total {}", - if me_allows_debug { - stats.get_me_d2c_batch_timeout_armed_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_d2c_batch_timeout_fired_total DC->Client max-delay timer fired events" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_d2c_batch_timeout_fired_total counter" - ); - let _ = writeln!( - out, - "telemt_me_d2c_batch_timeout_fired_total {}", - if me_allows_debug { - stats.get_me_d2c_batch_timeout_fired_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_writer_byte_budget_limit_bytes Configured resident-memory budget per ME writer" - ); - let _ = writeln!(out, "# TYPE telemt_me_writer_byte_budget_limit_bytes gauge"); - let _ = writeln!( - out, - "telemt_me_writer_byte_budget_limit_bytes {}", - if me_allows_normal { - stats.get_me_writer_byte_budget_limit_bytes_gauge() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_writer_byte_budget_reserved_bytes Aggregate ME writer memory reservations by lifecycle state" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_writer_byte_budget_reserved_bytes gauge" - ); - let _ = writeln!( - out, - "telemt_me_writer_byte_budget_reserved_bytes{{state=\"queued\"}} {}", - if me_allows_normal { - stats.get_me_writer_byte_budget_queued_bytes_gauge() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_writer_byte_budget_reserved_bytes{{state=\"inflight\"}} {}", - if me_allows_normal { - stats.get_me_writer_byte_budget_inflight_bytes_gauge() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_writer_byte_budget_events_total ME writer byte-budget outcomes" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_writer_byte_budget_events_total counter" - ); - let _ = writeln!( - out, - "telemt_me_writer_byte_budget_events_total{{result=\"wait\"}} {}", - if me_allows_normal { - stats.get_me_writer_byte_budget_wait_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_writer_byte_budget_events_total{{result=\"timeout\"}} {}", - if me_allows_normal { - stats.get_me_writer_byte_budget_timeout_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_writer_byte_budget_events_total{{result=\"oversize\"}} {}", - if me_allows_normal { - stats.get_me_writer_byte_budget_oversize_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_writer_pick_total ME writer-pick outcomes by mode and result" - ); - let _ = writeln!(out, "# TYPE telemt_me_writer_pick_total counter"); - let _ = writeln!( - out, - "telemt_me_writer_pick_total{{mode=\"sorted_rr\",result=\"success_try\"}} {}", - if me_allows_normal { - stats.get_me_writer_pick_sorted_rr_success_try_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_writer_pick_total{{mode=\"sorted_rr\",result=\"success_fallback\"}} {}", - if me_allows_normal { - stats.get_me_writer_pick_sorted_rr_success_fallback_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_writer_pick_total{{mode=\"sorted_rr\",result=\"full\"}} {}", - if me_allows_normal { - stats.get_me_writer_pick_sorted_rr_full_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_writer_pick_total{{mode=\"sorted_rr\",result=\"closed\"}} {}", - if me_allows_normal { - stats.get_me_writer_pick_sorted_rr_closed_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_writer_pick_total{{mode=\"sorted_rr\",result=\"no_candidate\"}} {}", - if me_allows_normal { - stats.get_me_writer_pick_sorted_rr_no_candidate_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_writer_pick_total{{mode=\"p2c\",result=\"success_try\"}} {}", - if me_allows_normal { - stats.get_me_writer_pick_p2c_success_try_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_writer_pick_total{{mode=\"p2c\",result=\"success_fallback\"}} {}", - if me_allows_normal { - stats.get_me_writer_pick_p2c_success_fallback_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_writer_pick_total{{mode=\"p2c\",result=\"full\"}} {}", - if me_allows_normal { - stats.get_me_writer_pick_p2c_full_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_writer_pick_total{{mode=\"p2c\",result=\"closed\"}} {}", - if me_allows_normal { - stats.get_me_writer_pick_p2c_closed_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_writer_pick_total{{mode=\"p2c\",result=\"no_candidate\"}} {}", - if me_allows_normal { - stats.get_me_writer_pick_p2c_no_candidate_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_writer_pick_blocking_fallback_total ME writer-pick blocking fallback attempts" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_writer_pick_blocking_fallback_total counter" - ); - let _ = writeln!( - out, - "telemt_me_writer_pick_blocking_fallback_total {}", - if me_allows_normal { - stats.get_me_writer_pick_blocking_fallback_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_writer_pick_mode_switch_total Writer-pick mode switches via runtime updates" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_writer_pick_mode_switch_total counter" - ); - let _ = writeln!( - out, - "telemt_me_writer_pick_mode_switch_total {}", - if me_allows_normal { - stats.get_me_writer_pick_mode_switch_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_socks_kdf_policy_total SOCKS KDF policy outcomes" - ); - let _ = writeln!(out, "# TYPE telemt_me_socks_kdf_policy_total counter"); - let _ = writeln!( - out, - "telemt_me_socks_kdf_policy_total{{policy=\"strict\",outcome=\"reject\"}} {}", - if me_allows_normal { - stats.get_me_socks_kdf_strict_reject() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_socks_kdf_policy_total{{policy=\"compat\",outcome=\"fallback\"}} {}", - if me_allows_debug { - stats.get_me_socks_kdf_compat_fallback() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_endpoint_quarantine_total ME endpoint quarantines due to rapid flaps" - ); - let _ = writeln!(out, "# TYPE telemt_me_endpoint_quarantine_total counter"); - let _ = writeln!( - out, - "telemt_me_endpoint_quarantine_total {}", - if me_allows_normal { - stats.get_me_endpoint_quarantine_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_endpoint_quarantine_unexpected_total ME endpoint quarantines caused by unexpected writer removals" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_endpoint_quarantine_unexpected_total counter" - ); - let _ = writeln!( - out, - "telemt_me_endpoint_quarantine_unexpected_total {}", - if me_allows_normal { - stats.get_me_endpoint_quarantine_unexpected_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_endpoint_quarantine_draining_suppressed_total Draining writer removals that skipped endpoint quarantine" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_endpoint_quarantine_draining_suppressed_total counter" - ); - let _ = writeln!( - out, - "telemt_me_endpoint_quarantine_draining_suppressed_total {}", - if me_allows_normal { - stats.get_me_endpoint_quarantine_draining_suppressed_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_kdf_drift_total ME KDF input drift detections" - ); - let _ = writeln!(out, "# TYPE telemt_me_kdf_drift_total counter"); - let _ = writeln!( - out, - "telemt_me_kdf_drift_total {}", - if me_allows_normal { - stats.get_me_kdf_drift_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_kdf_port_only_drift_total ME KDF client-port changes with stable non-port material" - ); - let _ = writeln!(out, "# TYPE telemt_me_kdf_port_only_drift_total counter"); - let _ = writeln!( - out, - "telemt_me_kdf_port_only_drift_total {}", - if me_allows_debug { - stats.get_me_kdf_port_only_drift_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_hardswap_pending_reuse_total Hardswap cycles that reused an existing pending generation" - ); - let _ = writeln!(out, "# TYPE telemt_me_hardswap_pending_reuse_total counter"); - let _ = writeln!( - out, - "telemt_me_hardswap_pending_reuse_total {}", - if me_allows_debug { - stats.get_me_hardswap_pending_reuse_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_hardswap_pending_ttl_expired_total Pending hardswap generations reset by TTL expiration" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_hardswap_pending_ttl_expired_total counter" - ); - let _ = writeln!( - out, - "telemt_me_hardswap_pending_ttl_expired_total {}", - if me_allows_normal { - stats.get_me_hardswap_pending_ttl_expired_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_single_endpoint_outage_enter_total Single-endpoint DC outage transitions to active state" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_single_endpoint_outage_enter_total counter" - ); - let _ = writeln!( - out, - "telemt_me_single_endpoint_outage_enter_total {}", - if me_allows_normal { - stats.get_me_single_endpoint_outage_enter_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_single_endpoint_outage_exit_total Single-endpoint DC outage recovery transitions" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_single_endpoint_outage_exit_total counter" - ); - let _ = writeln!( - out, - "telemt_me_single_endpoint_outage_exit_total {}", - if me_allows_normal { - stats.get_me_single_endpoint_outage_exit_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_single_endpoint_outage_reconnect_attempt_total Reconnect attempts performed during single-endpoint outages" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_single_endpoint_outage_reconnect_attempt_total counter" - ); - let _ = writeln!( - out, - "telemt_me_single_endpoint_outage_reconnect_attempt_total {}", - if me_allows_normal { - stats.get_me_single_endpoint_outage_reconnect_attempt_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_single_endpoint_outage_reconnect_success_total Successful reconnect attempts during single-endpoint outages" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_single_endpoint_outage_reconnect_success_total counter" - ); - let _ = writeln!( - out, - "telemt_me_single_endpoint_outage_reconnect_success_total {}", - if me_allows_normal { - stats.get_me_single_endpoint_outage_reconnect_success_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_single_endpoint_quarantine_bypass_total Outage reconnect attempts that bypassed quarantine" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_single_endpoint_quarantine_bypass_total counter" - ); - let _ = writeln!( - out, - "telemt_me_single_endpoint_quarantine_bypass_total {}", - if me_allows_normal { - stats.get_me_single_endpoint_quarantine_bypass_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_single_endpoint_shadow_rotate_total Successful periodic shadow rotations for single-endpoint DC groups" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_single_endpoint_shadow_rotate_total counter" - ); - let _ = writeln!( - out, - "telemt_me_single_endpoint_shadow_rotate_total {}", - if me_allows_normal { - stats.get_me_single_endpoint_shadow_rotate_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_single_endpoint_shadow_rotate_skipped_quarantine_total Shadow rotations skipped because endpoint is quarantined" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_single_endpoint_shadow_rotate_skipped_quarantine_total counter" - ); - let _ = writeln!( - out, - "telemt_me_single_endpoint_shadow_rotate_skipped_quarantine_total {}", - if me_allows_normal { - stats.get_me_single_endpoint_shadow_rotate_skipped_quarantine_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_floor_mode Runtime ME writer floor policy mode" - ); - let _ = writeln!(out, "# TYPE telemt_me_floor_mode gauge"); - let floor_mode = config.general.me_floor_mode; - let _ = writeln!( - out, - "telemt_me_floor_mode{{mode=\"static\"}} {}", - if matches!(floor_mode, crate::config::MeFloorMode::Static) { - 1 - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_floor_mode{{mode=\"adaptive\"}} {}", - if matches!(floor_mode, crate::config::MeFloorMode::Adaptive) { - 1 - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_floor_mode_switch_all_total Runtime ME floor mode switches" - ); - let _ = writeln!(out, "# TYPE telemt_me_floor_mode_switch_all_total counter"); - let _ = writeln!( - out, - "telemt_me_floor_mode_switch_all_total {}", - if me_allows_normal { - stats.get_me_floor_mode_switch_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_floor_mode_switch_total{{from=\"static\",to=\"adaptive\"}} {}", - if me_allows_normal { - stats.get_me_floor_mode_switch_static_to_adaptive_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_me_floor_mode_switch_total{{from=\"adaptive\",to=\"static\"}} {}", - if me_allows_normal { - stats.get_me_floor_mode_switch_adaptive_to_static_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_adaptive_floor_cpu_cores_detected Runtime detected logical CPU cores for adaptive floor" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_adaptive_floor_cpu_cores_detected gauge" - ); - let _ = writeln!( - out, - "telemt_me_adaptive_floor_cpu_cores_detected {}", - if me_allows_normal { - stats.get_me_floor_cpu_cores_detected_gauge() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_adaptive_floor_cpu_cores_effective Runtime effective logical CPU cores for adaptive floor" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_adaptive_floor_cpu_cores_effective gauge" - ); - let _ = writeln!( - out, - "telemt_me_adaptive_floor_cpu_cores_effective {}", - if me_allows_normal { - stats.get_me_floor_cpu_cores_effective_gauge() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_adaptive_floor_global_cap_raw Runtime raw global adaptive floor cap" - ); - let _ = writeln!(out, "# TYPE telemt_me_adaptive_floor_global_cap_raw gauge"); - let _ = writeln!( - out, - "telemt_me_adaptive_floor_global_cap_raw {}", - if me_allows_normal { - stats.get_me_floor_global_cap_raw_gauge() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_adaptive_floor_global_cap_effective Runtime effective global adaptive floor cap" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_adaptive_floor_global_cap_effective gauge" - ); - let _ = writeln!( - out, - "telemt_me_adaptive_floor_global_cap_effective {}", - if me_allows_normal { - stats.get_me_floor_global_cap_effective_gauge() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_adaptive_floor_target_writers_total Runtime adaptive floor target writers total" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_adaptive_floor_target_writers_total gauge" - ); - let _ = writeln!( - out, - "telemt_me_adaptive_floor_target_writers_total {}", - if me_allows_normal { - stats.get_me_floor_target_writers_total_gauge() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_adaptive_floor_active_cap_configured Runtime configured active writer cap" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_adaptive_floor_active_cap_configured gauge" - ); - let _ = writeln!( - out, - "telemt_me_adaptive_floor_active_cap_configured {}", - if me_allows_normal { - stats.get_me_floor_active_cap_configured_gauge() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_adaptive_floor_active_cap_effective Runtime effective active writer cap" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_adaptive_floor_active_cap_effective gauge" - ); - let _ = writeln!( - out, - "telemt_me_adaptive_floor_active_cap_effective {}", - if me_allows_normal { - stats.get_me_floor_active_cap_effective_gauge() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_adaptive_floor_warm_cap_configured Runtime configured warm writer cap" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_adaptive_floor_warm_cap_configured gauge" - ); - let _ = writeln!( - out, - "telemt_me_adaptive_floor_warm_cap_configured {}", - if me_allows_normal { - stats.get_me_floor_warm_cap_configured_gauge() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_adaptive_floor_warm_cap_effective Runtime effective warm writer cap" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_adaptive_floor_warm_cap_effective gauge" - ); - let _ = writeln!( - out, - "telemt_me_adaptive_floor_warm_cap_effective {}", - if me_allows_normal { - stats.get_me_floor_warm_cap_effective_gauge() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_writers_active_current Current non-draining active ME writers" - ); - let _ = writeln!(out, "# TYPE telemt_me_writers_active_current gauge"); - let _ = writeln!( - out, - "telemt_me_writers_active_current {}", - if me_allows_normal { - stats.get_me_writers_active_current_gauge() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_writers_warm_current Current non-draining warm ME writers" - ); - let _ = writeln!(out, "# TYPE telemt_me_writers_warm_current gauge"); - let _ = writeln!( - out, - "telemt_me_writers_warm_current {}", - if me_allows_normal { - stats.get_me_writers_warm_current_gauge() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_floor_cap_block_total Reconnect attempts blocked by adaptive floor caps" - ); - let _ = writeln!(out, "# TYPE telemt_me_floor_cap_block_total counter"); - let _ = writeln!( - out, - "telemt_me_floor_cap_block_total {}", - if me_allows_normal { - stats.get_me_floor_cap_block_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_floor_swap_idle_total Adaptive floor cap recovery via idle writer swap" - ); - let _ = writeln!(out, "# TYPE telemt_me_floor_swap_idle_total counter"); - let _ = writeln!( - out, - "telemt_me_floor_swap_idle_total {}", - if me_allows_normal { - stats.get_me_floor_swap_idle_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_floor_swap_idle_failed_total Failed idle swap attempts under adaptive floor caps" - ); - let _ = writeln!(out, "# TYPE telemt_me_floor_swap_idle_failed_total counter"); - let _ = writeln!( - out, - "telemt_me_floor_swap_idle_failed_total {}", - if me_allows_normal { - stats.get_me_floor_swap_idle_failed_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_secure_padding_invalid_total Invalid secure frame lengths" - ); - let _ = writeln!(out, "# TYPE telemt_secure_padding_invalid_total counter"); - let _ = writeln!( - out, - "telemt_secure_padding_invalid_total {}", - if me_allows_normal { - stats.get_secure_padding_invalid() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_desync_total Total crypto-desync detections" - ); - let _ = writeln!(out, "# TYPE telemt_desync_total counter"); - let _ = writeln!( - out, - "telemt_desync_total {}", - if me_allows_normal { - stats.get_desync_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_desync_full_logged_total Full forensic desync logs emitted" - ); - let _ = writeln!(out, "# TYPE telemt_desync_full_logged_total counter"); - let _ = writeln!( - out, - "telemt_desync_full_logged_total {}", - if me_allows_normal { - stats.get_desync_full_logged() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_desync_suppressed_total Suppressed desync forensic events" - ); - let _ = writeln!(out, "# TYPE telemt_desync_suppressed_total counter"); - let _ = writeln!( - out, - "telemt_desync_suppressed_total {}", - if me_allows_normal { - stats.get_desync_suppressed() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_desync_frames_bucket_total Desync count by frames_ok bucket" - ); - let _ = writeln!(out, "# TYPE telemt_desync_frames_bucket_total counter"); - let _ = writeln!( - out, - "telemt_desync_frames_bucket_total{{bucket=\"0\"}} {}", - if me_allows_normal { - stats.get_desync_frames_bucket_0() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_desync_frames_bucket_total{{bucket=\"1_2\"}} {}", - if me_allows_normal { - stats.get_desync_frames_bucket_1_2() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_desync_frames_bucket_total{{bucket=\"3_10\"}} {}", - if me_allows_normal { - stats.get_desync_frames_bucket_3_10() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_desync_frames_bucket_total{{bucket=\"gt_10\"}} {}", - if me_allows_normal { - stats.get_desync_frames_bucket_gt_10() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_pool_swap_total Successful ME pool swaps" - ); - let _ = writeln!(out, "# TYPE telemt_pool_swap_total counter"); - let _ = writeln!( - out, - "telemt_pool_swap_total {}", - if me_allows_normal { - stats.get_pool_swap_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_pool_drain_active Active draining ME writers" - ); - let _ = writeln!(out, "# TYPE telemt_pool_drain_active gauge"); - let _ = writeln!( - out, - "telemt_pool_drain_active {}", - if me_allows_debug { - stats.get_pool_drain_active() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_pool_force_close_total Forced close events for draining writers" - ); - let _ = writeln!(out, "# TYPE telemt_pool_force_close_total counter"); - let _ = writeln!( - out, - "telemt_pool_force_close_total {}", - if me_allows_normal { - stats.get_pool_force_close_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_pool_stale_pick_total Stale writer fallback picks for new binds" - ); - let _ = writeln!(out, "# TYPE telemt_pool_stale_pick_total counter"); - let _ = writeln!( - out, - "telemt_pool_stale_pick_total {}", - if me_allows_normal { - stats.get_pool_stale_pick_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_writer_removed_total Total ME writer removals" - ); - let _ = writeln!(out, "# TYPE telemt_me_writer_removed_total counter"); - let _ = writeln!( - out, - "telemt_me_writer_removed_total {}", - if me_allows_debug { - stats.get_me_writer_removed_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_writer_removed_unexpected_total Unexpected ME writer removals that triggered refill" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_writer_removed_unexpected_total counter" - ); - let _ = writeln!( - out, - "telemt_me_writer_removed_unexpected_total {}", - if me_allows_normal { - stats.get_me_writer_removed_unexpected_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_refill_triggered_total Immediate ME refill runs started" - ); - let _ = writeln!(out, "# TYPE telemt_me_refill_triggered_total counter"); - let _ = writeln!( - out, - "telemt_me_refill_triggered_total {}", - if me_allows_debug { - stats.get_me_refill_triggered_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_refill_skipped_inflight_total Immediate ME refill skips due to inflight dedup" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_refill_skipped_inflight_total counter" - ); - let _ = writeln!( - out, - "telemt_me_refill_skipped_inflight_total {}", - if me_allows_debug { - stats.get_me_refill_skipped_inflight_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_refill_failed_total Immediate ME refill failures" - ); - let _ = writeln!(out, "# TYPE telemt_me_refill_failed_total counter"); - let _ = writeln!( - out, - "telemt_me_refill_failed_total {}", - if me_allows_normal { - stats.get_me_refill_failed_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_writer_restored_same_endpoint_total Refilled ME writer restored on the same endpoint" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_writer_restored_same_endpoint_total counter" - ); - let _ = writeln!( - out, - "telemt_me_writer_restored_same_endpoint_total {}", - if me_allows_normal { - stats.get_me_writer_restored_same_endpoint_total() - } else { - 0 - } - ); - - let _ = writeln!( - out, - "# HELP telemt_me_writer_restored_fallback_total Refilled ME writer restored via fallback endpoint" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_writer_restored_fallback_total counter" - ); - let _ = writeln!( - out, - "telemt_me_writer_restored_fallback_total {}", - if me_allows_normal { - stats.get_me_writer_restored_fallback_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_no_writer_failfast_total ME route failfast errors due to missing writer in bounded wait window" - ); - let _ = writeln!(out, "# TYPE telemt_me_no_writer_failfast_total counter"); - let _ = writeln!( - out, - "telemt_me_no_writer_failfast_total {}", - if me_allows_normal { - stats.get_me_no_writer_failfast_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_hybrid_timeout_total ME hybrid route timeouts after bounded retry window" - ); - let _ = writeln!(out, "# TYPE telemt_me_hybrid_timeout_total counter"); - let _ = writeln!( - out, - "telemt_me_hybrid_timeout_total {}", - if me_allows_normal { - stats.get_me_hybrid_timeout_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_async_recovery_trigger_total Async ME recovery trigger attempts from route path" - ); - let _ = writeln!(out, "# TYPE telemt_me_async_recovery_trigger_total counter"); - let _ = writeln!( - out, - "telemt_me_async_recovery_trigger_total {}", - if me_allows_normal { - stats.get_me_async_recovery_trigger_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "# HELP telemt_me_inline_recovery_total Legacy inline ME recovery attempts from route path" - ); - let _ = writeln!(out, "# TYPE telemt_me_inline_recovery_total counter"); - let _ = writeln!( - out, - "telemt_me_inline_recovery_total {}", - if me_allows_normal { - stats.get_me_inline_recovery_total() - } else { - 0 - } - ); - - let unresolved_writer_losses = if me_allows_normal { - stats - .get_me_writer_removed_unexpected_total() - .saturating_sub( - stats - .get_me_writer_restored_same_endpoint_total() - .saturating_add(stats.get_me_writer_restored_fallback_total()), - ) - } else { - 0 - }; - let _ = writeln!( - out, - "# HELP telemt_me_writer_removed_unexpected_minus_restored_total Unexpected writer removals not yet compensated by restore" - ); - let _ = writeln!( - out, - "# TYPE telemt_me_writer_removed_unexpected_minus_restored_total gauge" - ); - let _ = writeln!( - out, - "telemt_me_writer_removed_unexpected_minus_restored_total {}", - unresolved_writer_losses - ); - - let _ = writeln!( - out, - "# HELP telemt_user_connections_total Per-user total connections" - ); - let _ = writeln!(out, "# TYPE telemt_user_connections_total counter"); - let _ = writeln!( - out, - "# HELP telemt_user_connections_current Per-user active connections" - ); - let _ = writeln!(out, "# TYPE telemt_user_connections_current gauge"); - let _ = writeln!( - out, - "# HELP telemt_user_octets_from_client_total Per-user total bytes received" - ); - let _ = writeln!(out, "# TYPE telemt_user_octets_from_client_total counter"); - let _ = writeln!( - out, - "# HELP telemt_user_octets_to_client_total Per-user total bytes sent" - ); - let _ = writeln!(out, "# TYPE telemt_user_octets_to_client_total counter"); - let _ = writeln!( - out, - "# HELP telemt_user_msgs_from_client_total Per-user total messages received" - ); - let _ = writeln!(out, "# TYPE telemt_user_msgs_from_client_total counter"); - let _ = writeln!( - out, - "# HELP telemt_user_msgs_to_client_total Per-user total messages sent" - ); - let _ = writeln!(out, "# TYPE telemt_user_msgs_to_client_total counter"); - let _ = writeln!( - out, - "# HELP telemt_ip_reservation_rollback_total IP reservation rollbacks caused by later limit checks" - ); - let _ = writeln!(out, "# TYPE telemt_ip_reservation_rollback_total counter"); - let _ = writeln!( - out, - "telemt_ip_reservation_rollback_total{{reason=\"tcp_limit\"}} {}", - if core_enabled { - stats.get_ip_reservation_rollback_tcp_limit_total() - } else { - 0 - } - ); - let _ = writeln!( - out, - "telemt_ip_reservation_rollback_total{{reason=\"quota_limit\"}} {}", - if core_enabled { - stats.get_ip_reservation_rollback_quota_limit_total() - } else { - 0 - } - ); - let ip_memory = ip_tracker.memory_stats().await; - let _ = writeln!( - out, - "# HELP telemt_ip_tracker_users Number of users tracked by IP limiter state" - ); - let _ = writeln!(out, "# TYPE telemt_ip_tracker_users gauge"); - let _ = writeln!( - out, - "telemt_ip_tracker_users{{scope=\"active\"}} {}", - ip_memory.active_users - ); - let _ = writeln!( - out, - "telemt_ip_tracker_users{{scope=\"recent\"}} {}", - ip_memory.recent_users - ); - let _ = writeln!( - out, - "# HELP telemt_ip_tracker_entries Number of IP entries tracked by limiter state" - ); - let _ = writeln!(out, "# TYPE telemt_ip_tracker_entries gauge"); - let _ = writeln!( - out, - "telemt_ip_tracker_entries{{scope=\"active\"}} {}", - ip_memory.active_entries - ); - let _ = writeln!( - out, - "telemt_ip_tracker_entries{{scope=\"recent\"}} {}", - ip_memory.recent_entries - ); - let _ = writeln!( - out, - "# HELP telemt_ip_tracker_cleanup_queue_len Deferred disconnect cleanup queue length" - ); - let _ = writeln!(out, "# TYPE telemt_ip_tracker_cleanup_queue_len gauge"); - let _ = writeln!( - out, - "telemt_ip_tracker_cleanup_queue_len {}", - ip_memory.cleanup_queue_len - ); - let _ = writeln!( - out, - "# HELP telemt_ip_tracker_cleanup_total Release cleanups deferred through the cleanup queue" - ); - let _ = writeln!(out, "# TYPE telemt_ip_tracker_cleanup_total counter"); - let _ = writeln!( - out, - "telemt_ip_tracker_cleanup_total{{path=\"deferred\"}} {}", - ip_memory.cleanup_deferred_releases - ); - let _ = writeln!( - out, - "# HELP telemt_ip_tracker_cap_rejects_total New connection rejects caused by global IP tracker caps" - ); - let _ = writeln!(out, "# TYPE telemt_ip_tracker_cap_rejects_total counter"); - let _ = writeln!( - out, - "telemt_ip_tracker_cap_rejects_total{{scope=\"active\"}} {}", - ip_memory.active_cap_rejects - ); - let _ = writeln!( - out, - "telemt_ip_tracker_cap_rejects_total{{scope=\"recent\"}} {}", - ip_memory.recent_cap_rejects - ); - - let mut user_stats_emitted = 0usize; - let mut user_stats_suppressed = 0usize; - let mut unique_ip_emitted = 0usize; - let mut unique_ip_suppressed = 0usize; - - if user_enabled { - for entry in stats.iter_user_stats() { - if user_stats_emitted >= USER_LABELED_METRICS_MAX_USERS { - user_stats_suppressed = user_stats_suppressed.saturating_add(1); - continue; - } - let user = entry.key(); - let s = entry.value(); - user_stats_emitted = user_stats_emitted.saturating_add(1); - let _ = writeln!( - out, - "telemt_user_connections_total{{user=\"{}\"}} {}", - user, - s.connects.load(std::sync::atomic::Ordering::Relaxed) - ); - let _ = writeln!( - out, - "telemt_user_connections_current{{user=\"{}\"}} {}", - user, - s.curr_connects.load(std::sync::atomic::Ordering::Relaxed) - ); - let _ = writeln!( - out, - "telemt_user_octets_from_client_total{{user=\"{}\"}} {}", - user, - s.octets_from_client - .load(std::sync::atomic::Ordering::Relaxed) - ); - let _ = writeln!( - out, - "telemt_user_octets_to_client_total{{user=\"{}\"}} {}", - user, - s.octets_to_client - .load(std::sync::atomic::Ordering::Relaxed) - ); - let _ = writeln!( - out, - "telemt_user_msgs_from_client_total{{user=\"{}\"}} {}", - user, - s.msgs_from_client - .load(std::sync::atomic::Ordering::Relaxed) - ); - let _ = writeln!( - out, - "telemt_user_msgs_to_client_total{{user=\"{}\"}} {}", - user, - s.msgs_to_client.load(std::sync::atomic::Ordering::Relaxed) - ); - } - - let ip_stats = ip_tracker.get_stats_snapshot().await; - let ip_counts: HashMap = ip_stats - .into_iter() - .map(|(user, count, _)| (user, count)) - .collect(); - - let mut unique_users = BTreeSet::new(); - unique_users.extend(config.access.users.keys().cloned()); - unique_users.extend(config.access.user_max_unique_ips.keys().cloned()); - unique_users.extend(ip_counts.keys().cloned()); - let unique_users_vec: Vec = unique_users.iter().cloned().collect(); - let recent_counts = ip_tracker - .get_recent_counts_for_users_snapshot(&unique_users_vec) - .await; - - let _ = writeln!( - out, - "# HELP telemt_user_unique_ips_current Per-user current number of unique active IPs" - ); - let _ = writeln!(out, "# TYPE telemt_user_unique_ips_current gauge"); - let _ = writeln!( - out, - "# HELP telemt_user_unique_ips_recent_window Per-user unique IPs seen in configured observation window" - ); - let _ = writeln!(out, "# TYPE telemt_user_unique_ips_recent_window gauge"); - let _ = writeln!( - out, - "# HELP telemt_user_unique_ips_limit Effective per-user unique IP limit (0 means unlimited)" - ); - let _ = writeln!(out, "# TYPE telemt_user_unique_ips_limit gauge"); - let _ = writeln!( - out, - "# HELP telemt_user_unique_ips_utilization Per-user unique IP usage ratio (0 for unlimited)" - ); - let _ = writeln!(out, "# TYPE telemt_user_unique_ips_utilization gauge"); - - for user in unique_users { - if unique_ip_emitted >= USER_LABELED_METRICS_MAX_USERS { - unique_ip_suppressed = unique_ip_suppressed.saturating_add(1); - continue; - } - unique_ip_emitted = unique_ip_emitted.saturating_add(1); - let current = ip_counts.get(&user).copied().unwrap_or(0); - let limit = config - .access - .user_max_unique_ips - .get(&user) - .copied() - .filter(|limit| *limit > 0) - .or((config.access.user_max_unique_ips_global_each > 0) - .then_some(config.access.user_max_unique_ips_global_each)) - .unwrap_or(0); - let utilization = if limit > 0 { - current as f64 / limit as f64 - } else { - 0.0 - }; - let _ = writeln!( - out, - "telemt_user_unique_ips_current{{user=\"{}\"}} {}", - user, current - ); - let _ = writeln!( - out, - "telemt_user_unique_ips_recent_window{{user=\"{}\"}} {}", - user, - recent_counts.get(&user).copied().unwrap_or(0) - ); - let _ = writeln!( - out, - "telemt_user_unique_ips_limit{{user=\"{}\"}} {}", - user, limit - ); - let _ = writeln!( - out, - "telemt_user_unique_ips_utilization{{user=\"{}\"}} {:.6}", - user, utilization - ); - } - } - - let _ = writeln!( - out, - "# HELP telemt_telemetry_user_series_suppressed User-labeled metric series suppression flag" - ); - let _ = writeln!(out, "# TYPE telemt_telemetry_user_series_suppressed gauge"); - let _ = writeln!( - out, - "telemt_telemetry_user_series_suppressed {}", - if user_enabled && user_stats_suppressed == 0 && unique_ip_suppressed == 0 { - 0 - } else { - 1 - } - ); - let _ = writeln!( - out, - "# HELP telemt_telemetry_user_series_users User-labeled metric users by export status" - ); - let _ = writeln!(out, "# TYPE telemt_telemetry_user_series_users gauge"); - let _ = writeln!( - out, - "telemt_telemetry_user_series_users{{family=\"stats\",status=\"emitted\"}} {}", - user_stats_emitted - ); - let _ = writeln!( - out, - "telemt_telemetry_user_series_users{{family=\"stats\",status=\"suppressed\"}} {}", - user_stats_suppressed - ); - let _ = writeln!( - out, - "telemt_telemetry_user_series_users{{family=\"unique_ip\",status=\"emitted\"}} {}", - unique_ip_emitted - ); - let _ = writeln!( - out, - "telemt_telemetry_user_series_users{{family=\"unique_ip\",status=\"suppressed\"}} {}", - unique_ip_suppressed - ); - - web::render(&mut out, web_publication, config); - out -} +// Ordered Prometheus text rendering split by bounded metric families. +mod render; +use render::render_metrics; #[cfg(test)] -mod tests { - use super::*; - use http_body_util::BodyExt; - use std::net::IpAddr; - use std::time::SystemTime; - - use crate::tls_front::types::{ - CachedTlsData, ParsedServerHello, TlsBehaviorProfile, TlsCertPayload, TlsProfileSource, - }; - - fn test_web_publication() -> crate::web::control::WebRuntimePublication { - let control = crate::web::control::WebRuntimeControl::new(); - control.subscribe().borrow().clone() - } - - #[tokio::test] - async fn test_render_metrics_format() { - let stats = Arc::new(Stats::new()); - let shared_state = ProxySharedState::new(); - let tracker = UserIpTracker::new(); - let mut config = ProxyConfig::default(); - config - .access - .user_max_unique_ips - .insert("alice".to_string(), 4); - - stats.increment_connects_all(); - stats.increment_connects_all(); - stats.increment_connects_bad_with_class("tls_handshake_bad_client"); - stats.increment_handshake_timeouts(); - stats.increment_handshake_failure_class("timeout"); - shared_state - .handshake - .auth_expensive_checks_total - .fetch_add(9, std::sync::atomic::Ordering::Relaxed); - shared_state - .handshake - .auth_budget_exhausted_total - .fetch_add(2, std::sync::atomic::Ordering::Relaxed); - stats.increment_upstream_connect_attempt_total(); - stats.increment_upstream_connect_attempt_total(); - stats.increment_upstream_connect_success_total(); - stats.increment_upstream_connect_fail_total(); - stats.increment_upstream_connect_failfast_hard_error_total(); - stats.observe_upstream_connect_attempts_per_request(2); - stats.observe_upstream_connect_duration_ms(220, true); - stats.observe_upstream_connect_duration_ms(1500, false); - stats.increment_me_rpc_proxy_req_signal_sent_total(); - stats.increment_me_rpc_proxy_req_signal_failed_total(); - stats.increment_me_rpc_proxy_req_signal_skipped_no_meta_total(); - stats.increment_me_rpc_proxy_req_signal_response_total(); - stats.increment_me_rpc_proxy_req_signal_close_sent_total(); - stats.increment_me_idle_close_by_peer_total(); - stats.increment_relay_idle_soft_mark_total(); - stats.increment_relay_idle_hard_close_total(); - stats.increment_relay_pressure_evict_total(); - stats.increment_relay_protocol_desync_close_total(); - stats.increment_me_d2c_batches_total(); - stats.add_me_d2c_batch_frames_total(3); - stats.add_me_d2c_batch_bytes_total(2048); - stats.increment_me_d2c_flush_reason(crate::stats::MeD2cFlushReason::AckImmediate); - stats.increment_me_d2c_data_frames_total(); - stats.increment_me_d2c_ack_frames_total(); - stats.add_me_d2c_payload_bytes_total(1800); - stats.increment_me_d2c_write_mode(crate::stats::MeD2cWriteMode::Coalesced); - stats.increment_me_d2c_quota_reject_total(crate::stats::MeD2cQuotaRejectStage::PostWrite); - stats.observe_me_d2c_frame_buf_shrink(4096); - stats.increment_me_endpoint_quarantine_total(); - stats.increment_me_endpoint_quarantine_unexpected_total(); - stats.increment_me_endpoint_quarantine_draining_suppressed_total(); - stats.increment_user_connects("alice"); - stats.increment_user_curr_connects("alice"); - stats.add_user_octets_from("alice", 1024); - stats.add_user_octets_to("alice", 2048); - stats.increment_user_msgs_from("alice"); - stats.increment_user_msgs_to("alice"); - stats.increment_user_msgs_to("alice"); - tracker - .check_and_add("alice", "203.0.113.10".parse().unwrap()) - .await - .unwrap(); - - let output = render_metrics( - &stats, - shared_state.as_ref(), - &config, - &tracker, - None, - &TlsFullCertBudget::new(), - &test_web_publication(), - ) - .await; - - assert!(output.contains(&format!( - "telemt_build_info{{version=\"{}\"}} 1", - env!("CARGO_PKG_VERSION") - ))); - assert!(output.contains("telemt_connections_total 2")); - assert!(output.contains("telemt_connections_bad_total 1")); - assert!(output.contains( - "telemt_connections_bad_by_class_total{class=\"tls_handshake_bad_client\"} 1" - )); - assert!(output.contains("telemt_handshake_timeouts_total 1")); - assert!(output.contains("telemt_handshake_failures_by_class_total{class=\"timeout\"} 1")); - assert!(output.contains("telemt_auth_expensive_checks_total 9")); - assert!(output.contains("telemt_auth_budget_exhausted_total 2")); - assert!(output.contains("telemt_upstream_connect_attempt_total 2")); - assert!(output.contains("telemt_upstream_connect_success_total 1")); - assert!(output.contains("telemt_upstream_connect_fail_total 1")); - assert!(output.contains("telemt_upstream_connect_failfast_hard_error_total 1")); - assert!(output.contains("telemt_upstream_connect_attempts_per_request{bucket=\"2\"} 1")); - assert!( - output - .contains("telemt_upstream_connect_duration_success_total{bucket=\"101_500ms\"} 1") - ); - assert!( - output.contains("telemt_upstream_connect_duration_fail_total{bucket=\"gt_1000ms\"} 1") - ); - assert!(output.contains("telemt_me_rpc_proxy_req_signal_sent_total 1")); - assert!(output.contains("telemt_me_rpc_proxy_req_signal_failed_total 1")); - assert!(output.contains("telemt_me_rpc_proxy_req_signal_skipped_no_meta_total 1")); - assert!(output.contains("telemt_me_rpc_proxy_req_signal_response_total 1")); - assert!(output.contains("telemt_me_rpc_proxy_req_signal_close_sent_total 1")); - assert!(output.contains("telemt_me_idle_close_by_peer_total 1")); - assert!(output.contains("telemt_relay_idle_soft_mark_total 1")); - assert!(output.contains("telemt_relay_idle_hard_close_total 1")); - assert!(output.contains("telemt_relay_pressure_evict_total 1")); - assert!(output.contains("telemt_relay_protocol_desync_close_total 1")); - assert!(output.contains("telemt_me_d2c_batches_total 1")); - assert!(output.contains("telemt_me_d2c_batch_frames_total 3")); - assert!(output.contains("telemt_me_d2c_batch_bytes_total 2048")); - assert!(output.contains("telemt_me_d2c_flush_reason_total{reason=\"ack_immediate\"} 1")); - assert!(output.contains("telemt_me_d2c_data_frames_total 1")); - assert!(output.contains("telemt_me_d2c_ack_frames_total 1")); - assert!(output.contains("telemt_me_d2c_payload_bytes_total 1800")); - assert!(output.contains("telemt_me_d2c_write_mode_total{mode=\"coalesced\"} 1")); - assert!(output.contains("telemt_me_d2c_quota_reject_total{stage=\"post_write\"} 1")); - assert!(output.contains("telemt_me_d2c_frame_buf_shrink_total 1")); - assert!(output.contains("telemt_me_d2c_frame_buf_shrink_bytes_total 4096")); - assert!(output.contains("telemt_me_endpoint_quarantine_total 1")); - assert!(output.contains("telemt_me_endpoint_quarantine_unexpected_total 1")); - assert!(output.contains("telemt_me_endpoint_quarantine_draining_suppressed_total 1")); - assert!(output.contains("telemt_user_connections_total{user=\"alice\"} 1")); - assert!(output.contains("telemt_user_connections_current{user=\"alice\"} 1")); - assert!(output.contains("telemt_user_octets_from_client_total{user=\"alice\"} 1024")); - assert!(output.contains("telemt_user_octets_to_client_total{user=\"alice\"} 2048")); - assert!(output.contains("telemt_user_msgs_from_client_total{user=\"alice\"} 1")); - assert!(output.contains("telemt_user_msgs_to_client_total{user=\"alice\"} 2")); - assert!(output.contains("telemt_user_unique_ips_current{user=\"alice\"} 1")); - assert!(output.contains("telemt_user_unique_ips_recent_window{user=\"alice\"} 1")); - assert!(output.contains("telemt_user_unique_ips_limit{user=\"alice\"} 4")); - assert!(output.contains("telemt_user_unique_ips_utilization{user=\"alice\"} 0.250000")); - assert!(output.contains("telemt_ip_tracker_users{scope=\"active\"} 1")); - assert!(output.contains("telemt_ip_tracker_entries{scope=\"active\"} 1")); - assert!(output.contains("telemt_ip_tracker_cleanup_queue_len 0")); - } - - #[tokio::test] - async fn test_render_tls_front_profile_health() { - let stats = Stats::new(); - let shared_state = ProxySharedState::new(); - let tracker = UserIpTracker::new(); - let mut config = ProxyConfig::default(); - config.censorship.tls_domain = "primary.example".to_string(); - config.censorship.tls_domains = vec!["fallback.example".to_string()]; - - let cache = TlsFrontCache::new( - &[ - "primary.example".to_string(), - "fallback.example".to_string(), - ], - 1024, - "tlsfront-profile-health-test", - ); - cache - .set( - "primary.example", - CachedTlsData { - server_hello_template: ParsedServerHello { - version: [0x03, 0x03], - random: [0u8; 32], - session_id: Vec::new(), - cipher_suite: [0x13, 0x01], - compression: 0, - extensions: { - let mut key_share = vec![0x00, 0x1d, 0x00, 0x20]; - key_share.resize(36, 0x42); - vec![ - crate::tls_front::types::TlsExtension { - ext_type: 0x002b, - data: vec![0x03, 0x04], - }, - crate::tls_front::types::TlsExtension { - ext_type: 0x0033, - data: key_share, - }, - ] - }, - }, - cert_info: None, - cert_payload: Some(TlsCertPayload { - cert_chain_der: vec![vec![0x30, 0x01]], - certificate_message: vec![0x0b, 0x00, 0x00, 0x00], - }), - app_data_records_sizes: vec![1024, 512], - total_app_data_len: 1536, - behavior_profile: TlsBehaviorProfile { - change_cipher_spec_count: 1, - app_data_record_sizes: vec![1024, 512], - ticket_record_sizes: vec![69], - source: TlsProfileSource::Merged, - ..TlsBehaviorProfile::default() - }, - fetched_at: SystemTime::now(), - domain: "primary.example".to_string(), - }, - ) - .await; - - let output = render_metrics( - &stats, - &shared_state, - &config, - &tracker, - Some(&cache), - &TlsFullCertBudget::new(), - &test_web_publication(), - ) - .await; - - assert!(output.contains("telemt_tls_front_profile_domains{status=\"configured\"} 2")); - assert!(output.contains("telemt_tls_front_profile_domains{status=\"emitted\"} 2")); - assert!(output.contains("telemt_tls_front_profile_domains{status=\"suppressed\"} 0")); - assert!( - output.contains("telemt_tls_front_profile_info{domain=\"primary.example\",source=\"merged\",is_default=\"false\",has_cert_info=\"false\",has_cert_payload=\"true\"} 1") - ); - assert!( - output.contains("telemt_tls_front_profile_info{domain=\"fallback.example\",source=\"default\",is_default=\"true\",has_cert_info=\"false\",has_cert_payload=\"false\"} 1") - ); - assert!( - output.contains("telemt_tls_front_profile_quality_info{domain=\"primary.example\",quality=\"raw_strict\",key_share_group=\"x25519\"} 1") - ); - assert!( - output.contains("telemt_tls_front_profile_quality_info{domain=\"fallback.example\",quality=\"fallback\",key_share_group=\"none\"} 1") - ); - assert!(output.contains( - "telemt_tls_front_profile_server_hello_bytes{domain=\"primary.example\"} 90" - )); - assert!(output.contains( - "telemt_tls_front_profile_server_hello_extensions{domain=\"primary.example\"} 2" - )); - assert!( - output.contains( - "telemt_tls_front_profile_app_data_records{domain=\"primary.example\"} 2" - ) - ); - assert!( - output - .contains("telemt_tls_front_profile_ticket_records{domain=\"primary.example\"} 1") - ); - assert!(output.contains( - "telemt_tls_front_profile_change_cipher_spec_records{domain=\"primary.example\"} 1" - )); - assert!( - output.contains( - "telemt_tls_front_profile_app_data_bytes{domain=\"primary.example\"} 1536" - ) - ); - } - - #[tokio::test] - async fn process_tls_budget_metrics_survive_a_generation_without_tls_cache() { - let stats = Stats::new(); - let shared_state = ProxySharedState::new(); - let tracker = UserIpTracker::new(); - let config = ProxyConfig::default(); - let budget = Arc::new(TlsFullCertBudget::new()); - let cache = TlsFrontCache::new_with_full_cert_budget( - &["example.com".to_string()], - 1024, - "tlsfront-test-cache", - Arc::clone(&budget), - ); - assert!( - cache - .take_full_cert_budget_for_ip( - "example.com", - "127.0.0.1".parse().unwrap(), - Duration::from_secs(60), - ) - .await - ); - - let output = render_metrics( - &stats, - &shared_state, - &config, - &tracker, - None, - budget.as_ref(), - &test_web_publication(), - ) - .await; - - assert!(output.contains("telemt_tls_front_full_cert_budget_entries 1")); - } - - #[tokio::test] - async fn test_render_empty_stats() { - let stats = Stats::new(); - let shared_state = ProxySharedState::new(); - let tracker = UserIpTracker::new(); - let config = ProxyConfig::default(); - let output = render_metrics( - &stats, - &shared_state, - &config, - &tracker, - None, - &TlsFullCertBudget::new(), - &test_web_publication(), - ) - .await; - assert!(output.contains("telemt_connections_total 0")); - assert!(output.contains("telemt_connections_bad_total 0")); - assert!(output.contains("telemt_handshake_timeouts_total 0")); - assert!(output.contains("telemt_auth_expensive_checks_total 0")); - assert!(output.contains("telemt_auth_budget_exhausted_total 0")); - assert!(output.contains("telemt_user_unique_ips_current{user=")); - assert!(output.contains("telemt_user_unique_ips_recent_window{user=")); - } - - #[tokio::test] - async fn test_render_uses_global_each_unique_ip_limit() { - let stats = Stats::new(); - let shared_state = ProxySharedState::new(); - stats.increment_user_connects("alice"); - stats.increment_user_curr_connects("alice"); - let tracker = UserIpTracker::new(); - tracker - .check_and_add("alice", "203.0.113.10".parse().unwrap()) - .await - .unwrap(); - let mut config = ProxyConfig::default(); - config.access.user_max_unique_ips_global_each = 2; - - let output = render_metrics( - &stats, - &shared_state, - &config, - &tracker, - None, - &TlsFullCertBudget::new(), - &test_web_publication(), - ) - .await; - - assert!(output.contains("telemt_user_unique_ips_limit{user=\"alice\"} 2")); - assert!(output.contains("telemt_user_unique_ips_utilization{user=\"alice\"} 0.500000")); - } - - #[tokio::test] - async fn test_render_has_type_annotations() { - let stats = Stats::new(); - let shared_state = ProxySharedState::new(); - let tracker = UserIpTracker::new(); - let config = ProxyConfig::default(); - let output = render_metrics( - &stats, - &shared_state, - &config, - &tracker, - None, - &TlsFullCertBudget::new(), - &test_web_publication(), - ) - .await; - assert!(output.contains("# TYPE telemt_uptime_seconds gauge")); - assert!(output.contains("# TYPE telemt_connections_total counter")); - assert!(output.contains("# TYPE telemt_connections_bad_total counter")); - assert!(output.contains("# TYPE telemt_connections_bad_by_class_total counter")); - assert!(output.contains("# TYPE telemt_handshake_timeouts_total counter")); - assert!(output.contains("# TYPE telemt_handshake_failures_by_class_total counter")); - assert!(output.contains("# TYPE telemt_auth_expensive_checks_total counter")); - assert!(output.contains("# TYPE telemt_auth_budget_exhausted_total counter")); - assert!(output.contains("# TYPE telemt_upstream_connect_attempt_total counter")); - assert!(output.contains("# TYPE telemt_me_rpc_proxy_req_signal_sent_total counter")); - assert!(output.contains("# TYPE telemt_me_idle_close_by_peer_total counter")); - assert!(output.contains("# TYPE telemt_relay_idle_soft_mark_total counter")); - assert!(output.contains("# TYPE telemt_relay_idle_hard_close_total counter")); - assert!(output.contains("# TYPE telemt_relay_pressure_evict_total counter")); - assert!(output.contains("# TYPE telemt_relay_protocol_desync_close_total counter")); - assert!(output.contains("# TYPE telemt_me_d2c_batches_total counter")); - assert!(output.contains("# TYPE telemt_me_d2c_flush_reason_total counter")); - assert!(output.contains("# TYPE telemt_me_d2c_write_mode_total counter")); - assert!(output.contains("# TYPE telemt_me_d2c_batch_frames_bucket_total counter")); - assert!(output.contains("# TYPE telemt_me_d2c_flush_duration_us_bucket_total counter")); - assert!(output.contains("# TYPE telemt_me_endpoint_quarantine_total counter")); - assert!(output.contains("# TYPE telemt_me_endpoint_quarantine_unexpected_total counter")); - assert!( - output - .contains("# TYPE telemt_me_endpoint_quarantine_draining_suppressed_total counter") - ); - assert!(output.contains("# TYPE telemt_me_writer_removed_total counter")); - assert!( - output - .contains("# TYPE telemt_me_writer_removed_unexpected_minus_restored_total gauge") - ); - assert!(output.contains("# TYPE telemt_user_unique_ips_current gauge")); - assert!(output.contains("# TYPE telemt_user_unique_ips_recent_window gauge")); - assert!(output.contains("# TYPE telemt_user_unique_ips_limit gauge")); - assert!(output.contains("# TYPE telemt_user_unique_ips_utilization gauge")); - assert!(output.contains("# TYPE telemt_stats_user_entries gauge")); - assert!(output.contains("# TYPE telemt_telemetry_user_series_users gauge")); - assert!(output.contains("# TYPE telemt_ip_tracker_users gauge")); - assert!(output.contains("# TYPE telemt_ip_tracker_entries gauge")); - assert!(output.contains("# TYPE telemt_ip_tracker_cleanup_queue_len gauge")); - assert!(output.contains("# TYPE telemt_ip_tracker_cleanup_total counter")); - assert!(output.contains("# TYPE telemt_ip_tracker_cap_rejects_total counter")); - assert!(output.contains("# TYPE telemt_tls_fetch_profile_cache_entries gauge")); - assert!(output.contains("# TYPE telemt_tls_fetch_profile_cache_cap_drops_total counter")); - assert!(output.contains("# TYPE telemt_tls_front_full_cert_budget_entries gauge")); - assert!( - output.contains("# TYPE telemt_tls_front_full_cert_budget_cap_drops_total counter") - ); - assert!(output.contains("# TYPE telemt_tls_front_profile_domains gauge")); - assert!(output.contains("# TYPE telemt_tls_front_profile_info gauge")); - assert!(output.contains("# TYPE telemt_tls_front_profile_quality_info gauge")); - assert!(output.contains("# TYPE telemt_tls_front_profile_age_seconds gauge")); - assert!(output.contains("# TYPE telemt_tls_front_profile_server_hello_bytes gauge")); - assert!(output.contains("# TYPE telemt_tls_front_profile_server_hello_extensions gauge")); - assert!(output.contains("# TYPE telemt_tls_front_profile_app_data_records gauge")); - assert!(output.contains("# TYPE telemt_tls_front_profile_ticket_records gauge")); - assert!( - output.contains("# TYPE telemt_tls_front_profile_change_cipher_spec_records gauge") - ); - assert!(output.contains("# TYPE telemt_tls_front_profile_app_data_bytes gauge")); - } - - #[tokio::test] - async fn test_endpoint_integration() { - let mut config = ProxyConfig::default(); - config.general.beobachten = true; - config.general.beobachten_minutes = 10; - let runtime = crate::maestro::generation::test_runtime_generation(1, config); - let web_publication = test_web_publication(); - let tls_full_cert_budget = TlsFullCertBudget::new(); - runtime.stats.increment_connects_all(); - runtime.stats.increment_connects_all(); - runtime.stats.increment_connects_all(); - - let req = Request::builder().uri("/metrics").body(()).unwrap(); - let resp = handle(req, &runtime, &web_publication, &tls_full_cert_budget) - .await - .unwrap(); - assert_eq!(resp.status(), StatusCode::OK); - let body = resp.into_body().collect().await.unwrap().to_bytes(); - assert!( - std::str::from_utf8(body.as_ref()) - .unwrap() - .contains("telemt_connections_total 3") - ); - assert!( - std::str::from_utf8(body.as_ref()) - .unwrap() - .contains(&format!( - "telemt_build_info{{version=\"{}\"}} 1", - env!("CARGO_PKG_VERSION") - )) - ); - - runtime.beobachten.record( - "TLS-scanner", - "203.0.113.10".parse::().unwrap(), - Duration::from_secs(600), - ); - let req_beob = Request::builder().uri("/beobachten").body(()).unwrap(); - let resp_beob = handle( - req_beob, - &runtime, - &web_publication, - &tls_full_cert_budget, - ) - .await - .unwrap(); - assert_eq!(resp_beob.status(), StatusCode::OK); - let body_beob = resp_beob.into_body().collect().await.unwrap().to_bytes(); - let beob_text = std::str::from_utf8(body_beob.as_ref()).unwrap(); - assert!(beob_text.contains("[TLS-scanner]")); - assert!(beob_text.contains("203.0.113.10-1")); - - let req404 = Request::builder().uri("/other").body(()).unwrap(); - let resp404 = handle(req404, &runtime, &web_publication, &tls_full_cert_budget) - .await - .unwrap(); - assert_eq!(resp404.status(), StatusCode::NOT_FOUND); - } -} +mod tests; diff --git a/src/metrics/render.rs b/src/metrics/render.rs new file mode 100644 index 0000000..e7af27d --- /dev/null +++ b/src/metrics/render.rs @@ -0,0 +1,78 @@ +use super::*; + +// Process, buffer, and TLS cache metrics. +mod process; +// Connection, quota, and conntrack metrics. +mod connections; +// Rate limiter, upstream, and initial ME metrics. +mod traffic; +// ME lifecycle and relay event metrics. +mod me_lifecycle; +// ME batching and resident-memory metrics. +mod me_buffers; +// ME writer selection, KDF, and hardswap metrics. +mod me_policy; +// Adaptive-floor and writer-cap metrics. +mod me_floor; +// Desync, pool recovery, and refill metrics. +mod me_recovery; +// Bounded per-user and IP-tracker metrics. +mod users; + +pub(super) async fn render_metrics( + stats: &Stats, + shared_state: &ProxySharedState, + config: &ProxyConfig, + ip_tracker: &UserIpTracker, + tls_cache: Option<&TlsFrontCache>, + tls_full_cert_budget: &TlsFullCertBudget, + web_publication: &crate::web::control::WebRuntimePublication, +) -> String { + let mut out = String::with_capacity(4096); + let telemetry = stats.telemetry_policy(); + let core_enabled = telemetry.core_enabled; + let user_enabled = telemetry.user_enabled; + let me_allows_normal = telemetry.me_level.allows_normal(); + let me_allows_debug = telemetry.me_level.allows_debug(); + + process::render( + &mut out, + stats, + shared_state, + telemetry, + tls_full_cert_budget, + ); + super::render_tls_front_profile_health(&mut out, config, tls_cache).await; + connections::render(&mut out, stats, shared_state, core_enabled); + traffic::render( + &mut out, + stats, + shared_state, + config, + core_enabled, + me_allows_normal, + me_allows_debug, + ); + me_lifecycle::render(&mut out, stats, me_allows_normal); + me_buffers::render( + &mut out, + stats, + core_enabled, + me_allows_normal, + me_allows_debug, + ); + me_policy::render(&mut out, stats, me_allows_normal, me_allows_debug); + me_floor::render(&mut out, stats, config, me_allows_normal); + me_recovery::render(&mut out, stats, me_allows_normal, me_allows_debug); + users::render( + &mut out, + stats, + config, + ip_tracker, + core_enabled, + user_enabled, + ) + .await; + super::web::render(&mut out, web_publication, config); + out +} diff --git a/src/metrics/render/connections.rs b/src/metrics/render/connections.rs new file mode 100644 index 0000000..f7c2ae0 --- /dev/null +++ b/src/metrics/render/connections.rs @@ -0,0 +1,339 @@ +use super::*; +use std::fmt::Write; + +pub(super) fn render( + out: &mut String, + stats: &Stats, + shared_state: &ProxySharedState, + core_enabled: bool, +) { + let _ = writeln!( + out, + "# HELP telemt_connections_total Total accepted connections" + ); + let _ = writeln!(out, "# TYPE telemt_connections_total counter"); + let _ = writeln!( + out, + "telemt_connections_total {}", + if core_enabled { + stats.get_connects_all() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_connections_bad_total Bad/rejected connections" + ); + let _ = writeln!(out, "# TYPE telemt_connections_bad_total counter"); + let _ = writeln!( + out, + "telemt_connections_bad_total {}", + if core_enabled { + stats.get_connects_bad() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_connections_bad_by_class_total Bad/rejected connections by class" + ); + let _ = writeln!(out, "# TYPE telemt_connections_bad_by_class_total counter"); + if core_enabled { + for (class, total) in stats.get_connects_bad_class_counts() { + let _ = writeln!( + out, + "telemt_connections_bad_by_class_total{{class=\"{}\"}} {}", + class, total + ); + } + } + + let _ = writeln!( + out, + "# HELP telemt_handshake_timeouts_total Handshake timeouts" + ); + let _ = writeln!(out, "# TYPE telemt_handshake_timeouts_total counter"); + let _ = writeln!( + out, + "telemt_handshake_timeouts_total {}", + if core_enabled { + stats.get_handshake_timeouts() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_handshake_failures_by_class_total Handshake failures by class" + ); + let _ = writeln!( + out, + "# TYPE telemt_handshake_failures_by_class_total counter" + ); + if core_enabled { + for (class, total) in stats.get_handshake_failure_class_counts() { + let _ = writeln!( + out, + "telemt_handshake_failures_by_class_total{{class=\"{}\"}} {}", + class, total + ); + } + } + + let _ = writeln!( + out, + "# HELP telemt_auth_expensive_checks_total Expensive authentication candidate checks executed during handshake validation" + ); + let _ = writeln!(out, "# TYPE telemt_auth_expensive_checks_total counter"); + let _ = writeln!( + out, + "telemt_auth_expensive_checks_total {}", + if core_enabled { + shared_state + .handshake + .auth_expensive_checks_total + .load(std::sync::atomic::Ordering::Relaxed) + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_auth_budget_exhausted_total Handshake validations that hit authentication candidate budget limits" + ); + let _ = writeln!(out, "# TYPE telemt_auth_budget_exhausted_total counter"); + let _ = writeln!( + out, + "telemt_auth_budget_exhausted_total {}", + if core_enabled { + shared_state + .handshake + .auth_budget_exhausted_total + .load(std::sync::atomic::Ordering::Relaxed) + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_accept_permit_timeout_total Accepted connections dropped due to permit wait timeout" + ); + let _ = writeln!(out, "# TYPE telemt_accept_permit_timeout_total counter"); + let _ = writeln!( + out, + "telemt_accept_permit_timeout_total {}", + if core_enabled { + stats.get_accept_permit_timeout_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_route_cutover_parked_current Sessions currently parked in route cutover stagger delay" + ); + let _ = writeln!(out, "# TYPE telemt_route_cutover_parked_current gauge"); + let _ = writeln!( + out, + "telemt_route_cutover_parked_current{{route=\"direct\"}} {}", + stats.get_route_cutover_parked_direct_current() + ); + let _ = writeln!( + out, + "telemt_route_cutover_parked_current{{route=\"middle\"}} {}", + stats.get_route_cutover_parked_middle_current() + ); + let _ = writeln!( + out, + "# HELP telemt_route_cutover_parked_total Sessions parked in route cutover stagger delay" + ); + let _ = writeln!(out, "# TYPE telemt_route_cutover_parked_total counter"); + let _ = writeln!( + out, + "telemt_route_cutover_parked_total{{route=\"direct\"}} {}", + stats.get_route_cutover_parked_direct_total() + ); + let _ = writeln!( + out, + "telemt_route_cutover_parked_total{{route=\"middle\"}} {}", + stats.get_route_cutover_parked_middle_total() + ); + + let _ = writeln!( + out, + "# HELP telemt_quota_refund_bytes_total Reserved quota bytes returned before commit" + ); + let _ = writeln!(out, "# TYPE telemt_quota_refund_bytes_total counter"); + let _ = writeln!( + out, + "telemt_quota_refund_bytes_total {}", + if core_enabled { + stats.get_quota_refund_bytes_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_quota_contention_total Quota reservation CAS contention events" + ); + let _ = writeln!(out, "# TYPE telemt_quota_contention_total counter"); + let _ = writeln!( + out, + "telemt_quota_contention_total {}", + if core_enabled { + stats.get_quota_contention_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_quota_contention_timeout_total Quota reservations that hit the bounded contention budget" + ); + let _ = writeln!(out, "# TYPE telemt_quota_contention_timeout_total counter"); + let _ = writeln!( + out, + "telemt_quota_contention_timeout_total {}", + if core_enabled { + stats.get_quota_contention_timeout_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_quota_acquire_cancelled_total Quota acquisitions cancelled before reservation completed" + ); + let _ = writeln!(out, "# TYPE telemt_quota_acquire_cancelled_total counter"); + let _ = writeln!( + out, + "telemt_quota_acquire_cancelled_total {}", + if core_enabled { + stats.get_quota_acquire_cancelled_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_conntrack_control_state Runtime conntrack control state flags" + ); + let _ = writeln!(out, "# TYPE telemt_conntrack_control_state gauge"); + let _ = writeln!( + out, + "telemt_conntrack_control_state{{flag=\"enabled\"}} {}", + if stats.get_conntrack_control_enabled() { + 1 + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_conntrack_control_state{{flag=\"available\"}} {}", + if stats.get_conntrack_control_available() { + 1 + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_conntrack_control_state{{flag=\"pressure_active\"}} {}", + if stats.get_conntrack_pressure_active() { + 1 + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_conntrack_control_state{{flag=\"rule_apply_ok\"}} {}", + if stats.get_conntrack_rule_apply_ok() { + 1 + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_conntrack_event_queue_depth Pending close events in conntrack control queue" + ); + let _ = writeln!(out, "# TYPE telemt_conntrack_event_queue_depth gauge"); + let _ = writeln!( + out, + "telemt_conntrack_event_queue_depth {}", + stats.get_conntrack_event_queue_depth() + ); + + let _ = writeln!( + out, + "# HELP telemt_conntrack_delete_total Conntrack delete attempts by outcome" + ); + let _ = writeln!(out, "# TYPE telemt_conntrack_delete_total counter"); + let _ = writeln!( + out, + "telemt_conntrack_delete_total{{result=\"attempt\"}} {}", + if core_enabled { + stats.get_conntrack_delete_attempt_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_conntrack_delete_total{{result=\"success\"}} {}", + if core_enabled { + stats.get_conntrack_delete_success_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_conntrack_delete_total{{result=\"not_found\"}} {}", + if core_enabled { + stats.get_conntrack_delete_not_found_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_conntrack_delete_total{{result=\"error\"}} {}", + if core_enabled { + stats.get_conntrack_delete_error_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_conntrack_close_event_drop_total Dropped conntrack close events due to queue pressure or unavailable sender" + ); + let _ = writeln!( + out, + "# TYPE telemt_conntrack_close_event_drop_total counter" + ); + let _ = writeln!( + out, + "telemt_conntrack_close_event_drop_total {}", + if core_enabled { + stats.get_conntrack_close_event_drop_total() + } else { + 0 + } + ); +} diff --git a/src/metrics/render/me_buffers.rs b/src/metrics/render/me_buffers.rs new file mode 100644 index 0000000..82fd147 --- /dev/null +++ b/src/metrics/render/me_buffers.rs @@ -0,0 +1,489 @@ +use super::*; +use std::fmt::Write; + +pub(super) fn render( + out: &mut String, + stats: &Stats, + core_enabled: bool, + me_allows_normal: bool, + me_allows_debug: bool, +) { + let _ = writeln!( + out, + "telemt_me_d2c_flush_reason_total{{reason=\"batch_bytes\"}} {}", + if me_allows_normal { + stats.get_me_d2c_flush_reason_batch_bytes_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_flush_reason_total{{reason=\"max_delay\"}} {}", + if me_allows_normal { + stats.get_me_d2c_flush_reason_max_delay_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_flush_reason_total{{reason=\"ack_immediate\"}} {}", + if me_allows_normal { + stats.get_me_d2c_flush_reason_ack_immediate_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_flush_reason_total{{reason=\"close\"}} {}", + if me_allows_normal { + stats.get_me_d2c_flush_reason_close_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_d2c_data_frames_total DC->Client data frames" + ); + let _ = writeln!(out, "# TYPE telemt_me_d2c_data_frames_total counter"); + let _ = writeln!( + out, + "telemt_me_d2c_data_frames_total {}", + if me_allows_normal { + stats.get_me_d2c_data_frames_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_d2c_ack_frames_total DC->Client quick-ack frames" + ); + let _ = writeln!(out, "# TYPE telemt_me_d2c_ack_frames_total counter"); + let _ = writeln!( + out, + "telemt_me_d2c_ack_frames_total {}", + if me_allows_normal { + stats.get_me_d2c_ack_frames_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_d2c_payload_bytes_total DC->Client payload bytes before transport framing" + ); + let _ = writeln!(out, "# TYPE telemt_me_d2c_payload_bytes_total counter"); + let _ = writeln!( + out, + "telemt_me_d2c_payload_bytes_total {}", + if me_allows_normal { + stats.get_me_d2c_payload_bytes_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_d2c_write_mode_total DC->Client writer mode selection" + ); + let _ = writeln!(out, "# TYPE telemt_me_d2c_write_mode_total counter"); + let _ = writeln!( + out, + "telemt_me_d2c_write_mode_total{{mode=\"coalesced\"}} {}", + if me_allows_normal { + stats.get_me_d2c_write_mode_coalesced_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_write_mode_total{{mode=\"split\"}} {}", + if me_allows_normal { + stats.get_me_d2c_write_mode_split_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_d2c_quota_reject_total DC->Client quota rejects" + ); + let _ = writeln!(out, "# TYPE telemt_me_d2c_quota_reject_total counter"); + let _ = writeln!( + out, + "telemt_me_d2c_quota_reject_total{{stage=\"pre_write\"}} {}", + if me_allows_normal { + stats.get_me_d2c_quota_reject_pre_write_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_quota_reject_total{{stage=\"post_write\"}} {}", + if me_allows_normal { + stats.get_me_d2c_quota_reject_post_write_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_child_join_timeout_total Middle relay child tasks that did not join before cleanup deadline" + ); + let _ = writeln!(out, "# TYPE telemt_me_child_join_timeout_total counter"); + let _ = writeln!( + out, + "telemt_me_child_join_timeout_total {}", + if core_enabled { + stats.get_me_child_join_timeout_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_child_abort_total Middle relay child tasks aborted after bounded cleanup timeout" + ); + let _ = writeln!(out, "# TYPE telemt_me_child_abort_total counter"); + let _ = writeln!( + out, + "telemt_me_child_abort_total {}", + if core_enabled { + stats.get_me_child_abort_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_flow_wait_events_total Flow wait events by reason, direction, and outcome" + ); + let _ = writeln!(out, "# TYPE telemt_flow_wait_events_total counter"); + let _ = writeln!( + out, + "telemt_flow_wait_events_total{{reason=\"middle_rate_limit\",direction=\"down\",outcome=\"waited\"}} {}", + if core_enabled { + stats.get_flow_wait_middle_rate_limit_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_flow_wait_events_total{{reason=\"middle_rate_limit\",direction=\"down\",outcome=\"cancelled\"}} {}", + if core_enabled { + stats.get_flow_wait_middle_rate_limit_cancelled_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_flow_wait_ms_total Flow wait time in milliseconds by reason and direction" + ); + let _ = writeln!(out, "# TYPE telemt_flow_wait_ms_total counter"); + let _ = writeln!( + out, + "telemt_flow_wait_ms_total{{reason=\"middle_rate_limit\",direction=\"down\"}} {}", + if core_enabled { + stats.get_flow_wait_middle_rate_limit_ms_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_session_drop_fallback_total Session reservations cleaned by Drop instead of explicit async release" + ); + let _ = writeln!(out, "# TYPE telemt_session_drop_fallback_total counter"); + let _ = writeln!( + out, + "telemt_session_drop_fallback_total {}", + if core_enabled { + stats.get_session_drop_fallback_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_d2c_frame_buf_shrink_total DC->Client reusable frame buffer shrink events" + ); + let _ = writeln!(out, "# TYPE telemt_me_d2c_frame_buf_shrink_total counter"); + let _ = writeln!( + out, + "telemt_me_d2c_frame_buf_shrink_total {}", + if me_allows_normal { + stats.get_me_d2c_frame_buf_shrink_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_d2c_frame_buf_shrink_bytes_total DC->Client reusable frame buffer bytes released" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_d2c_frame_buf_shrink_bytes_total counter" + ); + let _ = writeln!( + out, + "telemt_me_d2c_frame_buf_shrink_bytes_total {}", + if me_allows_normal { + stats.get_me_d2c_frame_buf_shrink_bytes_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_d2c_batch_frames_bucket_total DC->Client batch frame count buckets" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_d2c_batch_frames_bucket_total counter" + ); + let _ = writeln!( + out, + "telemt_me_d2c_batch_frames_bucket_total{{bucket=\"1\"}} {}", + if me_allows_debug { + stats.get_me_d2c_batch_frames_bucket_1() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_batch_frames_bucket_total{{bucket=\"2_4\"}} {}", + if me_allows_debug { + stats.get_me_d2c_batch_frames_bucket_2_4() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_batch_frames_bucket_total{{bucket=\"5_8\"}} {}", + if me_allows_debug { + stats.get_me_d2c_batch_frames_bucket_5_8() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_batch_frames_bucket_total{{bucket=\"9_16\"}} {}", + if me_allows_debug { + stats.get_me_d2c_batch_frames_bucket_9_16() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_batch_frames_bucket_total{{bucket=\"17_32\"}} {}", + if me_allows_debug { + stats.get_me_d2c_batch_frames_bucket_17_32() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_batch_frames_bucket_total{{bucket=\"gt_32\"}} {}", + if me_allows_debug { + stats.get_me_d2c_batch_frames_bucket_gt_32() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_d2c_batch_bytes_bucket_total DC->Client batch byte size buckets" + ); + let _ = writeln!(out, "# TYPE telemt_me_d2c_batch_bytes_bucket_total counter"); + let _ = writeln!( + out, + "telemt_me_d2c_batch_bytes_bucket_total{{bucket=\"0_1k\"}} {}", + if me_allows_debug { + stats.get_me_d2c_batch_bytes_bucket_0_1k() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_batch_bytes_bucket_total{{bucket=\"1k_4k\"}} {}", + if me_allows_debug { + stats.get_me_d2c_batch_bytes_bucket_1k_4k() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_batch_bytes_bucket_total{{bucket=\"4k_16k\"}} {}", + if me_allows_debug { + stats.get_me_d2c_batch_bytes_bucket_4k_16k() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_batch_bytes_bucket_total{{bucket=\"16k_64k\"}} {}", + if me_allows_debug { + stats.get_me_d2c_batch_bytes_bucket_16k_64k() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_batch_bytes_bucket_total{{bucket=\"64k_128k\"}} {}", + if me_allows_debug { + stats.get_me_d2c_batch_bytes_bucket_64k_128k() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_batch_bytes_bucket_total{{bucket=\"gt_128k\"}} {}", + if me_allows_debug { + stats.get_me_d2c_batch_bytes_bucket_gt_128k() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_d2c_flush_duration_us_bucket_total DC->Client flush duration buckets" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_d2c_flush_duration_us_bucket_total counter" + ); + let _ = writeln!( + out, + "telemt_me_d2c_flush_duration_us_bucket_total{{bucket=\"0_50\"}} {}", + if me_allows_debug { + stats.get_me_d2c_flush_duration_us_bucket_0_50() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_flush_duration_us_bucket_total{{bucket=\"51_200\"}} {}", + if me_allows_debug { + stats.get_me_d2c_flush_duration_us_bucket_51_200() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_flush_duration_us_bucket_total{{bucket=\"201_1000\"}} {}", + if me_allows_debug { + stats.get_me_d2c_flush_duration_us_bucket_201_1000() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_flush_duration_us_bucket_total{{bucket=\"1001_5000\"}} {}", + if me_allows_debug { + stats.get_me_d2c_flush_duration_us_bucket_1001_5000() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_flush_duration_us_bucket_total{{bucket=\"5001_20000\"}} {}", + if me_allows_debug { + stats.get_me_d2c_flush_duration_us_bucket_5001_20000() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_flush_duration_us_bucket_total{{bucket=\"gt_20000\"}} {}", + if me_allows_debug { + stats.get_me_d2c_flush_duration_us_bucket_gt_20000() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_d2c_batch_timeout_armed_total DC->Client max-delay timer armed events" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_d2c_batch_timeout_armed_total counter" + ); + let _ = writeln!( + out, + "telemt_me_d2c_batch_timeout_armed_total {}", + if me_allows_debug { + stats.get_me_d2c_batch_timeout_armed_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_d2c_batch_timeout_fired_total DC->Client max-delay timer fired events" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_d2c_batch_timeout_fired_total counter" + ); + let _ = writeln!( + out, + "telemt_me_d2c_batch_timeout_fired_total {}", + if me_allows_debug { + stats.get_me_d2c_batch_timeout_fired_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_writer_byte_budget_limit_bytes Configured resident-memory budget per ME writer" + ); + let _ = writeln!(out, "# TYPE telemt_me_writer_byte_budget_limit_bytes gauge"); + let _ = writeln!( + out, + "telemt_me_writer_byte_budget_limit_bytes {}", + if me_allows_normal { + stats.get_me_writer_byte_budget_limit_bytes_gauge() + } else { + 0 + } + ); +} diff --git a/src/metrics/render/me_floor.rs b/src/metrics/render/me_floor.rs new file mode 100644 index 0000000..4e1f213 --- /dev/null +++ b/src/metrics/render/me_floor.rs @@ -0,0 +1,282 @@ +use super::*; +use std::fmt::Write; + +pub(super) fn render( + out: &mut String, + stats: &Stats, + config: &ProxyConfig, + me_allows_normal: bool, +) { + let floor_mode = config.general.me_floor_mode; + let _ = writeln!( + out, + "telemt_me_floor_mode{{mode=\"static\"}} {}", + if matches!(floor_mode, crate::config::MeFloorMode::Static) { + 1 + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_floor_mode{{mode=\"adaptive\"}} {}", + if matches!(floor_mode, crate::config::MeFloorMode::Adaptive) { + 1 + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_floor_mode_switch_all_total Runtime ME floor mode switches" + ); + let _ = writeln!(out, "# TYPE telemt_me_floor_mode_switch_all_total counter"); + let _ = writeln!( + out, + "telemt_me_floor_mode_switch_all_total {}", + if me_allows_normal { + stats.get_me_floor_mode_switch_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_floor_mode_switch_total{{from=\"static\",to=\"adaptive\"}} {}", + if me_allows_normal { + stats.get_me_floor_mode_switch_static_to_adaptive_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_floor_mode_switch_total{{from=\"adaptive\",to=\"static\"}} {}", + if me_allows_normal { + stats.get_me_floor_mode_switch_adaptive_to_static_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_adaptive_floor_cpu_cores_detected Runtime detected logical CPU cores for adaptive floor" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_adaptive_floor_cpu_cores_detected gauge" + ); + let _ = writeln!( + out, + "telemt_me_adaptive_floor_cpu_cores_detected {}", + if me_allows_normal { + stats.get_me_floor_cpu_cores_detected_gauge() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_adaptive_floor_cpu_cores_effective Runtime effective logical CPU cores for adaptive floor" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_adaptive_floor_cpu_cores_effective gauge" + ); + let _ = writeln!( + out, + "telemt_me_adaptive_floor_cpu_cores_effective {}", + if me_allows_normal { + stats.get_me_floor_cpu_cores_effective_gauge() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_adaptive_floor_global_cap_raw Runtime raw global adaptive floor cap" + ); + let _ = writeln!(out, "# TYPE telemt_me_adaptive_floor_global_cap_raw gauge"); + let _ = writeln!( + out, + "telemt_me_adaptive_floor_global_cap_raw {}", + if me_allows_normal { + stats.get_me_floor_global_cap_raw_gauge() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_adaptive_floor_global_cap_effective Runtime effective global adaptive floor cap" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_adaptive_floor_global_cap_effective gauge" + ); + let _ = writeln!( + out, + "telemt_me_adaptive_floor_global_cap_effective {}", + if me_allows_normal { + stats.get_me_floor_global_cap_effective_gauge() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_adaptive_floor_target_writers_total Runtime adaptive floor target writers total" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_adaptive_floor_target_writers_total gauge" + ); + let _ = writeln!( + out, + "telemt_me_adaptive_floor_target_writers_total {}", + if me_allows_normal { + stats.get_me_floor_target_writers_total_gauge() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_adaptive_floor_active_cap_configured Runtime configured active writer cap" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_adaptive_floor_active_cap_configured gauge" + ); + let _ = writeln!( + out, + "telemt_me_adaptive_floor_active_cap_configured {}", + if me_allows_normal { + stats.get_me_floor_active_cap_configured_gauge() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_adaptive_floor_active_cap_effective Runtime effective active writer cap" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_adaptive_floor_active_cap_effective gauge" + ); + let _ = writeln!( + out, + "telemt_me_adaptive_floor_active_cap_effective {}", + if me_allows_normal { + stats.get_me_floor_active_cap_effective_gauge() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_adaptive_floor_warm_cap_configured Runtime configured warm writer cap" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_adaptive_floor_warm_cap_configured gauge" + ); + let _ = writeln!( + out, + "telemt_me_adaptive_floor_warm_cap_configured {}", + if me_allows_normal { + stats.get_me_floor_warm_cap_configured_gauge() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_adaptive_floor_warm_cap_effective Runtime effective warm writer cap" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_adaptive_floor_warm_cap_effective gauge" + ); + let _ = writeln!( + out, + "telemt_me_adaptive_floor_warm_cap_effective {}", + if me_allows_normal { + stats.get_me_floor_warm_cap_effective_gauge() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_writers_active_current Current non-draining active ME writers" + ); + let _ = writeln!(out, "# TYPE telemt_me_writers_active_current gauge"); + let _ = writeln!( + out, + "telemt_me_writers_active_current {}", + if me_allows_normal { + stats.get_me_writers_active_current_gauge() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_writers_warm_current Current non-draining warm ME writers" + ); + let _ = writeln!(out, "# TYPE telemt_me_writers_warm_current gauge"); + let _ = writeln!( + out, + "telemt_me_writers_warm_current {}", + if me_allows_normal { + stats.get_me_writers_warm_current_gauge() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_floor_cap_block_total Reconnect attempts blocked by adaptive floor caps" + ); + let _ = writeln!(out, "# TYPE telemt_me_floor_cap_block_total counter"); + let _ = writeln!( + out, + "telemt_me_floor_cap_block_total {}", + if me_allows_normal { + stats.get_me_floor_cap_block_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_floor_swap_idle_total Adaptive floor cap recovery via idle writer swap" + ); + let _ = writeln!(out, "# TYPE telemt_me_floor_swap_idle_total counter"); + let _ = writeln!( + out, + "telemt_me_floor_swap_idle_total {}", + if me_allows_normal { + stats.get_me_floor_swap_idle_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_floor_swap_idle_failed_total Failed idle swap attempts under adaptive floor caps" + ); + let _ = writeln!(out, "# TYPE telemt_me_floor_swap_idle_failed_total counter"); + let _ = writeln!( + out, + "telemt_me_floor_swap_idle_failed_total {}", + if me_allows_normal { + stats.get_me_floor_swap_idle_failed_total() + } else { + 0 + } + ); +} diff --git a/src/metrics/render/me_lifecycle.rs b/src/metrics/render/me_lifecycle.rs new file mode 100644 index 0000000..6674b9e --- /dev/null +++ b/src/metrics/render/me_lifecycle.rs @@ -0,0 +1,503 @@ +use super::*; +use std::fmt::Write; + +pub(super) fn render(out: &mut String, stats: &Stats, me_allows_normal: bool) { + let _ = writeln!( + out, + "# HELP telemt_me_reconnect_attempts_total ME reconnect attempts" + ); + let _ = writeln!(out, "# TYPE telemt_me_reconnect_attempts_total counter"); + let _ = writeln!( + out, + "telemt_me_reconnect_attempts_total {}", + if me_allows_normal { + stats.get_me_reconnect_attempts() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_reconnect_success_total ME reconnect successes" + ); + let _ = writeln!(out, "# TYPE telemt_me_reconnect_success_total counter"); + let _ = writeln!( + out, + "telemt_me_reconnect_success_total {}", + if me_allows_normal { + stats.get_me_reconnect_success() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_handshake_reject_total ME handshake rejects from upstream" + ); + let _ = writeln!(out, "# TYPE telemt_me_handshake_reject_total counter"); + let _ = writeln!( + out, + "telemt_me_handshake_reject_total {}", + if me_allows_normal { + stats.get_me_handshake_reject_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_handshake_error_code_total ME handshake reject errors by code" + ); + let _ = writeln!(out, "# TYPE telemt_me_handshake_error_code_total counter"); + if me_allows_normal { + for (error_code, count) in stats.get_me_handshake_error_code_counts() { + let _ = writeln!( + out, + "telemt_me_handshake_error_code_total{{error_code=\"{}\"}} {}", + error_code, count + ); + } + let _ = writeln!( + out, + "telemt_me_handshake_error_code_total{{error_code=\"overflow\"}} {}", + stats.get_me_handshake_error_code_overflow_total() + ); + } + + let _ = writeln!( + out, + "# HELP telemt_me_reader_eof_total ME reader EOF terminations" + ); + let _ = writeln!(out, "# TYPE telemt_me_reader_eof_total counter"); + let _ = writeln!( + out, + "telemt_me_reader_eof_total {}", + if me_allows_normal { + stats.get_me_reader_eof_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_idle_close_by_peer_total ME idle writers closed by peer" + ); + let _ = writeln!(out, "# TYPE telemt_me_idle_close_by_peer_total counter"); + let _ = writeln!( + out, + "telemt_me_idle_close_by_peer_total {}", + if me_allows_normal { + stats.get_me_idle_close_by_peer_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_relay_idle_soft_mark_total Middle-relay sessions marked as soft-idle candidates" + ); + let _ = writeln!(out, "# TYPE telemt_relay_idle_soft_mark_total counter"); + let _ = writeln!( + out, + "telemt_relay_idle_soft_mark_total {}", + if me_allows_normal { + stats.get_relay_idle_soft_mark_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_relay_idle_hard_close_total Middle-relay sessions closed by hard-idle policy" + ); + let _ = writeln!(out, "# TYPE telemt_relay_idle_hard_close_total counter"); + let _ = writeln!( + out, + "telemt_relay_idle_hard_close_total {}", + if me_allows_normal { + stats.get_relay_idle_hard_close_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_relay_pressure_evict_total Middle-relay sessions evicted under resource pressure" + ); + let _ = writeln!(out, "# TYPE telemt_relay_pressure_evict_total counter"); + let _ = writeln!( + out, + "telemt_relay_pressure_evict_total {}", + if me_allows_normal { + stats.get_relay_pressure_evict_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_relay_protocol_desync_close_total Middle-relay sessions closed due to protocol desync" + ); + let _ = writeln!( + out, + "# TYPE telemt_relay_protocol_desync_close_total counter" + ); + let _ = writeln!( + out, + "telemt_relay_protocol_desync_close_total {}", + if me_allows_normal { + stats.get_relay_protocol_desync_close_total() + } else { + 0 + } + ); + + let _ = writeln!(out, "# HELP telemt_me_crc_mismatch_total ME CRC mismatches"); + let _ = writeln!(out, "# TYPE telemt_me_crc_mismatch_total counter"); + let _ = writeln!( + out, + "telemt_me_crc_mismatch_total {}", + if me_allows_normal { + stats.get_me_crc_mismatch() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_seq_mismatch_total ME sequence mismatches" + ); + let _ = writeln!(out, "# TYPE telemt_me_seq_mismatch_total counter"); + let _ = writeln!( + out, + "telemt_me_seq_mismatch_total {}", + if me_allows_normal { + stats.get_me_seq_mismatch() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_route_drop_no_conn_total ME route drops: no conn" + ); + let _ = writeln!(out, "# TYPE telemt_me_route_drop_no_conn_total counter"); + let _ = writeln!( + out, + "telemt_me_route_drop_no_conn_total {}", + if me_allows_normal { + stats.get_me_route_drop_no_conn() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_route_drop_channel_closed_total ME route drops: channel closed" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_route_drop_channel_closed_total counter" + ); + let _ = writeln!( + out, + "telemt_me_route_drop_channel_closed_total {}", + if me_allows_normal { + stats.get_me_route_drop_channel_closed() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_route_drop_queue_full_total ME route drops: queue full" + ); + let _ = writeln!(out, "# TYPE telemt_me_route_drop_queue_full_total counter"); + let _ = writeln!( + out, + "telemt_me_route_drop_queue_full_total {}", + if me_allows_normal { + stats.get_me_route_drop_queue_full() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_route_drop_queue_full_profile_total ME route drops: queue full by adaptive profile" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_route_drop_queue_full_profile_total counter" + ); + let _ = writeln!( + out, + "telemt_me_route_drop_queue_full_profile_total{{profile=\"base\"}} {}", + if me_allows_normal { + stats.get_me_route_drop_queue_full_base() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_route_drop_queue_full_profile_total{{profile=\"high\"}} {}", + if me_allows_normal { + stats.get_me_route_drop_queue_full_high() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_fair_pressure_state Worker-local fairness pressure state" + ); + let _ = writeln!(out, "# TYPE telemt_me_fair_pressure_state gauge"); + let _ = writeln!( + out, + "telemt_me_fair_pressure_state {}", + if me_allows_normal { + stats.get_me_fair_pressure_state_gauge() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_fair_active_flows Fair-scheduler active flow count" + ); + let _ = writeln!(out, "# TYPE telemt_me_fair_active_flows gauge"); + let _ = writeln!( + out, + "telemt_me_fair_active_flows {}", + if me_allows_normal { + stats.get_me_fair_active_flows_gauge() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_fair_queued_bytes Fair-scheduler queued bytes" + ); + let _ = writeln!(out, "# TYPE telemt_me_fair_queued_bytes gauge"); + let _ = writeln!( + out, + "telemt_me_fair_queued_bytes {}", + if me_allows_normal { + stats.get_me_fair_queued_bytes_gauge() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_fair_flow_state_gauge Fair-scheduler flow health classes" + ); + let _ = writeln!(out, "# TYPE telemt_me_fair_flow_state_gauge gauge"); + let _ = writeln!( + out, + "telemt_me_fair_flow_state_gauge{{class=\"standing\"}} {}", + if me_allows_normal { + stats.get_me_fair_standing_flows_gauge() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_fair_flow_state_gauge{{class=\"backpressured\"}} {}", + if me_allows_normal { + stats.get_me_fair_backpressured_flows_gauge() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_fair_events_total Fair-scheduler event counters" + ); + let _ = writeln!(out, "# TYPE telemt_me_fair_events_total counter"); + let _ = writeln!( + out, + "telemt_me_fair_events_total{{event=\"scheduler_round\"}} {}", + if me_allows_normal { + stats.get_me_fair_scheduler_rounds_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_fair_events_total{{event=\"deficit_grant\"}} {}", + if me_allows_normal { + stats.get_me_fair_deficit_grants_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_fair_events_total{{event=\"deficit_skip\"}} {}", + if me_allows_normal { + stats.get_me_fair_deficit_skips_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_fair_events_total{{event=\"enqueue_reject\"}} {}", + if me_allows_normal { + stats.get_me_fair_enqueue_rejects_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_fair_events_total{{event=\"shed_drop\"}} {}", + if me_allows_normal { + stats.get_me_fair_shed_drops_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_fair_events_total{{event=\"penalty\"}} {}", + if me_allows_normal { + stats.get_me_fair_penalties_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_fair_events_total{{event=\"downstream_stall\"}} {}", + if me_allows_normal { + stats.get_me_fair_downstream_stalls_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_c2me_enqueue_events_total ME client->ME enqueue outcomes" + ); + let _ = writeln!(out, "# TYPE telemt_me_c2me_enqueue_events_total counter"); + let _ = writeln!( + out, + "telemt_me_c2me_enqueue_events_total{{event=\"full\"}} {}", + if me_allows_normal { + stats.get_me_c2me_send_full_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_c2me_enqueue_events_total{{event=\"high_water\"}} {}", + if me_allows_normal { + stats.get_me_c2me_send_high_water_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_c2me_enqueue_events_total{{event=\"timeout\"}} {}", + if me_allows_normal { + stats.get_me_c2me_send_timeout_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_d2c_batches_total Total DC->Client flush batches" + ); + let _ = writeln!(out, "# TYPE telemt_me_d2c_batches_total counter"); + let _ = writeln!( + out, + "telemt_me_d2c_batches_total {}", + if me_allows_normal { + stats.get_me_d2c_batches_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_d2c_batch_frames_total Total DC->Client frames flushed in batches" + ); + let _ = writeln!(out, "# TYPE telemt_me_d2c_batch_frames_total counter"); + let _ = writeln!( + out, + "telemt_me_d2c_batch_frames_total {}", + if me_allows_normal { + stats.get_me_d2c_batch_frames_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_d2c_batch_bytes_total Total DC->Client bytes flushed in batches" + ); + let _ = writeln!(out, "# TYPE telemt_me_d2c_batch_bytes_total counter"); + let _ = writeln!( + out, + "telemt_me_d2c_batch_bytes_total {}", + if me_allows_normal { + stats.get_me_d2c_batch_bytes_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_d2c_flush_reason_total DC->Client flush reasons" + ); + let _ = writeln!(out, "# TYPE telemt_me_d2c_flush_reason_total counter"); + let _ = writeln!( + out, + "telemt_me_d2c_flush_reason_total{{reason=\"queue_drain\"}} {}", + if me_allows_normal { + stats.get_me_d2c_flush_reason_queue_drain_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_d2c_flush_reason_total{{reason=\"batch_frames\"}} {}", + if me_allows_normal { + stats.get_me_d2c_flush_reason_batch_frames_total() + } else { + 0 + } + ); +} diff --git a/src/metrics/render/me_policy.rs b/src/metrics/render/me_policy.rs new file mode 100644 index 0000000..0cad1cb --- /dev/null +++ b/src/metrics/render/me_policy.rs @@ -0,0 +1,471 @@ +use super::*; +use std::fmt::Write; + +pub(super) fn render( + out: &mut String, + stats: &Stats, + me_allows_normal: bool, + me_allows_debug: bool, +) { + let _ = writeln!( + out, + "# HELP telemt_me_writer_byte_budget_reserved_bytes Aggregate ME writer memory reservations by lifecycle state" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_writer_byte_budget_reserved_bytes gauge" + ); + let _ = writeln!( + out, + "telemt_me_writer_byte_budget_reserved_bytes{{state=\"queued\"}} {}", + if me_allows_normal { + stats.get_me_writer_byte_budget_queued_bytes_gauge() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_writer_byte_budget_reserved_bytes{{state=\"inflight\"}} {}", + if me_allows_normal { + stats.get_me_writer_byte_budget_inflight_bytes_gauge() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_writer_byte_budget_events_total ME writer byte-budget outcomes" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_writer_byte_budget_events_total counter" + ); + let _ = writeln!( + out, + "telemt_me_writer_byte_budget_events_total{{result=\"wait\"}} {}", + if me_allows_normal { + stats.get_me_writer_byte_budget_wait_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_writer_byte_budget_events_total{{result=\"timeout\"}} {}", + if me_allows_normal { + stats.get_me_writer_byte_budget_timeout_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_writer_byte_budget_events_total{{result=\"oversize\"}} {}", + if me_allows_normal { + stats.get_me_writer_byte_budget_oversize_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_writer_pick_total ME writer-pick outcomes by mode and result" + ); + let _ = writeln!(out, "# TYPE telemt_me_writer_pick_total counter"); + let _ = writeln!( + out, + "telemt_me_writer_pick_total{{mode=\"sorted_rr\",result=\"success_try\"}} {}", + if me_allows_normal { + stats.get_me_writer_pick_sorted_rr_success_try_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_writer_pick_total{{mode=\"sorted_rr\",result=\"success_fallback\"}} {}", + if me_allows_normal { + stats.get_me_writer_pick_sorted_rr_success_fallback_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_writer_pick_total{{mode=\"sorted_rr\",result=\"full\"}} {}", + if me_allows_normal { + stats.get_me_writer_pick_sorted_rr_full_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_writer_pick_total{{mode=\"sorted_rr\",result=\"closed\"}} {}", + if me_allows_normal { + stats.get_me_writer_pick_sorted_rr_closed_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_writer_pick_total{{mode=\"sorted_rr\",result=\"no_candidate\"}} {}", + if me_allows_normal { + stats.get_me_writer_pick_sorted_rr_no_candidate_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_writer_pick_total{{mode=\"p2c\",result=\"success_try\"}} {}", + if me_allows_normal { + stats.get_me_writer_pick_p2c_success_try_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_writer_pick_total{{mode=\"p2c\",result=\"success_fallback\"}} {}", + if me_allows_normal { + stats.get_me_writer_pick_p2c_success_fallback_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_writer_pick_total{{mode=\"p2c\",result=\"full\"}} {}", + if me_allows_normal { + stats.get_me_writer_pick_p2c_full_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_writer_pick_total{{mode=\"p2c\",result=\"closed\"}} {}", + if me_allows_normal { + stats.get_me_writer_pick_p2c_closed_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_writer_pick_total{{mode=\"p2c\",result=\"no_candidate\"}} {}", + if me_allows_normal { + stats.get_me_writer_pick_p2c_no_candidate_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_writer_pick_blocking_fallback_total ME writer-pick blocking fallback attempts" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_writer_pick_blocking_fallback_total counter" + ); + let _ = writeln!( + out, + "telemt_me_writer_pick_blocking_fallback_total {}", + if me_allows_normal { + stats.get_me_writer_pick_blocking_fallback_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_writer_pick_mode_switch_total Writer-pick mode switches via runtime updates" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_writer_pick_mode_switch_total counter" + ); + let _ = writeln!( + out, + "telemt_me_writer_pick_mode_switch_total {}", + if me_allows_normal { + stats.get_me_writer_pick_mode_switch_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_socks_kdf_policy_total SOCKS KDF policy outcomes" + ); + let _ = writeln!(out, "# TYPE telemt_me_socks_kdf_policy_total counter"); + let _ = writeln!( + out, + "telemt_me_socks_kdf_policy_total{{policy=\"strict\",outcome=\"reject\"}} {}", + if me_allows_normal { + stats.get_me_socks_kdf_strict_reject() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_me_socks_kdf_policy_total{{policy=\"compat\",outcome=\"fallback\"}} {}", + if me_allows_debug { + stats.get_me_socks_kdf_compat_fallback() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_endpoint_quarantine_total ME endpoint quarantines due to rapid flaps" + ); + let _ = writeln!(out, "# TYPE telemt_me_endpoint_quarantine_total counter"); + let _ = writeln!( + out, + "telemt_me_endpoint_quarantine_total {}", + if me_allows_normal { + stats.get_me_endpoint_quarantine_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_endpoint_quarantine_unexpected_total ME endpoint quarantines caused by unexpected writer removals" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_endpoint_quarantine_unexpected_total counter" + ); + let _ = writeln!( + out, + "telemt_me_endpoint_quarantine_unexpected_total {}", + if me_allows_normal { + stats.get_me_endpoint_quarantine_unexpected_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_endpoint_quarantine_draining_suppressed_total Draining writer removals that skipped endpoint quarantine" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_endpoint_quarantine_draining_suppressed_total counter" + ); + let _ = writeln!( + out, + "telemt_me_endpoint_quarantine_draining_suppressed_total {}", + if me_allows_normal { + stats.get_me_endpoint_quarantine_draining_suppressed_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_kdf_drift_total ME KDF input drift detections" + ); + let _ = writeln!(out, "# TYPE telemt_me_kdf_drift_total counter"); + let _ = writeln!( + out, + "telemt_me_kdf_drift_total {}", + if me_allows_normal { + stats.get_me_kdf_drift_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_kdf_port_only_drift_total ME KDF client-port changes with stable non-port material" + ); + let _ = writeln!(out, "# TYPE telemt_me_kdf_port_only_drift_total counter"); + let _ = writeln!( + out, + "telemt_me_kdf_port_only_drift_total {}", + if me_allows_debug { + stats.get_me_kdf_port_only_drift_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_hardswap_pending_reuse_total Hardswap cycles that reused an existing pending generation" + ); + let _ = writeln!(out, "# TYPE telemt_me_hardswap_pending_reuse_total counter"); + let _ = writeln!( + out, + "telemt_me_hardswap_pending_reuse_total {}", + if me_allows_debug { + stats.get_me_hardswap_pending_reuse_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_hardswap_pending_ttl_expired_total Pending hardswap generations reset by TTL expiration" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_hardswap_pending_ttl_expired_total counter" + ); + let _ = writeln!( + out, + "telemt_me_hardswap_pending_ttl_expired_total {}", + if me_allows_normal { + stats.get_me_hardswap_pending_ttl_expired_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_single_endpoint_outage_enter_total Single-endpoint DC outage transitions to active state" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_single_endpoint_outage_enter_total counter" + ); + let _ = writeln!( + out, + "telemt_me_single_endpoint_outage_enter_total {}", + if me_allows_normal { + stats.get_me_single_endpoint_outage_enter_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_single_endpoint_outage_exit_total Single-endpoint DC outage recovery transitions" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_single_endpoint_outage_exit_total counter" + ); + let _ = writeln!( + out, + "telemt_me_single_endpoint_outage_exit_total {}", + if me_allows_normal { + stats.get_me_single_endpoint_outage_exit_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_single_endpoint_outage_reconnect_attempt_total Reconnect attempts performed during single-endpoint outages" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_single_endpoint_outage_reconnect_attempt_total counter" + ); + let _ = writeln!( + out, + "telemt_me_single_endpoint_outage_reconnect_attempt_total {}", + if me_allows_normal { + stats.get_me_single_endpoint_outage_reconnect_attempt_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_single_endpoint_outage_reconnect_success_total Successful reconnect attempts during single-endpoint outages" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_single_endpoint_outage_reconnect_success_total counter" + ); + let _ = writeln!( + out, + "telemt_me_single_endpoint_outage_reconnect_success_total {}", + if me_allows_normal { + stats.get_me_single_endpoint_outage_reconnect_success_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_single_endpoint_quarantine_bypass_total Outage reconnect attempts that bypassed quarantine" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_single_endpoint_quarantine_bypass_total counter" + ); + let _ = writeln!( + out, + "telemt_me_single_endpoint_quarantine_bypass_total {}", + if me_allows_normal { + stats.get_me_single_endpoint_quarantine_bypass_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_single_endpoint_shadow_rotate_total Successful periodic shadow rotations for single-endpoint DC groups" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_single_endpoint_shadow_rotate_total counter" + ); + let _ = writeln!( + out, + "telemt_me_single_endpoint_shadow_rotate_total {}", + if me_allows_normal { + stats.get_me_single_endpoint_shadow_rotate_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_single_endpoint_shadow_rotate_skipped_quarantine_total Shadow rotations skipped because endpoint is quarantined" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_single_endpoint_shadow_rotate_skipped_quarantine_total counter" + ); + let _ = writeln!( + out, + "telemt_me_single_endpoint_shadow_rotate_skipped_quarantine_total {}", + if me_allows_normal { + stats.get_me_single_endpoint_shadow_rotate_skipped_quarantine_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_floor_mode Runtime ME writer floor policy mode" + ); + let _ = writeln!(out, "# TYPE telemt_me_floor_mode gauge"); +} diff --git a/src/metrics/render/me_recovery.rs b/src/metrics/render/me_recovery.rs new file mode 100644 index 0000000..a54773f --- /dev/null +++ b/src/metrics/render/me_recovery.rs @@ -0,0 +1,369 @@ +use super::*; +use std::fmt::Write; + +pub(super) fn render( + out: &mut String, + stats: &Stats, + me_allows_normal: bool, + me_allows_debug: bool, +) { + let _ = writeln!( + out, + "# HELP telemt_secure_padding_invalid_total Invalid secure frame lengths" + ); + let _ = writeln!(out, "# TYPE telemt_secure_padding_invalid_total counter"); + let _ = writeln!( + out, + "telemt_secure_padding_invalid_total {}", + if me_allows_normal { + stats.get_secure_padding_invalid() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_desync_total Total crypto-desync detections" + ); + let _ = writeln!(out, "# TYPE telemt_desync_total counter"); + let _ = writeln!( + out, + "telemt_desync_total {}", + if me_allows_normal { + stats.get_desync_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_desync_full_logged_total Full forensic desync logs emitted" + ); + let _ = writeln!(out, "# TYPE telemt_desync_full_logged_total counter"); + let _ = writeln!( + out, + "telemt_desync_full_logged_total {}", + if me_allows_normal { + stats.get_desync_full_logged() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_desync_suppressed_total Suppressed desync forensic events" + ); + let _ = writeln!(out, "# TYPE telemt_desync_suppressed_total counter"); + let _ = writeln!( + out, + "telemt_desync_suppressed_total {}", + if me_allows_normal { + stats.get_desync_suppressed() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_desync_frames_bucket_total Desync count by frames_ok bucket" + ); + let _ = writeln!(out, "# TYPE telemt_desync_frames_bucket_total counter"); + let _ = writeln!( + out, + "telemt_desync_frames_bucket_total{{bucket=\"0\"}} {}", + if me_allows_normal { + stats.get_desync_frames_bucket_0() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_desync_frames_bucket_total{{bucket=\"1_2\"}} {}", + if me_allows_normal { + stats.get_desync_frames_bucket_1_2() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_desync_frames_bucket_total{{bucket=\"3_10\"}} {}", + if me_allows_normal { + stats.get_desync_frames_bucket_3_10() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_desync_frames_bucket_total{{bucket=\"gt_10\"}} {}", + if me_allows_normal { + stats.get_desync_frames_bucket_gt_10() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_pool_swap_total Successful ME pool swaps" + ); + let _ = writeln!(out, "# TYPE telemt_pool_swap_total counter"); + let _ = writeln!( + out, + "telemt_pool_swap_total {}", + if me_allows_normal { + stats.get_pool_swap_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_pool_drain_active Active draining ME writers" + ); + let _ = writeln!(out, "# TYPE telemt_pool_drain_active gauge"); + let _ = writeln!( + out, + "telemt_pool_drain_active {}", + if me_allows_debug { + stats.get_pool_drain_active() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_pool_force_close_total Forced close events for draining writers" + ); + let _ = writeln!(out, "# TYPE telemt_pool_force_close_total counter"); + let _ = writeln!( + out, + "telemt_pool_force_close_total {}", + if me_allows_normal { + stats.get_pool_force_close_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_pool_stale_pick_total Stale writer fallback picks for new binds" + ); + let _ = writeln!(out, "# TYPE telemt_pool_stale_pick_total counter"); + let _ = writeln!( + out, + "telemt_pool_stale_pick_total {}", + if me_allows_normal { + stats.get_pool_stale_pick_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_writer_removed_total Total ME writer removals" + ); + let _ = writeln!(out, "# TYPE telemt_me_writer_removed_total counter"); + let _ = writeln!( + out, + "telemt_me_writer_removed_total {}", + if me_allows_debug { + stats.get_me_writer_removed_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_writer_removed_unexpected_total Unexpected ME writer removals that triggered refill" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_writer_removed_unexpected_total counter" + ); + let _ = writeln!( + out, + "telemt_me_writer_removed_unexpected_total {}", + if me_allows_normal { + stats.get_me_writer_removed_unexpected_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_refill_triggered_total Immediate ME refill runs started" + ); + let _ = writeln!(out, "# TYPE telemt_me_refill_triggered_total counter"); + let _ = writeln!( + out, + "telemt_me_refill_triggered_total {}", + if me_allows_debug { + stats.get_me_refill_triggered_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_refill_skipped_inflight_total Immediate ME refill skips due to inflight dedup" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_refill_skipped_inflight_total counter" + ); + let _ = writeln!( + out, + "telemt_me_refill_skipped_inflight_total {}", + if me_allows_debug { + stats.get_me_refill_skipped_inflight_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_refill_failed_total Immediate ME refill failures" + ); + let _ = writeln!(out, "# TYPE telemt_me_refill_failed_total counter"); + let _ = writeln!( + out, + "telemt_me_refill_failed_total {}", + if me_allows_normal { + stats.get_me_refill_failed_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_writer_restored_same_endpoint_total Refilled ME writer restored on the same endpoint" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_writer_restored_same_endpoint_total counter" + ); + let _ = writeln!( + out, + "telemt_me_writer_restored_same_endpoint_total {}", + if me_allows_normal { + stats.get_me_writer_restored_same_endpoint_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_writer_restored_fallback_total Refilled ME writer restored via fallback endpoint" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_writer_restored_fallback_total counter" + ); + let _ = writeln!( + out, + "telemt_me_writer_restored_fallback_total {}", + if me_allows_normal { + stats.get_me_writer_restored_fallback_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_no_writer_failfast_total ME route failfast errors due to missing writer in bounded wait window" + ); + let _ = writeln!(out, "# TYPE telemt_me_no_writer_failfast_total counter"); + let _ = writeln!( + out, + "telemt_me_no_writer_failfast_total {}", + if me_allows_normal { + stats.get_me_no_writer_failfast_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_hybrid_timeout_total ME hybrid route timeouts after bounded retry window" + ); + let _ = writeln!(out, "# TYPE telemt_me_hybrid_timeout_total counter"); + let _ = writeln!( + out, + "telemt_me_hybrid_timeout_total {}", + if me_allows_normal { + stats.get_me_hybrid_timeout_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_async_recovery_trigger_total Async ME recovery trigger attempts from route path" + ); + let _ = writeln!(out, "# TYPE telemt_me_async_recovery_trigger_total counter"); + let _ = writeln!( + out, + "telemt_me_async_recovery_trigger_total {}", + if me_allows_normal { + stats.get_me_async_recovery_trigger_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_me_inline_recovery_total Legacy inline ME recovery attempts from route path" + ); + let _ = writeln!(out, "# TYPE telemt_me_inline_recovery_total counter"); + let _ = writeln!( + out, + "telemt_me_inline_recovery_total {}", + if me_allows_normal { + stats.get_me_inline_recovery_total() + } else { + 0 + } + ); + + let unresolved_writer_losses = if me_allows_normal { + stats + .get_me_writer_removed_unexpected_total() + .saturating_sub( + stats + .get_me_writer_restored_same_endpoint_total() + .saturating_add(stats.get_me_writer_restored_fallback_total()), + ) + } else { + 0 + }; + let _ = writeln!( + out, + "# HELP telemt_me_writer_removed_unexpected_minus_restored_total Unexpected writer removals not yet compensated by restore" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_writer_removed_unexpected_minus_restored_total gauge" + ); + let _ = writeln!( + out, + "telemt_me_writer_removed_unexpected_minus_restored_total {}", + unresolved_writer_losses + ); +} diff --git a/src/metrics/render/process.rs b/src/metrics/render/process.rs new file mode 100644 index 0000000..e71e2e8 --- /dev/null +++ b/src/metrics/render/process.rs @@ -0,0 +1,234 @@ +use super::*; +use std::fmt::Write; + +pub(super) fn render( + out: &mut String, + stats: &Stats, + shared_state: &ProxySharedState, + telemetry: crate::stats::telemetry::TelemetryPolicy, + tls_full_cert_budget: &TlsFullCertBudget, +) { + let core_enabled = telemetry.core_enabled; + let user_enabled = telemetry.user_enabled; + let _ = writeln!( + out, + "# HELP telemt_build_info Build information for the running telemt binary" + ); + let _ = writeln!(out, "# TYPE telemt_build_info gauge"); + let _ = writeln!( + out, + "telemt_build_info{{version=\"{}\"}} 1", + env!("CARGO_PKG_VERSION") + ); + + let _ = writeln!(out, "# HELP telemt_uptime_seconds Proxy uptime"); + let _ = writeln!(out, "# TYPE telemt_uptime_seconds gauge"); + let _ = writeln!(out, "telemt_uptime_seconds {:.1}", stats.uptime_secs()); + + let _ = writeln!( + out, + "# HELP telemt_telemetry_core_enabled Runtime core telemetry switch" + ); + let _ = writeln!(out, "# TYPE telemt_telemetry_core_enabled gauge"); + let _ = writeln!( + out, + "telemt_telemetry_core_enabled {}", + if core_enabled { 1 } else { 0 } + ); + + let _ = writeln!( + out, + "# HELP telemt_telemetry_user_enabled Runtime per-user telemetry switch" + ); + let _ = writeln!(out, "# TYPE telemt_telemetry_user_enabled gauge"); + let _ = writeln!( + out, + "telemt_telemetry_user_enabled {}", + if user_enabled { 1 } else { 0 } + ); + let _ = writeln!( + out, + "# HELP telemt_stats_user_entries Retained per-user stats entries" + ); + let _ = writeln!(out, "# TYPE telemt_stats_user_entries gauge"); + let _ = writeln!(out, "telemt_stats_user_entries {}", stats.user_stats_len()); + + let _ = writeln!( + out, + "# HELP telemt_telemetry_me_level Runtime ME telemetry level flag" + ); + let _ = writeln!(out, "# TYPE telemt_telemetry_me_level gauge"); + let _ = writeln!( + out, + "telemt_telemetry_me_level{{level=\"silent\"}} {}", + if matches!(telemetry.me_level, crate::config::MeTelemetryLevel::Silent) { + 1 + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_telemetry_me_level{{level=\"normal\"}} {}", + if matches!(telemetry.me_level, crate::config::MeTelemetryLevel::Normal) { + 1 + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_telemetry_me_level{{level=\"debug\"}} {}", + if matches!(telemetry.me_level, crate::config::MeTelemetryLevel::Debug) { + 1 + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_buffer_pool_buffers_total Snapshot of pooled and allocated buffers" + ); + let _ = writeln!(out, "# TYPE telemt_buffer_pool_buffers_total gauge"); + let _ = writeln!( + out, + "telemt_buffer_pool_buffers_total{{kind=\"pooled\"}} {}", + stats.get_buffer_pool_pooled_gauge() + ); + let _ = writeln!( + out, + "telemt_buffer_pool_buffers_total{{kind=\"allocated\"}} {}", + stats.get_buffer_pool_allocated_gauge() + ); + let _ = writeln!( + out, + "telemt_buffer_pool_buffers_total{{kind=\"in_use\"}} {}", + stats.get_buffer_pool_in_use_gauge() + ); + let _ = writeln!( + out, + "# HELP telemt_buffer_pool_events_total Buffer-pool allocation lifecycle events" + ); + let _ = writeln!(out, "# TYPE telemt_buffer_pool_events_total counter"); + let _ = writeln!( + out, + "telemt_buffer_pool_events_total{{event=\"replaced_nonstandard\"}} {}", + stats.get_buffer_pool_replaced_nonstandard_total() + ); + + let direct_budget = shared_state.direct_buffer_budget.snapshot(); + let _ = writeln!( + out, + "# HELP telemt_direct_relay_buffer_budget_bytes Direct relay copy-buffer budget and memory inputs" + ); + let _ = writeln!(out, "# TYPE telemt_direct_relay_buffer_budget_bytes gauge"); + for (kind, value) in [ + ("hard_limit", direct_budget.hard_limit_bytes), + ("target", direct_budget.target_bytes), + ("reserved", direct_budget.reserved_bytes), + ("memory_total", direct_budget.memory_total_bytes), + ("memory_available", direct_budget.memory_available_bytes), + ("process_rss", direct_budget.process_rss_bytes), + ] { + let _ = writeln!( + out, + "telemt_direct_relay_buffer_budget_bytes{{kind=\"{}\"}} {}", + kind, value + ); + } + let _ = writeln!( + out, + "# HELP telemt_direct_relay_buffer_budget_events_total Direct relay buffer-budget lifecycle events" + ); + let _ = writeln!( + out, + "# TYPE telemt_direct_relay_buffer_budget_events_total counter" + ); + for (result, value) in [ + ("promotion", direct_budget.promotion_total), + ("promotion_denied", direct_budget.promotion_denied_total), + ("minimum_fallback", direct_budget.minimum_fallback_total), + ("admission_rejected", direct_budget.admission_rejected_total), + ("quiet_demotion", direct_budget.quiet_demotion_total), + ( + "write_pressure_demotion", + direct_budget.write_pressure_demotion_total, + ), + ( + "global_pressure_demotion", + direct_budget.global_pressure_demotion_total, + ), + ] { + let _ = writeln!( + out, + "telemt_direct_relay_buffer_budget_events_total{{result=\"{}\"}} {}", + result, value + ); + } + let _ = writeln!( + out, + "# HELP telemt_direct_relay_buffer_sessions Current Direct relay sessions by adaptive tier" + ); + let _ = writeln!(out, "# TYPE telemt_direct_relay_buffer_sessions gauge"); + for (tier, value) in ["base", "tier1", "tier2", "tier3"] + .into_iter() + .zip(direct_budget.tier_sessions) + { + let _ = writeln!( + out, + "telemt_direct_relay_buffer_sessions{{tier=\"{}\"}} {}", + tier, value + ); + } + + let _ = writeln!( + out, + "# HELP telemt_tls_fetch_profile_cache_entries Current adaptive TLS fetch profile-cache entries" + ); + let _ = writeln!(out, "# TYPE telemt_tls_fetch_profile_cache_entries gauge"); + let _ = writeln!( + out, + "telemt_tls_fetch_profile_cache_entries {}", + fetcher::profile_cache_entries_for_metrics() + ); + let _ = writeln!( + out, + "# HELP telemt_tls_fetch_profile_cache_cap_drops_total Profile-cache winner inserts skipped because the cache cap was reached" + ); + let _ = writeln!( + out, + "# TYPE telemt_tls_fetch_profile_cache_cap_drops_total counter" + ); + let _ = writeln!( + out, + "telemt_tls_fetch_profile_cache_cap_drops_total {}", + fetcher::profile_cache_cap_drops_for_metrics() + ); + let _ = writeln!( + out, + "# HELP telemt_tls_front_full_cert_budget_entries Current domain and IP entries tracked by the process-owned TLS full-cert budget" + ); + let _ = writeln!( + out, + "# TYPE telemt_tls_front_full_cert_budget_entries gauge" + ); + let _ = writeln!( + out, + "telemt_tls_front_full_cert_budget_entries {}", + tls_full_cert_budget.entries_for_metrics() + ); + let _ = writeln!( + out, + "# HELP telemt_tls_front_full_cert_budget_cap_drops_total New domain and IP entries denied full-cert budget tracking because a bound was reached" + ); + let _ = writeln!( + out, + "# TYPE telemt_tls_front_full_cert_budget_cap_drops_total counter" + ); + let _ = writeln!( + out, + "telemt_tls_front_full_cert_budget_cap_drops_total {}", + tls_full_cert_budget.cap_drops_for_metrics() + ); +} diff --git a/src/metrics/render/traffic.rs b/src/metrics/render/traffic.rs new file mode 100644 index 0000000..29636f6 --- /dev/null +++ b/src/metrics/render/traffic.rs @@ -0,0 +1,516 @@ +use super::*; +use std::fmt::Write; + +pub(super) fn render( + out: &mut String, + stats: &Stats, + shared_state: &ProxySharedState, + config: &ProxyConfig, + core_enabled: bool, + me_allows_normal: bool, + me_allows_debug: bool, +) { + let limiter_metrics = shared_state.traffic_limiter.metrics_snapshot(); + let _ = writeln!( + out, + "# HELP telemt_rate_limiter_burst_bound_bytes Configured upper bound for one direct relay rate-limit burst" + ); + let _ = writeln!(out, "# TYPE telemt_rate_limiter_burst_bound_bytes gauge"); + let _ = writeln!( + out, + "telemt_rate_limiter_burst_bound_bytes{{direction=\"up\"}} {}", + if core_enabled { + config.general.direct_relay_copy_buf_c2s_bytes + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_rate_limiter_burst_bound_bytes{{direction=\"down\"}} {}", + if core_enabled { + config.general.direct_relay_copy_buf_s2c_bytes + } else { + 0 + } + ); + let _ = writeln!( + out, + "# HELP telemt_rate_limiter_throttle_total Traffic limiter throttle events by scope and direction" + ); + let _ = writeln!(out, "# TYPE telemt_rate_limiter_throttle_total counter"); + let _ = writeln!( + out, + "telemt_rate_limiter_throttle_total{{scope=\"user\",direction=\"up\"}} {}", + if core_enabled { + limiter_metrics.user_throttle_up_total + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_rate_limiter_throttle_total{{scope=\"user\",direction=\"down\"}} {}", + if core_enabled { + limiter_metrics.user_throttle_down_total + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_rate_limiter_throttle_total{{scope=\"cidr\",direction=\"up\"}} {}", + if core_enabled { + limiter_metrics.cidr_throttle_up_total + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_rate_limiter_throttle_total{{scope=\"cidr\",direction=\"down\"}} {}", + if core_enabled { + limiter_metrics.cidr_throttle_down_total + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_rate_limiter_wait_ms_total Traffic limiter accumulated wait time in milliseconds by scope and direction" + ); + let _ = writeln!(out, "# TYPE telemt_rate_limiter_wait_ms_total counter"); + let _ = writeln!( + out, + "telemt_rate_limiter_wait_ms_total{{scope=\"user\",direction=\"up\"}} {}", + if core_enabled { + limiter_metrics.user_wait_up_ms_total + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_rate_limiter_wait_ms_total{{scope=\"user\",direction=\"down\"}} {}", + if core_enabled { + limiter_metrics.user_wait_down_ms_total + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_rate_limiter_wait_ms_total{{scope=\"cidr\",direction=\"up\"}} {}", + if core_enabled { + limiter_metrics.cidr_wait_up_ms_total + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_rate_limiter_wait_ms_total{{scope=\"cidr\",direction=\"down\"}} {}", + if core_enabled { + limiter_metrics.cidr_wait_down_ms_total + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_rate_limiter_active_leases Active relay leases under rate limiting by scope" + ); + let _ = writeln!(out, "# TYPE telemt_rate_limiter_active_leases gauge"); + let _ = writeln!( + out, + "telemt_rate_limiter_active_leases{{scope=\"user\"}} {}", + if core_enabled { + limiter_metrics.user_active_leases + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_rate_limiter_active_leases{{scope=\"cidr\"}} {}", + if core_enabled { + limiter_metrics.cidr_active_leases + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_rate_limiter_policy_entries Active rate-limit policy entries by scope" + ); + let _ = writeln!(out, "# TYPE telemt_rate_limiter_policy_entries gauge"); + let _ = writeln!( + out, + "telemt_rate_limiter_policy_entries{{scope=\"user\"}} {}", + if core_enabled { + limiter_metrics.user_policy_entries + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_rate_limiter_policy_entries{{scope=\"cidr\"}} {}", + if core_enabled { + limiter_metrics.cidr_policy_entries + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_upstream_connect_attempt_total Upstream connect attempts across all requests" + ); + let _ = writeln!(out, "# TYPE telemt_upstream_connect_attempt_total counter"); + let _ = writeln!( + out, + "telemt_upstream_connect_attempt_total {}", + if core_enabled { + stats.get_upstream_connect_attempt_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_upstream_connect_success_total Successful upstream connect request cycles" + ); + let _ = writeln!(out, "# TYPE telemt_upstream_connect_success_total counter"); + let _ = writeln!( + out, + "telemt_upstream_connect_success_total {}", + if core_enabled { + stats.get_upstream_connect_success_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_upstream_connect_fail_total Failed upstream connect request cycles" + ); + let _ = writeln!(out, "# TYPE telemt_upstream_connect_fail_total counter"); + let _ = writeln!( + out, + "telemt_upstream_connect_fail_total {}", + if core_enabled { + stats.get_upstream_connect_fail_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_upstream_connect_failfast_hard_error_total Hard errors that triggered upstream connect failfast" + ); + let _ = writeln!( + out, + "# TYPE telemt_upstream_connect_failfast_hard_error_total counter" + ); + let _ = writeln!( + out, + "telemt_upstream_connect_failfast_hard_error_total {}", + if core_enabled { + stats.get_upstream_connect_failfast_hard_error_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_upstream_connect_attempts_per_request Histogram-like buckets for attempts per upstream connect request cycle" + ); + let _ = writeln!( + out, + "# TYPE telemt_upstream_connect_attempts_per_request counter" + ); + let _ = writeln!( + out, + "telemt_upstream_connect_attempts_per_request{{bucket=\"1\"}} {}", + if core_enabled { + stats.get_upstream_connect_attempts_bucket_1() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_upstream_connect_attempts_per_request{{bucket=\"2\"}} {}", + if core_enabled { + stats.get_upstream_connect_attempts_bucket_2() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_upstream_connect_attempts_per_request{{bucket=\"3_4\"}} {}", + if core_enabled { + stats.get_upstream_connect_attempts_bucket_3_4() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_upstream_connect_attempts_per_request{{bucket=\"gt_4\"}} {}", + if core_enabled { + stats.get_upstream_connect_attempts_bucket_gt_4() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_upstream_connect_duration_success_total Histogram-like buckets of successful upstream connect cycle duration" + ); + let _ = writeln!( + out, + "# TYPE telemt_upstream_connect_duration_success_total counter" + ); + let _ = writeln!( + out, + "telemt_upstream_connect_duration_success_total{{bucket=\"le_100ms\"}} {}", + if core_enabled { + stats.get_upstream_connect_duration_success_bucket_le_100ms() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_upstream_connect_duration_success_total{{bucket=\"101_500ms\"}} {}", + if core_enabled { + stats.get_upstream_connect_duration_success_bucket_101_500ms() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_upstream_connect_duration_success_total{{bucket=\"501_1000ms\"}} {}", + if core_enabled { + stats.get_upstream_connect_duration_success_bucket_501_1000ms() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_upstream_connect_duration_success_total{{bucket=\"gt_1000ms\"}} {}", + if core_enabled { + stats.get_upstream_connect_duration_success_bucket_gt_1000ms() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_upstream_connect_duration_fail_total Histogram-like buckets of failed upstream connect cycle duration" + ); + let _ = writeln!( + out, + "# TYPE telemt_upstream_connect_duration_fail_total counter" + ); + let _ = writeln!( + out, + "telemt_upstream_connect_duration_fail_total{{bucket=\"le_100ms\"}} {}", + if core_enabled { + stats.get_upstream_connect_duration_fail_bucket_le_100ms() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_upstream_connect_duration_fail_total{{bucket=\"101_500ms\"}} {}", + if core_enabled { + stats.get_upstream_connect_duration_fail_bucket_101_500ms() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_upstream_connect_duration_fail_total{{bucket=\"501_1000ms\"}} {}", + if core_enabled { + stats.get_upstream_connect_duration_fail_bucket_501_1000ms() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_upstream_connect_duration_fail_total{{bucket=\"gt_1000ms\"}} {}", + if core_enabled { + stats.get_upstream_connect_duration_fail_bucket_gt_1000ms() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_keepalive_sent_total ME keepalive frames sent" + ); + let _ = writeln!(out, "# TYPE telemt_me_keepalive_sent_total counter"); + let _ = writeln!( + out, + "telemt_me_keepalive_sent_total {}", + if me_allows_debug { + stats.get_me_keepalive_sent() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_keepalive_failed_total ME keepalive send failures" + ); + let _ = writeln!(out, "# TYPE telemt_me_keepalive_failed_total counter"); + let _ = writeln!( + out, + "telemt_me_keepalive_failed_total {}", + if me_allows_normal { + stats.get_me_keepalive_failed() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_keepalive_pong_total ME keepalive pong replies" + ); + let _ = writeln!(out, "# TYPE telemt_me_keepalive_pong_total counter"); + let _ = writeln!( + out, + "telemt_me_keepalive_pong_total {}", + if me_allows_debug { + stats.get_me_keepalive_pong() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_keepalive_timeout_total ME keepalive ping timeouts" + ); + let _ = writeln!(out, "# TYPE telemt_me_keepalive_timeout_total counter"); + let _ = writeln!( + out, + "telemt_me_keepalive_timeout_total {}", + if me_allows_normal { + stats.get_me_keepalive_timeout() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_rpc_proxy_req_signal_sent_total Service RPC_PROXY_REQ activity signals sent" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_rpc_proxy_req_signal_sent_total counter" + ); + let _ = writeln!( + out, + "telemt_me_rpc_proxy_req_signal_sent_total {}", + if me_allows_normal { + stats.get_me_rpc_proxy_req_signal_sent_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_rpc_proxy_req_signal_failed_total Service RPC_PROXY_REQ activity signal failures" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_rpc_proxy_req_signal_failed_total counter" + ); + let _ = writeln!( + out, + "telemt_me_rpc_proxy_req_signal_failed_total {}", + if me_allows_normal { + stats.get_me_rpc_proxy_req_signal_failed_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_rpc_proxy_req_signal_skipped_no_meta_total Service RPC_PROXY_REQ skipped due to missing writer metadata" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_rpc_proxy_req_signal_skipped_no_meta_total counter" + ); + let _ = writeln!( + out, + "telemt_me_rpc_proxy_req_signal_skipped_no_meta_total {}", + if me_allows_normal { + stats.get_me_rpc_proxy_req_signal_skipped_no_meta_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_rpc_proxy_req_signal_response_total Service RPC_PROXY_REQ responses observed" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_rpc_proxy_req_signal_response_total counter" + ); + let _ = writeln!( + out, + "telemt_me_rpc_proxy_req_signal_response_total {}", + if me_allows_normal { + stats.get_me_rpc_proxy_req_signal_response_total() + } else { + 0 + } + ); + + let _ = writeln!( + out, + "# HELP telemt_me_rpc_proxy_req_signal_close_sent_total Service RPC_CLOSE_EXT sent after activity signals" + ); + let _ = writeln!( + out, + "# TYPE telemt_me_rpc_proxy_req_signal_close_sent_total counter" + ); + let _ = writeln!( + out, + "telemt_me_rpc_proxy_req_signal_close_sent_total {}", + if me_allows_normal { + stats.get_me_rpc_proxy_req_signal_close_sent_total() + } else { + 0 + } + ); +} diff --git a/src/metrics/render/users.rs b/src/metrics/render/users.rs new file mode 100644 index 0000000..bdfd5e8 --- /dev/null +++ b/src/metrics/render/users.rs @@ -0,0 +1,307 @@ +use super::*; +use std::fmt::Write; + +pub(super) async fn render( + out: &mut String, + stats: &Stats, + config: &ProxyConfig, + ip_tracker: &UserIpTracker, + core_enabled: bool, + user_enabled: bool, +) { + let _ = writeln!( + out, + "# HELP telemt_user_connections_total Per-user total connections" + ); + let _ = writeln!(out, "# TYPE telemt_user_connections_total counter"); + let _ = writeln!( + out, + "# HELP telemt_user_connections_current Per-user active connections" + ); + let _ = writeln!(out, "# TYPE telemt_user_connections_current gauge"); + let _ = writeln!( + out, + "# HELP telemt_user_octets_from_client_total Per-user total bytes received" + ); + let _ = writeln!(out, "# TYPE telemt_user_octets_from_client_total counter"); + let _ = writeln!( + out, + "# HELP telemt_user_octets_to_client_total Per-user total bytes sent" + ); + let _ = writeln!(out, "# TYPE telemt_user_octets_to_client_total counter"); + let _ = writeln!( + out, + "# HELP telemt_user_msgs_from_client_total Per-user total messages received" + ); + let _ = writeln!(out, "# TYPE telemt_user_msgs_from_client_total counter"); + let _ = writeln!( + out, + "# HELP telemt_user_msgs_to_client_total Per-user total messages sent" + ); + let _ = writeln!(out, "# TYPE telemt_user_msgs_to_client_total counter"); + let _ = writeln!( + out, + "# HELP telemt_ip_reservation_rollback_total IP reservation rollbacks caused by later limit checks" + ); + let _ = writeln!(out, "# TYPE telemt_ip_reservation_rollback_total counter"); + let _ = writeln!( + out, + "telemt_ip_reservation_rollback_total{{reason=\"tcp_limit\"}} {}", + if core_enabled { + stats.get_ip_reservation_rollback_tcp_limit_total() + } else { + 0 + } + ); + let _ = writeln!( + out, + "telemt_ip_reservation_rollback_total{{reason=\"quota_limit\"}} {}", + if core_enabled { + stats.get_ip_reservation_rollback_quota_limit_total() + } else { + 0 + } + ); + let ip_memory = ip_tracker.memory_stats().await; + let _ = writeln!( + out, + "# HELP telemt_ip_tracker_users Number of users tracked by IP limiter state" + ); + let _ = writeln!(out, "# TYPE telemt_ip_tracker_users gauge"); + let _ = writeln!( + out, + "telemt_ip_tracker_users{{scope=\"active\"}} {}", + ip_memory.active_users + ); + let _ = writeln!( + out, + "telemt_ip_tracker_users{{scope=\"recent\"}} {}", + ip_memory.recent_users + ); + let _ = writeln!( + out, + "# HELP telemt_ip_tracker_entries Number of IP entries tracked by limiter state" + ); + let _ = writeln!(out, "# TYPE telemt_ip_tracker_entries gauge"); + let _ = writeln!( + out, + "telemt_ip_tracker_entries{{scope=\"active\"}} {}", + ip_memory.active_entries + ); + let _ = writeln!( + out, + "telemt_ip_tracker_entries{{scope=\"recent\"}} {}", + ip_memory.recent_entries + ); + let _ = writeln!( + out, + "# HELP telemt_ip_tracker_cleanup_queue_len Deferred disconnect cleanup queue length" + ); + let _ = writeln!(out, "# TYPE telemt_ip_tracker_cleanup_queue_len gauge"); + let _ = writeln!( + out, + "telemt_ip_tracker_cleanup_queue_len {}", + ip_memory.cleanup_queue_len + ); + let _ = writeln!( + out, + "# HELP telemt_ip_tracker_cleanup_total Release cleanups deferred through the cleanup queue" + ); + let _ = writeln!(out, "# TYPE telemt_ip_tracker_cleanup_total counter"); + let _ = writeln!( + out, + "telemt_ip_tracker_cleanup_total{{path=\"deferred\"}} {}", + ip_memory.cleanup_deferred_releases + ); + let _ = writeln!( + out, + "# HELP telemt_ip_tracker_cap_rejects_total New connection rejects caused by global IP tracker caps" + ); + let _ = writeln!(out, "# TYPE telemt_ip_tracker_cap_rejects_total counter"); + let _ = writeln!( + out, + "telemt_ip_tracker_cap_rejects_total{{scope=\"active\"}} {}", + ip_memory.active_cap_rejects + ); + let _ = writeln!( + out, + "telemt_ip_tracker_cap_rejects_total{{scope=\"recent\"}} {}", + ip_memory.recent_cap_rejects + ); + + let mut user_stats_emitted = 0usize; + let mut user_stats_suppressed = 0usize; + let mut unique_ip_emitted = 0usize; + let mut unique_ip_suppressed = 0usize; + + if user_enabled { + for entry in stats.iter_user_stats() { + if user_stats_emitted >= USER_LABELED_METRICS_MAX_USERS { + user_stats_suppressed = user_stats_suppressed.saturating_add(1); + continue; + } + let user = entry.key(); + let s = entry.value(); + user_stats_emitted = user_stats_emitted.saturating_add(1); + let _ = writeln!( + out, + "telemt_user_connections_total{{user=\"{}\"}} {}", + user, + s.connects.load(std::sync::atomic::Ordering::Relaxed) + ); + let _ = writeln!( + out, + "telemt_user_connections_current{{user=\"{}\"}} {}", + user, + s.curr_connects.load(std::sync::atomic::Ordering::Relaxed) + ); + let _ = writeln!( + out, + "telemt_user_octets_from_client_total{{user=\"{}\"}} {}", + user, + s.octets_from_client + .load(std::sync::atomic::Ordering::Relaxed) + ); + let _ = writeln!( + out, + "telemt_user_octets_to_client_total{{user=\"{}\"}} {}", + user, + s.octets_to_client + .load(std::sync::atomic::Ordering::Relaxed) + ); + let _ = writeln!( + out, + "telemt_user_msgs_from_client_total{{user=\"{}\"}} {}", + user, + s.msgs_from_client + .load(std::sync::atomic::Ordering::Relaxed) + ); + let _ = writeln!( + out, + "telemt_user_msgs_to_client_total{{user=\"{}\"}} {}", + user, + s.msgs_to_client.load(std::sync::atomic::Ordering::Relaxed) + ); + } + + let ip_stats = ip_tracker.get_stats_snapshot().await; + let ip_counts: HashMap = ip_stats + .into_iter() + .map(|(user, count, _)| (user, count)) + .collect(); + + let mut unique_users = BTreeSet::new(); + unique_users.extend(config.access.users.keys().cloned()); + unique_users.extend(config.access.user_max_unique_ips.keys().cloned()); + unique_users.extend(ip_counts.keys().cloned()); + let unique_users_vec: Vec = unique_users.iter().cloned().collect(); + let recent_counts = ip_tracker + .get_recent_counts_for_users_snapshot(&unique_users_vec) + .await; + + let _ = writeln!( + out, + "# HELP telemt_user_unique_ips_current Per-user current number of unique active IPs" + ); + let _ = writeln!(out, "# TYPE telemt_user_unique_ips_current gauge"); + let _ = writeln!( + out, + "# HELP telemt_user_unique_ips_recent_window Per-user unique IPs seen in configured observation window" + ); + let _ = writeln!(out, "# TYPE telemt_user_unique_ips_recent_window gauge"); + let _ = writeln!( + out, + "# HELP telemt_user_unique_ips_limit Effective per-user unique IP limit (0 means unlimited)" + ); + let _ = writeln!(out, "# TYPE telemt_user_unique_ips_limit gauge"); + let _ = writeln!( + out, + "# HELP telemt_user_unique_ips_utilization Per-user unique IP usage ratio (0 for unlimited)" + ); + let _ = writeln!(out, "# TYPE telemt_user_unique_ips_utilization gauge"); + + for user in unique_users { + if unique_ip_emitted >= USER_LABELED_METRICS_MAX_USERS { + unique_ip_suppressed = unique_ip_suppressed.saturating_add(1); + continue; + } + unique_ip_emitted = unique_ip_emitted.saturating_add(1); + let current = ip_counts.get(&user).copied().unwrap_or(0); + let limit = config + .access + .user_max_unique_ips + .get(&user) + .copied() + .filter(|limit| *limit > 0) + .or((config.access.user_max_unique_ips_global_each > 0) + .then_some(config.access.user_max_unique_ips_global_each)) + .unwrap_or(0); + let utilization = if limit > 0 { + current as f64 / limit as f64 + } else { + 0.0 + }; + let _ = writeln!( + out, + "telemt_user_unique_ips_current{{user=\"{}\"}} {}", + user, current + ); + let _ = writeln!( + out, + "telemt_user_unique_ips_recent_window{{user=\"{}\"}} {}", + user, + recent_counts.get(&user).copied().unwrap_or(0) + ); + let _ = writeln!( + out, + "telemt_user_unique_ips_limit{{user=\"{}\"}} {}", + user, limit + ); + let _ = writeln!( + out, + "telemt_user_unique_ips_utilization{{user=\"{}\"}} {:.6}", + user, utilization + ); + } + } + + let _ = writeln!( + out, + "# HELP telemt_telemetry_user_series_suppressed User-labeled metric series suppression flag" + ); + let _ = writeln!(out, "# TYPE telemt_telemetry_user_series_suppressed gauge"); + let _ = writeln!( + out, + "telemt_telemetry_user_series_suppressed {}", + if user_enabled && user_stats_suppressed == 0 && unique_ip_suppressed == 0 { + 0 + } else { + 1 + } + ); + let _ = writeln!( + out, + "# HELP telemt_telemetry_user_series_users User-labeled metric users by export status" + ); + let _ = writeln!(out, "# TYPE telemt_telemetry_user_series_users gauge"); + let _ = writeln!( + out, + "telemt_telemetry_user_series_users{{family=\"stats\",status=\"emitted\"}} {}", + user_stats_emitted + ); + let _ = writeln!( + out, + "telemt_telemetry_user_series_users{{family=\"stats\",status=\"suppressed\"}} {}", + user_stats_suppressed + ); + let _ = writeln!( + out, + "telemt_telemetry_user_series_users{{family=\"unique_ip\",status=\"emitted\"}} {}", + unique_ip_emitted + ); + let _ = writeln!( + out, + "telemt_telemetry_user_series_users{{family=\"unique_ip\",status=\"suppressed\"}} {}", + unique_ip_suppressed + ); +} diff --git a/src/metrics/tests.rs b/src/metrics/tests.rs new file mode 100644 index 0000000..04a0293 --- /dev/null +++ b/src/metrics/tests.rs @@ -0,0 +1,480 @@ +use super::*; +use http_body_util::BodyExt; +use std::net::IpAddr; +use std::time::SystemTime; + +use crate::tls_front::types::{ + CachedTlsData, ParsedServerHello, TlsBehaviorProfile, TlsCertPayload, TlsProfileSource, +}; + +fn test_web_publication() -> crate::web::control::WebRuntimePublication { + let control = crate::web::control::WebRuntimeControl::new(); + control.subscribe().borrow().clone() +} + +#[tokio::test] +async fn test_render_metrics_format() { + let stats = Arc::new(Stats::new()); + let shared_state = ProxySharedState::new(); + let tracker = UserIpTracker::new(); + let mut config = ProxyConfig::default(); + config + .access + .user_max_unique_ips + .insert("alice".to_string(), 4); + + stats.increment_connects_all(); + stats.increment_connects_all(); + stats.increment_connects_bad_with_class("tls_handshake_bad_client"); + stats.increment_handshake_timeouts(); + stats.increment_handshake_failure_class("timeout"); + shared_state + .handshake + .auth_expensive_checks_total + .fetch_add(9, std::sync::atomic::Ordering::Relaxed); + shared_state + .handshake + .auth_budget_exhausted_total + .fetch_add(2, std::sync::atomic::Ordering::Relaxed); + stats.increment_upstream_connect_attempt_total(); + stats.increment_upstream_connect_attempt_total(); + stats.increment_upstream_connect_success_total(); + stats.increment_upstream_connect_fail_total(); + stats.increment_upstream_connect_failfast_hard_error_total(); + stats.observe_upstream_connect_attempts_per_request(2); + stats.observe_upstream_connect_duration_ms(220, true); + stats.observe_upstream_connect_duration_ms(1500, false); + stats.increment_me_rpc_proxy_req_signal_sent_total(); + stats.increment_me_rpc_proxy_req_signal_failed_total(); + stats.increment_me_rpc_proxy_req_signal_skipped_no_meta_total(); + stats.increment_me_rpc_proxy_req_signal_response_total(); + stats.increment_me_rpc_proxy_req_signal_close_sent_total(); + stats.increment_me_idle_close_by_peer_total(); + stats.increment_relay_idle_soft_mark_total(); + stats.increment_relay_idle_hard_close_total(); + stats.increment_relay_pressure_evict_total(); + stats.increment_relay_protocol_desync_close_total(); + stats.increment_me_d2c_batches_total(); + stats.add_me_d2c_batch_frames_total(3); + stats.add_me_d2c_batch_bytes_total(2048); + stats.increment_me_d2c_flush_reason(crate::stats::MeD2cFlushReason::AckImmediate); + stats.increment_me_d2c_data_frames_total(); + stats.increment_me_d2c_ack_frames_total(); + stats.add_me_d2c_payload_bytes_total(1800); + stats.increment_me_d2c_write_mode(crate::stats::MeD2cWriteMode::Coalesced); + stats.increment_me_d2c_quota_reject_total(crate::stats::MeD2cQuotaRejectStage::PostWrite); + stats.observe_me_d2c_frame_buf_shrink(4096); + stats.increment_me_endpoint_quarantine_total(); + stats.increment_me_endpoint_quarantine_unexpected_total(); + stats.increment_me_endpoint_quarantine_draining_suppressed_total(); + stats.increment_user_connects("alice"); + stats.increment_user_curr_connects("alice"); + stats.add_user_octets_from("alice", 1024); + stats.add_user_octets_to("alice", 2048); + stats.increment_user_msgs_from("alice"); + stats.increment_user_msgs_to("alice"); + stats.increment_user_msgs_to("alice"); + tracker + .check_and_add("alice", "203.0.113.10".parse().unwrap()) + .await + .unwrap(); + + let output = render_metrics( + &stats, + shared_state.as_ref(), + &config, + &tracker, + None, + &TlsFullCertBudget::new(), + &test_web_publication(), + ) + .await; + + assert!(output.contains(&format!( + "telemt_build_info{{version=\"{}\"}} 1", + env!("CARGO_PKG_VERSION") + ))); + assert!(output.contains("telemt_connections_total 2")); + assert!(output.contains("telemt_connections_bad_total 1")); + assert!( + output.contains( + "telemt_connections_bad_by_class_total{class=\"tls_handshake_bad_client\"} 1" + ) + ); + assert!(output.contains("telemt_handshake_timeouts_total 1")); + assert!(output.contains("telemt_handshake_failures_by_class_total{class=\"timeout\"} 1")); + assert!(output.contains("telemt_auth_expensive_checks_total 9")); + assert!(output.contains("telemt_auth_budget_exhausted_total 2")); + assert!(output.contains("telemt_upstream_connect_attempt_total 2")); + assert!(output.contains("telemt_upstream_connect_success_total 1")); + assert!(output.contains("telemt_upstream_connect_fail_total 1")); + assert!(output.contains("telemt_upstream_connect_failfast_hard_error_total 1")); + assert!(output.contains("telemt_upstream_connect_attempts_per_request{bucket=\"2\"} 1")); + assert!( + output.contains("telemt_upstream_connect_duration_success_total{bucket=\"101_500ms\"} 1") + ); + assert!(output.contains("telemt_upstream_connect_duration_fail_total{bucket=\"gt_1000ms\"} 1")); + assert!(output.contains("telemt_me_rpc_proxy_req_signal_sent_total 1")); + assert!(output.contains("telemt_me_rpc_proxy_req_signal_failed_total 1")); + assert!(output.contains("telemt_me_rpc_proxy_req_signal_skipped_no_meta_total 1")); + assert!(output.contains("telemt_me_rpc_proxy_req_signal_response_total 1")); + assert!(output.contains("telemt_me_rpc_proxy_req_signal_close_sent_total 1")); + assert!(output.contains("telemt_me_idle_close_by_peer_total 1")); + assert!(output.contains("telemt_relay_idle_soft_mark_total 1")); + assert!(output.contains("telemt_relay_idle_hard_close_total 1")); + assert!(output.contains("telemt_relay_pressure_evict_total 1")); + assert!(output.contains("telemt_relay_protocol_desync_close_total 1")); + assert!(output.contains("telemt_me_d2c_batches_total 1")); + assert!(output.contains("telemt_me_d2c_batch_frames_total 3")); + assert!(output.contains("telemt_me_d2c_batch_bytes_total 2048")); + assert!(output.contains("telemt_me_d2c_flush_reason_total{reason=\"ack_immediate\"} 1")); + assert!(output.contains("telemt_me_d2c_data_frames_total 1")); + assert!(output.contains("telemt_me_d2c_ack_frames_total 1")); + assert!(output.contains("telemt_me_d2c_payload_bytes_total 1800")); + assert!(output.contains("telemt_me_d2c_write_mode_total{mode=\"coalesced\"} 1")); + assert!(output.contains("telemt_me_d2c_quota_reject_total{stage=\"post_write\"} 1")); + assert!(output.contains("telemt_me_d2c_frame_buf_shrink_total 1")); + assert!(output.contains("telemt_me_d2c_frame_buf_shrink_bytes_total 4096")); + assert!(output.contains("telemt_me_endpoint_quarantine_total 1")); + assert!(output.contains("telemt_me_endpoint_quarantine_unexpected_total 1")); + assert!(output.contains("telemt_me_endpoint_quarantine_draining_suppressed_total 1")); + assert!(output.contains("telemt_user_connections_total{user=\"alice\"} 1")); + assert!(output.contains("telemt_user_connections_current{user=\"alice\"} 1")); + assert!(output.contains("telemt_user_octets_from_client_total{user=\"alice\"} 1024")); + assert!(output.contains("telemt_user_octets_to_client_total{user=\"alice\"} 2048")); + assert!(output.contains("telemt_user_msgs_from_client_total{user=\"alice\"} 1")); + assert!(output.contains("telemt_user_msgs_to_client_total{user=\"alice\"} 2")); + assert!(output.contains("telemt_user_unique_ips_current{user=\"alice\"} 1")); + assert!(output.contains("telemt_user_unique_ips_recent_window{user=\"alice\"} 1")); + assert!(output.contains("telemt_user_unique_ips_limit{user=\"alice\"} 4")); + assert!(output.contains("telemt_user_unique_ips_utilization{user=\"alice\"} 0.250000")); + assert!(output.contains("telemt_ip_tracker_users{scope=\"active\"} 1")); + assert!(output.contains("telemt_ip_tracker_entries{scope=\"active\"} 1")); + assert!(output.contains("telemt_ip_tracker_cleanup_queue_len 0")); +} + +#[tokio::test] +async fn test_render_tls_front_profile_health() { + let stats = Stats::new(); + let shared_state = ProxySharedState::new(); + let tracker = UserIpTracker::new(); + let mut config = ProxyConfig::default(); + config.censorship.tls_domain = "primary.example".to_string(); + config.censorship.tls_domains = vec!["fallback.example".to_string()]; + + let cache = TlsFrontCache::new( + &[ + "primary.example".to_string(), + "fallback.example".to_string(), + ], + 1024, + "tlsfront-profile-health-test", + ); + cache + .set( + "primary.example", + CachedTlsData { + server_hello_template: ParsedServerHello { + version: [0x03, 0x03], + random: [0u8; 32], + session_id: Vec::new(), + cipher_suite: [0x13, 0x01], + compression: 0, + extensions: { + let mut key_share = vec![0x00, 0x1d, 0x00, 0x20]; + key_share.resize(36, 0x42); + vec![ + crate::tls_front::types::TlsExtension { + ext_type: 0x002b, + data: vec![0x03, 0x04], + }, + crate::tls_front::types::TlsExtension { + ext_type: 0x0033, + data: key_share, + }, + ] + }, + }, + cert_info: None, + cert_payload: Some(TlsCertPayload { + cert_chain_der: vec![vec![0x30, 0x01]], + certificate_message: vec![0x0b, 0x00, 0x00, 0x00], + }), + app_data_records_sizes: vec![1024, 512], + total_app_data_len: 1536, + behavior_profile: TlsBehaviorProfile { + change_cipher_spec_count: 1, + app_data_record_sizes: vec![1024, 512], + ticket_record_sizes: vec![69], + source: TlsProfileSource::Merged, + ..TlsBehaviorProfile::default() + }, + fetched_at: SystemTime::now(), + domain: "primary.example".to_string(), + }, + ) + .await; + + let output = render_metrics( + &stats, + &shared_state, + &config, + &tracker, + Some(&cache), + &TlsFullCertBudget::new(), + &test_web_publication(), + ) + .await; + + assert!(output.contains("telemt_tls_front_profile_domains{status=\"configured\"} 2")); + assert!(output.contains("telemt_tls_front_profile_domains{status=\"emitted\"} 2")); + assert!(output.contains("telemt_tls_front_profile_domains{status=\"suppressed\"} 0")); + assert!( + output.contains("telemt_tls_front_profile_info{domain=\"primary.example\",source=\"merged\",is_default=\"false\",has_cert_info=\"false\",has_cert_payload=\"true\"} 1") + ); + assert!( + output.contains("telemt_tls_front_profile_info{domain=\"fallback.example\",source=\"default\",is_default=\"true\",has_cert_info=\"false\",has_cert_payload=\"false\"} 1") + ); + assert!( + output.contains("telemt_tls_front_profile_quality_info{domain=\"primary.example\",quality=\"raw_strict\",key_share_group=\"x25519\"} 1") + ); + assert!( + output.contains("telemt_tls_front_profile_quality_info{domain=\"fallback.example\",quality=\"fallback\",key_share_group=\"none\"} 1") + ); + assert!( + output + .contains("telemt_tls_front_profile_server_hello_bytes{domain=\"primary.example\"} 90") + ); + assert!(output.contains( + "telemt_tls_front_profile_server_hello_extensions{domain=\"primary.example\"} 2" + )); + assert!( + output.contains("telemt_tls_front_profile_app_data_records{domain=\"primary.example\"} 2") + ); + assert!( + output.contains("telemt_tls_front_profile_ticket_records{domain=\"primary.example\"} 1") + ); + assert!(output.contains( + "telemt_tls_front_profile_change_cipher_spec_records{domain=\"primary.example\"} 1" + )); + assert!( + output.contains("telemt_tls_front_profile_app_data_bytes{domain=\"primary.example\"} 1536") + ); +} + +#[tokio::test] +async fn process_tls_budget_metrics_survive_a_generation_without_tls_cache() { + let stats = Stats::new(); + let shared_state = ProxySharedState::new(); + let tracker = UserIpTracker::new(); + let config = ProxyConfig::default(); + let budget = Arc::new(TlsFullCertBudget::new()); + let cache = TlsFrontCache::new_with_full_cert_budget( + &["example.com".to_string()], + 1024, + "tlsfront-test-cache", + Arc::clone(&budget), + ); + assert!( + cache + .take_full_cert_budget_for_ip( + "example.com", + "127.0.0.1".parse().unwrap(), + Duration::from_secs(60), + ) + .await + ); + + let output = render_metrics( + &stats, + &shared_state, + &config, + &tracker, + None, + budget.as_ref(), + &test_web_publication(), + ) + .await; + + assert!(output.contains("telemt_tls_front_full_cert_budget_entries 1")); +} + +#[tokio::test] +async fn test_render_empty_stats() { + let stats = Stats::new(); + let shared_state = ProxySharedState::new(); + let tracker = UserIpTracker::new(); + let config = ProxyConfig::default(); + let output = render_metrics( + &stats, + &shared_state, + &config, + &tracker, + None, + &TlsFullCertBudget::new(), + &test_web_publication(), + ) + .await; + assert!(output.contains("telemt_connections_total 0")); + assert!(output.contains("telemt_connections_bad_total 0")); + assert!(output.contains("telemt_handshake_timeouts_total 0")); + assert!(output.contains("telemt_auth_expensive_checks_total 0")); + assert!(output.contains("telemt_auth_budget_exhausted_total 0")); + assert!(output.contains("telemt_user_unique_ips_current{user=")); + assert!(output.contains("telemt_user_unique_ips_recent_window{user=")); +} + +#[tokio::test] +async fn test_render_uses_global_each_unique_ip_limit() { + let stats = Stats::new(); + let shared_state = ProxySharedState::new(); + stats.increment_user_connects("alice"); + stats.increment_user_curr_connects("alice"); + let tracker = UserIpTracker::new(); + tracker + .check_and_add("alice", "203.0.113.10".parse().unwrap()) + .await + .unwrap(); + let mut config = ProxyConfig::default(); + config.access.user_max_unique_ips_global_each = 2; + + let output = render_metrics( + &stats, + &shared_state, + &config, + &tracker, + None, + &TlsFullCertBudget::new(), + &test_web_publication(), + ) + .await; + + assert!(output.contains("telemt_user_unique_ips_limit{user=\"alice\"} 2")); + assert!(output.contains("telemt_user_unique_ips_utilization{user=\"alice\"} 0.500000")); +} + +#[tokio::test] +async fn test_render_has_type_annotations() { + let stats = Stats::new(); + let shared_state = ProxySharedState::new(); + let tracker = UserIpTracker::new(); + let config = ProxyConfig::default(); + let output = render_metrics( + &stats, + &shared_state, + &config, + &tracker, + None, + &TlsFullCertBudget::new(), + &test_web_publication(), + ) + .await; + assert!(output.contains("# TYPE telemt_uptime_seconds gauge")); + assert!(output.contains("# TYPE telemt_connections_total counter")); + assert!(output.contains("# TYPE telemt_connections_bad_total counter")); + assert!(output.contains("# TYPE telemt_connections_bad_by_class_total counter")); + assert!(output.contains("# TYPE telemt_handshake_timeouts_total counter")); + assert!(output.contains("# TYPE telemt_handshake_failures_by_class_total counter")); + assert!(output.contains("# TYPE telemt_auth_expensive_checks_total counter")); + assert!(output.contains("# TYPE telemt_auth_budget_exhausted_total counter")); + assert!(output.contains("# TYPE telemt_upstream_connect_attempt_total counter")); + assert!(output.contains("# TYPE telemt_me_rpc_proxy_req_signal_sent_total counter")); + assert!(output.contains("# TYPE telemt_me_idle_close_by_peer_total counter")); + assert!(output.contains("# TYPE telemt_relay_idle_soft_mark_total counter")); + assert!(output.contains("# TYPE telemt_relay_idle_hard_close_total counter")); + assert!(output.contains("# TYPE telemt_relay_pressure_evict_total counter")); + assert!(output.contains("# TYPE telemt_relay_protocol_desync_close_total counter")); + assert!(output.contains("# TYPE telemt_me_d2c_batches_total counter")); + assert!(output.contains("# TYPE telemt_me_d2c_flush_reason_total counter")); + assert!(output.contains("# TYPE telemt_me_d2c_write_mode_total counter")); + assert!(output.contains("# TYPE telemt_me_d2c_batch_frames_bucket_total counter")); + assert!(output.contains("# TYPE telemt_me_d2c_flush_duration_us_bucket_total counter")); + assert!(output.contains("# TYPE telemt_me_endpoint_quarantine_total counter")); + assert!(output.contains("# TYPE telemt_me_endpoint_quarantine_unexpected_total counter")); + assert!( + output.contains("# TYPE telemt_me_endpoint_quarantine_draining_suppressed_total counter") + ); + assert!(output.contains("# TYPE telemt_me_writer_removed_total counter")); + assert!( + output.contains("# TYPE telemt_me_writer_removed_unexpected_minus_restored_total gauge") + ); + assert!(output.contains("# TYPE telemt_user_unique_ips_current gauge")); + assert!(output.contains("# TYPE telemt_user_unique_ips_recent_window gauge")); + assert!(output.contains("# TYPE telemt_user_unique_ips_limit gauge")); + assert!(output.contains("# TYPE telemt_user_unique_ips_utilization gauge")); + assert!(output.contains("# TYPE telemt_stats_user_entries gauge")); + assert!(output.contains("# TYPE telemt_telemetry_user_series_users gauge")); + assert!(output.contains("# TYPE telemt_ip_tracker_users gauge")); + assert!(output.contains("# TYPE telemt_ip_tracker_entries gauge")); + assert!(output.contains("# TYPE telemt_ip_tracker_cleanup_queue_len gauge")); + assert!(output.contains("# TYPE telemt_ip_tracker_cleanup_total counter")); + assert!(output.contains("# TYPE telemt_ip_tracker_cap_rejects_total counter")); + assert!(output.contains("# TYPE telemt_tls_fetch_profile_cache_entries gauge")); + assert!(output.contains("# TYPE telemt_tls_fetch_profile_cache_cap_drops_total counter")); + assert!(output.contains("# TYPE telemt_tls_front_full_cert_budget_entries gauge")); + assert!(output.contains("# TYPE telemt_tls_front_full_cert_budget_cap_drops_total counter")); + assert!(output.contains("# TYPE telemt_tls_front_profile_domains gauge")); + assert!(output.contains("# TYPE telemt_tls_front_profile_info gauge")); + assert!(output.contains("# TYPE telemt_tls_front_profile_quality_info gauge")); + assert!(output.contains("# TYPE telemt_tls_front_profile_age_seconds gauge")); + assert!(output.contains("# TYPE telemt_tls_front_profile_server_hello_bytes gauge")); + assert!(output.contains("# TYPE telemt_tls_front_profile_server_hello_extensions gauge")); + assert!(output.contains("# TYPE telemt_tls_front_profile_app_data_records gauge")); + assert!(output.contains("# TYPE telemt_tls_front_profile_ticket_records gauge")); + assert!(output.contains("# TYPE telemt_tls_front_profile_change_cipher_spec_records gauge")); + assert!(output.contains("# TYPE telemt_tls_front_profile_app_data_bytes gauge")); +} + +#[tokio::test] +async fn test_endpoint_integration() { + let mut config = ProxyConfig::default(); + config.general.beobachten = true; + config.general.beobachten_minutes = 10; + let runtime = crate::maestro::generation::test_runtime_generation(1, config); + let web_publication = test_web_publication(); + let tls_full_cert_budget = TlsFullCertBudget::new(); + runtime.stats.increment_connects_all(); + runtime.stats.increment_connects_all(); + runtime.stats.increment_connects_all(); + + let req = Request::builder().uri("/metrics").body(()).unwrap(); + let resp = handle(req, &runtime, &web_publication, &tls_full_cert_budget) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + let body = resp.into_body().collect().await.unwrap().to_bytes(); + assert!( + std::str::from_utf8(body.as_ref()) + .unwrap() + .contains("telemt_connections_total 3") + ); + assert!( + std::str::from_utf8(body.as_ref()) + .unwrap() + .contains(&format!( + "telemt_build_info{{version=\"{}\"}} 1", + env!("CARGO_PKG_VERSION") + )) + ); + + runtime.beobachten.record( + "TLS-scanner", + "203.0.113.10".parse::().unwrap(), + Duration::from_secs(600), + ); + let req_beob = Request::builder().uri("/beobachten").body(()).unwrap(); + let resp_beob = handle(req_beob, &runtime, &web_publication, &tls_full_cert_budget) + .await + .unwrap(); + assert_eq!(resp_beob.status(), StatusCode::OK); + let body_beob = resp_beob.into_body().collect().await.unwrap().to_bytes(); + let beob_text = std::str::from_utf8(body_beob.as_ref()).unwrap(); + assert!(beob_text.contains("[TLS-scanner]")); + assert!(beob_text.contains("203.0.113.10-1")); + + let req404 = Request::builder().uri("/other").body(()).unwrap(); + let resp404 = handle(req404, &runtime, &web_publication, &tls_full_cert_budget) + .await + .unwrap(); + assert_eq!(resp404.status(), StatusCode::NOT_FOUND); +} diff --git a/src/network/probe.rs b/src/network/probe.rs index bf18029..f60c50d 100644 --- a/src/network/probe.rs +++ b/src/network/probe.rs @@ -70,9 +70,7 @@ pub async fn run_probe( ) -> Result { let mut probe = NetworkProbe::default(); let dns_resolver = Arc::new( - crate::network::dns_overrides::GenerationDnsResolver::from_entries( - &config.dns_overrides, - )?, + crate::network::dns_overrides::GenerationDnsResolver::from_entries(&config.dns_overrides)?, ); let servers = collect_stun_servers(config); let mut detected_ipv4 = detect_local_ip_v4(); @@ -426,185 +424,13 @@ pub fn decide_network_capabilities( } } +// Local interface discovery and bogon classification. +mod local; +pub use local::{ + detect_interface_ipv4, detect_interface_ipv6, is_bogon, is_bogon_v4, is_bogon_v6, + log_probe_result, +}; +use local::{detect_local_ip_v4, detect_local_ip_v6}; + #[cfg(test)] -mod tests { - use super::*; - use crate::config::NetworkConfig; - - #[test] - fn manual_nat_ip_enables_ipv4_me_without_reflection() { - let config = NetworkConfig { - ipv4: true, - ..Default::default() - }; - let probe = NetworkProbe { - detected_ipv4: Some(Ipv4Addr::new(10, 0, 0, 10)), - ipv4_is_bogon: true, - ..Default::default() - }; - - let decision = decide_network_capabilities( - &config, - &probe, - Some(IpAddr::V4(Ipv4Addr::new(1, 2, 3, 4))), - ); - - assert!(decision.ipv4_me); - } - - #[test] - fn manual_nat_ip_does_not_enable_other_family() { - let config = NetworkConfig { - ipv4: true, - ipv6: Some(true), - ..Default::default() - }; - let probe = NetworkProbe { - detected_ipv4: Some(Ipv4Addr::new(10, 0, 0, 10)), - detected_ipv6: Some(Ipv6Addr::LOCALHOST), - ipv4_is_bogon: true, - ipv6_is_bogon: true, - ..Default::default() - }; - - let decision = decide_network_capabilities( - &config, - &probe, - Some(IpAddr::V4(Ipv4Addr::new(1, 2, 3, 4))), - ); - - assert!(decision.ipv4_me); - assert!(!decision.ipv6_me); - } -} - -fn detect_local_ip_v4() -> Option { - let socket = UdpSocket::bind("0.0.0.0:0").ok()?; - socket.connect("8.8.8.8:80").ok()?; - match socket.local_addr().ok()?.ip() { - IpAddr::V4(v4) => Some(v4), - _ => None, - } -} - -fn detect_local_ip_v6() -> Option { - let socket = UdpSocket::bind("[::]:0").ok()?; - socket.connect("[2001:4860:4860::8888]:80").ok()?; - match socket.local_addr().ok()?.ip() { - IpAddr::V6(v6) => Some(v6), - _ => None, - } -} - -pub fn detect_interface_ipv4() -> Option { - detect_local_ip_v4() -} - -pub fn detect_interface_ipv6() -> Option { - detect_local_ip_v6() -} - -pub fn is_bogon(ip: IpAddr) -> bool { - match ip { - IpAddr::V4(v4) => is_bogon_v4(v4), - IpAddr::V6(v6) => is_bogon_v6(v6), - } -} - -pub fn is_bogon_v4(ip: Ipv4Addr) -> bool { - let octets = ip.octets(); - if ip.is_private() || ip.is_loopback() || ip.is_link_local() { - return true; - } - if octets[0] == 0 { - return true; - } - if octets[0] == 100 && (octets[1] & 0xC0) == 64 { - return true; - } - if octets[0] == 192 && octets[1] == 0 && octets[2] == 0 { - return true; - } - if octets[0] == 192 && octets[1] == 0 && octets[2] == 2 { - return true; - } - if octets[0] == 198 && (octets[1] & 0xFE) == 18 { - return true; - } - if octets[0] == 198 && octets[1] == 51 && octets[2] == 100 { - return true; - } - if octets[0] == 203 && octets[1] == 0 && octets[2] == 113 { - return true; - } - if ip.is_multicast() { - return true; - } - if octets[0] >= 240 { - return true; - } - if ip.is_broadcast() { - return true; - } - false -} - -pub fn is_bogon_v6(ip: Ipv6Addr) -> bool { - if ip.is_unspecified() || ip.is_loopback() || ip.is_unique_local() { - return true; - } - let segs = ip.segments(); - if (segs[0] & 0xFFC0) == 0xFE80 { - return true; - } - if segs[0..5] == [0, 0, 0, 0, 0] && segs[5] == 0xFFFF { - return true; - } - if segs[0] == 0x0100 && segs[1..4] == [0, 0, 0] { - return true; - } - if segs[0] == 0x2001 && segs[1] == 0x0db8 { - return true; - } - if segs[0] == 0x2002 { - return true; - } - if ip.is_multicast() { - return true; - } - false -} - -pub fn log_probe_result(probe: &NetworkProbe, decision: &NetworkDecision) { - info!( - ipv4 = probe - .detected_ipv4 - .as_ref() - .map(|v| v.to_string()) - .unwrap_or_else(|| "-".into()), - ipv6 = probe - .detected_ipv6 - .as_ref() - .map(|v| v.to_string()) - .unwrap_or_else(|| "-".into()), - reflected_v4 = probe - .reflected_ipv4 - .as_ref() - .map(|v| v.ip().to_string()) - .unwrap_or_else(|| "-".into()), - reflected_v6 = probe - .reflected_ipv6 - .as_ref() - .map(|v| v.ip().to_string()) - .unwrap_or_else(|| "-".into()), - ipv4_bogon = probe.ipv4_is_bogon, - ipv6_bogon = probe.ipv6_is_bogon, - ipv4_me = decision.ipv4_me, - ipv6_me = decision.ipv6_me, - ipv4_dc = decision.ipv4_dc, - ipv6_dc = decision.ipv6_dc, - prefer = decision.effective_prefer, - multipath = decision.effective_multipath, - "Network capabilities resolved" - ); -} +mod tests; diff --git a/src/network/probe/local.rs b/src/network/probe/local.rs new file mode 100644 index 0000000..e3d8928 --- /dev/null +++ b/src/network/probe/local.rs @@ -0,0 +1,132 @@ +use super::*; + +pub(super) fn detect_local_ip_v4() -> Option { + let socket = UdpSocket::bind("0.0.0.0:0").ok()?; + socket.connect("8.8.8.8:80").ok()?; + match socket.local_addr().ok()?.ip() { + IpAddr::V4(v4) => Some(v4), + _ => None, + } +} + +pub(super) fn detect_local_ip_v6() -> Option { + let socket = UdpSocket::bind("[::]:0").ok()?; + socket.connect("[2001:4860:4860::8888]:80").ok()?; + match socket.local_addr().ok()?.ip() { + IpAddr::V6(v6) => Some(v6), + _ => None, + } +} + +pub fn detect_interface_ipv4() -> Option { + detect_local_ip_v4() +} + +pub fn detect_interface_ipv6() -> Option { + detect_local_ip_v6() +} + +pub fn is_bogon(ip: IpAddr) -> bool { + match ip { + IpAddr::V4(v4) => is_bogon_v4(v4), + IpAddr::V6(v6) => is_bogon_v6(v6), + } +} + +pub fn is_bogon_v4(ip: Ipv4Addr) -> bool { + let octets = ip.octets(); + if ip.is_private() || ip.is_loopback() || ip.is_link_local() { + return true; + } + if octets[0] == 0 { + return true; + } + if octets[0] == 100 && (octets[1] & 0xC0) == 64 { + return true; + } + if octets[0] == 192 && octets[1] == 0 && octets[2] == 0 { + return true; + } + if octets[0] == 192 && octets[1] == 0 && octets[2] == 2 { + return true; + } + if octets[0] == 198 && (octets[1] & 0xFE) == 18 { + return true; + } + if octets[0] == 198 && octets[1] == 51 && octets[2] == 100 { + return true; + } + if octets[0] == 203 && octets[1] == 0 && octets[2] == 113 { + return true; + } + if ip.is_multicast() { + return true; + } + if octets[0] >= 240 { + return true; + } + if ip.is_broadcast() { + return true; + } + false +} + +pub fn is_bogon_v6(ip: Ipv6Addr) -> bool { + if ip.is_unspecified() || ip.is_loopback() || ip.is_unique_local() { + return true; + } + let segs = ip.segments(); + if (segs[0] & 0xFFC0) == 0xFE80 { + return true; + } + if segs[0..5] == [0, 0, 0, 0, 0] && segs[5] == 0xFFFF { + return true; + } + if segs[0] == 0x0100 && segs[1..4] == [0, 0, 0] { + return true; + } + if segs[0] == 0x2001 && segs[1] == 0x0db8 { + return true; + } + if segs[0] == 0x2002 { + return true; + } + if ip.is_multicast() { + return true; + } + false +} + +pub fn log_probe_result(probe: &NetworkProbe, decision: &NetworkDecision) { + info!( + ipv4 = probe + .detected_ipv4 + .as_ref() + .map(|v| v.to_string()) + .unwrap_or_else(|| "-".into()), + ipv6 = probe + .detected_ipv6 + .as_ref() + .map(|v| v.to_string()) + .unwrap_or_else(|| "-".into()), + reflected_v4 = probe + .reflected_ipv4 + .as_ref() + .map(|v| v.ip().to_string()) + .unwrap_or_else(|| "-".into()), + reflected_v6 = probe + .reflected_ipv6 + .as_ref() + .map(|v| v.ip().to_string()) + .unwrap_or_else(|| "-".into()), + ipv4_bogon = probe.ipv4_is_bogon, + ipv6_bogon = probe.ipv6_is_bogon, + ipv4_me = decision.ipv4_me, + ipv6_me = decision.ipv6_me, + ipv4_dc = decision.ipv4_dc, + ipv6_dc = decision.ipv6_dc, + prefer = decision.effective_prefer, + multipath = decision.effective_multipath, + "Network capabilities resolved" + ); +} diff --git a/src/network/probe/tests.rs b/src/network/probe/tests.rs new file mode 100644 index 0000000..890d702 --- /dev/null +++ b/src/network/probe/tests.rs @@ -0,0 +1,42 @@ +use super::*; +use crate::config::NetworkConfig; + +#[test] +fn manual_nat_ip_enables_ipv4_me_without_reflection() { + let config = NetworkConfig { + ipv4: true, + ..Default::default() + }; + let probe = NetworkProbe { + detected_ipv4: Some(Ipv4Addr::new(10, 0, 0, 10)), + ipv4_is_bogon: true, + ..Default::default() + }; + + let decision = + decide_network_capabilities(&config, &probe, Some(IpAddr::V4(Ipv4Addr::new(1, 2, 3, 4)))); + + assert!(decision.ipv4_me); +} + +#[test] +fn manual_nat_ip_does_not_enable_other_family() { + let config = NetworkConfig { + ipv4: true, + ipv6: Some(true), + ..Default::default() + }; + let probe = NetworkProbe { + detected_ipv4: Some(Ipv4Addr::new(10, 0, 0, 10)), + detected_ipv6: Some(Ipv6Addr::LOCALHOST), + ipv4_is_bogon: true, + ipv6_is_bogon: true, + ..Default::default() + }; + + let decision = + decide_network_capabilities(&config, &probe, Some(IpAddr::V4(Ipv4Addr::new(1, 2, 3, 4)))); + + assert!(decision.ipv4_me); + assert!(!decision.ipv6_me); +} diff --git a/src/network/stun.rs b/src/network/stun.rs index bda8eae..79de459 100644 --- a/src/network/stun.rs +++ b/src/network/stun.rs @@ -394,8 +394,8 @@ async fn resolve_stun_addr( } if let Some((host, port)) = split_host_port(stun_addr) - && let Some(addr) = dns_resolver - .and_then(|resolver| resolver.resolve_socket_addr(&host, port)) + && let Some(addr) = + dns_resolver.and_then(|resolver| resolver.resolve_socket_addr(&host, port)) { return Ok(match (addr.is_ipv4(), family) { (true, IpFamily::V4) | (false, IpFamily::V6) => Some(addr), diff --git a/src/proxy/client.rs b/src/proxy/client.rs index 2e30370..bfe6469 100644 --- a/src/proxy/client.rs +++ b/src/proxy/client.rs @@ -55,861 +55,25 @@ use crate::proxy::route_mode::RelayRouteMode; use crate::proxy::route_mode::RouteRuntimeController; use crate::proxy::shared_state::{ConntrackClosePolicy, ProxySharedState}; -fn beobachten_ttl(config: &ProxyConfig) -> Duration { - const BEOBACHTEN_TTL_MAX_MINUTES: u64 = 24 * 60; - let minutes = config.general.beobachten_minutes; - if minutes == 0 { - static BEOBACHTEN_ZERO_MINUTES_WARNED: OnceLock = OnceLock::new(); - let warned = BEOBACHTEN_ZERO_MINUTES_WARNED.get_or_init(|| AtomicBool::new(false)); - if !warned.swap(true, Ordering::Relaxed) { - warn!( - "general.beobachten_minutes=0 is insecure because entries expire immediately; forcing minimum TTL to 1 minute" - ); - } - return Duration::from_secs(60); - } +// Handshake classification, telemetry, and masking helpers. +mod handshake_support; +// Stream-level handshake timeout and relay dispatch entry points. +mod stream_entry; +// Running client socket setup and first-stage handshake. +mod running_lifecycle; +// TLS client handshake and authenticated dispatch. +mod tls_client; +// Direct MTProxy handshake and authenticated dispatch. +mod direct_client; +// Shared authenticated relay and user-admission helpers. +mod authenticated; - if minutes > BEOBACHTEN_TTL_MAX_MINUTES { - static BEOBACHTEN_OVERSIZED_MINUTES_WARNED: OnceLock = OnceLock::new(); - let warned = BEOBACHTEN_OVERSIZED_MINUTES_WARNED.get_or_init(|| AtomicBool::new(false)); - if !warned.swap(true, Ordering::Relaxed) { - warn!( - configured_minutes = minutes, - max_minutes = BEOBACHTEN_TTL_MAX_MINUTES, - "general.beobachten_minutes is too large; clamping to secure maximum" - ); - } - } - - Duration::from_secs(minutes.min(BEOBACHTEN_TTL_MAX_MINUTES).saturating_mul(60)) -} - -fn wrap_tls_application_record(payload: &[u8]) -> Vec { - let chunks = payload.len().div_ceil(u16::MAX as usize).max(1); - let mut record = Vec::with_capacity(payload.len() + 5 * chunks); - - if payload.is_empty() { - record.push(TLS_RECORD_APPLICATION); - record.extend_from_slice(&TLS_VERSION); - record.extend_from_slice(&0u16.to_be_bytes()); - return record; - } - - for chunk in payload.chunks(u16::MAX as usize) { - record.push(TLS_RECORD_APPLICATION); - record.extend_from_slice(&TLS_VERSION); - record.extend_from_slice(&(chunk.len() as u16).to_be_bytes()); - record.extend_from_slice(chunk); - } - - record -} - -fn tls_clienthello_len_in_bounds(tls_len: usize) -> bool { - (MIN_TLS_CLIENT_HELLO_SIZE..=MAX_TLS_PLAINTEXT_SIZE).contains(&tls_len) -} - -async fn read_with_progress( - reader: &mut R, - mut buf: &mut [u8], -) -> std::io::Result { - let mut total = 0usize; - while !buf.is_empty() { - match reader.read(buf).await { - Ok(0) => return Ok(total), - Ok(n) => { - total += n; - let (_, rest) = buf.split_at_mut(n); - buf = rest; - } - Err(e) => return Err(e), - } - } - Ok(total) -} - -async fn maybe_apply_mask_reject_delay(config: &ProxyConfig) { - let min = config.censorship.server_hello_delay_min_ms; - let max = config.censorship.server_hello_delay_max_ms; - if max == 0 { - return; - } - - let delay_ms = if min >= max { - max - } else { - rand::rng().random_range(min..=max) - }; - - if delay_ms > 0 { - tokio::time::sleep(Duration::from_millis(delay_ms)).await; - } -} - -fn handshake_timeout_with_mask_grace(config: &ProxyConfig) -> Duration { - let base = Duration::from_secs(config.timeouts.client_handshake); - if config.censorship.mask { - base.saturating_add(Duration::from_millis(750)) - } else { - base - } -} - -fn effective_client_first_byte_idle_secs(config: &ProxyConfig, shared: &ProxySharedState) -> u64 { - let idle_secs = config.timeouts.client_first_byte_idle_secs; - if idle_secs == 0 { - return 0; - } - if shared.conntrack_pressure_active() { - idle_secs.min( - config - .server - .conntrack_control - .profile - .client_first_byte_idle_cap_secs(), - ) - } else { - idle_secs - } -} - -const MASK_CLASSIFIER_PREFETCH_WINDOW: usize = 16; +use handshake_support::*; #[cfg(test)] -const MASK_CLASSIFIER_PREFETCH_TIMEOUT: Duration = Duration::from_millis(5); - -fn mask_classifier_prefetch_timeout(config: &ProxyConfig) -> Duration { - Duration::from_millis(config.censorship.mask_classifier_prefetch_timeout_ms) -} - -fn should_prefetch_mask_classifier_window(initial_data: &[u8]) -> bool { - if initial_data.len() >= MASK_CLASSIFIER_PREFETCH_WINDOW { - return false; - } - - if initial_data.is_empty() { - // Empty initial_data means there is no client probe prefix to refine. - // Prefetching in this case can consume fallback relay payload bytes and - // accidentally route them through shaping heuristics. - return false; - } - - if initial_data[0] == 0x16 || initial_data.starts_with(b"SSH-") { - return false; - } - - initial_data - .iter() - .all(|b| b.is_ascii_alphabetic() || *b == b' ') -} - -#[cfg(test)] -async fn extend_masking_initial_window(reader: &mut R, initial_data: &mut Vec) -where - R: AsyncRead + Unpin, -{ - extend_masking_initial_window_with_timeout( - reader, - initial_data, - MASK_CLASSIFIER_PREFETCH_TIMEOUT, - ) - .await; -} - -async fn extend_masking_initial_window_with_timeout( - reader: &mut R, - initial_data: &mut Vec, - prefetch_timeout: Duration, -) where - R: AsyncRead + Unpin, -{ - if !should_prefetch_mask_classifier_window(initial_data) { - return; - } - - let need = MASK_CLASSIFIER_PREFETCH_WINDOW.saturating_sub(initial_data.len()); - if need == 0 { - return; - } - - let mut extra = [0u8; MASK_CLASSIFIER_PREFETCH_WINDOW]; - if let Ok(Ok(n)) = timeout(prefetch_timeout, reader.read(&mut extra[..need])).await - && n > 0 - { - initial_data.extend_from_slice(&extra[..n]); - } -} - -fn masking_outcome( - reader: R, - writer: W, - initial_data: Vec, - peer: SocketAddr, - local_addr: SocketAddr, - config: Arc, - upstream_manager: Arc, - beobachten: Arc, - shared: Arc, -) -> HandshakeOutcome -where - R: AsyncRead + Unpin + Send + 'static, - W: AsyncWrite + Unpin + Send + 'static, -{ - HandshakeOutcome::NeedsMasking(Box::pin(async move { - let mut reader = reader; - let mut initial_data = initial_data; - extend_masking_initial_window_with_timeout( - &mut reader, - &mut initial_data, - mask_classifier_prefetch_timeout(&config), - ) - .await; - - crate::proxy::masking::handle_bad_client_with_shared_resolver( - reader, - writer, - &initial_data, - peer, - local_addr, - &config, - &beobachten, - shared.as_ref(), - Some(upstream_manager.as_ref()), - ) - .await; - Ok(()) - })) -} - -fn record_beobachten_class( - beobachten: &BeobachtenStore, - config: &ProxyConfig, - peer_ip: IpAddr, - class: &str, -) { - if !config.general.beobachten { - return; - } - beobachten.record(class, peer_ip, beobachten_ttl(config)); -} - -fn tls_fingerprint_collection_enabled(config: &ProxyConfig) -> bool { - config.general.beobachten || config.server.api.runtime_edge_enabled -} - -fn observe_tls_client_fingerprint( - stats: &Stats, - config: &ProxyConfig, - peer_ip: IpAddr, - handshake: &[u8], -) -> Option { - if !tls_fingerprint_collection_enabled(config) { - return None; - } - - match tls_fingerprint::fingerprint_client_hello(handshake) { - Some(fingerprint) => { - stats.record_tls_fingerprint_observed(&fingerprint, peer_ip, beobachten_ttl(config)); - Some(fingerprint) - } - None => { - stats.increment_tls_fingerprint_parse_error(); - None - } - } -} - -fn record_tls_fingerprint_auth_success( - stats: &Stats, - config: &ProxyConfig, - peer_ip: IpAddr, - fingerprint: Option<&TlsClientFingerprint>, - user: &str, -) { - if let Some(fingerprint) = fingerprint { - stats.record_tls_fingerprint_auth_success( - fingerprint, - peer_ip, - user, - beobachten_ttl(config), - ); - } -} - -fn record_tls_fingerprint_bad_or_probe( - stats: &Stats, - config: &ProxyConfig, - peer_ip: IpAddr, - fingerprint: Option<&TlsClientFingerprint>, -) { - if let Some(fingerprint) = fingerprint { - stats.record_tls_fingerprint_bad_or_probe(fingerprint, peer_ip, beobachten_ttl(config)); - } -} - -fn classify_expected_64_got_0(kind: std::io::ErrorKind) -> Option<&'static str> { - match kind { - std::io::ErrorKind::UnexpectedEof => Some("expected_64_got_0_unexpected_eof"), - std::io::ErrorKind::ConnectionReset => Some("expected_64_got_0_connection_reset"), - std::io::ErrorKind::ConnectionAborted => Some("expected_64_got_0_connection_aborted"), - std::io::ErrorKind::BrokenPipe => Some("expected_64_got_0_broken_pipe"), - std::io::ErrorKind::NotConnected => Some("expected_64_got_0_not_connected"), - _ => None, - } -} - -fn classify_handshake_failure_class(error: &ProxyError) -> &'static str { - match error { - ProxyError::Io(err) => classify_expected_64_got_0(err.kind()).unwrap_or("other"), - ProxyError::Stream(StreamError::UnexpectedEof) => "expected_64_got_0_unexpected_eof", - ProxyError::Stream(StreamError::Io(err)) => { - classify_expected_64_got_0(err.kind()).unwrap_or("other") - } - _ => "other", - } -} - -fn record_handshake_failure_class( - beobachten: &BeobachtenStore, - config: &ProxyConfig, - peer_ip: IpAddr, - error: &ProxyError, -) { - // Keep beobachten buckets stable while detailed per-kind classification - // is tracked in API counters. - let class = match classify_handshake_failure_class(error) { - value if value.starts_with("expected_64_got_0_") => "expected_64_got_0", - _ => "other", - }; - record_beobachten_class(beobachten, config, peer_ip, class); -} - -#[inline] -fn increment_bad_on_unknown_tls_sni(stats: &Stats, error: &ProxyError) { - if matches!(error, ProxyError::UnknownTlsSni) { - stats.increment_connects_bad_with_class("unknown_tls_sni"); - } -} - -fn is_trusted_proxy_source(peer_ip: IpAddr, trusted: &[IpNetwork]) -> bool { - if trusted.is_empty() { - static EMPTY_PROXY_TRUST_WARNED: OnceLock = OnceLock::new(); - let warned = EMPTY_PROXY_TRUST_WARNED.get_or_init(|| AtomicBool::new(false)); - if !warned.swap(true, Ordering::Relaxed) { - warn!( - "PROXY protocol enabled but server.proxy_protocol_trusted_cidrs is empty; rejecting all PROXY headers" - ); - } - return false; - } - trusted.iter().any(|cidr| cidr.contains(peer_ip)) -} - -fn synthetic_local_addr(port: u16) -> SocketAddr { - SocketAddr::from(([0, 0, 0, 0], port)) -} - -#[cfg(test)] -pub async fn handle_client_stream( - stream: S, - peer: SocketAddr, - config: Arc, - stats: Arc, - upstream_manager: Arc, - replay_checker: Arc, - buffer_pool: Arc, - rng: Arc, - me_pool: Option>, - route_runtime: Arc, - tls_cache: Option>, - ip_tracker: Arc, - beobachten: Arc, - proxy_protocol_enabled: bool, -) -> Result<()> -where - S: AsyncRead + AsyncWrite + Unpin + Send + 'static, -{ - handle_client_stream_with_shared( - stream, - peer, - config, - stats, - upstream_manager, - replay_checker, - buffer_pool, - rng, - me_pool, - route_runtime, - tls_cache, - ip_tracker, - beobachten, - ProxySharedState::new(), - proxy_protocol_enabled, - ) - .await -} - -#[allow(clippy::too_many_arguments)] -#[allow(dead_code)] -pub async fn handle_client_stream_with_shared( - stream: S, - peer: SocketAddr, - config: Arc, - stats: Arc, - upstream_manager: Arc, - replay_checker: Arc, - buffer_pool: Arc, - rng: Arc, - me_pool: Option>, - route_runtime: Arc, - tls_cache: Option>, - ip_tracker: Arc, - beobachten: Arc, - shared: Arc, - proxy_protocol_enabled: bool, -) -> Result<()> -where - S: AsyncRead + AsyncWrite + Unpin + Send + 'static, -{ - handle_client_stream_with_shared_and_pool_runtime( - stream, - peer, - config, - stats, - upstream_manager, - replay_checker, - buffer_pool, - rng, - me_pool, - None, - route_runtime, - tls_cache, - ip_tracker, - beobachten, - shared, - proxy_protocol_enabled, - ) - .await -} - -#[allow(clippy::too_many_arguments)] -pub async fn handle_client_stream_with_shared_and_pool_runtime( - mut stream: S, - peer: SocketAddr, - config: Arc, - stats: Arc, - upstream_manager: Arc, - replay_checker: Arc, - buffer_pool: Arc, - rng: Arc, - me_pool: Option>, - me_pool_runtime: Option>>>>, - route_runtime: Arc, - tls_cache: Option>, - ip_tracker: Arc, - beobachten: Arc, - shared: Arc, - proxy_protocol_enabled: bool, -) -> Result<()> -where - S: AsyncRead + AsyncWrite + Unpin + Send + 'static, -{ - stats.increment_connects_all(); - let mut real_peer = normalize_ip(peer); - - // For non-TCP streams, use a synthetic local address; may be overridden by PROXY protocol dst - let mut local_addr = synthetic_local_addr(config.server.port); - - if proxy_protocol_enabled { - if !is_trusted_proxy_source(peer.ip(), &config.server.proxy_protocol_trusted_cidrs) { - stats.increment_connects_bad_with_class("proxy_protocol_untrusted"); - warn!( - peer = %peer, - trusted = ?config.server.proxy_protocol_trusted_cidrs, - "Rejecting PROXY protocol header from untrusted source" - ); - record_beobachten_class(&beobachten, &config, peer.ip(), "other"); - return Err(ProxyError::InvalidProxyProtocol); - } - - let proxy_header_timeout = - Duration::from_millis(config.server.proxy_protocol_header_timeout_ms.max(1)); - match timeout( - proxy_header_timeout, - parse_proxy_protocol(&mut stream, peer), - ) - .await - { - Ok(Ok(info)) => { - debug!( - peer = %peer, - client = %info.src_addr, - version = info.version, - "PROXY protocol header parsed" - ); - real_peer = normalize_ip(info.src_addr); - if let Some(dst) = info.dst_addr { - local_addr = dst; - } - } - Ok(Err(e)) => { - stats.increment_connects_bad_with_class("proxy_protocol_invalid_header"); - warn!(peer = %peer, error = %e, "Invalid PROXY protocol header"); - record_beobachten_class(&beobachten, &config, peer.ip(), "other"); - return Err(e); - } - Err(_) => { - stats.increment_connects_bad_with_class("proxy_protocol_header_timeout"); - warn!(peer = %peer, timeout_ms = proxy_header_timeout.as_millis(), "PROXY protocol header timeout"); - record_beobachten_class(&beobachten, &config, peer.ip(), "other"); - return Err(ProxyError::InvalidProxyProtocol); - } - } - } - - debug!(peer = %real_peer, "New connection (generic stream)"); - - let first_byte_idle_secs = effective_client_first_byte_idle_secs(&config, shared.as_ref()); - let first_byte = if first_byte_idle_secs == 0 { - None - } else { - let idle_timeout = Duration::from_secs(first_byte_idle_secs); - let mut first_byte = [0u8; 1]; - match timeout(idle_timeout, stream.read(&mut first_byte)).await { - Ok(Ok(0)) => { - debug!(peer = %real_peer, "Connection closed before first client byte"); - return Ok(()); - } - Ok(Ok(_)) => Some(first_byte[0]), - Ok(Err(e)) - if matches!( - e.kind(), - std::io::ErrorKind::UnexpectedEof - | std::io::ErrorKind::ConnectionReset - | std::io::ErrorKind::ConnectionAborted - | std::io::ErrorKind::BrokenPipe - | std::io::ErrorKind::NotConnected - ) => - { - debug!( - peer = %real_peer, - error = %e, - "Connection closed before first client byte" - ); - return Ok(()); - } - Ok(Err(e)) => { - debug!( - peer = %real_peer, - error = %e, - "Failed while waiting for first client byte" - ); - return Err(ProxyError::Io(e)); - } - Err(_) => { - debug!( - peer = %real_peer, - idle_secs = first_byte_idle_secs, - "Closing idle pooled connection before first client byte" - ); - return Ok(()); - } - } - }; - - let handshake_timeout = handshake_timeout_with_mask_grace(&config); - let stats_for_timeout = stats.clone(); - let config_for_timeout = config.clone(); - let beobachten_for_timeout = beobachten.clone(); - let peer_for_timeout = real_peer.ip(); - - // Phase 2: active handshake (with timeout after the first client byte) - let outcome = match timeout(handshake_timeout, async { - let mut first_bytes = [0u8; 5]; - if let Some(first_byte) = first_byte { - first_bytes[0] = first_byte; - stream.read_exact(&mut first_bytes[1..]).await?; - } else { - stream.read_exact(&mut first_bytes).await?; - } - - let is_tls = tls::is_tls_handshake(&first_bytes[..3]); - debug!(peer = %real_peer, is_tls = is_tls, "Handshake type detected"); - - if is_tls { - let tls_len = u16::from_be_bytes([first_bytes[3], first_bytes[4]]) as usize; - - // RFC 8446 §5.1: TLS record payload MUST NOT exceed 2^14 (16_384) bytes. - // Lower bound is a structural minimum for a valid TLS 1.3 ClientHello - // (record header + handshake header + random + session_id + cipher_suites - // + compression + at least one extension with SNI). The previous value of - // 512 was implicitly coupled to TLS_REQUEST_LENGTH=517 from the official - // Telegram MTProxy reference server, leaving only a 5-byte margin and - // incorrectly rejecting compact but spec-compliant ClientHellos from - // third-party clients or future Telegram versions. - if !tls_clienthello_len_in_bounds(tls_len) { - debug!(peer = %real_peer, tls_len = tls_len, max_tls_len = MAX_TLS_PLAINTEXT_SIZE, "TLS handshake length out of bounds"); - stats.increment_connects_bad_with_class("tls_clienthello_len_out_of_bounds"); - maybe_apply_mask_reject_delay(&config).await; - let (reader, writer) = tokio::io::split(stream); - return Ok(masking_outcome( - reader, - writer, - first_bytes.to_vec(), - real_peer, - local_addr, - config.clone(), - upstream_manager.clone(), - beobachten.clone(), - shared.clone(), - )); - } - - let mut handshake = vec![0u8; 5 + tls_len]; - handshake[..5].copy_from_slice(&first_bytes); - let body_read = match read_with_progress(&mut stream, &mut handshake[5..]).await { - Ok(n) => n, - Err(e) => { - debug!(peer = %real_peer, error = %e, tls_len = tls_len, "TLS ClientHello body read failed; engaging masking fallback"); - stats.increment_connects_bad_with_class("tls_clienthello_read_error"); - maybe_apply_mask_reject_delay(&config).await; - let initial_len = 5; - let (reader, writer) = tokio::io::split(stream); - return Ok(masking_outcome( - reader, - writer, - handshake[..initial_len].to_vec(), - real_peer, - local_addr, - config.clone(), - upstream_manager.clone(), - beobachten.clone(), - shared.clone(), - )); - } - }; - - if body_read < tls_len { - debug!(peer = %real_peer, got = body_read, expected = tls_len, "Truncated in-range TLS ClientHello; engaging masking fallback"); - stats.increment_connects_bad_with_class("tls_clienthello_truncated"); - maybe_apply_mask_reject_delay(&config).await; - let initial_len = 5 + body_read; - let (reader, writer) = tokio::io::split(stream); - return Ok(masking_outcome( - reader, - writer, - handshake[..initial_len].to_vec(), - real_peer, - local_addr, - config.clone(), - upstream_manager.clone(), - beobachten.clone(), - shared.clone(), - )); - } - - let tls_fingerprint = - observe_tls_client_fingerprint(stats.as_ref(), &config, real_peer.ip(), &handshake); - - let (read_half, write_half) = tokio::io::split(stream); - - let (mut tls_reader, tls_writer, tls_user) = match handle_tls_handshake_with_shared( - &handshake, read_half, write_half, real_peer, - &config, &replay_checker, &rng, tls_cache.clone(), - shared.as_ref(), - ).await { - HandshakeResult::Success(result) => result, - HandshakeResult::BadClient { reader, writer } => { - stats.increment_connects_bad_with_class("tls_handshake_bad_client"); - record_tls_fingerprint_bad_or_probe( - stats.as_ref(), - &config, - real_peer.ip(), - tls_fingerprint.as_ref(), - ); - return Ok(masking_outcome( - reader, - writer, - handshake.clone(), - real_peer, - local_addr, - config.clone(), - upstream_manager.clone(), - beobachten.clone(), - shared.clone(), - )); - } - HandshakeResult::Error(e) => { - record_tls_fingerprint_bad_or_probe( - stats.as_ref(), - &config, - real_peer.ip(), - tls_fingerprint.as_ref(), - ); - increment_bad_on_unknown_tls_sni(stats.as_ref(), &e); - return Err(e); - } - }; - record_tls_fingerprint_auth_success( - stats.as_ref(), - &config, - real_peer.ip(), - tls_fingerprint.as_ref(), - tls_user.as_str(), - ); - - debug!(peer = %peer, "Reading MTProto handshake through TLS"); - let mtproto_data = tls_reader.read_exact(HANDSHAKE_LEN).await?; - let mtproto_handshake: [u8; HANDSHAKE_LEN] = mtproto_data[..].try_into() - .map_err(|_| ProxyError::InvalidHandshake("Short MTProto handshake".into()))?; - - let (crypto_reader, crypto_writer, success) = match handle_mtproto_handshake_with_shared( - &mtproto_handshake, tls_reader, tls_writer, real_peer, - &config, &replay_checker, true, Some(tls_user.as_str()), - shared.as_ref(), - ).await { - HandshakeResult::Success(result) => result, - HandshakeResult::BadClient { reader, writer } => { - // MTProto failed after TLS ServerHello was already sent. - // Switch fallback relay back to raw transport so the mask - // backend receives valid TLS records (not unwrapped payload). - let (reader, pending_plaintext) = reader.into_inner_with_pending_plaintext(); - let writer = writer.into_inner(); - let pending_record = if pending_plaintext.is_empty() { - Vec::new() - } else { - wrap_tls_application_record(&pending_plaintext) - }; - let reader = tokio::io::AsyncReadExt::chain(std::io::Cursor::new(pending_record), reader); - stats.increment_connects_bad_with_class("tls_mtproto_bad_client"); - debug!( - peer = %peer, - "Authenticated TLS session failed MTProto validation; engaging masking fallback" - ); - return Ok(masking_outcome( - reader, - writer, - Vec::new(), - real_peer, - local_addr, - config.clone(), - upstream_manager.clone(), - beobachten.clone(), - shared.clone(), - )); - } - HandshakeResult::Error(e) => return Err(e), - }; - - Ok(HandshakeOutcome::NeedsRelay(Box::pin( - RunningClientHandler::handle_authenticated_static_with_shared( - crypto_reader, crypto_writer, success, - upstream_manager, stats, config, buffer_pool, rng, me_pool, - me_pool_runtime, - route_runtime.clone(), - local_addr, real_peer, ip_tracker.clone(), - shared.clone(), - ), - ))) - } else { - if !config.general.modes.classic && !config.general.modes.secure { - debug!(peer = %real_peer, "Non-TLS modes disabled"); - stats.increment_connects_bad_with_class("direct_modes_disabled"); - maybe_apply_mask_reject_delay(&config).await; - let (reader, writer) = tokio::io::split(stream); - return Ok(masking_outcome( - reader, - writer, - first_bytes.to_vec(), - real_peer, - local_addr, - config.clone(), - upstream_manager.clone(), - beobachten.clone(), - shared.clone(), - )); - } - - let mut handshake = [0u8; HANDSHAKE_LEN]; - handshake[..5].copy_from_slice(&first_bytes); - stream.read_exact(&mut handshake[5..]).await?; - - let (read_half, write_half) = tokio::io::split(stream); - - let (crypto_reader, crypto_writer, success) = match handle_mtproto_handshake_with_shared( - &handshake, read_half, write_half, real_peer, - &config, &replay_checker, false, None, - shared.as_ref(), - ).await { - HandshakeResult::Success(result) => result, - HandshakeResult::BadClient { reader, writer } => { - stats.increment_connects_bad_with_class("direct_mtproto_bad_client"); - return Ok(masking_outcome( - reader, - writer, - handshake.to_vec(), - real_peer, - local_addr, - config.clone(), - upstream_manager.clone(), - beobachten.clone(), - shared.clone(), - )); - } - HandshakeResult::Error(e) => return Err(e), - }; - - Ok(HandshakeOutcome::NeedsRelay(Box::pin( - RunningClientHandler::handle_authenticated_static_with_shared( - crypto_reader, - crypto_writer, - success, - upstream_manager, - stats, - config, - buffer_pool, - rng, - me_pool, - me_pool_runtime, - route_runtime.clone(), - local_addr, - real_peer, - ip_tracker.clone(), - shared.clone(), - ) - ))) - } - }).await { - Ok(Ok(outcome)) => outcome, - Ok(Err(e)) => { - debug!(peer = %peer, error = %e, "Handshake failed"); - stats_for_timeout.increment_handshake_failure_class(classify_handshake_failure_class(&e)); - record_handshake_failure_class( - &beobachten_for_timeout, - &config_for_timeout, - peer_for_timeout, - &e, - ); - return Err(e); - } - Err(_) => { - stats_for_timeout.increment_handshake_timeouts(); - stats_for_timeout.increment_handshake_failure_class("timeout"); - debug!(peer = %peer, "Handshake timeout"); - record_beobachten_class( - &beobachten_for_timeout, - &config_for_timeout, - peer_for_timeout, - "other", - ); - return Err(ProxyError::TgHandshakeTimeout); - } - }; - - // Phase 2: relay (WITHOUT handshake timeout — relay has its own activity timeouts) - match outcome { - HandshakeOutcome::NeedsRelay(fut) | HandshakeOutcome::NeedsMasking(fut) => fut.await, - } -} +pub use stream_entry::handle_client_stream; +pub use stream_entry::{ + handle_client_stream_with_shared, handle_client_stream_with_shared_and_pool_runtime, +}; pub struct ClientHandler; @@ -1038,704 +202,6 @@ impl ClientHandler { } } -impl RunningClientHandler { - pub async fn run(self) -> Result<()> { - self.stats.increment_connects_all(); - let peer = self.peer; - debug!(peer = %peer, "New connection"); - - if let Err(e) = configure_client_socket( - &self.stream, - self.config.timeouts.client_keepalive, - self.config.timeouts.client_ack, - ) { - debug!(peer = %peer, error = %e, "Failed to configure client socket"); - } - - #[cfg(unix)] - let raw_fd = self.raw_fd; - let rst_on_close = self.rst_on_close; - - let outcome = match self.do_handshake().await? { - Some(outcome) => outcome, - None => return Ok(()), - }; - - // Phase 2: relay (WITHOUT handshake timeout — relay has its own activity timeouts) - match outcome { - HandshakeOutcome::NeedsRelay(fut) => { - #[cfg(unix)] - if matches!(rst_on_close, crate::config::RstOnCloseMode::Errors) { - let _ = crate::transport::socket::clear_linger_fd(raw_fd); - } - fut.await - } - HandshakeOutcome::NeedsMasking(fut) => fut.await, - } - } - - async fn do_handshake(mut self) -> Result> { - let mut local_addr = self.stream.local_addr().map_err(ProxyError::Io)?; - - if self.proxy_protocol_enabled { - if !is_trusted_proxy_source( - self.peer.ip(), - &self.config.server.proxy_protocol_trusted_cidrs, - ) { - self.stats - .increment_connects_bad_with_class("proxy_protocol_untrusted"); - warn!( - peer = %self.peer, - trusted = ?self.config.server.proxy_protocol_trusted_cidrs, - "Rejecting PROXY protocol header from untrusted source" - ); - record_beobachten_class(&self.beobachten, &self.config, self.peer.ip(), "other"); - return Err(ProxyError::InvalidProxyProtocol); - } - - let proxy_header_timeout = - Duration::from_millis(self.config.server.proxy_protocol_header_timeout_ms.max(1)); - match timeout( - proxy_header_timeout, - parse_proxy_protocol(&mut self.stream, self.peer), - ) - .await - { - Ok(Ok(info)) => { - debug!( - peer = %self.peer, - client = %info.src_addr, - version = info.version, - "PROXY protocol header parsed" - ); - self.peer = normalize_ip(info.src_addr); - self.real_peer_from_proxy = Some(self.peer); - if let Ok(mut slot) = self.real_peer_report.lock() { - *slot = Some(self.peer); - } - if let Some(dst) = info.dst_addr { - local_addr = dst; - } - } - Ok(Err(e)) => { - self.stats - .increment_connects_bad_with_class("proxy_protocol_invalid_header"); - warn!(peer = %self.peer, error = %e, "Invalid PROXY protocol header"); - record_beobachten_class( - &self.beobachten, - &self.config, - self.peer.ip(), - "other", - ); - return Err(e); - } - Err(_) => { - self.stats - .increment_connects_bad_with_class("proxy_protocol_header_timeout"); - warn!( - peer = %self.peer, - timeout_ms = proxy_header_timeout.as_millis(), - "PROXY protocol header timeout" - ); - record_beobachten_class( - &self.beobachten, - &self.config, - self.peer.ip(), - "other", - ); - return Err(ProxyError::InvalidProxyProtocol); - } - } - } - - let first_byte_idle_secs = - effective_client_first_byte_idle_secs(&self.config, self.shared.as_ref()); - let first_byte = if first_byte_idle_secs == 0 { - None - } else { - let idle_timeout = Duration::from_secs(first_byte_idle_secs); - let mut first_byte = [0u8; 1]; - match timeout(idle_timeout, self.stream.read(&mut first_byte)).await { - Ok(Ok(0)) => { - debug!(peer = %self.peer, "Connection closed before first client byte"); - return Ok(None); - } - Ok(Ok(_)) => Some(first_byte[0]), - Ok(Err(e)) - if matches!( - e.kind(), - std::io::ErrorKind::UnexpectedEof - | std::io::ErrorKind::ConnectionReset - | std::io::ErrorKind::ConnectionAborted - | std::io::ErrorKind::BrokenPipe - | std::io::ErrorKind::NotConnected - ) => - { - debug!( - peer = %self.peer, - error = %e, - "Connection closed before first client byte" - ); - return Ok(None); - } - Ok(Err(e)) => { - debug!( - peer = %self.peer, - error = %e, - "Failed while waiting for first client byte" - ); - return Err(ProxyError::Io(e)); - } - Err(_) => { - debug!( - peer = %self.peer, - idle_secs = first_byte_idle_secs, - "Closing idle pooled connection before first client byte" - ); - return Ok(None); - } - } - }; - - let handshake_timeout = handshake_timeout_with_mask_grace(&self.config); - let stats = self.stats.clone(); - let config_for_timeout = self.config.clone(); - let beobachten_for_timeout = self.beobachten.clone(); - let peer_for_timeout = self.peer.ip(); - let peer_for_log = self.peer; - - let outcome = match timeout(handshake_timeout, async { - let mut first_bytes = [0u8; 5]; - if let Some(first_byte) = first_byte { - first_bytes[0] = first_byte; - self.stream.read_exact(&mut first_bytes[1..]).await?; - } else { - self.stream.read_exact(&mut first_bytes).await?; - } - - let is_tls = tls::is_tls_handshake(&first_bytes[..3]); - let peer = self.peer; - - debug!(peer = %peer, is_tls = is_tls, "Handshake type detected"); - - if is_tls { - self.handle_tls_client(first_bytes, local_addr).await - } else { - self.handle_direct_client(first_bytes, local_addr).await - } - }) - .await - { - Ok(Ok(outcome)) => outcome, - Ok(Err(e)) => { - debug!(peer = %peer_for_log, error = %e, "Handshake failed"); - stats.increment_handshake_failure_class(classify_handshake_failure_class(&e)); - record_handshake_failure_class( - &beobachten_for_timeout, - &config_for_timeout, - peer_for_timeout, - &e, - ); - return Err(e); - } - Err(_) => { - stats.increment_handshake_timeouts(); - stats.increment_handshake_failure_class("timeout"); - debug!(peer = %peer_for_log, "Handshake timeout"); - record_beobachten_class( - &beobachten_for_timeout, - &config_for_timeout, - peer_for_timeout, - "other", - ); - return Err(ProxyError::TgHandshakeTimeout); - } - }; - - Ok(Some(outcome)) - } - - async fn handle_tls_client( - mut self, - first_bytes: [u8; 5], - local_addr: SocketAddr, - ) -> Result { - let peer = self.peer; - - let tls_len = u16::from_be_bytes([first_bytes[3], first_bytes[4]]) as usize; - - debug!(peer = %peer, tls_len = tls_len, "Reading TLS handshake"); - - // RFC 8446 §5.1: TLS record payload MUST NOT exceed 2^14 (16_384) bytes. - // Lower bound is a structural minimum for a valid TLS 1.3 ClientHello - // (record header + handshake header + random + session_id + cipher_suites - // + compression + at least one extension with SNI). The previous value of - // 512 was implicitly coupled to TLS_REQUEST_LENGTH=517 from the official - // Telegram MTProxy reference server, leaving only a 5-byte margin and - // incorrectly rejecting compact but spec-compliant ClientHellos from - // third-party clients or future Telegram versions. - if !tls_clienthello_len_in_bounds(tls_len) { - debug!(peer = %peer, tls_len = tls_len, max_tls_len = MAX_TLS_PLAINTEXT_SIZE, "TLS handshake length out of bounds"); - self.stats - .increment_connects_bad_with_class("tls_clienthello_len_out_of_bounds"); - maybe_apply_mask_reject_delay(&self.config).await; - let (reader, writer) = self.stream.into_split(); - return Ok(masking_outcome( - reader, - writer, - first_bytes.to_vec(), - peer, - local_addr, - self.config.clone(), - self.upstream_manager.clone(), - self.beobachten.clone(), - self.shared.clone(), - )); - } - - let mut handshake = vec![0u8; 5 + tls_len]; - handshake[..5].copy_from_slice(&first_bytes); - let body_read = match read_with_progress(&mut self.stream, &mut handshake[5..]).await { - Ok(n) => n, - Err(e) => { - debug!(peer = %peer, error = %e, tls_len = tls_len, "TLS ClientHello body read failed; engaging masking fallback"); - self.stats - .increment_connects_bad_with_class("tls_clienthello_read_error"); - maybe_apply_mask_reject_delay(&self.config).await; - let (reader, writer) = self.stream.into_split(); - return Ok(masking_outcome( - reader, - writer, - handshake[..5].to_vec(), - peer, - local_addr, - self.config.clone(), - self.upstream_manager.clone(), - self.beobachten.clone(), - self.shared.clone(), - )); - } - }; - - if body_read < tls_len { - debug!(peer = %peer, got = body_read, expected = tls_len, "Truncated in-range TLS ClientHello; engaging masking fallback"); - self.stats - .increment_connects_bad_with_class("tls_clienthello_truncated"); - maybe_apply_mask_reject_delay(&self.config).await; - let initial_len = 5 + body_read; - let (reader, writer) = self.stream.into_split(); - return Ok(masking_outcome( - reader, - writer, - handshake[..initial_len].to_vec(), - peer, - local_addr, - self.config.clone(), - self.upstream_manager.clone(), - self.beobachten.clone(), - self.shared.clone(), - )); - } - - let tls_fingerprint = observe_tls_client_fingerprint( - self.stats.as_ref(), - &self.config, - peer.ip(), - &handshake, - ); - - let config = self.config.clone(); - let replay_checker = self.replay_checker.clone(); - let stats = self.stats.clone(); - let buffer_pool = self.buffer_pool.clone(); - - let (read_half, write_half) = self.stream.into_split(); - - #[cfg(target_os = "linux")] - let response_write_options = - TlsResponseWriteOptions::tcp(self.raw_fd, self.tls_response_fragment_size); - #[cfg(not(target_os = "linux"))] - let response_write_options = TlsResponseWriteOptions::default(); - - let (mut tls_reader, tls_writer, tls_user) = - match handle_tls_handshake_with_shared_and_options( - &handshake, - read_half, - write_half, - peer, - &config, - &replay_checker, - &self.rng, - self.tls_cache.clone(), - self.shared.as_ref(), - response_write_options, - ) - .await - { - HandshakeResult::Success(result) => result, - HandshakeResult::BadClient { reader, writer } => { - stats.increment_connects_bad_with_class("tls_handshake_bad_client"); - record_tls_fingerprint_bad_or_probe( - stats.as_ref(), - &config, - peer.ip(), - tls_fingerprint.as_ref(), - ); - return Ok(masking_outcome( - reader, - writer, - handshake.clone(), - peer, - local_addr, - config.clone(), - self.upstream_manager.clone(), - self.beobachten.clone(), - self.shared.clone(), - )); - } - HandshakeResult::Error(e) => { - record_tls_fingerprint_bad_or_probe( - stats.as_ref(), - &config, - peer.ip(), - tls_fingerprint.as_ref(), - ); - increment_bad_on_unknown_tls_sni(stats.as_ref(), &e); - return Err(e); - } - }; - record_tls_fingerprint_auth_success( - stats.as_ref(), - &config, - peer.ip(), - tls_fingerprint.as_ref(), - tls_user.as_str(), - ); - - debug!(peer = %peer, "Reading MTProto handshake through TLS"); - let mtproto_data = tls_reader.read_exact(HANDSHAKE_LEN).await?; - let mtproto_handshake: [u8; HANDSHAKE_LEN] = mtproto_data[..] - .try_into() - .map_err(|_| ProxyError::InvalidHandshake("Short MTProto handshake".into()))?; - - let (crypto_reader, crypto_writer, success) = match handle_mtproto_handshake_with_shared( - &mtproto_handshake, - tls_reader, - tls_writer, - peer, - &config, - &replay_checker, - true, - Some(tls_user.as_str()), - self.shared.as_ref(), - ) - .await - { - HandshakeResult::Success(result) => result, - HandshakeResult::BadClient { reader, writer } => { - // MTProto failed after TLS ServerHello was already sent. - // Switch fallback relay back to raw transport so the mask - // backend receives valid TLS records (not unwrapped payload). - let (reader, pending_plaintext) = reader.into_inner_with_pending_plaintext(); - let writer = writer.into_inner(); - let pending_record = if pending_plaintext.is_empty() { - Vec::new() - } else { - wrap_tls_application_record(&pending_plaintext) - }; - let reader = - tokio::io::AsyncReadExt::chain(std::io::Cursor::new(pending_record), reader); - stats.increment_connects_bad_with_class("tls_mtproto_bad_client"); - debug!( - peer = %peer, - "Authenticated TLS session failed MTProto validation; engaging masking fallback" - ); - return Ok(masking_outcome( - reader, - writer, - Vec::new(), - peer, - local_addr, - config.clone(), - self.upstream_manager.clone(), - self.beobachten.clone(), - self.shared.clone(), - )); - } - HandshakeResult::Error(e) => return Err(e), - }; - - Ok(HandshakeOutcome::NeedsRelay(Box::pin( - Self::handle_authenticated_static_with_shared( - crypto_reader, - crypto_writer, - success, - self.upstream_manager, - self.stats, - self.config, - buffer_pool, - self.rng, - self.me_pool, - self.me_pool_runtime, - self.route_runtime.clone(), - local_addr, - peer, - self.ip_tracker, - self.shared, - ), - ))) - } - - async fn handle_direct_client( - mut self, - first_bytes: [u8; 5], - local_addr: SocketAddr, - ) -> Result { - let peer = self.peer; - - if !self.config.general.modes.classic && !self.config.general.modes.secure { - debug!(peer = %peer, "Non-TLS modes disabled"); - self.stats - .increment_connects_bad_with_class("direct_modes_disabled"); - maybe_apply_mask_reject_delay(&self.config).await; - let (reader, writer) = self.stream.into_split(); - return Ok(masking_outcome( - reader, - writer, - first_bytes.to_vec(), - peer, - local_addr, - self.config.clone(), - self.upstream_manager.clone(), - self.beobachten.clone(), - self.shared.clone(), - )); - } - - let mut handshake = [0u8; HANDSHAKE_LEN]; - handshake[..5].copy_from_slice(&first_bytes); - self.stream.read_exact(&mut handshake[5..]).await?; - - let config = self.config.clone(); - let replay_checker = self.replay_checker.clone(); - let stats = self.stats.clone(); - let buffer_pool = self.buffer_pool.clone(); - - let (read_half, write_half) = self.stream.into_split(); - - let (crypto_reader, crypto_writer, success) = match handle_mtproto_handshake_with_shared( - &handshake, - read_half, - write_half, - peer, - &config, - &replay_checker, - false, - None, - self.shared.as_ref(), - ) - .await - { - HandshakeResult::Success(result) => result, - HandshakeResult::BadClient { reader, writer } => { - stats.increment_connects_bad_with_class("direct_mtproto_bad_client"); - return Ok(masking_outcome( - reader, - writer, - handshake.to_vec(), - peer, - local_addr, - config.clone(), - self.upstream_manager.clone(), - self.beobachten.clone(), - self.shared.clone(), - )); - } - HandshakeResult::Error(e) => return Err(e), - }; - - Ok(HandshakeOutcome::NeedsRelay(Box::pin( - Self::handle_authenticated_static_with_shared( - crypto_reader, - crypto_writer, - success, - self.upstream_manager, - self.stats, - self.config, - buffer_pool, - self.rng, - self.me_pool, - self.me_pool_runtime, - self.route_runtime.clone(), - local_addr, - peer, - self.ip_tracker, - self.shared, - ), - ))) - } - - /// Main dispatch after successful handshake. - /// Two modes: - /// - Direct: TCP relay to TG DC (existing behavior) - /// - Middle Proxy: RPC multiplex through ME pool (supports CDN DCs) - #[cfg(test)] - async fn handle_authenticated_static( - client_reader: CryptoReader, - client_writer: CryptoWriter, - success: HandshakeSuccess, - upstream_manager: Arc, - stats: Arc, - config: Arc, - buffer_pool: Arc, - rng: Arc, - me_pool: Option>, - route_runtime: Arc, - local_addr: SocketAddr, - peer_addr: SocketAddr, - ip_tracker: Arc, - ) -> Result<()> - where - R: AsyncRead + Unpin + Send + 'static, - W: AsyncWrite + Unpin + Send + 'static, - { - Self::handle_authenticated_static_with_shared( - client_reader, - client_writer, - success, - upstream_manager, - stats, - config, - buffer_pool, - rng, - me_pool, - None, - route_runtime, - local_addr, - peer_addr, - ip_tracker, - ProxySharedState::new(), - ) - .await - } - - async fn handle_authenticated_static_with_shared( - client_reader: CryptoReader, - client_writer: CryptoWriter, - success: HandshakeSuccess, - upstream_manager: Arc, - stats: Arc, - config: Arc, - buffer_pool: Arc, - rng: Arc, - me_pool: Option>, - me_pool_runtime: Option>>>>, - route_runtime: Arc, - local_addr: SocketAddr, - peer_addr: SocketAddr, - ip_tracker: Arc, - shared: Arc, - ) -> Result<()> - where - R: AsyncRead + Unpin + Send + 'static, - W: AsyncWrite + Unpin + Send + 'static, - { - run_authenticated( - client_reader, - client_writer, - success, - ClientRuntimeDeps { - config, - stats, - upstream_manager, - buffer_pool, - rng, - 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, - stats: Arc, - peer_addr: SocketAddr, - ip_tracker: Arc, - ) -> Result { - acquire_user_connection_reservation(user, config, stats, peer_addr, ip_tracker).await - } - - #[cfg(test)] - async fn check_user_limits_static( - user: &str, - config: &ProxyConfig, - stats: &Stats, - peer_addr: SocketAddr, - ip_tracker: &UserIpTracker, - ) -> Result<()> { - 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(()) => { - ip_tracker.remove_ip(user, peer_addr.ip()).await; - } - 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(), - }); - } - } - - stats.decrement_user_curr_connects(user); - Ok(()) - } -} - #[cfg(test)] #[path = "tests/client_security_tests.rs"] mod security_tests; diff --git a/src/proxy/client/authenticated.rs b/src/proxy/client/authenticated.rs new file mode 100644 index 0000000..d131863 --- /dev/null +++ b/src/proxy/client/authenticated.rs @@ -0,0 +1,163 @@ +use super::*; + +impl RunningClientHandler { + /// Main dispatch after successful handshake. + /// Two modes: + /// - Direct: TCP relay to TG DC (existing behavior) + /// - Middle Proxy: RPC multiplex through ME pool (supports CDN DCs) + #[cfg(test)] + pub(super) async fn handle_authenticated_static( + client_reader: CryptoReader, + client_writer: CryptoWriter, + success: HandshakeSuccess, + upstream_manager: Arc, + stats: Arc, + config: Arc, + buffer_pool: Arc, + rng: Arc, + me_pool: Option>, + route_runtime: Arc, + local_addr: SocketAddr, + peer_addr: SocketAddr, + ip_tracker: Arc, + ) -> Result<()> + where + R: AsyncRead + Unpin + Send + 'static, + W: AsyncWrite + Unpin + Send + 'static, + { + Self::handle_authenticated_static_with_shared( + client_reader, + client_writer, + success, + upstream_manager, + stats, + config, + buffer_pool, + rng, + me_pool, + None, + route_runtime, + local_addr, + peer_addr, + ip_tracker, + ProxySharedState::new(), + ) + .await + } + + pub(super) async fn handle_authenticated_static_with_shared( + client_reader: CryptoReader, + client_writer: CryptoWriter, + success: HandshakeSuccess, + upstream_manager: Arc, + stats: Arc, + config: Arc, + buffer_pool: Arc, + rng: Arc, + me_pool: Option>, + me_pool_runtime: Option>>>>, + route_runtime: Arc, + local_addr: SocketAddr, + peer_addr: SocketAddr, + ip_tracker: Arc, + shared: Arc, + ) -> Result<()> + where + R: AsyncRead + Unpin + Send + 'static, + W: AsyncWrite + Unpin + Send + 'static, + { + run_authenticated( + client_reader, + client_writer, + success, + ClientRuntimeDeps { + config, + stats, + upstream_manager, + buffer_pool, + rng, + me_pool, + me_pool_runtime, + route_runtime, + ip_tracker, + shared, + }, + local_addr, + peer_addr, + ConntrackClosePolicy::Publish, + ) + .await + } + + #[cfg(test)] + pub(super) async fn acquire_user_connection_reservation_static( + user: &str, + config: &ProxyConfig, + stats: Arc, + peer_addr: SocketAddr, + ip_tracker: Arc, + ) -> Result { + acquire_user_connection_reservation(user, config, stats, peer_addr, ip_tracker).await + } + + #[cfg(test)] + pub(super) async fn check_user_limits_static( + user: &str, + config: &ProxyConfig, + stats: &Stats, + peer_addr: SocketAddr, + ip_tracker: &UserIpTracker, + ) -> Result<()> { + 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(()) => { + ip_tracker.remove_ip(user, peer_addr.ip()).await; + } + 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(), + }); + } + } + + stats.decrement_user_curr_connects(user); + Ok(()) + } +} diff --git a/src/proxy/client/direct_client.rs b/src/proxy/client/direct_client.rs new file mode 100644 index 0000000..6c2390d --- /dev/null +++ b/src/proxy/client/direct_client.rs @@ -0,0 +1,92 @@ +use super::*; + +impl RunningClientHandler { + pub(super) async fn handle_direct_client( + mut self, + first_bytes: [u8; 5], + local_addr: SocketAddr, + ) -> Result { + let peer = self.peer; + + if !self.config.general.modes.classic && !self.config.general.modes.secure { + debug!(peer = %peer, "Non-TLS modes disabled"); + self.stats + .increment_connects_bad_with_class("direct_modes_disabled"); + maybe_apply_mask_reject_delay(&self.config).await; + let (reader, writer) = self.stream.into_split(); + return Ok(masking_outcome( + reader, + writer, + first_bytes.to_vec(), + peer, + local_addr, + self.config.clone(), + self.upstream_manager.clone(), + self.beobachten.clone(), + self.shared.clone(), + )); + } + + let mut handshake = [0u8; HANDSHAKE_LEN]; + handshake[..5].copy_from_slice(&first_bytes); + self.stream.read_exact(&mut handshake[5..]).await?; + + let config = self.config.clone(); + let replay_checker = self.replay_checker.clone(); + let stats = self.stats.clone(); + let buffer_pool = self.buffer_pool.clone(); + + let (read_half, write_half) = self.stream.into_split(); + + let (crypto_reader, crypto_writer, success) = match handle_mtproto_handshake_with_shared( + &handshake, + read_half, + write_half, + peer, + &config, + &replay_checker, + false, + None, + self.shared.as_ref(), + ) + .await + { + HandshakeResult::Success(result) => result, + HandshakeResult::BadClient { reader, writer } => { + stats.increment_connects_bad_with_class("direct_mtproto_bad_client"); + return Ok(masking_outcome( + reader, + writer, + handshake.to_vec(), + peer, + local_addr, + config.clone(), + self.upstream_manager.clone(), + self.beobachten.clone(), + self.shared.clone(), + )); + } + HandshakeResult::Error(e) => return Err(e), + }; + + Ok(HandshakeOutcome::NeedsRelay(Box::pin( + Self::handle_authenticated_static_with_shared( + crypto_reader, + crypto_writer, + success, + self.upstream_manager, + self.stats, + self.config, + buffer_pool, + self.rng, + self.me_pool, + self.me_pool_runtime, + self.route_runtime.clone(), + local_addr, + peer, + self.ip_tracker, + self.shared, + ), + ))) + } +} diff --git a/src/proxy/client/handshake_support.rs b/src/proxy/client/handshake_support.rs new file mode 100644 index 0000000..a98504e --- /dev/null +++ b/src/proxy/client/handshake_support.rs @@ -0,0 +1,357 @@ +use super::*; + +pub(super) fn beobachten_ttl(config: &ProxyConfig) -> Duration { + const BEOBACHTEN_TTL_MAX_MINUTES: u64 = 24 * 60; + let minutes = config.general.beobachten_minutes; + if minutes == 0 { + static BEOBACHTEN_ZERO_MINUTES_WARNED: OnceLock = OnceLock::new(); + let warned = BEOBACHTEN_ZERO_MINUTES_WARNED.get_or_init(|| AtomicBool::new(false)); + if !warned.swap(true, Ordering::Relaxed) { + warn!( + "general.beobachten_minutes=0 is insecure because entries expire immediately; forcing minimum TTL to 1 minute" + ); + } + return Duration::from_secs(60); + } + + if minutes > BEOBACHTEN_TTL_MAX_MINUTES { + static BEOBACHTEN_OVERSIZED_MINUTES_WARNED: OnceLock = OnceLock::new(); + let warned = BEOBACHTEN_OVERSIZED_MINUTES_WARNED.get_or_init(|| AtomicBool::new(false)); + if !warned.swap(true, Ordering::Relaxed) { + warn!( + configured_minutes = minutes, + max_minutes = BEOBACHTEN_TTL_MAX_MINUTES, + "general.beobachten_minutes is too large; clamping to secure maximum" + ); + } + } + + Duration::from_secs(minutes.min(BEOBACHTEN_TTL_MAX_MINUTES).saturating_mul(60)) +} + +pub(super) fn wrap_tls_application_record(payload: &[u8]) -> Vec { + let chunks = payload.len().div_ceil(u16::MAX as usize).max(1); + let mut record = Vec::with_capacity(payload.len() + 5 * chunks); + + if payload.is_empty() { + record.push(TLS_RECORD_APPLICATION); + record.extend_from_slice(&TLS_VERSION); + record.extend_from_slice(&0u16.to_be_bytes()); + return record; + } + + for chunk in payload.chunks(u16::MAX as usize) { + record.push(TLS_RECORD_APPLICATION); + record.extend_from_slice(&TLS_VERSION); + record.extend_from_slice(&(chunk.len() as u16).to_be_bytes()); + record.extend_from_slice(chunk); + } + + record +} + +pub(super) fn tls_clienthello_len_in_bounds(tls_len: usize) -> bool { + (MIN_TLS_CLIENT_HELLO_SIZE..=MAX_TLS_PLAINTEXT_SIZE).contains(&tls_len) +} + +pub(super) async fn read_with_progress( + reader: &mut R, + mut buf: &mut [u8], +) -> std::io::Result { + let mut total = 0usize; + while !buf.is_empty() { + match reader.read(buf).await { + Ok(0) => return Ok(total), + Ok(n) => { + total += n; + let (_, rest) = buf.split_at_mut(n); + buf = rest; + } + Err(e) => return Err(e), + } + } + Ok(total) +} + +pub(super) async fn maybe_apply_mask_reject_delay(config: &ProxyConfig) { + let min = config.censorship.server_hello_delay_min_ms; + let max = config.censorship.server_hello_delay_max_ms; + if max == 0 { + return; + } + + let delay_ms = if min >= max { + max + } else { + rand::rng().random_range(min..=max) + }; + + if delay_ms > 0 { + tokio::time::sleep(Duration::from_millis(delay_ms)).await; + } +} + +pub(super) fn handshake_timeout_with_mask_grace(config: &ProxyConfig) -> Duration { + let base = Duration::from_secs(config.timeouts.client_handshake); + if config.censorship.mask { + base.saturating_add(Duration::from_millis(750)) + } else { + base + } +} + +pub(super) fn effective_client_first_byte_idle_secs( + config: &ProxyConfig, + shared: &ProxySharedState, +) -> u64 { + let idle_secs = config.timeouts.client_first_byte_idle_secs; + if idle_secs == 0 { + return 0; + } + if shared.conntrack_pressure_active() { + idle_secs.min( + config + .server + .conntrack_control + .profile + .client_first_byte_idle_cap_secs(), + ) + } else { + idle_secs + } +} + +const MASK_CLASSIFIER_PREFETCH_WINDOW: usize = 16; +#[cfg(test)] +pub(super) const MASK_CLASSIFIER_PREFETCH_TIMEOUT: Duration = Duration::from_millis(5); + +pub(super) fn mask_classifier_prefetch_timeout(config: &ProxyConfig) -> Duration { + Duration::from_millis(config.censorship.mask_classifier_prefetch_timeout_ms) +} + +pub(super) fn should_prefetch_mask_classifier_window(initial_data: &[u8]) -> bool { + if initial_data.len() >= MASK_CLASSIFIER_PREFETCH_WINDOW { + return false; + } + + if initial_data.is_empty() { + // Empty initial_data means there is no client probe prefix to refine. + // Prefetching in this case can consume fallback relay payload bytes and + // accidentally route them through shaping heuristics. + return false; + } + + if initial_data[0] == 0x16 || initial_data.starts_with(b"SSH-") { + return false; + } + + initial_data + .iter() + .all(|b| b.is_ascii_alphabetic() || *b == b' ') +} + +#[cfg(test)] +pub(super) async fn extend_masking_initial_window(reader: &mut R, initial_data: &mut Vec) +where + R: AsyncRead + Unpin, +{ + extend_masking_initial_window_with_timeout( + reader, + initial_data, + MASK_CLASSIFIER_PREFETCH_TIMEOUT, + ) + .await; +} + +pub(super) async fn extend_masking_initial_window_with_timeout( + reader: &mut R, + initial_data: &mut Vec, + prefetch_timeout: Duration, +) where + R: AsyncRead + Unpin, +{ + if !should_prefetch_mask_classifier_window(initial_data) { + return; + } + + let need = MASK_CLASSIFIER_PREFETCH_WINDOW.saturating_sub(initial_data.len()); + if need == 0 { + return; + } + + let mut extra = [0u8; MASK_CLASSIFIER_PREFETCH_WINDOW]; + if let Ok(Ok(n)) = timeout(prefetch_timeout, reader.read(&mut extra[..need])).await + && n > 0 + { + initial_data.extend_from_slice(&extra[..n]); + } +} + +pub(super) fn masking_outcome( + reader: R, + writer: W, + initial_data: Vec, + peer: SocketAddr, + local_addr: SocketAddr, + config: Arc, + upstream_manager: Arc, + beobachten: Arc, + shared: Arc, +) -> HandshakeOutcome +where + R: AsyncRead + Unpin + Send + 'static, + W: AsyncWrite + Unpin + Send + 'static, +{ + HandshakeOutcome::NeedsMasking(Box::pin(async move { + let mut reader = reader; + let mut initial_data = initial_data; + extend_masking_initial_window_with_timeout( + &mut reader, + &mut initial_data, + mask_classifier_prefetch_timeout(&config), + ) + .await; + + crate::proxy::masking::handle_bad_client_with_shared_resolver( + reader, + writer, + &initial_data, + peer, + local_addr, + &config, + &beobachten, + shared.as_ref(), + Some(upstream_manager.as_ref()), + ) + .await; + Ok(()) + })) +} + +pub(super) fn record_beobachten_class( + beobachten: &BeobachtenStore, + config: &ProxyConfig, + peer_ip: IpAddr, + class: &str, +) { + if !config.general.beobachten { + return; + } + beobachten.record(class, peer_ip, beobachten_ttl(config)); +} + +pub(super) fn tls_fingerprint_collection_enabled(config: &ProxyConfig) -> bool { + config.general.beobachten || config.server.api.runtime_edge_enabled +} + +pub(super) fn observe_tls_client_fingerprint( + stats: &Stats, + config: &ProxyConfig, + peer_ip: IpAddr, + handshake: &[u8], +) -> Option { + if !tls_fingerprint_collection_enabled(config) { + return None; + } + + match tls_fingerprint::fingerprint_client_hello(handshake) { + Some(fingerprint) => { + stats.record_tls_fingerprint_observed(&fingerprint, peer_ip, beobachten_ttl(config)); + Some(fingerprint) + } + None => { + stats.increment_tls_fingerprint_parse_error(); + None + } + } +} + +pub(super) fn record_tls_fingerprint_auth_success( + stats: &Stats, + config: &ProxyConfig, + peer_ip: IpAddr, + fingerprint: Option<&TlsClientFingerprint>, + user: &str, +) { + if let Some(fingerprint) = fingerprint { + stats.record_tls_fingerprint_auth_success( + fingerprint, + peer_ip, + user, + beobachten_ttl(config), + ); + } +} + +pub(super) fn record_tls_fingerprint_bad_or_probe( + stats: &Stats, + config: &ProxyConfig, + peer_ip: IpAddr, + fingerprint: Option<&TlsClientFingerprint>, +) { + if let Some(fingerprint) = fingerprint { + stats.record_tls_fingerprint_bad_or_probe(fingerprint, peer_ip, beobachten_ttl(config)); + } +} + +pub(super) fn classify_expected_64_got_0(kind: std::io::ErrorKind) -> Option<&'static str> { + match kind { + std::io::ErrorKind::UnexpectedEof => Some("expected_64_got_0_unexpected_eof"), + std::io::ErrorKind::ConnectionReset => Some("expected_64_got_0_connection_reset"), + std::io::ErrorKind::ConnectionAborted => Some("expected_64_got_0_connection_aborted"), + std::io::ErrorKind::BrokenPipe => Some("expected_64_got_0_broken_pipe"), + std::io::ErrorKind::NotConnected => Some("expected_64_got_0_not_connected"), + _ => None, + } +} + +pub(super) fn classify_handshake_failure_class(error: &ProxyError) -> &'static str { + match error { + ProxyError::Io(err) => classify_expected_64_got_0(err.kind()).unwrap_or("other"), + ProxyError::Stream(StreamError::UnexpectedEof) => "expected_64_got_0_unexpected_eof", + ProxyError::Stream(StreamError::Io(err)) => { + classify_expected_64_got_0(err.kind()).unwrap_or("other") + } + _ => "other", + } +} + +pub(super) fn record_handshake_failure_class( + beobachten: &BeobachtenStore, + config: &ProxyConfig, + peer_ip: IpAddr, + error: &ProxyError, +) { + // Keep beobachten buckets stable while detailed per-kind classification + // is tracked in API counters. + let class = match classify_handshake_failure_class(error) { + value if value.starts_with("expected_64_got_0_") => "expected_64_got_0", + _ => "other", + }; + record_beobachten_class(beobachten, config, peer_ip, class); +} + +#[inline] +pub(super) fn increment_bad_on_unknown_tls_sni(stats: &Stats, error: &ProxyError) { + if matches!(error, ProxyError::UnknownTlsSni) { + stats.increment_connects_bad_with_class("unknown_tls_sni"); + } +} + +pub(super) fn is_trusted_proxy_source(peer_ip: IpAddr, trusted: &[IpNetwork]) -> bool { + if trusted.is_empty() { + static EMPTY_PROXY_TRUST_WARNED: OnceLock = OnceLock::new(); + let warned = EMPTY_PROXY_TRUST_WARNED.get_or_init(|| AtomicBool::new(false)); + if !warned.swap(true, Ordering::Relaxed) { + warn!( + "PROXY protocol enabled but server.proxy_protocol_trusted_cidrs is empty; rejecting all PROXY headers" + ); + } + return false; + } + trusted.iter().any(|cidr| cidr.contains(peer_ip)) +} + +pub(super) fn synthetic_local_addr(port: u16) -> SocketAddr { + SocketAddr::from(([0, 0, 0, 0], port)) +} diff --git a/src/proxy/client/running_lifecycle.rs b/src/proxy/client/running_lifecycle.rs new file mode 100644 index 0000000..8c3ed61 --- /dev/null +++ b/src/proxy/client/running_lifecycle.rs @@ -0,0 +1,219 @@ +use super::*; + +impl RunningClientHandler { + pub async fn run(self) -> Result<()> { + self.stats.increment_connects_all(); + let peer = self.peer; + debug!(peer = %peer, "New connection"); + + if let Err(e) = configure_client_socket( + &self.stream, + self.config.timeouts.client_keepalive, + self.config.timeouts.client_ack, + ) { + debug!(peer = %peer, error = %e, "Failed to configure client socket"); + } + + #[cfg(unix)] + let raw_fd = self.raw_fd; + let rst_on_close = self.rst_on_close; + + let outcome = match self.do_handshake().await? { + Some(outcome) => outcome, + None => return Ok(()), + }; + + // Phase 2: relay (WITHOUT handshake timeout — relay has its own activity timeouts) + match outcome { + HandshakeOutcome::NeedsRelay(fut) => { + #[cfg(unix)] + if matches!(rst_on_close, crate::config::RstOnCloseMode::Errors) { + let _ = crate::transport::socket::clear_linger_fd(raw_fd); + } + fut.await + } + HandshakeOutcome::NeedsMasking(fut) => fut.await, + } + } + + pub(super) async fn do_handshake(mut self) -> Result> { + let mut local_addr = self.stream.local_addr().map_err(ProxyError::Io)?; + + if self.proxy_protocol_enabled { + if !is_trusted_proxy_source( + self.peer.ip(), + &self.config.server.proxy_protocol_trusted_cidrs, + ) { + self.stats + .increment_connects_bad_with_class("proxy_protocol_untrusted"); + warn!( + peer = %self.peer, + trusted = ?self.config.server.proxy_protocol_trusted_cidrs, + "Rejecting PROXY protocol header from untrusted source" + ); + record_beobachten_class(&self.beobachten, &self.config, self.peer.ip(), "other"); + return Err(ProxyError::InvalidProxyProtocol); + } + + let proxy_header_timeout = + Duration::from_millis(self.config.server.proxy_protocol_header_timeout_ms.max(1)); + match timeout( + proxy_header_timeout, + parse_proxy_protocol(&mut self.stream, self.peer), + ) + .await + { + Ok(Ok(info)) => { + debug!( + peer = %self.peer, + client = %info.src_addr, + version = info.version, + "PROXY protocol header parsed" + ); + self.peer = normalize_ip(info.src_addr); + self.real_peer_from_proxy = Some(self.peer); + if let Ok(mut slot) = self.real_peer_report.lock() { + *slot = Some(self.peer); + } + if let Some(dst) = info.dst_addr { + local_addr = dst; + } + } + Ok(Err(e)) => { + self.stats + .increment_connects_bad_with_class("proxy_protocol_invalid_header"); + warn!(peer = %self.peer, error = %e, "Invalid PROXY protocol header"); + record_beobachten_class( + &self.beobachten, + &self.config, + self.peer.ip(), + "other", + ); + return Err(e); + } + Err(_) => { + self.stats + .increment_connects_bad_with_class("proxy_protocol_header_timeout"); + warn!( + peer = %self.peer, + timeout_ms = proxy_header_timeout.as_millis(), + "PROXY protocol header timeout" + ); + record_beobachten_class( + &self.beobachten, + &self.config, + self.peer.ip(), + "other", + ); + return Err(ProxyError::InvalidProxyProtocol); + } + } + } + + let first_byte_idle_secs = + effective_client_first_byte_idle_secs(&self.config, self.shared.as_ref()); + let first_byte = if first_byte_idle_secs == 0 { + None + } else { + let idle_timeout = Duration::from_secs(first_byte_idle_secs); + let mut first_byte = [0u8; 1]; + match timeout(idle_timeout, self.stream.read(&mut first_byte)).await { + Ok(Ok(0)) => { + debug!(peer = %self.peer, "Connection closed before first client byte"); + return Ok(None); + } + Ok(Ok(_)) => Some(first_byte[0]), + Ok(Err(e)) + if matches!( + e.kind(), + std::io::ErrorKind::UnexpectedEof + | std::io::ErrorKind::ConnectionReset + | std::io::ErrorKind::ConnectionAborted + | std::io::ErrorKind::BrokenPipe + | std::io::ErrorKind::NotConnected + ) => + { + debug!( + peer = %self.peer, + error = %e, + "Connection closed before first client byte" + ); + return Ok(None); + } + Ok(Err(e)) => { + debug!( + peer = %self.peer, + error = %e, + "Failed while waiting for first client byte" + ); + return Err(ProxyError::Io(e)); + } + Err(_) => { + debug!( + peer = %self.peer, + idle_secs = first_byte_idle_secs, + "Closing idle pooled connection before first client byte" + ); + return Ok(None); + } + } + }; + + let handshake_timeout = handshake_timeout_with_mask_grace(&self.config); + let stats = self.stats.clone(); + let config_for_timeout = self.config.clone(); + let beobachten_for_timeout = self.beobachten.clone(); + let peer_for_timeout = self.peer.ip(); + let peer_for_log = self.peer; + + let outcome = match timeout(handshake_timeout, async { + let mut first_bytes = [0u8; 5]; + if let Some(first_byte) = first_byte { + first_bytes[0] = first_byte; + self.stream.read_exact(&mut first_bytes[1..]).await?; + } else { + self.stream.read_exact(&mut first_bytes).await?; + } + + let is_tls = tls::is_tls_handshake(&first_bytes[..3]); + let peer = self.peer; + + debug!(peer = %peer, is_tls = is_tls, "Handshake type detected"); + + if is_tls { + self.handle_tls_client(first_bytes, local_addr).await + } else { + self.handle_direct_client(first_bytes, local_addr).await + } + }) + .await + { + Ok(Ok(outcome)) => outcome, + Ok(Err(e)) => { + debug!(peer = %peer_for_log, error = %e, "Handshake failed"); + stats.increment_handshake_failure_class(classify_handshake_failure_class(&e)); + record_handshake_failure_class( + &beobachten_for_timeout, + &config_for_timeout, + peer_for_timeout, + &e, + ); + return Err(e); + } + Err(_) => { + stats.increment_handshake_timeouts(); + stats.increment_handshake_failure_class("timeout"); + debug!(peer = %peer_for_log, "Handshake timeout"); + record_beobachten_class( + &beobachten_for_timeout, + &config_for_timeout, + peer_for_timeout, + "other", + ); + return Err(ProxyError::TgHandshakeTimeout); + } + }; + + Ok(Some(outcome)) + } +} diff --git a/src/proxy/client/stream_entry.rs b/src/proxy/client/stream_entry.rs new file mode 100644 index 0000000..34f68d8 --- /dev/null +++ b/src/proxy/client/stream_entry.rs @@ -0,0 +1,504 @@ +use super::*; + +#[cfg(test)] +pub async fn handle_client_stream( + stream: S, + peer: SocketAddr, + config: Arc, + stats: Arc, + upstream_manager: Arc, + replay_checker: Arc, + buffer_pool: Arc, + rng: Arc, + me_pool: Option>, + route_runtime: Arc, + tls_cache: Option>, + ip_tracker: Arc, + beobachten: Arc, + proxy_protocol_enabled: bool, +) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send + 'static, +{ + handle_client_stream_with_shared( + stream, + peer, + config, + stats, + upstream_manager, + replay_checker, + buffer_pool, + rng, + me_pool, + route_runtime, + tls_cache, + ip_tracker, + beobachten, + ProxySharedState::new(), + proxy_protocol_enabled, + ) + .await +} + +#[allow(clippy::too_many_arguments)] +#[allow(dead_code)] +pub async fn handle_client_stream_with_shared( + stream: S, + peer: SocketAddr, + config: Arc, + stats: Arc, + upstream_manager: Arc, + replay_checker: Arc, + buffer_pool: Arc, + rng: Arc, + me_pool: Option>, + route_runtime: Arc, + tls_cache: Option>, + ip_tracker: Arc, + beobachten: Arc, + shared: Arc, + proxy_protocol_enabled: bool, +) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send + 'static, +{ + handle_client_stream_with_shared_and_pool_runtime( + stream, + peer, + config, + stats, + upstream_manager, + replay_checker, + buffer_pool, + rng, + me_pool, + None, + route_runtime, + tls_cache, + ip_tracker, + beobachten, + shared, + proxy_protocol_enabled, + ) + .await +} + +#[allow(clippy::too_many_arguments)] +pub async fn handle_client_stream_with_shared_and_pool_runtime( + mut stream: S, + peer: SocketAddr, + config: Arc, + stats: Arc, + upstream_manager: Arc, + replay_checker: Arc, + buffer_pool: Arc, + rng: Arc, + me_pool: Option>, + me_pool_runtime: Option>>>>, + route_runtime: Arc, + tls_cache: Option>, + ip_tracker: Arc, + beobachten: Arc, + shared: Arc, + proxy_protocol_enabled: bool, +) -> Result<()> +where + S: AsyncRead + AsyncWrite + Unpin + Send + 'static, +{ + stats.increment_connects_all(); + let mut real_peer = normalize_ip(peer); + + // For non-TCP streams, use a synthetic local address; may be overridden by PROXY protocol dst + let mut local_addr = synthetic_local_addr(config.server.port); + + if proxy_protocol_enabled { + if !is_trusted_proxy_source(peer.ip(), &config.server.proxy_protocol_trusted_cidrs) { + stats.increment_connects_bad_with_class("proxy_protocol_untrusted"); + warn!( + peer = %peer, + trusted = ?config.server.proxy_protocol_trusted_cidrs, + "Rejecting PROXY protocol header from untrusted source" + ); + record_beobachten_class(&beobachten, &config, peer.ip(), "other"); + return Err(ProxyError::InvalidProxyProtocol); + } + + let proxy_header_timeout = + Duration::from_millis(config.server.proxy_protocol_header_timeout_ms.max(1)); + match timeout( + proxy_header_timeout, + parse_proxy_protocol(&mut stream, peer), + ) + .await + { + Ok(Ok(info)) => { + debug!( + peer = %peer, + client = %info.src_addr, + version = info.version, + "PROXY protocol header parsed" + ); + real_peer = normalize_ip(info.src_addr); + if let Some(dst) = info.dst_addr { + local_addr = dst; + } + } + Ok(Err(e)) => { + stats.increment_connects_bad_with_class("proxy_protocol_invalid_header"); + warn!(peer = %peer, error = %e, "Invalid PROXY protocol header"); + record_beobachten_class(&beobachten, &config, peer.ip(), "other"); + return Err(e); + } + Err(_) => { + stats.increment_connects_bad_with_class("proxy_protocol_header_timeout"); + warn!(peer = %peer, timeout_ms = proxy_header_timeout.as_millis(), "PROXY protocol header timeout"); + record_beobachten_class(&beobachten, &config, peer.ip(), "other"); + return Err(ProxyError::InvalidProxyProtocol); + } + } + } + + debug!(peer = %real_peer, "New connection (generic stream)"); + + let first_byte_idle_secs = effective_client_first_byte_idle_secs(&config, shared.as_ref()); + let first_byte = if first_byte_idle_secs == 0 { + None + } else { + let idle_timeout = Duration::from_secs(first_byte_idle_secs); + let mut first_byte = [0u8; 1]; + match timeout(idle_timeout, stream.read(&mut first_byte)).await { + Ok(Ok(0)) => { + debug!(peer = %real_peer, "Connection closed before first client byte"); + return Ok(()); + } + Ok(Ok(_)) => Some(first_byte[0]), + Ok(Err(e)) + if matches!( + e.kind(), + std::io::ErrorKind::UnexpectedEof + | std::io::ErrorKind::ConnectionReset + | std::io::ErrorKind::ConnectionAborted + | std::io::ErrorKind::BrokenPipe + | std::io::ErrorKind::NotConnected + ) => + { + debug!( + peer = %real_peer, + error = %e, + "Connection closed before first client byte" + ); + return Ok(()); + } + Ok(Err(e)) => { + debug!( + peer = %real_peer, + error = %e, + "Failed while waiting for first client byte" + ); + return Err(ProxyError::Io(e)); + } + Err(_) => { + debug!( + peer = %real_peer, + idle_secs = first_byte_idle_secs, + "Closing idle pooled connection before first client byte" + ); + return Ok(()); + } + } + }; + + let handshake_timeout = handshake_timeout_with_mask_grace(&config); + let stats_for_timeout = stats.clone(); + let config_for_timeout = config.clone(); + let beobachten_for_timeout = beobachten.clone(); + let peer_for_timeout = real_peer.ip(); + + // Phase 2: active handshake (with timeout after the first client byte) + let outcome = match timeout(handshake_timeout, async { + let mut first_bytes = [0u8; 5]; + if let Some(first_byte) = first_byte { + first_bytes[0] = first_byte; + stream.read_exact(&mut first_bytes[1..]).await?; + } else { + stream.read_exact(&mut first_bytes).await?; + } + + let is_tls = tls::is_tls_handshake(&first_bytes[..3]); + debug!(peer = %real_peer, is_tls = is_tls, "Handshake type detected"); + + if is_tls { + let tls_len = u16::from_be_bytes([first_bytes[3], first_bytes[4]]) as usize; + + // RFC 8446 §5.1: TLS record payload MUST NOT exceed 2^14 (16_384) bytes. + // Lower bound is a structural minimum for a valid TLS 1.3 ClientHello + // (record header + handshake header + random + session_id + cipher_suites + // + compression + at least one extension with SNI). The previous value of + // 512 was implicitly coupled to TLS_REQUEST_LENGTH=517 from the official + // Telegram MTProxy reference server, leaving only a 5-byte margin and + // incorrectly rejecting compact but spec-compliant ClientHellos from + // third-party clients or future Telegram versions. + if !tls_clienthello_len_in_bounds(tls_len) { + debug!(peer = %real_peer, tls_len = tls_len, max_tls_len = MAX_TLS_PLAINTEXT_SIZE, "TLS handshake length out of bounds"); + stats.increment_connects_bad_with_class("tls_clienthello_len_out_of_bounds"); + maybe_apply_mask_reject_delay(&config).await; + let (reader, writer) = tokio::io::split(stream); + return Ok(masking_outcome( + reader, + writer, + first_bytes.to_vec(), + real_peer, + local_addr, + config.clone(), + upstream_manager.clone(), + beobachten.clone(), + shared.clone(), + )); + } + + let mut handshake = vec![0u8; 5 + tls_len]; + handshake[..5].copy_from_slice(&first_bytes); + let body_read = match read_with_progress(&mut stream, &mut handshake[5..]).await { + Ok(n) => n, + Err(e) => { + debug!(peer = %real_peer, error = %e, tls_len = tls_len, "TLS ClientHello body read failed; engaging masking fallback"); + stats.increment_connects_bad_with_class("tls_clienthello_read_error"); + maybe_apply_mask_reject_delay(&config).await; + let initial_len = 5; + let (reader, writer) = tokio::io::split(stream); + return Ok(masking_outcome( + reader, + writer, + handshake[..initial_len].to_vec(), + real_peer, + local_addr, + config.clone(), + upstream_manager.clone(), + beobachten.clone(), + shared.clone(), + )); + } + }; + + if body_read < tls_len { + debug!(peer = %real_peer, got = body_read, expected = tls_len, "Truncated in-range TLS ClientHello; engaging masking fallback"); + stats.increment_connects_bad_with_class("tls_clienthello_truncated"); + maybe_apply_mask_reject_delay(&config).await; + let initial_len = 5 + body_read; + let (reader, writer) = tokio::io::split(stream); + return Ok(masking_outcome( + reader, + writer, + handshake[..initial_len].to_vec(), + real_peer, + local_addr, + config.clone(), + upstream_manager.clone(), + beobachten.clone(), + shared.clone(), + )); + } + + let tls_fingerprint = + observe_tls_client_fingerprint(stats.as_ref(), &config, real_peer.ip(), &handshake); + + let (read_half, write_half) = tokio::io::split(stream); + + let (mut tls_reader, tls_writer, tls_user) = match handle_tls_handshake_with_shared( + &handshake, read_half, write_half, real_peer, + &config, &replay_checker, &rng, tls_cache.clone(), + shared.as_ref(), + ).await { + HandshakeResult::Success(result) => result, + HandshakeResult::BadClient { reader, writer } => { + stats.increment_connects_bad_with_class("tls_handshake_bad_client"); + record_tls_fingerprint_bad_or_probe( + stats.as_ref(), + &config, + real_peer.ip(), + tls_fingerprint.as_ref(), + ); + return Ok(masking_outcome( + reader, + writer, + handshake.clone(), + real_peer, + local_addr, + config.clone(), + upstream_manager.clone(), + beobachten.clone(), + shared.clone(), + )); + } + HandshakeResult::Error(e) => { + record_tls_fingerprint_bad_or_probe( + stats.as_ref(), + &config, + real_peer.ip(), + tls_fingerprint.as_ref(), + ); + increment_bad_on_unknown_tls_sni(stats.as_ref(), &e); + return Err(e); + } + }; + record_tls_fingerprint_auth_success( + stats.as_ref(), + &config, + real_peer.ip(), + tls_fingerprint.as_ref(), + tls_user.as_str(), + ); + + debug!(peer = %peer, "Reading MTProto handshake through TLS"); + let mtproto_data = tls_reader.read_exact(HANDSHAKE_LEN).await?; + let mtproto_handshake: [u8; HANDSHAKE_LEN] = mtproto_data[..].try_into() + .map_err(|_| ProxyError::InvalidHandshake("Short MTProto handshake".into()))?; + + let (crypto_reader, crypto_writer, success) = match handle_mtproto_handshake_with_shared( + &mtproto_handshake, tls_reader, tls_writer, real_peer, + &config, &replay_checker, true, Some(tls_user.as_str()), + shared.as_ref(), + ).await { + HandshakeResult::Success(result) => result, + HandshakeResult::BadClient { reader, writer } => { + // MTProto failed after TLS ServerHello was already sent. + // Switch fallback relay back to raw transport so the mask + // backend receives valid TLS records (not unwrapped payload). + let (reader, pending_plaintext) = reader.into_inner_with_pending_plaintext(); + let writer = writer.into_inner(); + let pending_record = if pending_plaintext.is_empty() { + Vec::new() + } else { + wrap_tls_application_record(&pending_plaintext) + }; + let reader = tokio::io::AsyncReadExt::chain(std::io::Cursor::new(pending_record), reader); + stats.increment_connects_bad_with_class("tls_mtproto_bad_client"); + debug!( + peer = %peer, + "Authenticated TLS session failed MTProto validation; engaging masking fallback" + ); + return Ok(masking_outcome( + reader, + writer, + Vec::new(), + real_peer, + local_addr, + config.clone(), + upstream_manager.clone(), + beobachten.clone(), + shared.clone(), + )); + } + HandshakeResult::Error(e) => return Err(e), + }; + + Ok(HandshakeOutcome::NeedsRelay(Box::pin( + RunningClientHandler::handle_authenticated_static_with_shared( + crypto_reader, crypto_writer, success, + upstream_manager, stats, config, buffer_pool, rng, me_pool, + me_pool_runtime, + route_runtime.clone(), + local_addr, real_peer, ip_tracker.clone(), + shared.clone(), + ), + ))) + } else { + if !config.general.modes.classic && !config.general.modes.secure { + debug!(peer = %real_peer, "Non-TLS modes disabled"); + stats.increment_connects_bad_with_class("direct_modes_disabled"); + maybe_apply_mask_reject_delay(&config).await; + let (reader, writer) = tokio::io::split(stream); + return Ok(masking_outcome( + reader, + writer, + first_bytes.to_vec(), + real_peer, + local_addr, + config.clone(), + upstream_manager.clone(), + beobachten.clone(), + shared.clone(), + )); + } + + let mut handshake = [0u8; HANDSHAKE_LEN]; + handshake[..5].copy_from_slice(&first_bytes); + stream.read_exact(&mut handshake[5..]).await?; + + let (read_half, write_half) = tokio::io::split(stream); + + let (crypto_reader, crypto_writer, success) = match handle_mtproto_handshake_with_shared( + &handshake, read_half, write_half, real_peer, + &config, &replay_checker, false, None, + shared.as_ref(), + ).await { + HandshakeResult::Success(result) => result, + HandshakeResult::BadClient { reader, writer } => { + stats.increment_connects_bad_with_class("direct_mtproto_bad_client"); + return Ok(masking_outcome( + reader, + writer, + handshake.to_vec(), + real_peer, + local_addr, + config.clone(), + upstream_manager.clone(), + beobachten.clone(), + shared.clone(), + )); + } + HandshakeResult::Error(e) => return Err(e), + }; + + Ok(HandshakeOutcome::NeedsRelay(Box::pin( + RunningClientHandler::handle_authenticated_static_with_shared( + crypto_reader, + crypto_writer, + success, + upstream_manager, + stats, + config, + buffer_pool, + rng, + me_pool, + me_pool_runtime, + route_runtime.clone(), + local_addr, + real_peer, + ip_tracker.clone(), + shared.clone(), + ) + ))) + } + }).await { + Ok(Ok(outcome)) => outcome, + Ok(Err(e)) => { + debug!(peer = %peer, error = %e, "Handshake failed"); + stats_for_timeout.increment_handshake_failure_class(classify_handshake_failure_class(&e)); + record_handshake_failure_class( + &beobachten_for_timeout, + &config_for_timeout, + peer_for_timeout, + &e, + ); + return Err(e); + } + Err(_) => { + stats_for_timeout.increment_handshake_timeouts(); + stats_for_timeout.increment_handshake_failure_class("timeout"); + debug!(peer = %peer, "Handshake timeout"); + record_beobachten_class( + &beobachten_for_timeout, + &config_for_timeout, + peer_for_timeout, + "other", + ); + return Err(ProxyError::TgHandshakeTimeout); + } + }; + + // Phase 2: relay (WITHOUT handshake timeout — relay has its own activity timeouts) + match outcome { + HandshakeOutcome::NeedsRelay(fut) | HandshakeOutcome::NeedsMasking(fut) => fut.await, + } +} diff --git a/src/proxy/client/tls_client.rs b/src/proxy/client/tls_client.rs new file mode 100644 index 0000000..65f76ae --- /dev/null +++ b/src/proxy/client/tls_client.rs @@ -0,0 +1,234 @@ +use super::*; + +impl RunningClientHandler { + pub(super) async fn handle_tls_client( + mut self, + first_bytes: [u8; 5], + local_addr: SocketAddr, + ) -> Result { + let peer = self.peer; + + let tls_len = u16::from_be_bytes([first_bytes[3], first_bytes[4]]) as usize; + + debug!(peer = %peer, tls_len = tls_len, "Reading TLS handshake"); + + // RFC 8446 §5.1: TLS record payload MUST NOT exceed 2^14 (16_384) bytes. + // Lower bound is a structural minimum for a valid TLS 1.3 ClientHello + // (record header + handshake header + random + session_id + cipher_suites + // + compression + at least one extension with SNI). The previous value of + // 512 was implicitly coupled to TLS_REQUEST_LENGTH=517 from the official + // Telegram MTProxy reference server, leaving only a 5-byte margin and + // incorrectly rejecting compact but spec-compliant ClientHellos from + // third-party clients or future Telegram versions. + if !tls_clienthello_len_in_bounds(tls_len) { + debug!(peer = %peer, tls_len = tls_len, max_tls_len = MAX_TLS_PLAINTEXT_SIZE, "TLS handshake length out of bounds"); + self.stats + .increment_connects_bad_with_class("tls_clienthello_len_out_of_bounds"); + maybe_apply_mask_reject_delay(&self.config).await; + let (reader, writer) = self.stream.into_split(); + return Ok(masking_outcome( + reader, + writer, + first_bytes.to_vec(), + peer, + local_addr, + self.config.clone(), + self.upstream_manager.clone(), + self.beobachten.clone(), + self.shared.clone(), + )); + } + + let mut handshake = vec![0u8; 5 + tls_len]; + handshake[..5].copy_from_slice(&first_bytes); + let body_read = match read_with_progress(&mut self.stream, &mut handshake[5..]).await { + Ok(n) => n, + Err(e) => { + debug!(peer = %peer, error = %e, tls_len = tls_len, "TLS ClientHello body read failed; engaging masking fallback"); + self.stats + .increment_connects_bad_with_class("tls_clienthello_read_error"); + maybe_apply_mask_reject_delay(&self.config).await; + let (reader, writer) = self.stream.into_split(); + return Ok(masking_outcome( + reader, + writer, + handshake[..5].to_vec(), + peer, + local_addr, + self.config.clone(), + self.upstream_manager.clone(), + self.beobachten.clone(), + self.shared.clone(), + )); + } + }; + + if body_read < tls_len { + debug!(peer = %peer, got = body_read, expected = tls_len, "Truncated in-range TLS ClientHello; engaging masking fallback"); + self.stats + .increment_connects_bad_with_class("tls_clienthello_truncated"); + maybe_apply_mask_reject_delay(&self.config).await; + let initial_len = 5 + body_read; + let (reader, writer) = self.stream.into_split(); + return Ok(masking_outcome( + reader, + writer, + handshake[..initial_len].to_vec(), + peer, + local_addr, + self.config.clone(), + self.upstream_manager.clone(), + self.beobachten.clone(), + self.shared.clone(), + )); + } + + let tls_fingerprint = observe_tls_client_fingerprint( + self.stats.as_ref(), + &self.config, + peer.ip(), + &handshake, + ); + + let config = self.config.clone(); + let replay_checker = self.replay_checker.clone(); + let stats = self.stats.clone(); + let buffer_pool = self.buffer_pool.clone(); + + let (read_half, write_half) = self.stream.into_split(); + + #[cfg(target_os = "linux")] + let response_write_options = + TlsResponseWriteOptions::tcp(self.raw_fd, self.tls_response_fragment_size); + #[cfg(not(target_os = "linux"))] + let response_write_options = TlsResponseWriteOptions::default(); + + let (mut tls_reader, tls_writer, tls_user) = + match handle_tls_handshake_with_shared_and_options( + &handshake, + read_half, + write_half, + peer, + &config, + &replay_checker, + &self.rng, + self.tls_cache.clone(), + self.shared.as_ref(), + response_write_options, + ) + .await + { + HandshakeResult::Success(result) => result, + HandshakeResult::BadClient { reader, writer } => { + stats.increment_connects_bad_with_class("tls_handshake_bad_client"); + record_tls_fingerprint_bad_or_probe( + stats.as_ref(), + &config, + peer.ip(), + tls_fingerprint.as_ref(), + ); + return Ok(masking_outcome( + reader, + writer, + handshake.clone(), + peer, + local_addr, + config.clone(), + self.upstream_manager.clone(), + self.beobachten.clone(), + self.shared.clone(), + )); + } + HandshakeResult::Error(e) => { + record_tls_fingerprint_bad_or_probe( + stats.as_ref(), + &config, + peer.ip(), + tls_fingerprint.as_ref(), + ); + increment_bad_on_unknown_tls_sni(stats.as_ref(), &e); + return Err(e); + } + }; + record_tls_fingerprint_auth_success( + stats.as_ref(), + &config, + peer.ip(), + tls_fingerprint.as_ref(), + tls_user.as_str(), + ); + + debug!(peer = %peer, "Reading MTProto handshake through TLS"); + let mtproto_data = tls_reader.read_exact(HANDSHAKE_LEN).await?; + let mtproto_handshake: [u8; HANDSHAKE_LEN] = mtproto_data[..] + .try_into() + .map_err(|_| ProxyError::InvalidHandshake("Short MTProto handshake".into()))?; + + let (crypto_reader, crypto_writer, success) = match handle_mtproto_handshake_with_shared( + &mtproto_handshake, + tls_reader, + tls_writer, + peer, + &config, + &replay_checker, + true, + Some(tls_user.as_str()), + self.shared.as_ref(), + ) + .await + { + HandshakeResult::Success(result) => result, + HandshakeResult::BadClient { reader, writer } => { + // MTProto failed after TLS ServerHello was already sent. + // Switch fallback relay back to raw transport so the mask + // backend receives valid TLS records (not unwrapped payload). + let (reader, pending_plaintext) = reader.into_inner_with_pending_plaintext(); + let writer = writer.into_inner(); + let pending_record = if pending_plaintext.is_empty() { + Vec::new() + } else { + wrap_tls_application_record(&pending_plaintext) + }; + let reader = + tokio::io::AsyncReadExt::chain(std::io::Cursor::new(pending_record), reader); + stats.increment_connects_bad_with_class("tls_mtproto_bad_client"); + debug!( + peer = %peer, + "Authenticated TLS session failed MTProto validation; engaging masking fallback" + ); + return Ok(masking_outcome( + reader, + writer, + Vec::new(), + peer, + local_addr, + config.clone(), + self.upstream_manager.clone(), + self.beobachten.clone(), + self.shared.clone(), + )); + } + HandshakeResult::Error(e) => return Err(e), + }; + + Ok(HandshakeOutcome::NeedsRelay(Box::pin( + Self::handle_authenticated_static_with_shared( + crypto_reader, + crypto_writer, + success, + self.upstream_manager, + self.stats, + self.config, + buffer_pool, + self.rng, + self.me_pool, + self.me_pool_runtime, + self.route_runtime.clone(), + local_addr, + peer, + self.ip_tracker, + self.shared, + ), + ))) + } +} diff --git a/src/proxy/direct_relay.rs b/src/proxy/direct_relay.rs index 9215912..ef65298 100644 --- a/src/proxy/direct_relay.rs +++ b/src/proxy/direct_relay.rs @@ -36,6 +36,15 @@ use nix::sys::stat::Mode; #[cfg(unix)] use std::os::unix::fs::OpenOptionsExt; +// Direct relay lifecycle and conntrack publication. +mod relay; +// Telegram DC resolution and upstream handshake. +mod routing; + +pub(crate) use relay::{ + handle_via_direct, handle_via_direct_with_shared, handle_via_direct_with_shared_and_conntrack, +}; +use routing::*; const UNKNOWN_DC_LOG_DISTINCT_LIMIT: usize = 1024; static LOGGED_UNKNOWN_DCS: OnceLock>> = OnceLock::new(); const MAX_SCOPE_HINT_LEN: usize = 64; @@ -224,401 +233,9 @@ fn clear_unknown_dc_log_cache_for_testing() { } #[cfg(test)] -fn unknown_dc_test_lock() -> &'static Mutex<()> { - static TEST_LOCK: OnceLock> = OnceLock::new(); - TEST_LOCK.get_or_init(|| Mutex::new(())) -} - -#[allow(dead_code)] -/// Runs Direct relay with standalone cancellation and shared-state defaults. -pub(crate) async fn handle_via_direct( - client_reader: CryptoReader, - client_writer: CryptoWriter, - success: HandshakeSuccess, - upstream_manager: Arc, - stats: Arc, - config: Arc, - buffer_pool: Arc, - rng: Arc, - route_rx: watch::Receiver, - route_snapshot: RouteCutoverState, - session_id: u64, -) -> Result<()> -where - R: AsyncRead + Unpin + Send + 'static, - W: AsyncWrite + Unpin + Send + 'static, -{ - handle_via_direct_with_shared( - client_reader, - client_writer, - success, - upstream_manager, - stats, - config.clone(), - buffer_pool, - rng, - route_rx, - route_snapshot, - session_id, - SocketAddr::from(([0, 0, 0, 0], config.server.port)), - CancellationToken::new(), - ProxySharedState::new(), - ) - .await -} - -/// Runs Direct relay for a kernel-backed TCP client tuple. -pub(crate) async fn handle_via_direct_with_shared( - client_reader: CryptoReader, - client_writer: CryptoWriter, - success: HandshakeSuccess, - upstream_manager: Arc, - stats: Arc, - config: Arc, - buffer_pool: Arc, - rng: Arc, - route_rx: watch::Receiver, - route_snapshot: RouteCutoverState, - session_id: u64, - local_addr: SocketAddr, - session_cancel: CancellationToken, - shared: Arc, -) -> 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( - client_reader: CryptoReader, - client_writer: CryptoWriter, - success: HandshakeSuccess, - upstream_manager: Arc, - stats: Arc, - config: Arc, - buffer_pool: Arc, - rng: Arc, - mut route_rx: watch::Receiver, - route_snapshot: RouteCutoverState, - session_id: u64, - local_addr: SocketAddr, - session_cancel: CancellationToken, - shared: Arc, - conntrack_close_policy: ConntrackClosePolicy, -) -> Result<()> -where - R: AsyncRead + Unpin + Send + 'static, - W: AsyncWrite + Unpin + Send + 'static, -{ - let user = &success.user; - let dc_addr = get_dc_addr_static(success.dc_idx, &config)?; - - debug!( - user = %user, - peer = %success.peer, - dc = success.dc_idx, - dc_addr = %dc_addr, - proto = ?success.proto_tag, - mode = "direct", - "Connecting to Telegram DC" - ); - - let scope_hint = validated_scope_hint(user); - if user.starts_with("scope_") && scope_hint.is_none() { - warn!( - user = %user, - "Ignoring invalid scope hint and falling back to default upstream selection" - ); - } - let tg_stream = tokio::select! { - result = upstream_manager.connect(dc_addr, Some(success.dc_idx), scope_hint) => result?, - _ = session_cancel.cancelled() => { - return Err(ProxyError::UserDisabled { - user: user.to_string(), - }); - } - }; - - debug!(peer = %success.peer, dc_addr = %dc_addr, "Connected, performing TG handshake"); - - let (tg_reader, tg_writer) = tokio::select! { - result = do_tg_handshake_static(tg_stream, &success, &config, rng.as_ref()) => result?, - _ = session_cancel.cancelled() => { - return Err(ProxyError::UserDisabled { - user: user.to_string(), - }); - } - }; - - debug!(peer = %success.peer, "TG handshake complete, starting relay"); - - stats.increment_user_connects(user); - let _direct_connection_lease = stats.acquire_direct_connection_lease(); - let traffic_lease = shared - .traffic_limiter - .acquire_lease(user, success.peer.ip()); - - let buffer_pool_trim = Arc::clone(&buffer_pool); - let relay_activity_timeout = if shared.conntrack_pressure_active() { - Duration::from_secs( - config - .server - .conntrack_control - .profile - .direct_activity_timeout_secs(), - ) - } else { - Duration::from_secs(1800) - }; - let relay_result = crate::proxy::relay::relay_direct_adaptive( - client_reader, - client_writer, - tg_reader, - tg_writer, - config.general.direct_relay_copy_buf_c2s_bytes, - config.general.direct_relay_copy_buf_s2c_bytes, - config.server.max_connections, - user, - Arc::clone(&stats), - config.access.user_data_quota.get(user).copied(), - traffic_lease, - relay_activity_timeout, - session_cancel.clone(), - Arc::clone(&shared.direct_buffer_budget), - ); - tokio::pin!(relay_result); - let relay_result = loop { - if let Some(cutover) = - affected_cutover_state(&route_rx, RelayRouteMode::Direct, route_snapshot.generation) - { - let delay = cutover_stagger_delay(session_id, cutover.generation); - warn!( - user = %user, - target_mode = cutover.mode.as_str(), - cutover_generation = cutover.generation, - delay_ms = delay.as_millis() as u64, - "Cutover affected direct session, closing client connection" - ); - let _cutover_park_lease = stats.acquire_direct_cutover_park_lease(); - tokio::time::sleep(delay).await; - break Err(ProxyError::RouteSwitched); - } - tokio::select! { - result = &mut relay_result => { - break result; - } - changed = route_rx.changed() => { - if changed.is_err() { - break relay_result.await; - } - } - _ = session_cancel.cancelled() => { - break Err(ProxyError::UserDisabled { - user: user.to_string(), - }); - } - } - }; - - match &relay_result { - Ok(()) => debug!(user = %user, "Direct relay completed"), - Err(e) => debug!(user = %user, error = %e, "Direct relay ended with error"), - } - - let pool_snapshot = buffer_pool_trim.stats(); - stats.set_buffer_pool_gauges( - pool_snapshot.pooled, - pool_snapshot.allocated, - pool_snapshot.allocated.saturating_sub(pool_snapshot.pooled), - ); - - 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 -} - -fn classify_conntrack_close_reason(result: &Result<()>) -> ConntrackCloseReason { - match result { - Ok(()) => ConntrackCloseReason::NormalEof, - Err(crate::error::ProxyError::Io(error)) - if matches!(error.kind(), std::io::ErrorKind::TimedOut) => - { - ConntrackCloseReason::Timeout - } - Err(crate::error::ProxyError::Io(error)) - if matches!( - error.kind(), - std::io::ErrorKind::ConnectionReset - | std::io::ErrorKind::ConnectionAborted - | std::io::ErrorKind::BrokenPipe - | std::io::ErrorKind::NotConnected - | std::io::ErrorKind::UnexpectedEof - ) => - { - ConntrackCloseReason::Reset - } - Err(crate::error::ProxyError::Proxy(message)) - if message.contains("pressure") || message.contains("evicted") => - { - ConntrackCloseReason::Pressure - } - Err(_) => ConntrackCloseReason::Other, - } -} - -fn get_dc_addr_static(dc_idx: i16, config: &ProxyConfig) -> Result { - let prefer_v6 = config.network.prefer == 6 && config.network.ipv6.unwrap_or(true); - let datacenters = if prefer_v6 { - &*TG_DATACENTERS_V6 - } else { - &*TG_DATACENTERS_V4 - }; - - let num_dcs = datacenters.len(); - - let dc_key = dc_idx.to_string(); - if let Some(addrs) = config.dc_overrides.get(&dc_key) { - let mut parsed = Vec::new(); - for addr_str in addrs { - match addr_str.parse::() { - Ok(addr) => parsed.push(addr), - Err(_) => { - warn!(dc_idx = dc_idx, addr_str = %addr_str, "Invalid DC override address in config, ignoring") - } - } - } - - if let Some(addr) = parsed - .iter() - .find(|a| a.is_ipv6() == prefer_v6) - .or_else(|| parsed.first()) - .copied() - { - debug!(dc_idx = dc_idx, addr = %addr, count = parsed.len(), "Using DC override from config"); - return Ok(addr); - } - } - - let abs_dc = dc_idx.unsigned_abs() as usize; - if abs_dc >= 1 && abs_dc <= num_dcs { - return Ok(SocketAddr::new(datacenters[abs_dc - 1], TG_DATACENTER_PORT)); - } - - // Unknown DC requested by client without override: log and fall back. - if !config.dc_overrides.contains_key(&dc_key) { - warn!( - dc_idx = dc_idx, - "Requested non-standard DC with no override; falling back to default cluster" - ); - if config.general.unknown_dc_file_log_enabled - && let Some(path) = &config.general.unknown_dc_log_path - && let Ok(handle) = tokio::runtime::Handle::try_current() - { - if let Some(path) = sanitize_unknown_dc_log_path(path) { - if should_log_unknown_dc(dc_idx) { - handle.spawn_blocking(move || { - if unknown_dc_log_path_is_still_safe(&path) - && let Ok(mut file) = open_unknown_dc_log_append_anchored(&path) - { - let _ = append_unknown_dc_line(&mut file, dc_idx); - } - }); - } - } else { - warn!(dc_idx = dc_idx, raw_path = %path, "Rejected unsafe unknown DC log path"); - } - } - } - - let default_dc = config.default_dc.unwrap_or(2) as usize; - let fallback_idx = if default_dc >= 1 && default_dc <= num_dcs { - default_dc - 1 - } else { - 0 - }; - - info!( - original_dc = dc_idx, - fallback_dc = (fallback_idx + 1) as u16, - fallback_addr = %datacenters[fallback_idx], - "Special DC ---> default_cluster" - ); - - Ok(SocketAddr::new( - datacenters[fallback_idx], - TG_DATACENTER_PORT, - )) -} - -async fn do_tg_handshake_static( - mut stream: S, - success: &HandshakeSuccess, - config: &ProxyConfig, - rng: &SecureRandom, -) -> Result<(CryptoReader>, CryptoWriter>)> -where - S: AsyncRead + AsyncWrite + Unpin, -{ - let (nonce, _tg_enc_key, _tg_enc_iv, _tg_dec_key, _tg_dec_iv) = generate_tg_nonce( - success.proto_tag, - success.dc_idx, - &success.enc_key, - success.enc_iv, - rng, - config.general.fast_mode, - ); - - let (encrypted_nonce, tg_encryptor, tg_decryptor) = encrypt_tg_nonce_with_ciphers(&nonce); - - debug!( - peer = %success.peer, - nonce_head = %hex::encode(&nonce[..16]), - "Sending nonce to Telegram" - ); - - stream.write_all(&encrypted_nonce).await?; - stream.flush().await?; - - let (read_half, write_half) = split(stream); - - let max_pending = config.general.crypto_pending_buffer; - Ok(( - CryptoReader::new(read_half, tg_decryptor), - CryptoWriter::new(write_half, tg_encryptor, max_pending), - )) +fn unknown_dc_test_lock() -> &'static tokio::sync::Mutex<()> { + static TEST_LOCK: OnceLock> = OnceLock::new(); + TEST_LOCK.get_or_init(|| tokio::sync::Mutex::new(())) } #[cfg(test)] diff --git a/src/proxy/direct_relay/relay.rs b/src/proxy/direct_relay/relay.rs new file mode 100644 index 0000000..b70420a --- /dev/null +++ b/src/proxy/direct_relay/relay.rs @@ -0,0 +1,271 @@ +use super::*; + +#[allow(dead_code)] +/// Runs Direct relay with standalone cancellation and shared-state defaults. +pub(crate) async fn handle_via_direct( + client_reader: CryptoReader, + client_writer: CryptoWriter, + success: HandshakeSuccess, + upstream_manager: Arc, + stats: Arc, + config: Arc, + buffer_pool: Arc, + rng: Arc, + route_rx: watch::Receiver, + route_snapshot: RouteCutoverState, + session_id: u64, +) -> Result<()> +where + R: AsyncRead + Unpin + Send + 'static, + W: AsyncWrite + Unpin + Send + 'static, +{ + handle_via_direct_with_shared( + client_reader, + client_writer, + success, + upstream_manager, + stats, + config.clone(), + buffer_pool, + rng, + route_rx, + route_snapshot, + session_id, + SocketAddr::from(([0, 0, 0, 0], config.server.port)), + CancellationToken::new(), + ProxySharedState::new(), + ) + .await +} + +/// Runs Direct relay for a kernel-backed TCP client tuple. +pub(crate) async fn handle_via_direct_with_shared( + client_reader: CryptoReader, + client_writer: CryptoWriter, + success: HandshakeSuccess, + upstream_manager: Arc, + stats: Arc, + config: Arc, + buffer_pool: Arc, + rng: Arc, + route_rx: watch::Receiver, + route_snapshot: RouteCutoverState, + session_id: u64, + local_addr: SocketAddr, + session_cancel: CancellationToken, + shared: Arc, +) -> 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( + client_reader: CryptoReader, + client_writer: CryptoWriter, + success: HandshakeSuccess, + upstream_manager: Arc, + stats: Arc, + config: Arc, + buffer_pool: Arc, + rng: Arc, + mut route_rx: watch::Receiver, + route_snapshot: RouteCutoverState, + session_id: u64, + local_addr: SocketAddr, + session_cancel: CancellationToken, + shared: Arc, + conntrack_close_policy: ConntrackClosePolicy, +) -> Result<()> +where + R: AsyncRead + Unpin + Send + 'static, + W: AsyncWrite + Unpin + Send + 'static, +{ + let user = &success.user; + let dc_addr = get_dc_addr_static(success.dc_idx, &config)?; + + debug!( + user = %user, + peer = %success.peer, + dc = success.dc_idx, + dc_addr = %dc_addr, + proto = ?success.proto_tag, + mode = "direct", + "Connecting to Telegram DC" + ); + + let scope_hint = validated_scope_hint(user); + if user.starts_with("scope_") && scope_hint.is_none() { + warn!( + user = %user, + "Ignoring invalid scope hint and falling back to default upstream selection" + ); + } + let tg_stream = tokio::select! { + result = upstream_manager.connect(dc_addr, Some(success.dc_idx), scope_hint) => result?, + _ = session_cancel.cancelled() => { + return Err(ProxyError::UserDisabled { + user: user.to_string(), + }); + } + }; + + debug!(peer = %success.peer, dc_addr = %dc_addr, "Connected, performing TG handshake"); + + let (tg_reader, tg_writer) = tokio::select! { + result = do_tg_handshake_static(tg_stream, &success, &config, rng.as_ref()) => result?, + _ = session_cancel.cancelled() => { + return Err(ProxyError::UserDisabled { + user: user.to_string(), + }); + } + }; + + debug!(peer = %success.peer, "TG handshake complete, starting relay"); + + stats.increment_user_connects(user); + let _direct_connection_lease = stats.acquire_direct_connection_lease(); + let traffic_lease = shared + .traffic_limiter + .acquire_lease(user, success.peer.ip()); + + let buffer_pool_trim = Arc::clone(&buffer_pool); + let relay_activity_timeout = if shared.conntrack_pressure_active() { + Duration::from_secs( + config + .server + .conntrack_control + .profile + .direct_activity_timeout_secs(), + ) + } else { + Duration::from_secs(1800) + }; + let relay_result = crate::proxy::relay::relay_direct_adaptive( + client_reader, + client_writer, + tg_reader, + tg_writer, + config.general.direct_relay_copy_buf_c2s_bytes, + config.general.direct_relay_copy_buf_s2c_bytes, + config.server.max_connections, + user, + Arc::clone(&stats), + config.access.user_data_quota.get(user).copied(), + traffic_lease, + relay_activity_timeout, + session_cancel.clone(), + Arc::clone(&shared.direct_buffer_budget), + ); + tokio::pin!(relay_result); + let relay_result = loop { + if let Some(cutover) = + affected_cutover_state(&route_rx, RelayRouteMode::Direct, route_snapshot.generation) + { + let delay = cutover_stagger_delay(session_id, cutover.generation); + warn!( + user = %user, + target_mode = cutover.mode.as_str(), + cutover_generation = cutover.generation, + delay_ms = delay.as_millis() as u64, + "Cutover affected direct session, closing client connection" + ); + let _cutover_park_lease = stats.acquire_direct_cutover_park_lease(); + tokio::time::sleep(delay).await; + break Err(ProxyError::RouteSwitched); + } + tokio::select! { + result = &mut relay_result => { + break result; + } + changed = route_rx.changed() => { + if changed.is_err() { + break relay_result.await; + } + } + _ = session_cancel.cancelled() => { + break Err(ProxyError::UserDisabled { + user: user.to_string(), + }); + } + } + }; + + match &relay_result { + Ok(()) => debug!(user = %user, "Direct relay completed"), + Err(e) => debug!(user = %user, error = %e, "Direct relay ended with error"), + } + + let pool_snapshot = buffer_pool_trim.stats(); + stats.set_buffer_pool_gauges( + pool_snapshot.pooled, + pool_snapshot.allocated, + pool_snapshot.allocated.saturating_sub(pool_snapshot.pooled), + ); + + 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 +} + +fn classify_conntrack_close_reason(result: &Result<()>) -> ConntrackCloseReason { + match result { + Ok(()) => ConntrackCloseReason::NormalEof, + Err(crate::error::ProxyError::Io(error)) + if matches!(error.kind(), std::io::ErrorKind::TimedOut) => + { + ConntrackCloseReason::Timeout + } + Err(crate::error::ProxyError::Io(error)) + if matches!( + error.kind(), + std::io::ErrorKind::ConnectionReset + | std::io::ErrorKind::ConnectionAborted + | std::io::ErrorKind::BrokenPipe + | std::io::ErrorKind::NotConnected + | std::io::ErrorKind::UnexpectedEof + ) => + { + ConntrackCloseReason::Reset + } + Err(crate::error::ProxyError::Proxy(message)) + if message.contains("pressure") || message.contains("evicted") => + { + ConntrackCloseReason::Pressure + } + Err(_) => ConntrackCloseReason::Other, + } +} diff --git a/src/proxy/direct_relay/routing.rs b/src/proxy/direct_relay/routing.rs new file mode 100644 index 0000000..db6de42 --- /dev/null +++ b/src/proxy/direct_relay/routing.rs @@ -0,0 +1,123 @@ +use super::*; + +pub(super) fn get_dc_addr_static(dc_idx: i16, config: &ProxyConfig) -> Result { + let prefer_v6 = config.network.prefer == 6 && config.network.ipv6.unwrap_or(true); + let datacenters = if prefer_v6 { + &*TG_DATACENTERS_V6 + } else { + &*TG_DATACENTERS_V4 + }; + + let num_dcs = datacenters.len(); + + let dc_key = dc_idx.to_string(); + if let Some(addrs) = config.dc_overrides.get(&dc_key) { + let mut parsed = Vec::new(); + for addr_str in addrs { + match addr_str.parse::() { + Ok(addr) => parsed.push(addr), + Err(_) => { + warn!(dc_idx = dc_idx, addr_str = %addr_str, "Invalid DC override address in config, ignoring") + } + } + } + + if let Some(addr) = parsed + .iter() + .find(|a| a.is_ipv6() == prefer_v6) + .or_else(|| parsed.first()) + .copied() + { + debug!(dc_idx = dc_idx, addr = %addr, count = parsed.len(), "Using DC override from config"); + return Ok(addr); + } + } + + let abs_dc = dc_idx.unsigned_abs() as usize; + if abs_dc >= 1 && abs_dc <= num_dcs { + return Ok(SocketAddr::new(datacenters[abs_dc - 1], TG_DATACENTER_PORT)); + } + + // Unknown DC requested by client without override: log and fall back. + if !config.dc_overrides.contains_key(&dc_key) { + warn!( + dc_idx = dc_idx, + "Requested non-standard DC with no override; falling back to default cluster" + ); + if config.general.unknown_dc_file_log_enabled + && let Some(path) = &config.general.unknown_dc_log_path + && let Ok(handle) = tokio::runtime::Handle::try_current() + { + if let Some(path) = sanitize_unknown_dc_log_path(path) { + if should_log_unknown_dc(dc_idx) { + handle.spawn_blocking(move || { + if unknown_dc_log_path_is_still_safe(&path) + && let Ok(mut file) = open_unknown_dc_log_append_anchored(&path) + { + let _ = append_unknown_dc_line(&mut file, dc_idx); + } + }); + } + } else { + warn!(dc_idx = dc_idx, raw_path = %path, "Rejected unsafe unknown DC log path"); + } + } + } + + let default_dc = config.default_dc.unwrap_or(2) as usize; + let fallback_idx = if default_dc >= 1 && default_dc <= num_dcs { + default_dc - 1 + } else { + 0 + }; + + info!( + original_dc = dc_idx, + fallback_dc = (fallback_idx + 1) as u16, + fallback_addr = %datacenters[fallback_idx], + "Special DC ---> default_cluster" + ); + + Ok(SocketAddr::new( + datacenters[fallback_idx], + TG_DATACENTER_PORT, + )) +} + +pub(super) async fn do_tg_handshake_static( + mut stream: S, + success: &HandshakeSuccess, + config: &ProxyConfig, + rng: &SecureRandom, +) -> Result<(CryptoReader>, CryptoWriter>)> +where + S: AsyncRead + AsyncWrite + Unpin, +{ + let (nonce, _tg_enc_key, _tg_enc_iv, _tg_dec_key, _tg_dec_iv) = generate_tg_nonce( + success.proto_tag, + success.dc_idx, + &success.enc_key, + success.enc_iv, + rng, + config.general.fast_mode, + ); + + let (encrypted_nonce, tg_encryptor, tg_decryptor) = encrypt_tg_nonce_with_ciphers(&nonce); + + debug!( + peer = %success.peer, + nonce_head = %hex::encode(&nonce[..16]), + "Sending nonce to Telegram" + ); + + stream.write_all(&encrypted_nonce).await?; + stream.flush().await?; + + let (read_half, write_half) = split(stream); + + let max_pending = config.general.crypto_pending_buffer; + Ok(( + CryptoReader::new(read_half, tg_decryptor), + CryptoWriter::new(write_half, tg_encryptor, max_pending), + )) +} diff --git a/src/proxy/handshake/auth_probe.rs b/src/proxy/handshake/auth_probe.rs index 73e6aac..13e0428 100644 --- a/src/proxy/handshake/auth_probe.rs +++ b/src/proxy/handshake/auth_probe.rs @@ -98,7 +98,9 @@ pub(super) fn auth_probe_is_throttled_in( }; if auth_probe_state_expired(&entry, now) { drop(entry); - state.remove_if(&peer_ip, |_, current| auth_probe_state_expired(current, now)); + state.remove_if(&peer_ip, |_, current| { + auth_probe_state_expired(current, now) + }); return false; } now < entry.blocked_until @@ -116,7 +118,9 @@ pub(super) fn auth_probe_saturation_grace_exhausted_in( }; if auth_probe_state_expired(&entry, now) { drop(entry); - state.remove_if(&peer_ip, |_, current| auth_probe_state_expired(current, now)); + state.remove_if(&peer_ip, |_, current| { + auth_probe_state_expired(current, now) + }); return false; } @@ -264,7 +268,8 @@ pub(super) fn auth_probe_record_failure_with_state_in( } } - let Some((evict_key, evict_fail_streak, evict_last_seen)) = eviction_candidate else { + let Some((evict_key, evict_fail_streak, evict_last_seen)) = eviction_candidate + else { return; }; if state diff --git a/src/proxy/handshake/tls_handshake.rs b/src/proxy/handshake/tls_handshake.rs index ccc877a..c98dc17 100644 --- a/src/proxy/handshake/tls_handshake.rs +++ b/src/proxy/handshake/tls_handshake.rs @@ -266,8 +266,7 @@ where return HandshakeResult::BadClient { reader, writer }; } - let selected_tls_domain = - matched_tls_domain.unwrap_or(config.censorship.tls_domain.as_str()); + let selected_tls_domain = matched_tls_domain.unwrap_or(config.censorship.tls_domain.as_str()); let cached_entry = if config.censorship.tls_emulation { if let Some(cache) = tls_cache.as_ref() { let cached_entry = cache.get(selected_tls_domain).await; diff --git a/src/proxy/masking.rs b/src/proxy/masking.rs index cb60d21..e9c3284 100644 --- a/src/proxy/masking.rs +++ b/src/proxy/masking.rs @@ -58,1307 +58,34 @@ struct MaskTcpTarget<'a> { port: u16, } -fn mask_copy_read_len(total: usize, byte_cap: usize) -> usize { - // Keep short scanner probes on the small baseline buffer and grow only - // after the session has proven to be sustained masking relay traffic. - let active_buffer_size = if total >= MASK_BUFFER_GROW_AFTER_BYTES { - MASK_BUFFER_MAX_SIZE - } else { - MASK_BUFFER_SIZE - }; - - if byte_cap == 0 { - return active_buffer_size; - } - - let remaining_budget = byte_cap.saturating_sub(total); - if remaining_budget == 0 { - return 0; - } - - remaining_budget.min(active_buffer_size) -} - -async fn copy_with_idle_timeout( - reader: &mut R, - writer: &mut W, - byte_cap: usize, - shutdown_on_eof: bool, - idle_timeout: Duration, -) -> CopyOutcome -where - R: AsyncRead + Unpin, - W: AsyncWrite + Unpin, -{ - let mut buf = vec![0u8; MASK_BUFFER_SIZE]; - let mut total = 0usize; - let mut ended_by_eof = false; - - loop { - let read_len = mask_copy_read_len(total, byte_cap); - if read_len == 0 { - break; - } - if buf.len() < read_len { - buf.resize(read_len, 0); - } - let read_res = timeout(idle_timeout, reader.read(&mut buf[..read_len])).await; - let n = match read_res { - Ok(Ok(n)) => n, - Ok(Err(_)) | Err(_) => break, - }; - if n == 0 { - ended_by_eof = true; - if shutdown_on_eof { - let _ = timeout(idle_timeout, writer.shutdown()).await; - } - break; - } - total = total.saturating_add(n); - - let write_res = timeout(idle_timeout, writer.write_all(&buf[..n])).await; - match write_res { - Ok(Ok(())) => {} - Ok(Err(_)) | Err(_) => break, - } - } - CopyOutcome { - total, - ended_by_eof, - } -} - -fn is_http_probe(data: &[u8]) -> bool { - // RFC 7540 section 3.5: HTTP/2 client preface starts with "PRI ". - const HTTP_METHODS: [&[u8]; 10] = [ - b"GET ", b"POST", b"HEAD", b"PUT ", b"DELETE", b"OPTIONS", b"CONNECT", b"TRACE", b"PATCH", - b"PRI ", - ]; - - if data.is_empty() { - return false; - } - - let window = &data[..data.len().min(16)]; - for method in HTTP_METHODS { - if data.len() >= method.len() && window.starts_with(method) { - return true; - } - - if (2..=3).contains(&window.len()) && method.starts_with(window) { - return true; - } - } - - false -} - -fn next_mask_shape_bucket(total: usize, floor: usize, cap: usize) -> usize { - if total == 0 || floor == 0 || cap < floor { - return total; - } - - if total >= cap { - return total; - } - - let mut bucket = floor; - while bucket < total { - match bucket.checked_mul(2) { - Some(next) => bucket = next, - None => return total, - } - if bucket > cap { - return cap; - } - } - bucket -} - -async fn maybe_write_shape_padding( - mask_write: &mut W, - total_sent: usize, - enabled: bool, - floor: usize, - cap: usize, - above_cap_blur: bool, - above_cap_blur_max_bytes: usize, - aggressive_mode: bool, -) where - W: AsyncWrite + Unpin, -{ - if !enabled { - return; - } - - let target_total = if total_sent >= cap && above_cap_blur && above_cap_blur_max_bytes > 0 { - let mut rng = rand::rng(); - let extra = if aggressive_mode { - rng.random_range(1..=above_cap_blur_max_bytes) - } else { - rng.random_range(0..=above_cap_blur_max_bytes) - }; - total_sent.saturating_add(extra) - } else { - next_mask_shape_bucket(total_sent, floor, cap) - }; - - if target_total <= total_sent { - return; - } - - let mut remaining = target_total - total_sent; - let mut pad_chunk = [0u8; 1024]; - let deadline = Instant::now() + MASK_TIMEOUT; - // Use a Send RNG so relay futures remain spawn-safe under Tokio. - let mut rng = { - let mut seed_source = rand::rng(); - StdRng::from_rng(&mut seed_source) - }; - - while remaining > 0 { - let now = Instant::now(); - if now >= deadline { - return; - } - - let write_len = remaining.min(pad_chunk.len()); - rng.fill_bytes(&mut pad_chunk[..write_len]); - let write_budget = deadline.saturating_duration_since(now); - match timeout(write_budget, mask_write.write_all(&pad_chunk[..write_len])).await { - Ok(Ok(())) => {} - Ok(Err(_)) | Err(_) => return, - } - remaining -= write_len; - } - - let now = Instant::now(); - if now >= deadline { - return; - } - let flush_budget = deadline.saturating_duration_since(now); - let _ = timeout(flush_budget, mask_write.flush()).await; -} - -async fn write_proxy_header_with_timeout(mask_write: &mut W, header: &[u8]) -> bool -where - W: AsyncWrite + Unpin, -{ - match timeout(MASK_TIMEOUT, mask_write.write_all(header)).await { - Ok(Ok(())) => true, - Ok(Err(_)) => false, - Err(_) => { - debug!("Timeout writing proxy protocol header to mask backend"); - false - } - } -} - -async fn consume_client_data_with_timeout_and_cap( - reader: R, - byte_cap: usize, - relay_timeout: Duration, - idle_timeout: Duration, -) where - R: AsyncRead + Unpin, -{ - if timeout( - relay_timeout, - consume_client_data(reader, byte_cap, idle_timeout), - ) - .await - .is_err() - { - debug!("Timed out while consuming client data on masking fallback path"); - } -} - -fn mask_failure_drain_cap(config: &ProxyConfig) -> usize { - let configured_cap = config.censorship.mask_relay_max_bytes; - if configured_cap == 0 { - return MASK_BUFFER_SIZE; - } - - configured_cap.min(MASK_BUFFER_SIZE) -} - -async fn consume_mask_failure_path( - reader: R, - config: &ProxyConfig, - relay_timeout: Duration, - idle_timeout: Duration, -) where - R: AsyncRead + Unpin, -{ - consume_client_data_with_timeout_and_cap( - reader, - mask_failure_drain_cap(config), - relay_timeout, - idle_timeout, - ) - .await; -} - -async fn wait_mask_connect_budget(started: Instant) { - let elapsed = started.elapsed(); - if elapsed < MASK_TIMEOUT { - tokio::time::sleep(MASK_TIMEOUT - elapsed).await; - } -} - -// Log-normal sample bounded to [floor, ceiling]. Median = sqrt(floor * ceiling). -// Implements Box-Muller transform for standard normal sampling — no external -// dependency on rand_distr (which is incompatible with rand 0.10). -// sigma is chosen so ~99% of raw samples land inside [floor, ceiling] before clamp. -// When floor > ceiling (misconfiguration), returns ceiling (the smaller value). -// When floor == ceiling, returns that value. When both are 0, returns 0. -pub(crate) fn sample_lognormal_percentile_bounded( - floor: u64, - ceiling: u64, - rng: &mut impl Rng, -) -> u64 { - if ceiling == 0 && floor == 0 { - return 0; - } - if floor > ceiling { - return ceiling; - } - if floor == ceiling { - return floor; - } - let floor_f = floor.max(1) as f64; - let ceiling_f = ceiling.max(1) as f64; - let mu = (floor_f.ln() + ceiling_f.ln()) / 2.0; - // 4.65 ≈ 2 * 2.326 (double-sided z-score for 99th percentile) - let sigma = ((ceiling_f / floor_f).ln() / 4.65).max(0.01); - // Box-Muller transform: two uniform samples → one standard normal sample - let u1: f64 = rng.random_range(f64::MIN_POSITIVE..1.0); - let u2: f64 = rng.random_range(0.0_f64..std::f64::consts::TAU); - let normal_sample = (-2.0_f64 * u1.ln()).sqrt() * u2.cos(); - let raw = (mu + sigma * normal_sample).exp(); - if raw.is_finite() { - (raw as u64).clamp(floor, ceiling) - } else { - ((floor_f * ceiling_f).sqrt()) as u64 - } -} - -fn mask_outcome_target_budget(config: &ProxyConfig) -> Duration { - if config.censorship.mask_timing_normalization_enabled { - let floor = config.censorship.mask_timing_normalization_floor_ms; - let ceiling = config.censorship.mask_timing_normalization_ceiling_ms; - if floor == 0 { - if ceiling == 0 { - return Duration::from_millis(0); - } - // floor=0 stays uniform: log-normal cannot model distribution anchored at zero - let mut rng = rand::rng(); - return Duration::from_millis(rng.random_range(0..=ceiling)); - } - if ceiling > floor { - let mut rng = rand::rng(); - return Duration::from_millis(sample_lognormal_percentile_bounded( - floor, ceiling, &mut rng, - )); - } - // ceiling <= floor: use the larger value (fail-closed: preserve longer delay) - return Duration::from_millis(floor.max(ceiling)); - } - - MASK_TIMEOUT -} - -async fn wait_mask_connect_budget_if_needed(started: Instant, config: &ProxyConfig) { - if config.censorship.mask_timing_normalization_enabled { - return; - } - - wait_mask_connect_budget(started).await; -} - -async fn wait_mask_outcome_budget(started: Instant, config: &ProxyConfig) { - let target = mask_outcome_target_budget(config); - let elapsed = started.elapsed(); - if elapsed < target { - tokio::time::sleep(target - elapsed).await; - } -} +// Bounded relay-copy, probe classification, and failure draining. +mod copy; +// Masking delay distribution and outcome normalization. +mod timing; +// SNI-aware mask-target selection and bounded DNS resolution. +mod target; +// Local-interface snapshots and self-target rejection. +mod interfaces; +// Beobachten TTL, PROXY header, and backend socket setup. +mod backend_setup; +// Mask backend connection orchestration. +mod handler; +// Bidirectional masking relay. +mod relay; #[cfg(test)] -mod tls_domain_mask_host_tests { - use super::{ - mask_host_for_initial_data, mask_tcp_target_for_initial_data, matching_tls_domain_for_sni, - }; - use crate::config::ProxyConfig; - - fn client_hello_with_sni(sni_host: &str) -> Vec { - let mut body = Vec::new(); - body.extend_from_slice(&[0x03, 0x03]); - body.extend_from_slice(&[0u8; 32]); - body.push(32); - body.extend_from_slice(&[0x42u8; 32]); - body.extend_from_slice(&2u16.to_be_bytes()); - body.extend_from_slice(&[0x13, 0x01]); - body.push(1); - body.push(0); - - let host_bytes = sni_host.as_bytes(); - let mut sni_payload = Vec::new(); - sni_payload.extend_from_slice(&((host_bytes.len() + 3) as u16).to_be_bytes()); - sni_payload.push(0); - sni_payload.extend_from_slice(&(host_bytes.len() as u16).to_be_bytes()); - sni_payload.extend_from_slice(host_bytes); - - let mut extensions = Vec::new(); - extensions.extend_from_slice(&0x0000u16.to_be_bytes()); - extensions.extend_from_slice(&(sni_payload.len() as u16).to_be_bytes()); - extensions.extend_from_slice(&sni_payload); - body.extend_from_slice(&(extensions.len() as u16).to_be_bytes()); - body.extend_from_slice(&extensions); - - let mut handshake = Vec::new(); - handshake.push(0x01); - let body_len = (body.len() as u32).to_be_bytes(); - handshake.extend_from_slice(&body_len[1..4]); - handshake.extend_from_slice(&body); - - let mut record = Vec::new(); - record.push(0x16); - record.extend_from_slice(&[0x03, 0x01]); - record.extend_from_slice(&(handshake.len() as u16).to_be_bytes()); - record.extend_from_slice(&handshake); - record - } - - fn config_with_tls_domains() -> ProxyConfig { - let mut config = ProxyConfig::default(); - config.censorship.tls_domain = "a.com".to_string(); - config.censorship.tls_domains = vec!["b.com".to_string(), "c.com".to_string()]; - config.censorship.mask_host = None; - config - } - - #[test] - fn matching_tls_domain_accepts_primary_and_extra_domains_case_insensitively() { - let config = config_with_tls_domains(); - - assert_eq!(matching_tls_domain_for_sni(&config, "A.COM"), Some("a.com")); - assert_eq!(matching_tls_domain_for_sni(&config, "B.COM"), Some("b.com")); - assert_eq!(matching_tls_domain_for_sni(&config, "unknown.com"), None); - } - - #[test] - fn mask_host_preserves_explicit_non_primary_origin() { - let mut config = config_with_tls_domains(); - config.censorship.mask_host = Some("origin.example".to_string()); - - let initial_data = client_hello_with_sni("b.com"); - - assert_eq!( - mask_host_for_initial_data(&config, &initial_data), - "origin.example" - ); - } - - #[test] - fn mask_host_uses_matching_tls_domain_when_mask_host_is_primary_default() { - let config = config_with_tls_domains(); - let initial_data = client_hello_with_sni("b.com"); - - assert_eq!(mask_host_for_initial_data(&config, &initial_data), "b.com"); - } - - #[test] - fn mask_host_uses_primary_domain_when_dynamic_masking_is_disabled() { - let mut config = config_with_tls_domains(); - config.censorship.mask_dynamic = false; - let initial_data = client_hello_with_sni("b.com"); - - assert_eq!(mask_host_for_initial_data(&config, &initial_data), "a.com"); - } - - #[test] - fn exclusive_mask_target_overrides_only_matching_sni() { - let mut config = config_with_tls_domains(); - config - .censorship - .exclusive_mask - .insert("b.com".to_string(), "origin-b.example:8443".to_string()); - let b_initial_data = client_hello_with_sni("B.COM"); - let c_initial_data = client_hello_with_sni("c.com"); - - let b_target = mask_tcp_target_for_initial_data(&config, &b_initial_data); - let c_target = mask_tcp_target_for_initial_data(&config, &c_initial_data); - - assert_eq!(b_target.host, "origin-b.example"); - assert_eq!(b_target.port, 8443); - assert_eq!(c_target.host, "c.com"); - assert_eq!(c_target.port, config.censorship.mask_port); - } -} - -/// Detect client type based on initial data -fn detect_client_type(data: &[u8]) -> &'static str { - // Check for HTTP request - if is_http_probe(data) { - return "HTTP"; - } - - // Check for TLS ClientHello (0x16 = handshake, 0x03 0x01-0x03 = TLS version) - if data.len() > 3 && data[0] == 0x16 && data[1] == 0x03 { - return "TLS-scanner"; - } - - // Check for SSH - if data.starts_with(b"SSH-") { - return "SSH"; - } - - // Port scanner (very short data) - if data.len() < 10 { - return "port-scanner"; - } - - "unknown" -} - -fn parse_mask_host_ip_literal(host: &str) -> Option { - if host.starts_with('[') && host.ends_with(']') { - return host[1..host.len() - 1].parse::().ok(); - } - host.parse::().ok() -} - -async fn resolve_mask_target_addrs( - mask_host: &str, - mask_port: u16, - upstream_manager: Option<&crate::transport::UpstreamManager>, -) -> std::io::Result> { - if let Some(ip) = parse_mask_host_ip_literal(mask_host) { - return Ok(vec![SocketAddr::new(ip, mask_port)]); - } - - if let Some(upstream_manager) = upstream_manager { - return upstream_manager - .resolve_all(mask_host, mask_port) - .await - .map_err(|error| IoError::new(ErrorKind::NotFound, error.to_string())); - } - - let addrs = timeout(MASK_TIMEOUT, lookup_host((mask_host, mask_port))) - .await - .map_err(|_| IoError::new(ErrorKind::TimedOut, "mask target DNS lookup timed out"))??; - let addrs = addrs - .take(MASK_DNS_RESULT_MAX_ADDRESSES) - .collect::>(); - if addrs.is_empty() { - return Err(IoError::new( - ErrorKind::NotFound, - "mask target DNS lookup returned no addresses", - )); - } - - Ok(addrs) -} - -fn matching_tls_domain_for_sni<'a>(config: &'a ProxyConfig, sni: &str) -> Option<&'a str> { - if config.censorship.tls_domain.eq_ignore_ascii_case(sni) { - return Some(config.censorship.tls_domain.as_str()); - } - - for domain in &config.censorship.tls_domains { - if domain.eq_ignore_ascii_case(sni) { - return Some(domain.as_str()); - } - } - - None -} - -fn parse_exclusive_mask_target(target: &str) -> Option> { - let target = target.trim(); - if target.is_empty() { - return None; - } - - if target.starts_with('[') { - let end = target.find(']')?; - if target.get(end + 1..end + 2)? != ":" { - return None; - } - let port = target[end + 2..].parse::().ok()?; - return (port > 0).then_some(MaskTcpTarget { - host: &target[..=end], - port, - }); - } - - let (host, port) = target.rsplit_once(':')?; - if host.is_empty() || host.contains(':') { - return None; - } - let port = port.parse::().ok()?; - (port > 0).then_some(MaskTcpTarget { host, port }) -} - -fn exclusive_mask_target_for_sni<'a>( - config: &'a ProxyConfig, - sni: &str, -) -> Option> { - if let Some(target) = config.censorship.exclusive_mask_targets.get(sni) { - return Some(MaskTcpTarget { - host: target.host.as_str(), - port: target.port, - }); - } - if let Some(target) = config.censorship.exclusive_mask.get(sni) { - return parse_exclusive_mask_target(target); - } - - if sni.bytes().any(|byte| byte.is_ascii_uppercase()) { - let normalized_sni = sni.to_ascii_lowercase(); - if let Some(target) = config - .censorship - .exclusive_mask_targets - .get(&normalized_sni) - { - return Some(MaskTcpTarget { - host: target.host.as_str(), - port: target.port, - }); - } - if let Some(target) = config.censorship.exclusive_mask.get(&normalized_sni) { - return parse_exclusive_mask_target(target); - } - } - - None -} - -#[cfg(test)] -fn mask_host_for_initial_data<'a>(config: &'a ProxyConfig, initial_data: &[u8]) -> &'a str { - mask_tcp_target_for_initial_data(config, initial_data).host -} - -#[cfg(test)] -fn mask_tcp_target_for_initial_data<'a>( - config: &'a ProxyConfig, - initial_data: &[u8], -) -> MaskTcpTarget<'a> { - let sni = tls::extract_sni_from_client_hello(initial_data); - if let Some(target) = sni - .as_deref() - .and_then(|sni| exclusive_mask_target_for_sni(config, sni)) - { - return target; - } - - default_mask_tcp_target_for_initial_data(config, initial_data, sni.as_deref()) -} - -fn default_mask_tcp_target_for_initial_data<'a>( - config: &'a ProxyConfig, - initial_data: &[u8], - sni: Option<&str>, -) -> MaskTcpTarget<'a> { - let configured_mask_host = config - .censorship - .mask_host - .as_deref() - .unwrap_or(&config.censorship.tls_domain); - - if config.censorship.mask_host.is_none() && config.censorship.mask_dynamic { - let extracted_sni = if sni.is_none() { - tls::extract_sni_from_client_hello(initial_data) - } else { - None - }; - if let Some(host) = sni - .or(extracted_sni.as_deref()) - .and_then(|sni| matching_tls_domain_for_sni(config, sni)) - { - return MaskTcpTarget { - host, - port: config.censorship.mask_port, - }; - } - } - - if let Some(mask_host) = config.censorship.mask_host.as_deref() { - return MaskTcpTarget { - host: mask_host, - port: config.censorship.mask_port, - }; - } - - MaskTcpTarget { - host: configured_mask_host, - port: config.censorship.mask_port, - } -} - -fn canonical_ip(ip: IpAddr) -> IpAddr { - match ip { - IpAddr::V6(v6) => v6 - .to_ipv4_mapped() - .map(IpAddr::V4) - .unwrap_or(IpAddr::V6(v6)), - IpAddr::V4(v4) => IpAddr::V4(v4), - } -} - -#[cfg(unix)] -fn collect_local_interface_ips() -> Vec { - #[cfg(test)] - LOCAL_INTERFACE_ENUMERATIONS.fetch_add(1, Ordering::Relaxed); - - let mut out = Vec::new(); - if let Ok(addrs) = getifaddrs() { - for iface in addrs { - if let Some(address) = iface.address { - if let Some(v4) = address.as_sockaddr_in() { - out.push(canonical_ip(IpAddr::V4(v4.ip()))); - } else if let Some(v6) = address.as_sockaddr_in6() { - out.push(canonical_ip(IpAddr::V6(v6.ip()))); - } - } - } - } - out -} - -fn choose_interface_snapshot(previous: &[IpAddr], refreshed: Vec) -> Vec { - if refreshed.is_empty() && !previous.is_empty() { - return previous.to_vec(); - } - - refreshed -} - -#[cfg(unix)] -#[derive(Default)] -struct LocalInterfaceCache { - ips: Vec, - refreshed_at: Option, -} - -#[cfg(unix)] -static LOCAL_INTERFACE_CACHE: OnceLock> = OnceLock::new(); - -#[cfg(unix)] -static LOCAL_INTERFACE_REFRESH_LOCK: OnceLock> = OnceLock::new(); - -#[cfg(all(unix, test))] -fn local_interface_ips() -> Vec { - let cache = LOCAL_INTERFACE_CACHE.get_or_init(|| Mutex::new(LocalInterfaceCache::default())); - let mut guard = cache.lock().unwrap_or_else(|poison| poison.into_inner()); - - let stale = guard - .refreshed_at - .is_none_or(|at| at.elapsed() >= LOCAL_INTERFACE_CACHE_TTL); - if stale { - let refreshed = collect_local_interface_ips(); - guard.ips = choose_interface_snapshot(&guard.ips, refreshed); - guard.refreshed_at = Some(StdInstant::now()); - } - - guard.ips.clone() -} - -#[cfg(unix)] -async fn local_interface_ips_async() -> Vec { - let cache = LOCAL_INTERFACE_CACHE.get_or_init(|| Mutex::new(LocalInterfaceCache::default())); - - { - let guard = cache.lock().unwrap_or_else(|poison| poison.into_inner()); - let stale = guard - .refreshed_at - .is_none_or(|at| at.elapsed() >= LOCAL_INTERFACE_CACHE_TTL); - if !stale { - return guard.ips.clone(); - } - } - - let refresh_lock = LOCAL_INTERFACE_REFRESH_LOCK.get_or_init(|| AsyncMutex::new(())); - let _refresh_guard = refresh_lock.lock().await; - - { - let guard = cache.lock().unwrap_or_else(|poison| poison.into_inner()); - let stale = guard - .refreshed_at - .is_none_or(|at| at.elapsed() >= LOCAL_INTERFACE_CACHE_TTL); - if !stale { - return guard.ips.clone(); - } - } - - let refreshed = tokio::task::spawn_blocking(collect_local_interface_ips) - .await - .unwrap_or_default(); - - let mut guard = cache.lock().unwrap_or_else(|poison| poison.into_inner()); - let stale = guard - .refreshed_at - .is_none_or(|at| at.elapsed() >= LOCAL_INTERFACE_CACHE_TTL); - if stale { - guard.ips = choose_interface_snapshot(&guard.ips, refreshed); - guard.refreshed_at = Some(StdInstant::now()); - } - - guard.ips.clone() -} - -#[cfg(all(not(unix), test))] -fn local_interface_ips() -> Vec { - Vec::new() -} - -#[cfg(not(unix))] -async fn local_interface_ips_async() -> Vec { - Vec::new() -} - -#[cfg(test)] -static LOCAL_INTERFACE_ENUMERATIONS: AtomicUsize = AtomicUsize::new(0); - -#[cfg(test)] -fn reset_local_interface_enumerations_for_tests() { - LOCAL_INTERFACE_ENUMERATIONS.store(0, Ordering::Relaxed); - - #[cfg(unix)] - if let Some(cache) = LOCAL_INTERFACE_CACHE.get() { - let mut guard = cache.lock().unwrap_or_else(|poison| poison.into_inner()); - guard.ips.clear(); - guard.refreshed_at = None; - } -} - -#[cfg(test)] -fn local_interface_enumerations_for_tests() -> usize { - LOCAL_INTERFACE_ENUMERATIONS.load(Ordering::Relaxed) -} - -fn is_mask_target_local_listener_with_interfaces( - mask_host: &str, - mask_port: u16, - local_addr: SocketAddr, - resolved_addrs: &[SocketAddr], - interface_ips: &[IpAddr], -) -> bool { - if mask_port != local_addr.port() { - return false; - } - - let local_ip = canonical_ip(local_addr.ip()); - let literal_mask_ip = parse_mask_host_ip_literal(mask_host).map(canonical_ip); - - for addr in resolved_addrs { - let resolved_ip = canonical_ip(addr.ip()); - if resolved_ip == local_ip { - return true; - } - - if local_ip.is_unspecified() - && (resolved_ip.is_loopback() - || resolved_ip.is_unspecified() - || interface_ips.contains(&resolved_ip)) - { - return true; - } - } - - if let Some(mask_ip) = literal_mask_ip { - if mask_ip == local_ip { - return true; - } - - if local_ip.is_unspecified() - && (mask_ip.is_loopback() - || mask_ip.is_unspecified() - || interface_ips.contains(&mask_ip)) - { - return true; - } - } - - false -} - -#[cfg(test)] -fn is_mask_target_local_listener( - mask_host: &str, - mask_port: u16, - local_addr: SocketAddr, - resolved_addrs: &[SocketAddr], -) -> bool { - if mask_port != local_addr.port() { - return false; - } - - let interfaces = local_interface_ips(); - is_mask_target_local_listener_with_interfaces( - mask_host, - mask_port, - local_addr, - resolved_addrs, - &interfaces, - ) -} - -async fn is_mask_target_local_listener_async( - mask_host: &str, - mask_port: u16, - local_addr: SocketAddr, - resolved_addrs: &[SocketAddr], -) -> bool { - if mask_port != local_addr.port() { - return false; - } - - let interfaces = local_interface_ips_async().await; - is_mask_target_local_listener_with_interfaces( - mask_host, - mask_port, - local_addr, - resolved_addrs, - &interfaces, - ) -} - -fn masking_beobachten_ttl(config: &ProxyConfig) -> Duration { - let minutes = config.general.beobachten_minutes; - let clamped = minutes.clamp(1, 24 * 60); - Duration::from_secs(clamped.saturating_mul(60)) -} - -fn build_mask_proxy_header( - version: u8, - peer: SocketAddr, - local_addr: SocketAddr, -) -> Option> { - match version { - 0 => None, - 2 => Some( - ProxyProtocolV2Builder::new() - .with_addrs(peer, local_addr) - .build(), - ), - _ => { - let header = match (peer, local_addr) { - (SocketAddr::V4(src), SocketAddr::V4(dst)) => ProxyProtocolV1Builder::new() - .tcp4(src.into(), dst.into()) - .build(), - (SocketAddr::V6(src), SocketAddr::V6(dst)) => ProxyProtocolV1Builder::new() - .tcp6(src.into(), dst.into()) - .build(), - _ => ProxyProtocolV1Builder::new().build(), - }; - Some(header) - } - } -} - -fn configure_mask_backend_socket(stream: &TcpStream) { - if let Err(e) = configure_tcp_socket(stream, false, Duration::from_secs(0)) { - debug!(error = %e, "Failed to configure mask backend socket"); - } -} - -/// Handles a bad client by forwarding it to the configured mask target. -#[cfg(test)] -pub async fn handle_bad_client( - reader: R, - writer: W, - initial_data: &[u8], - peer: SocketAddr, - local_addr: SocketAddr, - config: &ProxyConfig, - beobachten: &BeobachtenStore, -) where - R: AsyncRead + Unpin + Send + 'static, - W: AsyncWrite + Unpin + Send + 'static, -{ - let shared = ProxySharedState::new(); - handle_bad_client_with_shared( - reader, - writer, - initial_data, - peer, - local_addr, - config, - beobachten, - shared.as_ref(), - ) - .await; -} - -/// Handles a bad client with shared pre-auth fallback admission state. -pub(crate) async fn handle_bad_client_with_shared( - reader: R, - writer: W, - initial_data: &[u8], - peer: SocketAddr, - local_addr: SocketAddr, - config: &ProxyConfig, - beobachten: &BeobachtenStore, - shared: &ProxySharedState, -) where - R: AsyncRead + Unpin + Send + 'static, - W: AsyncWrite + Unpin + Send + 'static, -{ - handle_bad_client_with_shared_resolver( - reader, - writer, - initial_data, - peer, - local_addr, - config, - beobachten, - shared, - None, - ) - .await; -} - -pub(super) async fn handle_bad_client_with_shared_resolver( - reader: R, - writer: W, - initial_data: &[u8], - peer: SocketAddr, - local_addr: SocketAddr, - config: &ProxyConfig, - beobachten: &BeobachtenStore, - shared: &ProxySharedState, - upstream_manager: Option<&crate::transport::UpstreamManager>, -) where - R: AsyncRead + Unpin + Send + 'static, - W: AsyncWrite + Unpin + Send + 'static, -{ - let client_type = detect_client_type(initial_data); - if config.general.beobachten { - let ttl = masking_beobachten_ttl(config); - beobachten.record(client_type, peer.ip(), ttl); - } - - let relay_timeout = Duration::from_millis(config.censorship.mask_relay_timeout_ms); - let idle_timeout = Duration::from_millis(config.censorship.mask_relay_idle_timeout_ms); - - if !config.censorship.mask { - // Masking disabled, just consume data - consume_client_data_with_timeout_and_cap( - reader, - config.censorship.mask_relay_max_bytes, - relay_timeout, - idle_timeout, - ) - .await; - return; - } - - let Some(_masking_permit) = shared.try_acquire_masking_fallback_permit() else { - let outcome_started = Instant::now(); - debug!( - client_type = client_type, - "Masking fallback concurrency limit reached" - ); - consume_mask_failure_path(reader, config, relay_timeout, idle_timeout).await; - wait_mask_outcome_budget(outcome_started, config).await; - return; - }; - - let client_sni = tls::extract_sni_from_client_hello(initial_data); - let exclusive_tcp_target = client_sni - .as_deref() - .and_then(|sni| exclusive_mask_target_for_sni(config, sni)); - - // Connect via Unix socket or TCP - #[cfg(unix)] - if exclusive_tcp_target.is_none() - && let Some(ref sock_path) = config.censorship.mask_unix_sock - { - let outcome_started = Instant::now(); - let connect_started = Instant::now(); - debug!( - client_type = client_type, - sock = %sock_path, - data_len = initial_data.len(), - "Forwarding bad client to mask unix socket" - ); - - let connect_result = timeout(MASK_TIMEOUT, UnixStream::connect(sock_path)).await; - match connect_result { - Ok(Ok(stream)) => { - let (mask_read, mut mask_write) = stream.into_split(); - let proxy_header = build_mask_proxy_header( - config.censorship.mask_proxy_protocol, - peer, - local_addr, - ); - if let Some(header) = proxy_header - && !write_proxy_header_with_timeout(&mut mask_write, &header).await - { - wait_mask_outcome_budget(outcome_started, config).await; - return; - } - if timeout( - relay_timeout, - relay_to_mask( - reader, - writer, - mask_read, - mask_write, - initial_data, - config.censorship.mask_shape_hardening, - config.censorship.mask_shape_bucket_floor_bytes, - config.censorship.mask_shape_bucket_cap_bytes, - config.censorship.mask_shape_above_cap_blur, - config.censorship.mask_shape_above_cap_blur_max_bytes, - config.censorship.mask_shape_hardening_aggressive_mode, - config.censorship.mask_relay_max_bytes, - idle_timeout, - ), - ) - .await - .is_err() - { - debug!("Mask relay timed out (unix socket)"); - } - wait_mask_outcome_budget(outcome_started, config).await; - } - Ok(Err(e)) => { - wait_mask_connect_budget_if_needed(connect_started, config).await; - debug!(error = %e, "Failed to connect to mask unix socket"); - consume_mask_failure_path(reader, config, relay_timeout, idle_timeout).await; - wait_mask_outcome_budget(outcome_started, config).await; - } - Err(_) => { - debug!("Timeout connecting to mask unix socket"); - consume_mask_failure_path(reader, config, relay_timeout, idle_timeout).await; - wait_mask_outcome_budget(outcome_started, config).await; - } - } - return; - } - - let mask_target = exclusive_tcp_target.unwrap_or_else(|| { - default_mask_tcp_target_for_initial_data(config, initial_data, client_sni.as_deref()) - }); - let mask_host = mask_target.host; - let mask_port = mask_target.port; - - let resolved_mask_addrs = - match resolve_mask_target_addrs(mask_host, mask_port, upstream_manager).await { - Ok(addrs) => addrs, - Err(e) => { - let outcome_started = Instant::now(); - debug!( - client_type = client_type, - host = %mask_host, - port = mask_port, - error = %e, - "Failed to resolve mask target" - ); - consume_mask_failure_path(reader, config, relay_timeout, idle_timeout).await; - wait_mask_outcome_budget(outcome_started, config).await; - return; - } - }; - - // Fail closed when fallback points at our own listener endpoint. - // Self-referential masking can create recursive proxy loops under - // misconfiguration and leak distinguishable load spikes to adversaries. - if is_mask_target_local_listener_async(mask_host, mask_port, local_addr, &resolved_mask_addrs) - .await - { - let outcome_started = Instant::now(); - debug!( - client_type = client_type, - host = %mask_host, - port = mask_port, - local = %local_addr, - "Mask target resolves to local listener; refusing self-referential masking fallback" - ); - consume_mask_failure_path(reader, config, relay_timeout, idle_timeout).await; - wait_mask_outcome_budget(outcome_started, config).await; - return; - } - - let outcome_started = Instant::now(); - - debug!( - client_type = client_type, - host = %mask_host, - port = mask_port, - data_len = initial_data.len(), - "Forwarding bad client to mask host" - ); - - let connect_started = Instant::now(); - let connect_result = timeout( - MASK_TIMEOUT, - TcpStream::connect(resolved_mask_addrs.as_slice()), - ) - .await; - match connect_result { - Ok(Ok(stream)) => { - configure_mask_backend_socket(&stream); - let proxy_header = - build_mask_proxy_header(config.censorship.mask_proxy_protocol, peer, local_addr); - - let (mask_read, mut mask_write) = stream.into_split(); - if let Some(header) = proxy_header - && !write_proxy_header_with_timeout(&mut mask_write, &header).await - { - wait_mask_outcome_budget(outcome_started, config).await; - return; - } - if timeout( - relay_timeout, - relay_to_mask( - reader, - writer, - mask_read, - mask_write, - initial_data, - config.censorship.mask_shape_hardening, - config.censorship.mask_shape_bucket_floor_bytes, - config.censorship.mask_shape_bucket_cap_bytes, - config.censorship.mask_shape_above_cap_blur, - config.censorship.mask_shape_above_cap_blur_max_bytes, - config.censorship.mask_shape_hardening_aggressive_mode, - config.censorship.mask_relay_max_bytes, - idle_timeout, - ), - ) - .await - .is_err() - { - debug!("Mask relay timed out"); - } - wait_mask_outcome_budget(outcome_started, config).await; - } - Ok(Err(e)) => { - wait_mask_connect_budget_if_needed(connect_started, config).await; - debug!(error = %e, "Failed to connect to mask host"); - consume_mask_failure_path(reader, config, relay_timeout, idle_timeout).await; - wait_mask_outcome_budget(outcome_started, config).await; - } - Err(_) => { - debug!("Timeout connecting to mask host"); - consume_mask_failure_path(reader, config, relay_timeout, idle_timeout).await; - wait_mask_outcome_budget(outcome_started, config).await; - } - } -} - -/// Relay traffic between client and mask backend -async fn relay_to_mask( - mut reader: R, - mut writer: W, - mut mask_read: MR, - mut mask_write: MW, - initial_data: &[u8], - shape_hardening_enabled: bool, - shape_bucket_floor_bytes: usize, - shape_bucket_cap_bytes: usize, - shape_above_cap_blur: bool, - shape_above_cap_blur_max_bytes: usize, - shape_hardening_aggressive_mode: bool, - mask_relay_max_bytes: usize, - idle_timeout: Duration, -) where - R: AsyncRead + Unpin + Send + 'static, - W: AsyncWrite + Unpin + Send + 'static, - MR: AsyncRead + Unpin + Send + 'static, - MW: AsyncWrite + Unpin + Send + 'static, -{ - // Send initial data to mask host - if mask_write.write_all(initial_data).await.is_err() { - return; - } - if mask_write.flush().await.is_err() { - return; - } - - let (upstream_copy, downstream_copy) = tokio::join!( - async { - copy_with_idle_timeout( - &mut reader, - &mut mask_write, - mask_relay_max_bytes, - !shape_hardening_enabled, - idle_timeout, - ) - .await - }, - async { - copy_with_idle_timeout( - &mut mask_read, - &mut writer, - mask_relay_max_bytes, - true, - idle_timeout, - ) - .await - } - ); - - let total_sent = initial_data.len().saturating_add(upstream_copy.total); - - let should_shape = shape_hardening_enabled - && !initial_data.is_empty() - && (upstream_copy.ended_by_eof - || (shape_hardening_aggressive_mode && downstream_copy.total == 0)); - - maybe_write_shape_padding( - &mut mask_write, - total_sent, - should_shape, - shape_bucket_floor_bytes, - shape_bucket_cap_bytes, - shape_above_cap_blur, - shape_above_cap_blur_max_bytes, - shape_hardening_aggressive_mode, - ) - .await; - - let _ = mask_write.shutdown().await; - let _ = writer.shutdown().await; -} - -/// Just consume all data from client without responding. -async fn consume_client_data( - mut reader: R, - byte_cap: usize, - idle_timeout: Duration, -) { - // Keep drain path fail-closed under slow-loris stalls. - let mut buf = vec![0u8; MASK_BUFFER_SIZE]; - let mut total = 0usize; - - loop { - let read_len = mask_copy_read_len(total, byte_cap); - if read_len == 0 { - break; - } - if buf.len() < read_len { - buf.resize(read_len, 0); - } - let n = match timeout(idle_timeout, reader.read(&mut buf[..read_len])).await { - Ok(Ok(n)) => n, - Ok(Err(_)) | Err(_) => break, - }; - - if n == 0 { - break; - } - - total = total.saturating_add(n); - if byte_cap != 0 && total >= byte_cap { - break; - } - } -} +pub use backend_setup::handle_bad_client; +use backend_setup::*; +use copy::*; +pub(crate) use handler::handle_bad_client_with_shared; +pub(in crate::proxy) use handler::handle_bad_client_with_shared_resolver; +use interfaces::*; +use relay::*; +use target::*; +pub(crate) use timing::sample_lognormal_percentile_bounded; +use timing::{ + mask_outcome_target_budget, wait_mask_connect_budget_if_needed, wait_mask_outcome_budget, +}; #[cfg(test)] #[path = "tests/masking_security_tests.rs"] diff --git a/src/proxy/masking/backend_setup.rs b/src/proxy/masking/backend_setup.rs new file mode 100644 index 0000000..1476788 --- /dev/null +++ b/src/proxy/masking/backend_setup.rs @@ -0,0 +1,68 @@ +use super::*; + +pub(super) fn masking_beobachten_ttl(config: &ProxyConfig) -> Duration { + let minutes = config.general.beobachten_minutes; + let clamped = minutes.clamp(1, 24 * 60); + Duration::from_secs(clamped.saturating_mul(60)) +} + +pub(super) fn build_mask_proxy_header( + version: u8, + peer: SocketAddr, + local_addr: SocketAddr, +) -> Option> { + match version { + 0 => None, + 2 => Some( + ProxyProtocolV2Builder::new() + .with_addrs(peer, local_addr) + .build(), + ), + _ => { + let header = match (peer, local_addr) { + (SocketAddr::V4(src), SocketAddr::V4(dst)) => ProxyProtocolV1Builder::new() + .tcp4(src.into(), dst.into()) + .build(), + (SocketAddr::V6(src), SocketAddr::V6(dst)) => ProxyProtocolV1Builder::new() + .tcp6(src.into(), dst.into()) + .build(), + _ => ProxyProtocolV1Builder::new().build(), + }; + Some(header) + } + } +} + +pub(super) fn configure_mask_backend_socket(stream: &TcpStream) { + if let Err(e) = configure_tcp_socket(stream, false, Duration::from_secs(0)) { + debug!(error = %e, "Failed to configure mask backend socket"); + } +} + +/// Handles a bad client by forwarding it to the configured mask target. +#[cfg(test)] +pub async fn handle_bad_client( + reader: R, + writer: W, + initial_data: &[u8], + peer: SocketAddr, + local_addr: SocketAddr, + config: &ProxyConfig, + beobachten: &BeobachtenStore, +) where + R: AsyncRead + Unpin + Send + 'static, + W: AsyncWrite + Unpin + Send + 'static, +{ + let shared = ProxySharedState::new(); + handle_bad_client_with_shared( + reader, + writer, + initial_data, + peer, + local_addr, + config, + beobachten, + shared.as_ref(), + ) + .await; +} diff --git a/src/proxy/masking/copy.rs b/src/proxy/masking/copy.rs new file mode 100644 index 0000000..d33f12e --- /dev/null +++ b/src/proxy/masking/copy.rs @@ -0,0 +1,256 @@ +use super::*; + +pub(super) fn mask_copy_read_len(total: usize, byte_cap: usize) -> usize { + // Keep short scanner probes on the small baseline buffer and grow only + // after the session has proven to be sustained masking relay traffic. + let active_buffer_size = if total >= MASK_BUFFER_GROW_AFTER_BYTES { + MASK_BUFFER_MAX_SIZE + } else { + MASK_BUFFER_SIZE + }; + + if byte_cap == 0 { + return active_buffer_size; + } + + let remaining_budget = byte_cap.saturating_sub(total); + if remaining_budget == 0 { + return 0; + } + + remaining_budget.min(active_buffer_size) +} + +pub(super) async fn copy_with_idle_timeout( + reader: &mut R, + writer: &mut W, + byte_cap: usize, + shutdown_on_eof: bool, + idle_timeout: Duration, +) -> CopyOutcome +where + R: AsyncRead + Unpin, + W: AsyncWrite + Unpin, +{ + let mut buf = vec![0u8; MASK_BUFFER_SIZE]; + let mut total = 0usize; + let mut ended_by_eof = false; + + loop { + let read_len = mask_copy_read_len(total, byte_cap); + if read_len == 0 { + break; + } + if buf.len() < read_len { + buf.resize(read_len, 0); + } + let read_res = timeout(idle_timeout, reader.read(&mut buf[..read_len])).await; + let n = match read_res { + Ok(Ok(n)) => n, + Ok(Err(_)) | Err(_) => break, + }; + if n == 0 { + ended_by_eof = true; + if shutdown_on_eof { + let _ = timeout(idle_timeout, writer.shutdown()).await; + } + break; + } + total = total.saturating_add(n); + + let write_res = timeout(idle_timeout, writer.write_all(&buf[..n])).await; + match write_res { + Ok(Ok(())) => {} + Ok(Err(_)) | Err(_) => break, + } + } + CopyOutcome { + total, + ended_by_eof, + } +} + +pub(super) fn is_http_probe(data: &[u8]) -> bool { + // RFC 7540 section 3.5: HTTP/2 client preface starts with "PRI ". + const HTTP_METHODS: [&[u8]; 10] = [ + b"GET ", b"POST", b"HEAD", b"PUT ", b"DELETE", b"OPTIONS", b"CONNECT", b"TRACE", b"PATCH", + b"PRI ", + ]; + + if data.is_empty() { + return false; + } + + let window = &data[..data.len().min(16)]; + for method in HTTP_METHODS { + if data.len() >= method.len() && window.starts_with(method) { + return true; + } + + if (2..=3).contains(&window.len()) && method.starts_with(window) { + return true; + } + } + + false +} + +pub(super) fn next_mask_shape_bucket(total: usize, floor: usize, cap: usize) -> usize { + if total == 0 || floor == 0 || cap < floor { + return total; + } + + if total >= cap { + return total; + } + + let mut bucket = floor; + while bucket < total { + match bucket.checked_mul(2) { + Some(next) => bucket = next, + None => return total, + } + if bucket > cap { + return cap; + } + } + bucket +} + +pub(super) async fn maybe_write_shape_padding( + mask_write: &mut W, + total_sent: usize, + enabled: bool, + floor: usize, + cap: usize, + above_cap_blur: bool, + above_cap_blur_max_bytes: usize, + aggressive_mode: bool, +) where + W: AsyncWrite + Unpin, +{ + if !enabled { + return; + } + + let target_total = if total_sent >= cap && above_cap_blur && above_cap_blur_max_bytes > 0 { + let mut rng = rand::rng(); + let extra = if aggressive_mode { + rng.random_range(1..=above_cap_blur_max_bytes) + } else { + rng.random_range(0..=above_cap_blur_max_bytes) + }; + total_sent.saturating_add(extra) + } else { + next_mask_shape_bucket(total_sent, floor, cap) + }; + + if target_total <= total_sent { + return; + } + + let mut remaining = target_total - total_sent; + let mut pad_chunk = [0u8; 1024]; + let deadline = Instant::now() + MASK_TIMEOUT; + // Use a Send RNG so relay futures remain spawn-safe under Tokio. + let mut rng = { + let mut seed_source = rand::rng(); + StdRng::from_rng(&mut seed_source) + }; + + while remaining > 0 { + let now = Instant::now(); + if now >= deadline { + return; + } + + let write_len = remaining.min(pad_chunk.len()); + rng.fill_bytes(&mut pad_chunk[..write_len]); + let write_budget = deadline.saturating_duration_since(now); + match timeout(write_budget, mask_write.write_all(&pad_chunk[..write_len])).await { + Ok(Ok(())) => {} + Ok(Err(_)) | Err(_) => return, + } + remaining -= write_len; + } + + let now = Instant::now(); + if now >= deadline { + return; + } + let flush_budget = deadline.saturating_duration_since(now); + let _ = timeout(flush_budget, mask_write.flush()).await; +} + +pub(super) async fn write_proxy_header_with_timeout(mask_write: &mut W, header: &[u8]) -> bool +where + W: AsyncWrite + Unpin, +{ + match timeout(MASK_TIMEOUT, mask_write.write_all(header)).await { + Ok(Ok(())) => true, + Ok(Err(_)) => false, + Err(_) => { + debug!("Timeout writing proxy protocol header to mask backend"); + false + } + } +} + +pub(super) async fn consume_client_data_with_timeout_and_cap( + reader: R, + byte_cap: usize, + relay_timeout: Duration, + idle_timeout: Duration, +) where + R: AsyncRead + Unpin, +{ + if timeout( + relay_timeout, + consume_client_data(reader, byte_cap, idle_timeout), + ) + .await + .is_err() + { + debug!("Timed out while consuming client data on masking fallback path"); + } +} + +pub(super) fn mask_failure_drain_cap(config: &ProxyConfig) -> usize { + let configured_cap = config.censorship.mask_relay_max_bytes; + if configured_cap == 0 { + return MASK_BUFFER_SIZE; + } + + configured_cap.min(MASK_BUFFER_SIZE) +} + +pub(super) async fn consume_mask_failure_path( + reader: R, + config: &ProxyConfig, + relay_timeout: Duration, + idle_timeout: Duration, +) where + R: AsyncRead + Unpin, +{ + consume_client_data_with_timeout_and_cap( + reader, + mask_failure_drain_cap(config), + relay_timeout, + idle_timeout, + ) + .await; +} + +pub(super) async fn wait_mask_connect_budget(started: Instant) { + let elapsed = started.elapsed(); + if elapsed < MASK_TIMEOUT { + tokio::time::sleep(MASK_TIMEOUT - elapsed).await; + } +} + +// Log-normal sample bounded to [floor, ceiling]. Median = sqrt(floor * ceiling). +// Implements Box-Muller transform for standard normal sampling — no external +// dependency on rand_distr (which is incompatible with rand 0.10). +// sigma is chosen so ~99% of raw samples land inside [floor, ceiling] before clamp. +// When floor > ceiling (misconfiguration), returns ceiling (the smaller value). +// When floor == ceiling, returns that value. When both are 0, returns 0. diff --git a/src/proxy/masking/handler.rs b/src/proxy/masking/handler.rs new file mode 100644 index 0000000..1c1e3a6 --- /dev/null +++ b/src/proxy/masking/handler.rs @@ -0,0 +1,260 @@ +use super::*; + +/// Handles a bad client with shared pre-auth fallback admission state. +pub(crate) async fn handle_bad_client_with_shared( + reader: R, + writer: W, + initial_data: &[u8], + peer: SocketAddr, + local_addr: SocketAddr, + config: &ProxyConfig, + beobachten: &BeobachtenStore, + shared: &ProxySharedState, +) where + R: AsyncRead + Unpin + Send + 'static, + W: AsyncWrite + Unpin + Send + 'static, +{ + handle_bad_client_with_shared_resolver( + reader, + writer, + initial_data, + peer, + local_addr, + config, + beobachten, + shared, + None, + ) + .await; +} + +pub(in crate::proxy) async fn handle_bad_client_with_shared_resolver( + reader: R, + writer: W, + initial_data: &[u8], + peer: SocketAddr, + local_addr: SocketAddr, + config: &ProxyConfig, + beobachten: &BeobachtenStore, + shared: &ProxySharedState, + upstream_manager: Option<&crate::transport::UpstreamManager>, +) where + R: AsyncRead + Unpin + Send + 'static, + W: AsyncWrite + Unpin + Send + 'static, +{ + let client_type = detect_client_type(initial_data); + if config.general.beobachten { + let ttl = masking_beobachten_ttl(config); + beobachten.record(client_type, peer.ip(), ttl); + } + + let relay_timeout = Duration::from_millis(config.censorship.mask_relay_timeout_ms); + let idle_timeout = Duration::from_millis(config.censorship.mask_relay_idle_timeout_ms); + + if !config.censorship.mask { + // Masking disabled, just consume data + consume_client_data_with_timeout_and_cap( + reader, + config.censorship.mask_relay_max_bytes, + relay_timeout, + idle_timeout, + ) + .await; + return; + } + + let Some(_masking_permit) = shared.try_acquire_masking_fallback_permit() else { + let outcome_started = Instant::now(); + debug!( + client_type = client_type, + "Masking fallback concurrency limit reached" + ); + consume_mask_failure_path(reader, config, relay_timeout, idle_timeout).await; + wait_mask_outcome_budget(outcome_started, config).await; + return; + }; + + let client_sni = tls::extract_sni_from_client_hello(initial_data); + let exclusive_tcp_target = client_sni + .as_deref() + .and_then(|sni| exclusive_mask_target_for_sni(config, sni)); + + // Connect via Unix socket or TCP + #[cfg(unix)] + if exclusive_tcp_target.is_none() + && let Some(ref sock_path) = config.censorship.mask_unix_sock + { + let outcome_started = Instant::now(); + let connect_started = Instant::now(); + debug!( + client_type = client_type, + sock = %sock_path, + data_len = initial_data.len(), + "Forwarding bad client to mask unix socket" + ); + + let connect_result = timeout(MASK_TIMEOUT, UnixStream::connect(sock_path)).await; + match connect_result { + Ok(Ok(stream)) => { + let (mask_read, mut mask_write) = stream.into_split(); + let proxy_header = build_mask_proxy_header( + config.censorship.mask_proxy_protocol, + peer, + local_addr, + ); + if let Some(header) = proxy_header + && !write_proxy_header_with_timeout(&mut mask_write, &header).await + { + wait_mask_outcome_budget(outcome_started, config).await; + return; + } + if timeout( + relay_timeout, + relay_to_mask( + reader, + writer, + mask_read, + mask_write, + initial_data, + config.censorship.mask_shape_hardening, + config.censorship.mask_shape_bucket_floor_bytes, + config.censorship.mask_shape_bucket_cap_bytes, + config.censorship.mask_shape_above_cap_blur, + config.censorship.mask_shape_above_cap_blur_max_bytes, + config.censorship.mask_shape_hardening_aggressive_mode, + config.censorship.mask_relay_max_bytes, + idle_timeout, + ), + ) + .await + .is_err() + { + debug!("Mask relay timed out (unix socket)"); + } + wait_mask_outcome_budget(outcome_started, config).await; + } + Ok(Err(e)) => { + wait_mask_connect_budget_if_needed(connect_started, config).await; + debug!(error = %e, "Failed to connect to mask unix socket"); + consume_mask_failure_path(reader, config, relay_timeout, idle_timeout).await; + wait_mask_outcome_budget(outcome_started, config).await; + } + Err(_) => { + debug!("Timeout connecting to mask unix socket"); + consume_mask_failure_path(reader, config, relay_timeout, idle_timeout).await; + wait_mask_outcome_budget(outcome_started, config).await; + } + } + return; + } + + let mask_target = exclusive_tcp_target.unwrap_or_else(|| { + default_mask_tcp_target_for_initial_data(config, initial_data, client_sni.as_deref()) + }); + let mask_host = mask_target.host; + let mask_port = mask_target.port; + + let resolved_mask_addrs = + match resolve_mask_target_addrs(mask_host, mask_port, upstream_manager).await { + Ok(addrs) => addrs, + Err(e) => { + let outcome_started = Instant::now(); + debug!( + client_type = client_type, + host = %mask_host, + port = mask_port, + error = %e, + "Failed to resolve mask target" + ); + consume_mask_failure_path(reader, config, relay_timeout, idle_timeout).await; + wait_mask_outcome_budget(outcome_started, config).await; + return; + } + }; + + // Fail closed when fallback points at our own listener endpoint. + // Self-referential masking can create recursive proxy loops under + // misconfiguration and leak distinguishable load spikes to adversaries. + if is_mask_target_local_listener_async(mask_host, mask_port, local_addr, &resolved_mask_addrs) + .await + { + let outcome_started = Instant::now(); + debug!( + client_type = client_type, + host = %mask_host, + port = mask_port, + local = %local_addr, + "Mask target resolves to local listener; refusing self-referential masking fallback" + ); + consume_mask_failure_path(reader, config, relay_timeout, idle_timeout).await; + wait_mask_outcome_budget(outcome_started, config).await; + return; + } + + let outcome_started = Instant::now(); + + debug!( + client_type = client_type, + host = %mask_host, + port = mask_port, + data_len = initial_data.len(), + "Forwarding bad client to mask host" + ); + + let connect_started = Instant::now(); + let connect_result = timeout( + MASK_TIMEOUT, + TcpStream::connect(resolved_mask_addrs.as_slice()), + ) + .await; + match connect_result { + Ok(Ok(stream)) => { + configure_mask_backend_socket(&stream); + let proxy_header = + build_mask_proxy_header(config.censorship.mask_proxy_protocol, peer, local_addr); + + let (mask_read, mut mask_write) = stream.into_split(); + if let Some(header) = proxy_header + && !write_proxy_header_with_timeout(&mut mask_write, &header).await + { + wait_mask_outcome_budget(outcome_started, config).await; + return; + } + if timeout( + relay_timeout, + relay_to_mask( + reader, + writer, + mask_read, + mask_write, + initial_data, + config.censorship.mask_shape_hardening, + config.censorship.mask_shape_bucket_floor_bytes, + config.censorship.mask_shape_bucket_cap_bytes, + config.censorship.mask_shape_above_cap_blur, + config.censorship.mask_shape_above_cap_blur_max_bytes, + config.censorship.mask_shape_hardening_aggressive_mode, + config.censorship.mask_relay_max_bytes, + idle_timeout, + ), + ) + .await + .is_err() + { + debug!("Mask relay timed out"); + } + wait_mask_outcome_budget(outcome_started, config).await; + } + Ok(Err(e)) => { + wait_mask_connect_budget_if_needed(connect_started, config).await; + debug!(error = %e, "Failed to connect to mask host"); + consume_mask_failure_path(reader, config, relay_timeout, idle_timeout).await; + wait_mask_outcome_budget(outcome_started, config).await; + } + Err(_) => { + debug!("Timeout connecting to mask host"); + consume_mask_failure_path(reader, config, relay_timeout, idle_timeout).await; + wait_mask_outcome_budget(outcome_started, config).await; + } + } +} diff --git a/src/proxy/masking/interfaces.rs b/src/proxy/masking/interfaces.rs new file mode 100644 index 0000000..ae843b7 --- /dev/null +++ b/src/proxy/masking/interfaces.rs @@ -0,0 +1,238 @@ +use super::*; + +pub(super) fn canonical_ip(ip: IpAddr) -> IpAddr { + match ip { + IpAddr::V6(v6) => v6 + .to_ipv4_mapped() + .map(IpAddr::V4) + .unwrap_or(IpAddr::V6(v6)), + IpAddr::V4(v4) => IpAddr::V4(v4), + } +} + +#[cfg(unix)] +pub(super) fn collect_local_interface_ips() -> Vec { + #[cfg(test)] + LOCAL_INTERFACE_ENUMERATIONS.fetch_add(1, Ordering::Relaxed); + + let mut out = Vec::new(); + if let Ok(addrs) = getifaddrs() { + for iface in addrs { + if let Some(address) = iface.address { + if let Some(v4) = address.as_sockaddr_in() { + out.push(canonical_ip(IpAddr::V4(v4.ip()))); + } else if let Some(v6) = address.as_sockaddr_in6() { + out.push(canonical_ip(IpAddr::V6(v6.ip()))); + } + } + } + } + out +} + +pub(super) fn choose_interface_snapshot( + previous: &[IpAddr], + refreshed: Vec, +) -> Vec { + if refreshed.is_empty() && !previous.is_empty() { + return previous.to_vec(); + } + + refreshed +} + +#[cfg(unix)] +#[derive(Default)] +struct LocalInterfaceCache { + ips: Vec, + refreshed_at: Option, +} + +#[cfg(unix)] +static LOCAL_INTERFACE_CACHE: OnceLock> = OnceLock::new(); + +#[cfg(unix)] +pub(super) static LOCAL_INTERFACE_REFRESH_LOCK: OnceLock> = OnceLock::new(); + +#[cfg(all(unix, test))] +pub(super) fn local_interface_ips() -> Vec { + let cache = LOCAL_INTERFACE_CACHE.get_or_init(|| Mutex::new(LocalInterfaceCache::default())); + let mut guard = cache.lock().unwrap_or_else(|poison| poison.into_inner()); + + let stale = guard + .refreshed_at + .is_none_or(|at| at.elapsed() >= LOCAL_INTERFACE_CACHE_TTL); + if stale { + let refreshed = collect_local_interface_ips(); + guard.ips = choose_interface_snapshot(&guard.ips, refreshed); + guard.refreshed_at = Some(StdInstant::now()); + } + + guard.ips.clone() +} + +#[cfg(unix)] +pub(super) async fn local_interface_ips_async() -> Vec { + let cache = LOCAL_INTERFACE_CACHE.get_or_init(|| Mutex::new(LocalInterfaceCache::default())); + + { + let guard = cache.lock().unwrap_or_else(|poison| poison.into_inner()); + let stale = guard + .refreshed_at + .is_none_or(|at| at.elapsed() >= LOCAL_INTERFACE_CACHE_TTL); + if !stale { + return guard.ips.clone(); + } + } + + let refresh_lock = LOCAL_INTERFACE_REFRESH_LOCK.get_or_init(|| AsyncMutex::new(())); + let _refresh_guard = refresh_lock.lock().await; + + { + let guard = cache.lock().unwrap_or_else(|poison| poison.into_inner()); + let stale = guard + .refreshed_at + .is_none_or(|at| at.elapsed() >= LOCAL_INTERFACE_CACHE_TTL); + if !stale { + return guard.ips.clone(); + } + } + + let refreshed = tokio::task::spawn_blocking(collect_local_interface_ips) + .await + .unwrap_or_default(); + + let mut guard = cache.lock().unwrap_or_else(|poison| poison.into_inner()); + let stale = guard + .refreshed_at + .is_none_or(|at| at.elapsed() >= LOCAL_INTERFACE_CACHE_TTL); + if stale { + guard.ips = choose_interface_snapshot(&guard.ips, refreshed); + guard.refreshed_at = Some(StdInstant::now()); + } + + guard.ips.clone() +} + +#[cfg(all(not(unix), test))] +pub(super) fn local_interface_ips() -> Vec { + Vec::new() +} + +#[cfg(not(unix))] +pub(super) async fn local_interface_ips_async() -> Vec { + Vec::new() +} + +#[cfg(test)] +static LOCAL_INTERFACE_ENUMERATIONS: AtomicUsize = AtomicUsize::new(0); + +#[cfg(test)] +pub(super) fn reset_local_interface_enumerations_for_tests() { + LOCAL_INTERFACE_ENUMERATIONS.store(0, Ordering::Relaxed); + + #[cfg(unix)] + if let Some(cache) = LOCAL_INTERFACE_CACHE.get() { + let mut guard = cache.lock().unwrap_or_else(|poison| poison.into_inner()); + guard.ips.clear(); + guard.refreshed_at = None; + } +} + +#[cfg(test)] +pub(super) fn local_interface_enumerations_for_tests() -> usize { + LOCAL_INTERFACE_ENUMERATIONS.load(Ordering::Relaxed) +} + +#[cfg(test)] +pub(super) fn interface_cache_test_lock() -> &'static tokio::sync::Mutex<()> { + static LOCK: OnceLock> = OnceLock::new(); + LOCK.get_or_init(|| tokio::sync::Mutex::new(())) +} + +pub(super) fn is_mask_target_local_listener_with_interfaces( + mask_host: &str, + mask_port: u16, + local_addr: SocketAddr, + resolved_addrs: &[SocketAddr], + interface_ips: &[IpAddr], +) -> bool { + if mask_port != local_addr.port() { + return false; + } + + let local_ip = canonical_ip(local_addr.ip()); + let literal_mask_ip = parse_mask_host_ip_literal(mask_host).map(canonical_ip); + + for addr in resolved_addrs { + let resolved_ip = canonical_ip(addr.ip()); + if resolved_ip == local_ip { + return true; + } + + if local_ip.is_unspecified() + && (resolved_ip.is_loopback() + || resolved_ip.is_unspecified() + || interface_ips.contains(&resolved_ip)) + { + return true; + } + } + + if let Some(mask_ip) = literal_mask_ip { + if mask_ip == local_ip { + return true; + } + + if local_ip.is_unspecified() + && (mask_ip.is_loopback() + || mask_ip.is_unspecified() + || interface_ips.contains(&mask_ip)) + { + return true; + } + } + + false +} + +#[cfg(test)] +pub(super) fn is_mask_target_local_listener( + mask_host: &str, + mask_port: u16, + local_addr: SocketAddr, + resolved_addrs: &[SocketAddr], +) -> bool { + if mask_port != local_addr.port() { + return false; + } + + let interfaces = local_interface_ips(); + is_mask_target_local_listener_with_interfaces( + mask_host, + mask_port, + local_addr, + resolved_addrs, + &interfaces, + ) +} + +pub(super) async fn is_mask_target_local_listener_async( + mask_host: &str, + mask_port: u16, + local_addr: SocketAddr, + resolved_addrs: &[SocketAddr], +) -> bool { + if mask_port != local_addr.port() { + return false; + } + + let interfaces = local_interface_ips_async().await; + is_mask_target_local_listener_with_interfaces( + mask_host, + mask_port, + local_addr, + resolved_addrs, + &interfaces, + ) +} diff --git a/src/proxy/masking/relay.rs b/src/proxy/masking/relay.rs new file mode 100644 index 0000000..f8a74b0 --- /dev/null +++ b/src/proxy/masking/relay.rs @@ -0,0 +1,110 @@ +use super::*; + +/// Relays traffic between the client and mask backend. +pub(super) async fn relay_to_mask( + mut reader: R, + mut writer: W, + mut mask_read: MR, + mut mask_write: MW, + initial_data: &[u8], + shape_hardening_enabled: bool, + shape_bucket_floor_bytes: usize, + shape_bucket_cap_bytes: usize, + shape_above_cap_blur: bool, + shape_above_cap_blur_max_bytes: usize, + shape_hardening_aggressive_mode: bool, + mask_relay_max_bytes: usize, + idle_timeout: Duration, +) where + R: AsyncRead + Unpin + Send + 'static, + W: AsyncWrite + Unpin + Send + 'static, + MR: AsyncRead + Unpin + Send + 'static, + MW: AsyncWrite + Unpin + Send + 'static, +{ + // Send initial data to mask host + if mask_write.write_all(initial_data).await.is_err() { + return; + } + if mask_write.flush().await.is_err() { + return; + } + + let (upstream_copy, downstream_copy) = tokio::join!( + async { + copy_with_idle_timeout( + &mut reader, + &mut mask_write, + mask_relay_max_bytes, + !shape_hardening_enabled, + idle_timeout, + ) + .await + }, + async { + copy_with_idle_timeout( + &mut mask_read, + &mut writer, + mask_relay_max_bytes, + true, + idle_timeout, + ) + .await + } + ); + + let total_sent = initial_data.len().saturating_add(upstream_copy.total); + + let should_shape = shape_hardening_enabled + && !initial_data.is_empty() + && (upstream_copy.ended_by_eof + || (shape_hardening_aggressive_mode && downstream_copy.total == 0)); + + maybe_write_shape_padding( + &mut mask_write, + total_sent, + should_shape, + shape_bucket_floor_bytes, + shape_bucket_cap_bytes, + shape_above_cap_blur, + shape_above_cap_blur_max_bytes, + shape_hardening_aggressive_mode, + ) + .await; + + let _ = mask_write.shutdown().await; + let _ = writer.shutdown().await; +} + +/// Just consume all data from client without responding. +pub(super) async fn consume_client_data( + mut reader: R, + byte_cap: usize, + idle_timeout: Duration, +) { + // Keep drain path fail-closed under slow-loris stalls. + let mut buf = vec![0u8; MASK_BUFFER_SIZE]; + let mut total = 0usize; + + loop { + let read_len = mask_copy_read_len(total, byte_cap); + if read_len == 0 { + break; + } + if buf.len() < read_len { + buf.resize(read_len, 0); + } + let n = match timeout(idle_timeout, reader.read(&mut buf[..read_len])).await { + Ok(Ok(n)) => n, + Ok(Err(_)) | Err(_) => break, + }; + + if n == 0 { + break; + } + + total = total.saturating_add(n); + if byte_cap != 0 && total >= byte_cap { + break; + } + } +} diff --git a/src/proxy/masking/target.rs b/src/proxy/masking/target.rs new file mode 100644 index 0000000..096cebb --- /dev/null +++ b/src/proxy/masking/target.rs @@ -0,0 +1,207 @@ +use super::*; + +/// Detect client type based on initial data. +pub(super) fn detect_client_type(data: &[u8]) -> &'static str { + // Check for HTTP request + if is_http_probe(data) { + return "HTTP"; + } + + // Check for TLS ClientHello (0x16 = handshake, 0x03 0x01-0x03 = TLS version) + if data.len() > 3 && data[0] == 0x16 && data[1] == 0x03 { + return "TLS-scanner"; + } + + // Check for SSH + if data.starts_with(b"SSH-") { + return "SSH"; + } + + // Port scanner (very short data) + if data.len() < 10 { + return "port-scanner"; + } + + "unknown" +} + +pub(super) fn parse_mask_host_ip_literal(host: &str) -> Option { + if host.starts_with('[') && host.ends_with(']') { + return host[1..host.len() - 1].parse::().ok(); + } + host.parse::().ok() +} + +pub(super) async fn resolve_mask_target_addrs( + mask_host: &str, + mask_port: u16, + upstream_manager: Option<&crate::transport::UpstreamManager>, +) -> std::io::Result> { + if let Some(ip) = parse_mask_host_ip_literal(mask_host) { + return Ok(vec![SocketAddr::new(ip, mask_port)]); + } + + if let Some(upstream_manager) = upstream_manager { + return upstream_manager + .resolve_all(mask_host, mask_port) + .await + .map_err(|error| IoError::new(ErrorKind::NotFound, error.to_string())); + } + + let addrs = timeout(MASK_TIMEOUT, lookup_host((mask_host, mask_port))) + .await + .map_err(|_| IoError::new(ErrorKind::TimedOut, "mask target DNS lookup timed out"))??; + let addrs = addrs + .take(MASK_DNS_RESULT_MAX_ADDRESSES) + .collect::>(); + if addrs.is_empty() { + return Err(IoError::new( + ErrorKind::NotFound, + "mask target DNS lookup returned no addresses", + )); + } + + Ok(addrs) +} + +pub(super) fn matching_tls_domain_for_sni<'a>( + config: &'a ProxyConfig, + sni: &str, +) -> Option<&'a str> { + if config.censorship.tls_domain.eq_ignore_ascii_case(sni) { + return Some(config.censorship.tls_domain.as_str()); + } + + for domain in &config.censorship.tls_domains { + if domain.eq_ignore_ascii_case(sni) { + return Some(domain.as_str()); + } + } + + None +} + +pub(super) fn parse_exclusive_mask_target(target: &str) -> Option> { + let target = target.trim(); + if target.is_empty() { + return None; + } + + if target.starts_with('[') { + let end = target.find(']')?; + if target.get(end + 1..end + 2)? != ":" { + return None; + } + let port = target[end + 2..].parse::().ok()?; + return (port > 0).then_some(MaskTcpTarget { + host: &target[..=end], + port, + }); + } + + let (host, port) = target.rsplit_once(':')?; + if host.is_empty() || host.contains(':') { + return None; + } + let port = port.parse::().ok()?; + (port > 0).then_some(MaskTcpTarget { host, port }) +} + +pub(super) fn exclusive_mask_target_for_sni<'a>( + config: &'a ProxyConfig, + sni: &str, +) -> Option> { + if let Some(target) = config.censorship.exclusive_mask_targets.get(sni) { + return Some(MaskTcpTarget { + host: target.host.as_str(), + port: target.port, + }); + } + if let Some(target) = config.censorship.exclusive_mask.get(sni) { + return parse_exclusive_mask_target(target); + } + + if sni.bytes().any(|byte| byte.is_ascii_uppercase()) { + let normalized_sni = sni.to_ascii_lowercase(); + if let Some(target) = config + .censorship + .exclusive_mask_targets + .get(&normalized_sni) + { + return Some(MaskTcpTarget { + host: target.host.as_str(), + port: target.port, + }); + } + if let Some(target) = config.censorship.exclusive_mask.get(&normalized_sni) { + return parse_exclusive_mask_target(target); + } + } + + None +} + +#[cfg(test)] +pub(super) fn mask_host_for_initial_data<'a>( + config: &'a ProxyConfig, + initial_data: &[u8], +) -> &'a str { + mask_tcp_target_for_initial_data(config, initial_data).host +} + +#[cfg(test)] +pub(super) fn mask_tcp_target_for_initial_data<'a>( + config: &'a ProxyConfig, + initial_data: &[u8], +) -> MaskTcpTarget<'a> { + let sni = tls::extract_sni_from_client_hello(initial_data); + if let Some(target) = sni + .as_deref() + .and_then(|sni| exclusive_mask_target_for_sni(config, sni)) + { + return target; + } + + default_mask_tcp_target_for_initial_data(config, initial_data, sni.as_deref()) +} + +pub(super) fn default_mask_tcp_target_for_initial_data<'a>( + config: &'a ProxyConfig, + initial_data: &[u8], + sni: Option<&str>, +) -> MaskTcpTarget<'a> { + let configured_mask_host = config + .censorship + .mask_host + .as_deref() + .unwrap_or(&config.censorship.tls_domain); + + if config.censorship.mask_host.is_none() && config.censorship.mask_dynamic { + let extracted_sni = if sni.is_none() { + tls::extract_sni_from_client_hello(initial_data) + } else { + None + }; + if let Some(host) = sni + .or(extracted_sni.as_deref()) + .and_then(|sni| matching_tls_domain_for_sni(config, sni)) + { + return MaskTcpTarget { + host, + port: config.censorship.mask_port, + }; + } + } + + if let Some(mask_host) = config.censorship.mask_host.as_deref() { + return MaskTcpTarget { + host: mask_host, + port: config.censorship.mask_port, + }; + } + + MaskTcpTarget { + host: configured_mask_host, + port: config.censorship.mask_port, + } +} diff --git a/src/proxy/masking/timing.rs b/src/proxy/masking/timing.rs new file mode 100644 index 0000000..bf7487b --- /dev/null +++ b/src/proxy/masking/timing.rs @@ -0,0 +1,186 @@ +use super::*; + +pub(crate) fn sample_lognormal_percentile_bounded( + floor: u64, + ceiling: u64, + rng: &mut impl Rng, +) -> u64 { + if ceiling == 0 && floor == 0 { + return 0; + } + if floor > ceiling { + return ceiling; + } + if floor == ceiling { + return floor; + } + let floor_f = floor.max(1) as f64; + let ceiling_f = ceiling.max(1) as f64; + let mu = (floor_f.ln() + ceiling_f.ln()) / 2.0; + // 4.65 ≈ 2 * 2.326 (double-sided z-score for 99th percentile) + let sigma = ((ceiling_f / floor_f).ln() / 4.65).max(0.01); + // Box-Muller transform: two uniform samples → one standard normal sample + let u1: f64 = rng.random_range(f64::MIN_POSITIVE..1.0); + let u2: f64 = rng.random_range(0.0_f64..std::f64::consts::TAU); + let normal_sample = (-2.0_f64 * u1.ln()).sqrt() * u2.cos(); + let raw = (mu + sigma * normal_sample).exp(); + if raw.is_finite() { + (raw as u64).clamp(floor, ceiling) + } else { + ((floor_f * ceiling_f).sqrt()) as u64 + } +} + +pub(super) fn mask_outcome_target_budget(config: &ProxyConfig) -> Duration { + if config.censorship.mask_timing_normalization_enabled { + let floor = config.censorship.mask_timing_normalization_floor_ms; + let ceiling = config.censorship.mask_timing_normalization_ceiling_ms; + if floor == 0 { + if ceiling == 0 { + return Duration::from_millis(0); + } + // floor=0 stays uniform: log-normal cannot model distribution anchored at zero + let mut rng = rand::rng(); + return Duration::from_millis(rng.random_range(0..=ceiling)); + } + if ceiling > floor { + let mut rng = rand::rng(); + return Duration::from_millis(sample_lognormal_percentile_bounded( + floor, ceiling, &mut rng, + )); + } + // ceiling <= floor: use the larger value (fail-closed: preserve longer delay) + return Duration::from_millis(floor.max(ceiling)); + } + + MASK_TIMEOUT +} + +pub(super) async fn wait_mask_connect_budget_if_needed(started: Instant, config: &ProxyConfig) { + if config.censorship.mask_timing_normalization_enabled { + return; + } + + wait_mask_connect_budget(started).await; +} + +pub(super) async fn wait_mask_outcome_budget(started: Instant, config: &ProxyConfig) { + let target = mask_outcome_target_budget(config); + let elapsed = started.elapsed(); + if elapsed < target { + tokio::time::sleep(target - elapsed).await; + } +} + +#[cfg(test)] +mod tls_domain_mask_host_tests { + use super::{ + mask_host_for_initial_data, mask_tcp_target_for_initial_data, matching_tls_domain_for_sni, + }; + use crate::config::ProxyConfig; + + fn client_hello_with_sni(sni_host: &str) -> Vec { + let mut body = Vec::new(); + body.extend_from_slice(&[0x03, 0x03]); + body.extend_from_slice(&[0u8; 32]); + body.push(32); + body.extend_from_slice(&[0x42u8; 32]); + body.extend_from_slice(&2u16.to_be_bytes()); + body.extend_from_slice(&[0x13, 0x01]); + body.push(1); + body.push(0); + + let host_bytes = sni_host.as_bytes(); + let mut sni_payload = Vec::new(); + sni_payload.extend_from_slice(&((host_bytes.len() + 3) as u16).to_be_bytes()); + sni_payload.push(0); + sni_payload.extend_from_slice(&(host_bytes.len() as u16).to_be_bytes()); + sni_payload.extend_from_slice(host_bytes); + + let mut extensions = Vec::new(); + extensions.extend_from_slice(&0x0000u16.to_be_bytes()); + extensions.extend_from_slice(&(sni_payload.len() as u16).to_be_bytes()); + extensions.extend_from_slice(&sni_payload); + body.extend_from_slice(&(extensions.len() as u16).to_be_bytes()); + body.extend_from_slice(&extensions); + + let mut handshake = Vec::new(); + handshake.push(0x01); + let body_len = (body.len() as u32).to_be_bytes(); + handshake.extend_from_slice(&body_len[1..4]); + handshake.extend_from_slice(&body); + + let mut record = Vec::new(); + record.push(0x16); + record.extend_from_slice(&[0x03, 0x01]); + record.extend_from_slice(&(handshake.len() as u16).to_be_bytes()); + record.extend_from_slice(&handshake); + record + } + + fn config_with_tls_domains() -> ProxyConfig { + let mut config = ProxyConfig::default(); + config.censorship.tls_domain = "a.com".to_string(); + config.censorship.tls_domains = vec!["b.com".to_string(), "c.com".to_string()]; + config.censorship.mask_host = None; + config + } + + #[test] + fn matching_tls_domain_accepts_primary_and_extra_domains_case_insensitively() { + let config = config_with_tls_domains(); + + assert_eq!(matching_tls_domain_for_sni(&config, "A.COM"), Some("a.com")); + assert_eq!(matching_tls_domain_for_sni(&config, "B.COM"), Some("b.com")); + assert_eq!(matching_tls_domain_for_sni(&config, "unknown.com"), None); + } + + #[test] + fn mask_host_preserves_explicit_non_primary_origin() { + let mut config = config_with_tls_domains(); + config.censorship.mask_host = Some("origin.example".to_string()); + + let initial_data = client_hello_with_sni("b.com"); + + assert_eq!( + mask_host_for_initial_data(&config, &initial_data), + "origin.example" + ); + } + + #[test] + fn mask_host_uses_matching_tls_domain_when_mask_host_is_primary_default() { + let config = config_with_tls_domains(); + let initial_data = client_hello_with_sni("b.com"); + + assert_eq!(mask_host_for_initial_data(&config, &initial_data), "b.com"); + } + + #[test] + fn mask_host_uses_primary_domain_when_dynamic_masking_is_disabled() { + let mut config = config_with_tls_domains(); + config.censorship.mask_dynamic = false; + let initial_data = client_hello_with_sni("b.com"); + + assert_eq!(mask_host_for_initial_data(&config, &initial_data), "a.com"); + } + + #[test] + fn exclusive_mask_target_overrides_only_matching_sni() { + let mut config = config_with_tls_domains(); + config + .censorship + .exclusive_mask + .insert("b.com".to_string(), "origin-b.example:8443".to_string()); + let b_initial_data = client_hello_with_sni("B.COM"); + let c_initial_data = client_hello_with_sni("c.com"); + + let b_target = mask_tcp_target_for_initial_data(&config, &b_initial_data); + let c_target = mask_tcp_target_for_initial_data(&config, &c_initial_data); + + assert_eq!(b_target.host, "origin-b.example"); + assert_eq!(b_target.port, 8443); + assert_eq!(c_target.host, "c.com"); + assert_eq!(c_target.port, config.censorship.mask_port); + } +} diff --git a/src/proxy/middle_relay/session.rs b/src/proxy/middle_relay/session.rs index 7c704f1..bee6675 100644 --- a/src/proxy/middle_relay/session.rs +++ b/src/proxy/middle_relay/session.rs @@ -1,5 +1,12 @@ use super::*; +// Bounded C2ME sender and downstream writer tasks. +mod tasks; +// Conntrack close classification. +mod close_reason; + +use close_reason::classify_conntrack_close_reason; +use tasks::{run_c2me_sender, run_me_writer}; struct RelayConnLease { connection: Option, conn_id: u64, @@ -174,47 +181,21 @@ where }; let c2me_byte_budget = c2me_queued_permit_budget(c2me_channel_capacity, frame_limit); let c2me_byte_semaphore = Arc::new(Semaphore::new(c2me_byte_budget)); - let (c2me_tx, mut c2me_rx) = mpsc::channel::(c2me_channel_capacity); + let (c2me_tx, c2me_rx) = mpsc::channel::(c2me_channel_capacity); let me_pool_c2me = me_pool.clone(); - let mut c2me_sender = tokio::spawn(async move { - let mut sent_since_yield = 0usize; - while let Some(cmd) = c2me_rx.recv().await { - match cmd { - C2MeCommand::Data { - payload, - flags, - _permit, - } => { - me_pool_c2me - .send_proxy_req_pooled( - conn_id, - success.dc_idx, - peer, - translated_local_addr, - payload, - _permit, - flags, - effective_tag_array, - ) - .await?; - sent_since_yield = sent_since_yield.saturating_add(1); - if should_yield_c2me_sender(sent_since_yield, !c2me_rx.is_empty()) { - sent_since_yield = 0; - tokio::task::yield_now().await; - } - } - C2MeCommand::Close => { - let _ = me_pool_c2me.send_close(conn_id).await; - return Ok(()); - } - } - } - Ok(()) - }); + let mut c2me_sender = tokio::spawn(run_c2me_sender( + c2me_rx, + me_pool_c2me, + conn_id, + success, + peer, + translated_local_addr, + effective_tag_array, + )); - let (stop_tx, mut stop_rx) = oneshot::channel::<()>(); + let (stop_tx, stop_rx) = oneshot::channel::<()>(); let flow_cancel = CancellationToken::new(); - let mut me_rx_task = me_rx; + let me_rx_task = me_rx; let stats_clone = stats.clone(); let rng_clone = rng.clone(); let user_clone = user.clone(); @@ -224,361 +205,24 @@ where let last_downstream_activity_ms_clone = last_downstream_activity_ms.clone(); let bytes_me2c_clone = bytes_me2c.clone(); let d2c_flush_policy = MeD2cFlushPolicy::from_config(&config); - let mut me_writer = tokio::spawn(async move { - let mut writer = crypto_writer; - let mut frame_buf = Vec::with_capacity(16 * 1024); - let shrink_threshold = d2c_flush_policy.frame_buf_shrink_threshold_bytes; - - fn shrink_session_vec(buf: &mut Vec, threshold: usize) { - if buf.capacity() > threshold { - buf.clear(); - buf.shrink_to(threshold); - } else { - buf.clear(); - } - } - - loop { - tokio::select! { - msg = me_rx_task.recv() => { - let Some(first) = msg else { - debug!(conn_id, "ME channel closed"); - shrink_session_vec(&mut frame_buf, shrink_threshold); - return Err(ProxyError::MiddleConnectionLost); - }; - - let mut batch_frames = 0usize; - let mut batch_bytes = 0usize; - let mut flush_immediately; - let mut max_delay_fired = false; - - let first_is_downstream_activity = - matches!(&first, MeResponse::Data { .. } | MeResponse::Ack(_)); - match process_me_writer_response_with_traffic_lease( - first, - &mut writer, - proto_tag, - rng_clone.as_ref(), - &mut frame_buf, - stats_clone.as_ref(), - &user_clone, - quota_user_stats_me_writer.as_deref(), - quota_limit, - d2c_flush_policy.quota_soft_overshoot_bytes, - traffic_lease_me_writer.as_ref(), - &flow_cancel_me_writer, - bytes_me2c_clone.as_ref(), - conn_id, - d2c_flush_policy.ack_flush_immediate, - false, - ).await? { - MeWriterResponseOutcome::Continue { frames, bytes, flush_immediately: immediate } => { - if first_is_downstream_activity { - last_downstream_activity_ms_clone - .store(session_started_at.elapsed().as_millis() as u64, Ordering::Relaxed); - } - batch_frames = batch_frames.saturating_add(frames); - batch_bytes = batch_bytes.saturating_add(bytes); - flush_immediately = immediate; - } - MeWriterResponseOutcome::Close => { - let flush_started_at = if stats_clone.telemetry_policy().me_level.allows_debug() { - Some(Instant::now()) - } else { - None - }; - let _ = flush_client_or_cancel(&mut writer, &flow_cancel_me_writer).await; - let flush_duration_us = flush_started_at.map(|started| { - started - .elapsed() - .as_micros() - .min(u128::from(u64::MAX)) as u64 - }); - observe_me_d2c_flush_event( - stats_clone.as_ref(), - MeD2cFlushReason::Close, - batch_frames, - batch_bytes, - flush_duration_us, - ); - shrink_session_vec(&mut frame_buf, shrink_threshold); - return Ok(()); - } - } - - while !flush_immediately - && batch_frames < d2c_flush_policy.max_frames - && batch_bytes < d2c_flush_policy.max_bytes - { - let Ok(next) = me_rx_task.try_recv() else { - break; - }; - - let next_is_downstream_activity = - matches!(&next, MeResponse::Data { .. } | MeResponse::Ack(_)); - match process_me_writer_response_with_traffic_lease( - next, - &mut writer, - proto_tag, - rng_clone.as_ref(), - &mut frame_buf, - stats_clone.as_ref(), - &user_clone, - quota_user_stats_me_writer.as_deref(), - quota_limit, - d2c_flush_policy.quota_soft_overshoot_bytes, - traffic_lease_me_writer.as_ref(), - &flow_cancel_me_writer, - bytes_me2c_clone.as_ref(), - conn_id, - d2c_flush_policy.ack_flush_immediate, - true, - ).await? { - MeWriterResponseOutcome::Continue { frames, bytes, flush_immediately: immediate } => { - if next_is_downstream_activity { - last_downstream_activity_ms_clone - .store(session_started_at.elapsed().as_millis() as u64, Ordering::Relaxed); - } - batch_frames = batch_frames.saturating_add(frames); - batch_bytes = batch_bytes.saturating_add(bytes); - flush_immediately |= immediate; - } - MeWriterResponseOutcome::Close => { - let flush_started_at = - if stats_clone.telemetry_policy().me_level.allows_debug() { - Some(Instant::now()) - } else { - None - }; - let _ = - flush_client_or_cancel(&mut writer, &flow_cancel_me_writer).await; - let flush_duration_us = flush_started_at.map(|started| { - started - .elapsed() - .as_micros() - .min(u128::from(u64::MAX)) - as u64 - }); - observe_me_d2c_flush_event( - stats_clone.as_ref(), - MeD2cFlushReason::Close, - batch_frames, - batch_bytes, - flush_duration_us, - ); - shrink_session_vec(&mut frame_buf, shrink_threshold); - return Ok(()); - } - } - } - - if !flush_immediately - && !d2c_flush_policy.max_delay.is_zero() - && batch_frames < d2c_flush_policy.max_frames - && batch_bytes < d2c_flush_policy.max_bytes - { - stats_clone.increment_me_d2c_batch_timeout_armed_total(); - match tokio::time::timeout(d2c_flush_policy.max_delay, me_rx_task.recv()).await { - Ok(Some(next)) => { - let next_is_downstream_activity = - matches!(&next, MeResponse::Data { .. } | MeResponse::Ack(_)); - match process_me_writer_response_with_traffic_lease( - next, - &mut writer, - proto_tag, - rng_clone.as_ref(), - &mut frame_buf, - stats_clone.as_ref(), - &user_clone, - quota_user_stats_me_writer.as_deref(), - quota_limit, - d2c_flush_policy.quota_soft_overshoot_bytes, - traffic_lease_me_writer.as_ref(), - &flow_cancel_me_writer, - bytes_me2c_clone.as_ref(), - conn_id, - d2c_flush_policy.ack_flush_immediate, - true, - ).await? { - MeWriterResponseOutcome::Continue { frames, bytes, flush_immediately: immediate } => { - if next_is_downstream_activity { - last_downstream_activity_ms_clone - .store(session_started_at.elapsed().as_millis() as u64, Ordering::Relaxed); - } - batch_frames = batch_frames.saturating_add(frames); - batch_bytes = batch_bytes.saturating_add(bytes); - flush_immediately |= immediate; - } - MeWriterResponseOutcome::Close => { - let flush_started_at = if stats_clone - .telemetry_policy() - .me_level - .allows_debug() - { - Some(Instant::now()) - } else { - None - }; - let _ = flush_client_or_cancel( - &mut writer, - &flow_cancel_me_writer, - ) - .await; - let flush_duration_us = flush_started_at.map(|started| { - started - .elapsed() - .as_micros() - .min(u128::from(u64::MAX)) - as u64 - }); - observe_me_d2c_flush_event( - stats_clone.as_ref(), - MeD2cFlushReason::Close, - batch_frames, - batch_bytes, - flush_duration_us, - ); - shrink_session_vec(&mut frame_buf, shrink_threshold); - return Ok(()); - } - } - - while !flush_immediately - && batch_frames < d2c_flush_policy.max_frames - && batch_bytes < d2c_flush_policy.max_bytes - { - let Ok(extra) = me_rx_task.try_recv() else { - break; - }; - - let extra_is_downstream_activity = - matches!(&extra, MeResponse::Data { .. } | MeResponse::Ack(_)); - match process_me_writer_response_with_traffic_lease( - extra, - &mut writer, - proto_tag, - rng_clone.as_ref(), - &mut frame_buf, - stats_clone.as_ref(), - &user_clone, - quota_user_stats_me_writer.as_deref(), - quota_limit, - d2c_flush_policy.quota_soft_overshoot_bytes, - traffic_lease_me_writer.as_ref(), - &flow_cancel_me_writer, - bytes_me2c_clone.as_ref(), - conn_id, - d2c_flush_policy.ack_flush_immediate, - true, - ).await? { - MeWriterResponseOutcome::Continue { frames, bytes, flush_immediately: immediate } => { - if extra_is_downstream_activity { - last_downstream_activity_ms_clone - .store(session_started_at.elapsed().as_millis() as u64, Ordering::Relaxed); - } - batch_frames = batch_frames.saturating_add(frames); - batch_bytes = batch_bytes.saturating_add(bytes); - flush_immediately |= immediate; - } - MeWriterResponseOutcome::Close => { - let flush_started_at = if stats_clone - .telemetry_policy() - .me_level - .allows_debug() - { - Some(Instant::now()) - } else { - None - }; - let _ = flush_client_or_cancel( - &mut writer, - &flow_cancel_me_writer, - ) - .await; - let flush_duration_us = flush_started_at.map(|started| { - started - .elapsed() - .as_micros() - .min(u128::from(u64::MAX)) - as u64 - }); - observe_me_d2c_flush_event( - stats_clone.as_ref(), - MeD2cFlushReason::Close, - batch_frames, - batch_bytes, - flush_duration_us, - ); - shrink_session_vec(&mut frame_buf, shrink_threshold); - return Ok(()); - } - } - } - } - Ok(None) => { - debug!(conn_id, "ME channel closed"); - shrink_session_vec(&mut frame_buf, shrink_threshold); - return Err(ProxyError::MiddleConnectionLost); - } - Err(_) => { - max_delay_fired = true; - stats_clone.increment_me_d2c_batch_timeout_fired_total(); - } - } - } - - let flush_reason = classify_me_d2c_flush_reason( - flush_immediately, - batch_frames, - d2c_flush_policy.max_frames, - batch_bytes, - d2c_flush_policy.max_bytes, - max_delay_fired, - ); - let physical_flush = - me_d2c_flush_reason_requires_client_flush(flush_reason); - let flush_started_at = if physical_flush - && stats_clone.telemetry_policy().me_level.allows_debug() - { - Some(Instant::now()) - } else { - None - }; - if physical_flush { - flush_client_or_cancel(&mut writer, &flow_cancel_me_writer).await?; - } - let flush_duration_us = flush_started_at.map(|started| { - started - .elapsed() - .as_micros() - .min(u128::from(u64::MAX)) as u64 - }); - observe_me_d2c_flush_event( - stats_clone.as_ref(), - flush_reason, - batch_frames, - batch_bytes, - flush_duration_us, - ); - let shrink_threshold = d2c_flush_policy.frame_buf_shrink_threshold_bytes; - let shrink_trigger = shrink_threshold - .saturating_mul(ME_D2C_FRAME_BUF_SHRINK_HYSTERESIS_FACTOR); - if frame_buf.capacity() > shrink_trigger { - let cap_before = frame_buf.capacity(); - frame_buf.shrink_to(shrink_threshold); - let cap_after = frame_buf.capacity(); - let bytes_freed = cap_before.saturating_sub(cap_after) as u64; - stats_clone.observe_me_d2c_frame_buf_shrink(bytes_freed); - } - } - _ = &mut stop_rx => { - debug!(conn_id, "ME writer stop signal"); - shrink_session_vec(&mut frame_buf, shrink_threshold); - return Ok(()); - } - } - } - }); + let mut me_writer = tokio::spawn(run_me_writer( + crypto_writer, + me_rx_task, + stats_clone, + rng_clone, + user_clone, + quota_user_stats_me_writer, + quota_limit, + traffic_lease_me_writer, + flow_cancel_me_writer, + last_downstream_activity_ms_clone, + bytes_me2c_clone, + d2c_flush_policy, + proto_tag, + session_started_at, + conn_id, + stop_rx, + )); let mut main_result: Result<()> = Ok(()); let mut client_closed = false; @@ -875,30 +519,3 @@ where ); result } - -fn classify_conntrack_close_reason(result: &Result<()>) -> ConntrackCloseReason { - match result { - Ok(()) => ConntrackCloseReason::NormalEof, - Err(ProxyError::Io(error)) if matches!(error.kind(), std::io::ErrorKind::TimedOut) => { - ConntrackCloseReason::Timeout - } - Err(ProxyError::Io(error)) - if matches!( - error.kind(), - std::io::ErrorKind::ConnectionReset - | std::io::ErrorKind::ConnectionAborted - | std::io::ErrorKind::BrokenPipe - | std::io::ErrorKind::NotConnected - | std::io::ErrorKind::UnexpectedEof - ) => - { - ConntrackCloseReason::Reset - } - Err(ProxyError::Proxy(message)) - if message.contains("pressure") || message.contains("evicted") => - { - ConntrackCloseReason::Pressure - } - Err(_) => ConntrackCloseReason::Other, - } -} diff --git a/src/proxy/middle_relay/session/close_reason.rs b/src/proxy/middle_relay/session/close_reason.rs new file mode 100644 index 0000000..024ab2c --- /dev/null +++ b/src/proxy/middle_relay/session/close_reason.rs @@ -0,0 +1,28 @@ +use super::*; + +pub(super) fn classify_conntrack_close_reason(result: &Result<()>) -> ConntrackCloseReason { + match result { + Ok(()) => ConntrackCloseReason::NormalEof, + Err(ProxyError::Io(error)) if matches!(error.kind(), std::io::ErrorKind::TimedOut) => { + ConntrackCloseReason::Timeout + } + Err(ProxyError::Io(error)) + if matches!( + error.kind(), + std::io::ErrorKind::ConnectionReset + | std::io::ErrorKind::ConnectionAborted + | std::io::ErrorKind::BrokenPipe + | std::io::ErrorKind::NotConnected + | std::io::ErrorKind::UnexpectedEof + ) => + { + ConntrackCloseReason::Reset + } + Err(ProxyError::Proxy(message)) + if message.contains("pressure") || message.contains("evicted") => + { + ConntrackCloseReason::Pressure + } + Err(_) => ConntrackCloseReason::Other, + } +} diff --git a/src/proxy/middle_relay/session/tasks.rs b/src/proxy/middle_relay/session/tasks.rs new file mode 100644 index 0000000..56f52dd --- /dev/null +++ b/src/proxy/middle_relay/session/tasks.rs @@ -0,0 +1,423 @@ +use super::*; + +#[allow(clippy::too_many_arguments)] +pub(super) async fn run_c2me_sender( + mut c2me_rx: mpsc::Receiver, + me_pool_c2me: Arc, + conn_id: u64, + success: HandshakeSuccess, + peer: SocketAddr, + translated_local_addr: SocketAddr, + effective_tag_array: Option<[u8; 16]>, +) -> Result<()> { + let mut sent_since_yield = 0usize; + while let Some(cmd) = c2me_rx.recv().await { + match cmd { + C2MeCommand::Data { + payload, + flags, + _permit, + } => { + me_pool_c2me + .send_proxy_req_pooled( + conn_id, + success.dc_idx, + peer, + translated_local_addr, + payload, + _permit, + flags, + effective_tag_array, + ) + .await?; + sent_since_yield = sent_since_yield.saturating_add(1); + if should_yield_c2me_sender(sent_since_yield, !c2me_rx.is_empty()) { + sent_since_yield = 0; + tokio::task::yield_now().await; + } + } + C2MeCommand::Close => { + let _ = me_pool_c2me.send_close(conn_id).await; + return Ok(()); + } + } + } + Ok(()) +} + +#[allow(clippy::too_many_arguments)] +pub(super) async fn run_me_writer( + crypto_writer: CryptoWriter, + mut me_rx_task: mpsc::Receiver, + stats_clone: Arc, + rng_clone: Arc, + user_clone: String, + quota_user_stats_me_writer: Option>, + quota_limit: Option, + traffic_lease_me_writer: Option>, + flow_cancel_me_writer: CancellationToken, + last_downstream_activity_ms_clone: Arc, + bytes_me2c_clone: Arc, + d2c_flush_policy: MeD2cFlushPolicy, + proto_tag: ProtoTag, + session_started_at: Instant, + conn_id: u64, + mut stop_rx: oneshot::Receiver<()>, +) -> Result<()> +where + W: AsyncWrite + Unpin + Send + 'static, +{ + let mut writer = crypto_writer; + let mut frame_buf = Vec::with_capacity(16 * 1024); + let shrink_threshold = d2c_flush_policy.frame_buf_shrink_threshold_bytes; + + fn shrink_session_vec(buf: &mut Vec, threshold: usize) { + if buf.capacity() > threshold { + buf.clear(); + buf.shrink_to(threshold); + } else { + buf.clear(); + } + } + + loop { + tokio::select! { + msg = me_rx_task.recv() => { + let Some(first) = msg else { + debug!(conn_id, "ME channel closed"); + shrink_session_vec(&mut frame_buf, shrink_threshold); + return Err(ProxyError::MiddleConnectionLost); + }; + + let mut batch_frames = 0usize; + let mut batch_bytes = 0usize; + let mut flush_immediately; + let mut max_delay_fired = false; + + let first_is_downstream_activity = + matches!(&first, MeResponse::Data { .. } | MeResponse::Ack(_)); + match process_me_writer_response_with_traffic_lease( + first, + &mut writer, + proto_tag, + rng_clone.as_ref(), + &mut frame_buf, + stats_clone.as_ref(), + &user_clone, + quota_user_stats_me_writer.as_deref(), + quota_limit, + d2c_flush_policy.quota_soft_overshoot_bytes, + traffic_lease_me_writer.as_ref(), + &flow_cancel_me_writer, + bytes_me2c_clone.as_ref(), + conn_id, + d2c_flush_policy.ack_flush_immediate, + false, + ).await? { + MeWriterResponseOutcome::Continue { frames, bytes, flush_immediately: immediate } => { + if first_is_downstream_activity { + last_downstream_activity_ms_clone + .store(session_started_at.elapsed().as_millis() as u64, Ordering::Relaxed); + } + batch_frames = batch_frames.saturating_add(frames); + batch_bytes = batch_bytes.saturating_add(bytes); + flush_immediately = immediate; + } + MeWriterResponseOutcome::Close => { + let flush_started_at = if stats_clone.telemetry_policy().me_level.allows_debug() { + Some(Instant::now()) + } else { + None + }; + let _ = flush_client_or_cancel(&mut writer, &flow_cancel_me_writer).await; + let flush_duration_us = flush_started_at.map(|started| { + started + .elapsed() + .as_micros() + .min(u128::from(u64::MAX)) as u64 + }); + observe_me_d2c_flush_event( + stats_clone.as_ref(), + MeD2cFlushReason::Close, + batch_frames, + batch_bytes, + flush_duration_us, + ); + shrink_session_vec(&mut frame_buf, shrink_threshold); + return Ok(()); + } + } + + while !flush_immediately + && batch_frames < d2c_flush_policy.max_frames + && batch_bytes < d2c_flush_policy.max_bytes + { + let Ok(next) = me_rx_task.try_recv() else { + break; + }; + + let next_is_downstream_activity = + matches!(&next, MeResponse::Data { .. } | MeResponse::Ack(_)); + match process_me_writer_response_with_traffic_lease( + next, + &mut writer, + proto_tag, + rng_clone.as_ref(), + &mut frame_buf, + stats_clone.as_ref(), + &user_clone, + quota_user_stats_me_writer.as_deref(), + quota_limit, + d2c_flush_policy.quota_soft_overshoot_bytes, + traffic_lease_me_writer.as_ref(), + &flow_cancel_me_writer, + bytes_me2c_clone.as_ref(), + conn_id, + d2c_flush_policy.ack_flush_immediate, + true, + ).await? { + MeWriterResponseOutcome::Continue { frames, bytes, flush_immediately: immediate } => { + if next_is_downstream_activity { + last_downstream_activity_ms_clone + .store(session_started_at.elapsed().as_millis() as u64, Ordering::Relaxed); + } + batch_frames = batch_frames.saturating_add(frames); + batch_bytes = batch_bytes.saturating_add(bytes); + flush_immediately |= immediate; + } + MeWriterResponseOutcome::Close => { + let flush_started_at = + if stats_clone.telemetry_policy().me_level.allows_debug() { + Some(Instant::now()) + } else { + None + }; + let _ = + flush_client_or_cancel(&mut writer, &flow_cancel_me_writer).await; + let flush_duration_us = flush_started_at.map(|started| { + started + .elapsed() + .as_micros() + .min(u128::from(u64::MAX)) + as u64 + }); + observe_me_d2c_flush_event( + stats_clone.as_ref(), + MeD2cFlushReason::Close, + batch_frames, + batch_bytes, + flush_duration_us, + ); + shrink_session_vec(&mut frame_buf, shrink_threshold); + return Ok(()); + } + } + } + + if !flush_immediately + && !d2c_flush_policy.max_delay.is_zero() + && batch_frames < d2c_flush_policy.max_frames + && batch_bytes < d2c_flush_policy.max_bytes + { + stats_clone.increment_me_d2c_batch_timeout_armed_total(); + match tokio::time::timeout(d2c_flush_policy.max_delay, me_rx_task.recv()).await { + Ok(Some(next)) => { + let next_is_downstream_activity = + matches!(&next, MeResponse::Data { .. } | MeResponse::Ack(_)); + match process_me_writer_response_with_traffic_lease( + next, + &mut writer, + proto_tag, + rng_clone.as_ref(), + &mut frame_buf, + stats_clone.as_ref(), + &user_clone, + quota_user_stats_me_writer.as_deref(), + quota_limit, + d2c_flush_policy.quota_soft_overshoot_bytes, + traffic_lease_me_writer.as_ref(), + &flow_cancel_me_writer, + bytes_me2c_clone.as_ref(), + conn_id, + d2c_flush_policy.ack_flush_immediate, + true, + ).await? { + MeWriterResponseOutcome::Continue { frames, bytes, flush_immediately: immediate } => { + if next_is_downstream_activity { + last_downstream_activity_ms_clone + .store(session_started_at.elapsed().as_millis() as u64, Ordering::Relaxed); + } + batch_frames = batch_frames.saturating_add(frames); + batch_bytes = batch_bytes.saturating_add(bytes); + flush_immediately |= immediate; + } + MeWriterResponseOutcome::Close => { + let flush_started_at = if stats_clone + .telemetry_policy() + .me_level + .allows_debug() + { + Some(Instant::now()) + } else { + None + }; + let _ = flush_client_or_cancel( + &mut writer, + &flow_cancel_me_writer, + ) + .await; + let flush_duration_us = flush_started_at.map(|started| { + started + .elapsed() + .as_micros() + .min(u128::from(u64::MAX)) + as u64 + }); + observe_me_d2c_flush_event( + stats_clone.as_ref(), + MeD2cFlushReason::Close, + batch_frames, + batch_bytes, + flush_duration_us, + ); + shrink_session_vec(&mut frame_buf, shrink_threshold); + return Ok(()); + } + } + + while !flush_immediately + && batch_frames < d2c_flush_policy.max_frames + && batch_bytes < d2c_flush_policy.max_bytes + { + let Ok(extra) = me_rx_task.try_recv() else { + break; + }; + + let extra_is_downstream_activity = + matches!(&extra, MeResponse::Data { .. } | MeResponse::Ack(_)); + match process_me_writer_response_with_traffic_lease( + extra, + &mut writer, + proto_tag, + rng_clone.as_ref(), + &mut frame_buf, + stats_clone.as_ref(), + &user_clone, + quota_user_stats_me_writer.as_deref(), + quota_limit, + d2c_flush_policy.quota_soft_overshoot_bytes, + traffic_lease_me_writer.as_ref(), + &flow_cancel_me_writer, + bytes_me2c_clone.as_ref(), + conn_id, + d2c_flush_policy.ack_flush_immediate, + true, + ).await? { + MeWriterResponseOutcome::Continue { frames, bytes, flush_immediately: immediate } => { + if extra_is_downstream_activity { + last_downstream_activity_ms_clone + .store(session_started_at.elapsed().as_millis() as u64, Ordering::Relaxed); + } + batch_frames = batch_frames.saturating_add(frames); + batch_bytes = batch_bytes.saturating_add(bytes); + flush_immediately |= immediate; + } + MeWriterResponseOutcome::Close => { + let flush_started_at = if stats_clone + .telemetry_policy() + .me_level + .allows_debug() + { + Some(Instant::now()) + } else { + None + }; + let _ = flush_client_or_cancel( + &mut writer, + &flow_cancel_me_writer, + ) + .await; + let flush_duration_us = flush_started_at.map(|started| { + started + .elapsed() + .as_micros() + .min(u128::from(u64::MAX)) + as u64 + }); + observe_me_d2c_flush_event( + stats_clone.as_ref(), + MeD2cFlushReason::Close, + batch_frames, + batch_bytes, + flush_duration_us, + ); + shrink_session_vec(&mut frame_buf, shrink_threshold); + return Ok(()); + } + } + } + } + Ok(None) => { + debug!(conn_id, "ME channel closed"); + shrink_session_vec(&mut frame_buf, shrink_threshold); + return Err(ProxyError::MiddleConnectionLost); + } + Err(_) => { + max_delay_fired = true; + stats_clone.increment_me_d2c_batch_timeout_fired_total(); + } + } + } + + let flush_reason = classify_me_d2c_flush_reason( + flush_immediately, + batch_frames, + d2c_flush_policy.max_frames, + batch_bytes, + d2c_flush_policy.max_bytes, + max_delay_fired, + ); + let physical_flush = + me_d2c_flush_reason_requires_client_flush(flush_reason); + let flush_started_at = if physical_flush + && stats_clone.telemetry_policy().me_level.allows_debug() + { + Some(Instant::now()) + } else { + None + }; + if physical_flush { + flush_client_or_cancel(&mut writer, &flow_cancel_me_writer).await?; + } + let flush_duration_us = flush_started_at.map(|started| { + started + .elapsed() + .as_micros() + .min(u128::from(u64::MAX)) as u64 + }); + observe_me_d2c_flush_event( + stats_clone.as_ref(), + flush_reason, + batch_frames, + batch_bytes, + flush_duration_us, + ); + let shrink_threshold = d2c_flush_policy.frame_buf_shrink_threshold_bytes; + let shrink_trigger = shrink_threshold + .saturating_mul(ME_D2C_FRAME_BUF_SHRINK_HYSTERESIS_FACTOR); + if frame_buf.capacity() > shrink_trigger { + let cap_before = frame_buf.capacity(); + frame_buf.shrink_to(shrink_threshold); + let cap_after = frame_buf.capacity(); + let bytes_freed = cap_before.saturating_sub(cap_after) as u64; + stats_clone.observe_me_d2c_frame_buf_shrink(bytes_freed); + } + } + _ = &mut stop_rx => { + debug!(conn_id, "ME writer stop signal"); + shrink_session_vec(&mut frame_buf, shrink_threshold); + return Ok(()); + } + } + } +} diff --git a/src/proxy/tests/direct_relay_security_tests.rs b/src/proxy/tests/direct_relay_security_tests.rs index f66397a..f5db06f 100644 --- a/src/proxy/tests/direct_relay_security_tests.rs +++ b/src/proxy/tests/direct_relay_security_tests.rs @@ -39,1932 +39,21 @@ fn nonempty_line_count(text: &str) -> usize { text.lines().filter(|line| !line.trim().is_empty()).count() } -#[test] -fn unknown_dc_log_is_deduplicated_per_dc_idx() { - let _guard = unknown_dc_test_lock() - .lock() - .expect("unknown dc test lock must be available"); - clear_unknown_dc_log_cache_for_testing(); - - assert!(should_log_unknown_dc(777)); - assert!( - !should_log_unknown_dc(777), - "same unknown dc_idx must not be logged repeatedly" - ); - assert!( - should_log_unknown_dc(778), - "different unknown dc_idx must still be loggable" - ); -} - -#[test] -fn unknown_dc_log_respects_distinct_limit() { - let _guard = unknown_dc_test_lock() - .lock() - .expect("unknown dc test lock must be available"); - clear_unknown_dc_log_cache_for_testing(); - - for dc in 1..=UNKNOWN_DC_LOG_DISTINCT_LIMIT { - assert!( - should_log_unknown_dc(dc as i16), - "expected first-time unknown dc_idx to be loggable" - ); - } - - assert!( - !should_log_unknown_dc(i16::MAX), - "distinct unknown dc_idx entries above limit must not be logged" - ); -} - -#[test] -fn unknown_dc_log_fails_closed_when_dedup_lock_is_poisoned() { - let poisoned = Arc::new(std::sync::Mutex::new( - std::collections::HashSet::::new(), - )); - let poisoned_for_thread = poisoned.clone(); - - let _ = std::thread::spawn(move || { - let _guard = poisoned_for_thread - .lock() - .expect("poison setup lock must be available"); - panic!("intentional poison for fail-closed regression"); - }) - .join(); - - assert!( - !should_log_unknown_dc_with_set(poisoned.as_ref(), 4242), - "poisoned unknown-DC dedup lock must fail closed" - ); -} - -#[test] -fn unsafe_unknown_dc_log_path_does_not_consume_dedup_slot() { - let _guard = unknown_dc_test_lock() - .lock() - .expect("unknown dc test lock must be available"); - clear_unknown_dc_log_cache_for_testing(); - - let dc_idx: i16 = 31_123; - let mut cfg = ProxyConfig::default(); - cfg.general.unknown_dc_file_log_enabled = true; - cfg.general.unknown_dc_log_path = Some("../telemt-unknown-dc-unsafe.log".to_string()); - - let _ = get_dc_addr_static(dc_idx, &cfg).expect("fallback routing must still work"); - - assert!( - should_log_unknown_dc(dc_idx), - "rejected unsafe log path must not consume unknown-dc dedup entry" - ); -} - -#[test] -fn stress_unknown_dc_log_concurrent_unique_churn_respects_cap() { - let _guard = unknown_dc_test_lock() - .lock() - .expect("unknown dc test lock must be available"); - clear_unknown_dc_log_cache_for_testing(); - - let accepted = Arc::new(AtomicUsize::new(0)); - let mut workers = Vec::new(); - - // Adversarial model: many concurrent peers rotate dc_idx values rapidly. - for worker in 0..16usize { - let accepted = Arc::clone(&accepted); - workers.push(std::thread::spawn(move || { - let base = (worker * 2048) as i32; - for offset in 0..512i32 { - let raw = base + offset; - let dc = (raw % i16::MAX as i32) as i16; - if should_log_unknown_dc(dc) { - accepted.fetch_add(1, Ordering::Relaxed); - } - } - })); - } - - for worker in workers { - worker.join().expect("worker thread must not panic"); - } - - assert_eq!( - accepted.load(Ordering::Relaxed), - UNKNOWN_DC_LOG_DISTINCT_LIMIT, - "concurrent unique churn must never admit more than the configured distinct cap" - ); -} - -#[test] -fn light_fuzz_unknown_dc_log_mixed_duplicates_never_exceeds_cap() { - let _guard = unknown_dc_test_lock() - .lock() - .expect("unknown dc test lock must be available"); - clear_unknown_dc_log_cache_for_testing(); - - // Deterministic xorshift sequence for reproducible mixed duplicate fuzzing. - let mut s: u64 = 0xA5A5_5A5A_C3C3_3C3C; - let mut admitted = 0usize; - - for _ in 0..20_000 { - s ^= s << 7; - s ^= s >> 9; - s ^= s << 8; - - let dc = (s as i16).wrapping_sub(i16::MAX / 2); - if should_log_unknown_dc(dc) { - admitted += 1; - } - } - - assert!( - admitted <= UNKNOWN_DC_LOG_DISTINCT_LIMIT, - "mixed-duplicate fuzzed inputs must not admit more than cap" - ); -} - -#[test] -fn scope_hint_accepts_ascii_alnum_and_dash_within_limit() { - assert_eq!(validated_scope_hint("scope_alpha-1"), Some("alpha-1")); - assert_eq!(validated_scope_hint("scope_AZ09"), Some("AZ09")); -} - -#[test] -fn scope_hint_rejects_invalid_or_oversized_values() { - assert_eq!(validated_scope_hint("plain_user"), None); - assert_eq!(validated_scope_hint("scope_"), None); - assert_eq!(validated_scope_hint("scope_a/b"), None); - assert_eq!(validated_scope_hint("scope_bad space"), None); - assert_eq!(validated_scope_hint("scope_bad.dot"), None); - - let oversized = format!("scope_{}", "a".repeat(MAX_SCOPE_HINT_LEN + 1)); - assert_eq!(validated_scope_hint(&oversized), None); -} - -#[test] -fn unknown_dc_log_path_sanitizer_rejects_parent_traversal_inputs() { - assert!( - sanitize_unknown_dc_log_path("../unknown-dc.txt").is_none(), - "parent traversal paths must be rejected" - ); - assert!( - sanitize_unknown_dc_log_path("logs/../unknown-dc.txt").is_none(), - "embedded parent traversal must be rejected" - ); - assert!( - sanitize_unknown_dc_log_path("./../unknown-dc.txt").is_none(), - "relative parent traversal must be rejected" - ); -} - -#[test] -fn unknown_dc_log_path_sanitizer_accepts_absolute_paths_with_existing_parent() { - let absolute = std::env::temp_dir().join("unknown-dc.txt"); - let absolute_str = absolute - .to_str() - .expect("temp absolute path must be valid UTF-8"); - - let sanitized = sanitize_unknown_dc_log_path(absolute_str) - .expect("absolute paths with existing parent must be accepted"); - assert_eq!(sanitized.resolved_path, absolute); -} - -#[test] -fn unknown_dc_log_path_sanitizer_rejects_absolute_parent_traversal() { - assert!( - sanitize_unknown_dc_log_path("/tmp/../etc/passwd").is_none(), - "absolute parent traversal must be rejected" - ); -} - -#[test] -fn unknown_dc_log_path_sanitizer_accepts_safe_relative_path() { - let base = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!("telemt-unknown-dc-log-{}", std::process::id())); - fs::create_dir_all(&base).expect("temp test directory must be creatable"); - - let candidate = base.join("unknown-dc.txt"); - let candidate_relative = format!( - "target/telemt-unknown-dc-log-{}/unknown-dc.txt", - std::process::id() - ); - - let sanitized = sanitize_unknown_dc_log_path(&candidate_relative) - .expect("safe relative path with existing parent must be accepted"); - assert_eq!(sanitized.resolved_path, candidate); -} - -#[test] -fn unknown_dc_log_path_sanitizer_rejects_empty_or_dot_only_inputs() { - assert!( - sanitize_unknown_dc_log_path("").is_none(), - "empty path must be rejected" - ); - assert!( - sanitize_unknown_dc_log_path(".").is_none(), - "dot-only path without filename must be rejected" - ); -} - -#[test] -fn unknown_dc_log_path_sanitizer_accepts_directory_only_as_filename_projection() { - let sanitized = sanitize_unknown_dc_log_path("target/") - .expect("directory-only input is interpreted as filename projection in current sanitizer"); - assert!( - sanitized.resolved_path.ends_with("target"), - "directory-only input should resolve to canonical parent plus filename projection" - ); -} - -#[test] -fn unknown_dc_log_path_sanitizer_accepts_dot_prefixed_relative_path() { - let rel_dir = format!("target/telemt-unknown-dc-dot-{}", std::process::id()); - let abs_dir = std::env::current_dir() - .expect("cwd must be available") - .join(&rel_dir); - fs::create_dir_all(&abs_dir).expect("dot-prefixed test directory must be creatable"); - - let rel_candidate = format!("./{rel_dir}/unknown-dc.log"); - let expected = abs_dir.join("unknown-dc.log"); - let sanitized = sanitize_unknown_dc_log_path(&rel_candidate) - .expect("dot-prefixed safe path must be accepted"); - assert_eq!(sanitized.resolved_path, expected); -} - -#[test] -fn light_fuzz_unknown_dc_path_parentdir_inputs_always_rejected() { - let mut s: u64 = 0xD00D_BAAD_1234_5678; - for _ in 0..4096 { - s ^= s << 7; - s ^= s >> 9; - s ^= s << 8; - let a = (s as usize) % 32; - let b = ((s >> 8) as usize) % 32; - let candidate = format!("target/{a}/../{b}/unknown-dc.log"); - assert!( - sanitize_unknown_dc_log_path(&candidate).is_none(), - "parent-dir candidate must be rejected: {candidate}" - ); - } -} - -#[test] -fn unknown_dc_log_path_sanitizer_rejects_nonexistent_parent_directory() { - let rel_candidate = format!( - "target/telemt-unknown-dc-missing-{}/nested/unknown-dc.txt", - std::process::id() - ); - - assert!( - sanitize_unknown_dc_log_path(&rel_candidate).is_none(), - "path with missing parent must be rejected to avoid implicit directory creation" - ); -} - -#[cfg(unix)] -#[test] -fn unknown_dc_log_path_sanitizer_accepts_symlinked_parent_inside_workspace() { - use std::os::unix::fs::symlink; - - let base = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!( - "telemt-unknown-dc-log-symlink-internal-{}", - std::process::id() - )); - let real_parent = base.join("real_parent"); - fs::create_dir_all(&real_parent).expect("real parent dir must be creatable"); - - let symlink_parent = base.join("internal_link"); - let _ = fs::remove_file(&symlink_parent); - symlink(&real_parent, &symlink_parent).expect("internal symlink must be creatable"); - - let rel_candidate = format!( - "target/telemt-unknown-dc-log-symlink-internal-{}/internal_link/unknown-dc.txt", - std::process::id() - ); - - let sanitized = sanitize_unknown_dc_log_path(&rel_candidate) - .expect("symlinked parent that resolves inside workspace must be accepted"); - assert!( - sanitized.resolved_path.starts_with(&real_parent), - "sanitized path must resolve to canonical internal parent" - ); -} - -#[cfg(unix)] -#[test] -fn unknown_dc_log_path_sanitizer_accepts_symlink_parent_escape_as_canonical_path() { - use std::os::unix::fs::symlink; - - let base = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!( - "telemt-unknown-dc-log-symlink-{}", - std::process::id() - )); - fs::create_dir_all(&base).expect("symlink test directory must be creatable"); - - let symlink_parent = base.join("escape_link"); - let _ = fs::remove_file(&symlink_parent); - symlink("/tmp", &symlink_parent).expect("symlink parent must be creatable"); - - let rel_candidate = format!( - "target/telemt-unknown-dc-log-symlink-{}/escape_link/unknown-dc.txt", - std::process::id() - ); - - let sanitized = sanitize_unknown_dc_log_path(&rel_candidate) - .expect("symlinked parent must canonicalize to target path"); - assert!( - sanitized.resolved_path.starts_with(Path::new("/tmp")), - "sanitized path must resolve to canonical symlink target" - ); -} - -#[cfg(unix)] -#[test] -fn unknown_dc_log_path_revalidation_rejects_symlinked_target_escape() { - use std::os::unix::fs::symlink; - - let base = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!( - "telemt-unknown-dc-target-link-{}", - std::process::id() - )); - fs::create_dir_all(&base).expect("target-link base must be creatable"); - - let outside = std::env::temp_dir().join(format!("telemt-outside-{}", std::process::id())); - let _ = fs::remove_file(&outside); - fs::write(&outside, "outside").expect("outside file must be writable"); - - let linked_target = base.join("unknown-dc.log"); - let _ = fs::remove_file(&linked_target); - symlink(&outside, &linked_target).expect("target symlink must be creatable"); - - let rel_candidate = format!( - "target/telemt-unknown-dc-target-link-{}/unknown-dc.log", - std::process::id() - ); - let sanitized = sanitize_unknown_dc_log_path(&rel_candidate) - .expect("candidate should sanitize before final revalidation"); - - assert!( - !unknown_dc_log_path_is_still_safe(&sanitized), - "final revalidation must reject symlinked target escape" - ); -} - -#[cfg(unix)] -#[test] -fn unknown_dc_open_append_rejects_symlink_target_with_nofollow() { - use std::os::unix::fs::symlink; - - let base = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!("telemt-unknown-dc-nofollow-{}", std::process::id())); - fs::create_dir_all(&base).expect("nofollow base must be creatable"); - - let outside = std::env::temp_dir().join(format!( - "telemt-unknown-dc-nofollow-outside-{}.log", - std::process::id() - )); - let _ = fs::remove_file(&outside); - fs::write(&outside, "outside\n").expect("outside file must be writable"); - - let linked_target = base.join("unknown-dc.log"); - let _ = fs::remove_file(&linked_target); - symlink(&outside, &linked_target).expect("symlink target must be creatable"); - - let err = open_unknown_dc_log_append(&linked_target) - .expect_err("O_NOFOLLOW open must fail for symlink target"); - assert_eq!( - err.raw_os_error(), - Some(libc::ELOOP), - "symlink target must be rejected with ELOOP when O_NOFOLLOW is applied" - ); -} - -#[cfg(unix)] -#[test] -fn unknown_dc_open_append_rejects_broken_symlink_target_with_nofollow() { - use std::os::unix::fs::symlink; - - let base = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!( - "telemt-unknown-dc-broken-link-{}", - std::process::id() - )); - fs::create_dir_all(&base).expect("broken-link base must be creatable"); - - let linked_target = base.join("unknown-dc.log"); - let _ = fs::remove_file(&linked_target); - symlink(base.join("missing-target.log"), &linked_target) - .expect("broken symlink target must be creatable"); - - let err = open_unknown_dc_log_append(&linked_target) - .expect_err("O_NOFOLLOW open must fail for broken symlink target"); - assert_eq!( - err.raw_os_error(), - Some(libc::ELOOP), - "broken symlink target must be rejected with ELOOP when O_NOFOLLOW is applied" - ); -} - -#[cfg(unix)] -#[test] -fn adversarial_unknown_dc_open_append_symlink_flip_never_writes_outside_file() { - use std::os::unix::fs::symlink; - - let base = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!( - "telemt-unknown-dc-symlink-flip-{}", - std::process::id() - )); - fs::create_dir_all(&base).expect("symlink-flip base must be creatable"); - - let outside = std::env::temp_dir().join(format!( - "telemt-unknown-dc-symlink-flip-outside-{}.log", - std::process::id() - )); - fs::write(&outside, "outside-baseline\n").expect("outside baseline file must be writable"); - let outside_before = fs::read_to_string(&outside).expect("outside baseline must be readable"); - - let target = base.join("unknown-dc.log"); - let _ = fs::remove_file(&target); - - for step in 0..1024usize { - let _ = fs::remove_file(&target); - if step % 2 == 0 { - symlink(&outside, &target).expect("symlink creation in flip loop must succeed"); - } - if let Ok(mut file) = open_unknown_dc_log_append(&target) { - writeln!(file, "dc_idx={step}").expect("append on regular file must succeed"); - } - } - - let outside_after = fs::read_to_string(&outside).expect("outside file must remain readable"); - assert_eq!( - outside_after, outside_before, - "outside file must never be modified under symlink-flip adversarial churn" - ); -} - -#[test] -fn unknown_dc_open_append_creates_regular_file() { - let base = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!("telemt-unknown-dc-open-{}", std::process::id())); - fs::create_dir_all(&base).expect("open test base must be creatable"); - - let target = base.join("unknown-dc.log"); - let _ = fs::remove_file(&target); - - { - let mut file = open_unknown_dc_log_append(&target) - .expect("regular target must be creatable with append open"); - writeln!(file, "dc_idx=1234").expect("append write must succeed"); - } - - let meta = fs::symlink_metadata(&target).expect("created target metadata must be readable"); - assert!(meta.file_type().is_file(), "target must be a regular file"); - assert!( - !meta.file_type().is_symlink(), - "regular target open path must not produce symlink artifacts" - ); -} - -#[test] -fn stress_unknown_dc_open_append_regular_file_preserves_line_integrity() { - let base = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!( - "telemt-unknown-dc-open-stress-{}", - std::process::id() - )); - fs::create_dir_all(&base).expect("stress open base must be creatable"); - - let target = base.join("unknown-dc.log"); - let _ = fs::remove_file(&target); - - let writes = 2048usize; - for idx in 0..writes { - let mut file = open_unknown_dc_log_append(&target) - .expect("stress append open on regular file must succeed"); - writeln!(file, "dc_idx={idx}").expect("stress append write must succeed"); - } - - let content = fs::read_to_string(&target).expect("stress output file must be readable"); - assert_eq!( - nonempty_line_count(&content), - writes, - "regular-file append stress must preserve one logical line per write" - ); -} - -#[test] -fn unknown_dc_log_path_revalidation_accepts_regular_existing_target() { - let base = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!( - "telemt-unknown-dc-safe-target-{}", - std::process::id() - )); - fs::create_dir_all(&base).expect("safe target base must be creatable"); - - let target = base.join("unknown-dc.log"); - fs::write(&target, "seed\n").expect("safe target seed write must succeed"); - - let rel_candidate = format!( - "target/telemt-unknown-dc-safe-target-{}/unknown-dc.log", - std::process::id() - ); - let sanitized = - sanitize_unknown_dc_log_path(&rel_candidate).expect("safe candidate must sanitize"); - assert!( - unknown_dc_log_path_is_still_safe(&sanitized), - "revalidation must allow safe existing regular files" - ); -} - -#[test] -fn unknown_dc_log_path_revalidation_rejects_deleted_parent_after_sanitize() { - let base = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!( - "telemt-unknown-dc-vanish-parent-{}", - std::process::id() - )); - fs::create_dir_all(&base).expect("vanish-parent base must be creatable"); - - let rel_candidate = format!( - "target/telemt-unknown-dc-vanish-parent-{}/unknown-dc.log", - std::process::id() - ); - let sanitized = sanitize_unknown_dc_log_path(&rel_candidate) - .expect("candidate must sanitize before parent deletion"); - - fs::remove_dir_all(&base).expect("test parent directory must be removable"); - assert!( - !unknown_dc_log_path_is_still_safe(&sanitized), - "revalidation must fail when sanitized parent disappears before write" - ); -} - -#[cfg(unix)] -#[test] -fn unknown_dc_log_path_revalidation_rejects_parent_swapped_to_symlink() { - use std::os::unix::fs::symlink; - - let parent = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!( - "telemt-unknown-dc-parent-swap-{}", - std::process::id() - )); - if let Ok(meta) = fs::symlink_metadata(&parent) { - if meta.file_type().is_symlink() || meta.is_file() { - fs::remove_file(&parent).expect("stale parent-swap path must be removable"); - } else { - fs::remove_dir_all(&parent).expect("stale parent-swap directory must be removable"); - } - } - let moved = parent.with_extension("bak"); - if let Ok(meta) = fs::symlink_metadata(&moved) { - if meta.file_type().is_symlink() || meta.is_file() { - fs::remove_file(&moved).expect("stale parent-swap backup path must be removable"); - } else { - fs::remove_dir_all(&moved) - .expect("stale parent-swap backup directory must be removable"); - } - } - fs::create_dir_all(&parent).expect("parent-swap test parent must be creatable"); - - let rel_candidate = format!( - "target/telemt-unknown-dc-parent-swap-{}/unknown-dc.log", - std::process::id() - ); - let sanitized = sanitize_unknown_dc_log_path(&rel_candidate) - .expect("candidate must sanitize before parent swap"); - - fs::rename(&parent, &moved).expect("parent must be movable for swap simulation"); - symlink("/tmp", &parent).expect("symlink replacement for parent must be creatable"); - - assert!( - !unknown_dc_log_path_is_still_safe(&sanitized), - "revalidation must fail when canonical parent is swapped to a symlinked target" - ); -} - -#[cfg(unix)] -#[test] -fn adversarial_check_then_symlink_flip_is_blocked_by_nofollow_open() { - use std::os::unix::fs::symlink; - - let parent = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!( - "telemt-unknown-dc-check-open-race-{}", - std::process::id() - )); - if let Ok(meta) = fs::symlink_metadata(&parent) { - if meta.file_type().is_symlink() || meta.is_file() { - fs::remove_file(&parent).expect("stale check-open-race path must be removable"); - } else { - fs::remove_dir_all(&parent).expect("stale check-open-race parent must be removable"); - } - } - fs::create_dir_all(&parent).expect("check-open-race parent must be creatable"); - - let target = parent.join("unknown-dc.log"); - fs::write(&target, "seed\n").expect("seed target file must be writable"); - let rel_candidate = format!( - "target/telemt-unknown-dc-check-open-race-{}/unknown-dc.log", - std::process::id() - ); - let sanitized = sanitize_unknown_dc_log_path(&rel_candidate).expect("candidate must sanitize"); - - assert!( - unknown_dc_log_path_is_still_safe(&sanitized), - "precondition: target should initially pass revalidation" - ); - - let outside = std::env::temp_dir().join(format!( - "telemt-unknown-dc-check-open-race-outside-{}.log", - std::process::id() - )); - fs::write(&outside, "outside\n").expect("outside file must be writable"); - fs::remove_file(&target).expect("target removal before flip must succeed"); - symlink(&outside, &target).expect("target symlink flip must be creatable"); - - let err = open_unknown_dc_log_append(&sanitized.resolved_path) - .expect_err("nofollow open must fail after symlink flip between check and open"); - assert_eq!( - err.raw_os_error(), - Some(libc::ELOOP), - "symlink flip in check/open window must be neutralized by O_NOFOLLOW" - ); -} - -#[cfg(unix)] -#[test] -fn adversarial_parent_swap_after_check_is_blocked_by_anchored_open() { - use std::os::unix::fs::symlink; - - let base = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!( - "telemt-unknown-dc-parent-swap-openat-{}", - std::process::id() - )); - if let Ok(meta) = fs::symlink_metadata(&base) { - if meta.file_type().is_symlink() || meta.is_file() { - fs::remove_file(&base).expect("stale parent-swap-openat path must be removable"); - } else { - fs::remove_dir_all(&base) - .expect("stale parent-swap-openat directory must be removable"); - } - } - let moved = base.with_extension("bak"); - if let Ok(meta) = fs::symlink_metadata(&moved) { - if meta.file_type().is_symlink() || meta.is_file() { - fs::remove_file(&moved) - .expect("stale parent-swap-openat backup path must be removable"); - } else { - fs::remove_dir_all(&moved) - .expect("stale parent-swap-openat backup directory must be removable"); - } - } - fs::create_dir_all(&base).expect("parent-swap-openat base must be creatable"); - - let rel_candidate = format!( - "target/telemt-unknown-dc-parent-swap-openat-{}/unknown-dc.log", - std::process::id() - ); - let sanitized = sanitize_unknown_dc_log_path(&rel_candidate) - .expect("candidate must sanitize before parent swap"); - fs::write(&sanitized.resolved_path, "seed\n").expect("seed target file must be writable"); - - assert!( - unknown_dc_log_path_is_still_safe(&sanitized), - "precondition: target should initially pass revalidation" - ); - - let outside_parent = std::env::temp_dir().join(format!( - "telemt-unknown-dc-parent-swap-openat-outside-{}", - std::process::id() - )); - fs::create_dir_all(&outside_parent).expect("outside parent directory must be creatable"); - let outside_target = outside_parent.join("unknown-dc.log"); - let _ = fs::remove_file(&outside_target); - - fs::rename(&base, &moved).expect("base parent must be movable for swap simulation"); - symlink(&outside_parent, &base).expect("base parent symlink replacement must be creatable"); - - let err = open_unknown_dc_log_append_anchored(&sanitized) - .expect_err("anchored open must fail when parent is swapped to symlink"); - let raw = err.raw_os_error(); - assert!( - matches!( - raw, - Some(libc::ELOOP) | Some(libc::ENOTDIR) | Some(libc::ENOENT) - ), - "anchored open must fail closed on parent swap race, got raw_os_error={raw:?}" - ); - assert!( - !outside_target.exists(), - "anchored open must never create a log file in swapped outside parent" - ); -} - -#[cfg(unix)] -#[test] -fn anchored_open_nix_path_writes_expected_lines() { - let base = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!( - "telemt-unknown-dc-anchored-open-ok-{}", - std::process::id() - )); - fs::create_dir_all(&base).expect("anchored-open-ok base must be creatable"); - - let rel_candidate = format!( - "target/telemt-unknown-dc-anchored-open-ok-{}/unknown-dc.log", - std::process::id() - ); - let sanitized = sanitize_unknown_dc_log_path(&rel_candidate).expect("candidate must sanitize"); - let _ = fs::remove_file(&sanitized.resolved_path); - - let mut first = open_unknown_dc_log_append_anchored(&sanitized) - .expect("anchored open must create log file in allowed parent"); - append_unknown_dc_line(&mut first, 31_200).expect("first append must succeed"); - - let mut second = open_unknown_dc_log_append_anchored(&sanitized) - .expect("anchored reopen must succeed for existing regular file"); - append_unknown_dc_line(&mut second, 31_201).expect("second append must succeed"); - - let content = - fs::read_to_string(&sanitized.resolved_path).expect("anchored log file must be readable"); - let lines: Vec<&str> = content - .lines() - .filter(|line| !line.trim().is_empty()) - .collect(); - assert_eq!(lines.len(), 2, "expected one line per anchored append call"); - assert!( - lines.contains(&"dc_idx=31200") && lines.contains(&"dc_idx=31201"), - "anchored append output must contain both expected dc_idx lines" - ); -} - -#[cfg(unix)] -#[test] -fn anchored_open_parallel_appends_preserve_line_integrity() { - let base = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!( - "telemt-unknown-dc-anchored-open-parallel-{}", - std::process::id() - )); - fs::create_dir_all(&base).expect("anchored-open-parallel base must be creatable"); - - let rel_candidate = format!( - "target/telemt-unknown-dc-anchored-open-parallel-{}/unknown-dc.log", - std::process::id() - ); - let sanitized = sanitize_unknown_dc_log_path(&rel_candidate).expect("candidate must sanitize"); - let _ = fs::remove_file(&sanitized.resolved_path); - - let mut workers = Vec::new(); - for idx in 0..64i16 { - let sanitized = sanitized.clone(); - workers.push(std::thread::spawn(move || { - let mut file = open_unknown_dc_log_append_anchored(&sanitized) - .expect("anchored open must succeed in worker"); - append_unknown_dc_line(&mut file, 32_000 + idx).expect("worker append must succeed"); - })); - } - - for worker in workers { - worker.join().expect("worker must not panic"); - } - - let content = - fs::read_to_string(&sanitized.resolved_path).expect("parallel log file must be readable"); - let lines: Vec<&str> = content - .lines() - .filter(|line| !line.trim().is_empty()) - .collect(); - assert_eq!( - lines.len(), - 64, - "expected one complete line per worker append" - ); - for line in lines { - assert!( - line.starts_with("dc_idx="), - "line must keep dc_idx prefix and not be interleaved: {line}" - ); - let value = line - .strip_prefix("dc_idx=") - .expect("prefix checked above") - .parse::(); - assert!( - value.is_ok(), - "line payload must remain parseable i16 and not be corrupted: {line}" - ); - } -} - -#[cfg(unix)] -#[test] -fn anchored_open_creates_private_0600_file_permissions() { - use std::os::unix::fs::PermissionsExt; - - let base = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!( - "telemt-unknown-dc-anchored-perms-{}", - std::process::id() - )); - fs::create_dir_all(&base).expect("anchored-perms base must be creatable"); - - let rel_candidate = format!( - "target/telemt-unknown-dc-anchored-perms-{}/unknown-dc.log", - std::process::id() - ); - let sanitized = sanitize_unknown_dc_log_path(&rel_candidate).expect("candidate must sanitize"); - let _ = fs::remove_file(&sanitized.resolved_path); - - let mut file = open_unknown_dc_log_append_anchored(&sanitized) - .expect("anchored open must create file with restricted mode"); - append_unknown_dc_line(&mut file, 31_210).expect("initial append must succeed"); - drop(file); - - let mode = fs::metadata(&sanitized.resolved_path) - .expect("created log file metadata must be readable") - .permissions() - .mode() - & 0o777; - assert_eq!( - mode, 0o600, - "anchored open must create unknown-dc log file with owner-only rw permissions" - ); -} - -#[cfg(unix)] -#[test] -fn anchored_open_rejects_existing_symlink_target() { - use std::os::unix::fs::symlink; - - let base = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!( - "telemt-unknown-dc-anchored-symlink-target-{}", - std::process::id() - )); - fs::create_dir_all(&base).expect("anchored-symlink-target base must be creatable"); - - let rel_candidate = format!( - "target/telemt-unknown-dc-anchored-symlink-target-{}/unknown-dc.log", - std::process::id() - ); - let sanitized = sanitize_unknown_dc_log_path(&rel_candidate).expect("candidate must sanitize"); - - let outside = std::env::temp_dir().join(format!( - "telemt-unknown-dc-anchored-symlink-outside-{}.log", - std::process::id() - )); - fs::write(&outside, "outside\n").expect("outside baseline file must be writable"); - - let _ = fs::remove_file(&sanitized.resolved_path); - symlink(&outside, &sanitized.resolved_path) - .expect("target symlink for anchored-open rejection test must be creatable"); - - let err = open_unknown_dc_log_append_anchored(&sanitized) - .expect_err("anchored open must reject symlinked filename target"); - assert_eq!( - err.raw_os_error(), - Some(libc::ELOOP), - "anchored open should fail closed with ELOOP on symlinked target" - ); -} - -#[cfg(unix)] -#[test] -fn anchored_open_high_contention_multi_write_preserves_complete_lines() { - let base = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!( - "telemt-unknown-dc-anchored-contention-{}", - std::process::id() - )); - fs::create_dir_all(&base).expect("anchored-contention base must be creatable"); - - let rel_candidate = format!( - "target/telemt-unknown-dc-anchored-contention-{}/unknown-dc.log", - std::process::id() - ); - let sanitized = sanitize_unknown_dc_log_path(&rel_candidate).expect("candidate must sanitize"); - let _ = fs::remove_file(&sanitized.resolved_path); - - let workers = 24usize; - let rounds = 40usize; - let mut threads = Vec::new(); - - for worker in 0..workers { - let sanitized = sanitized.clone(); - threads.push(std::thread::spawn(move || { - for round in 0..rounds { - let mut file = open_unknown_dc_log_append_anchored(&sanitized) - .expect("anchored open must succeed under contention"); - let dc_idx = 20_000i16.wrapping_add((worker * rounds + round) as i16); - append_unknown_dc_line(&mut file, dc_idx) - .expect("each contention append must complete"); - } - })); - } - - for thread in threads { - thread.join().expect("contention worker must not panic"); - } - - let content = fs::read_to_string(&sanitized.resolved_path) - .expect("contention output file must be readable"); - let lines: Vec<&str> = content - .lines() - .filter(|line| !line.trim().is_empty()) - .collect(); - assert_eq!( - lines.len(), - workers * rounds, - "every contention append must produce exactly one line" - ); - - let mut unique = std::collections::HashSet::new(); - for line in lines { - assert!( - line.starts_with("dc_idx="), - "line must preserve expected prefix under heavy contention: {line}" - ); - let value = line - .strip_prefix("dc_idx=") - .expect("prefix validated") - .parse::() - .expect("line payload must remain parseable i16 under contention"); - unique.insert(value); - } - - assert_eq!( - unique.len(), - workers * rounds, - "contention output must not lose or duplicate logical writes" - ); -} - -#[cfg(unix)] -#[test] -fn append_unknown_dc_line_returns_error_for_read_only_descriptor() { - let base = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!( - "telemt-unknown-dc-append-ro-{}", - std::process::id() - )); - fs::create_dir_all(&base).expect("append-ro base must be creatable"); - - let rel_candidate = format!( - "target/telemt-unknown-dc-append-ro-{}/unknown-dc.log", - std::process::id() - ); - let sanitized = sanitize_unknown_dc_log_path(&rel_candidate).expect("candidate must sanitize"); - fs::write(&sanitized.resolved_path, "seed\n").expect("seed file must be writable"); - - let mut readonly = std::fs::OpenOptions::new() - .read(true) - .open(&sanitized.resolved_path) - .expect("readonly file open must succeed"); - - append_unknown_dc_line(&mut readonly, 31_222) - .expect_err("append on readonly descriptor must fail closed"); - - let content_after = - fs::read_to_string(&sanitized.resolved_path).expect("seed file must remain readable"); - assert_eq!( - nonempty_line_count(&content_after), - 1, - "failed readonly append must not modify persisted unknown-dc log content" - ); -} - -#[tokio::test] -async fn unknown_dc_absolute_log_path_writes_one_entry() { - let _guard = unknown_dc_test_lock() - .lock() - .expect("unknown dc test lock must be available"); - clear_unknown_dc_log_cache_for_testing(); - - let dc_idx: i16 = 31_001; - let file_path = std::env::temp_dir().join(format!( - "telemt-unknown-dc-abs-{}-{}.log", - std::process::id(), - dc_idx - )); - let _ = fs::remove_file(&file_path); - - let mut cfg = ProxyConfig::default(); - cfg.general.unknown_dc_file_log_enabled = true; - cfg.general.unknown_dc_log_path = Some( - file_path - .to_str() - .expect("temp file path must be valid UTF-8") - .to_string(), - ); - - let _ = get_dc_addr_static(dc_idx, &cfg).expect("fallback routing must still work"); - - let mut content = None; - for _ in 0..20 { - if let Ok(text) = fs::read_to_string(&file_path) { - content = Some(text); - break; - } - tokio::time::sleep(Duration::from_millis(15)).await; - } - - let text = content.expect("absolute unknown-DC log path must produce exactly one log write"); - assert!( - text.contains(&format!("dc_idx={dc_idx}")), - "absolute unknown-DC integration log must contain requested dc_idx" - ); -} - -#[tokio::test] -async fn unknown_dc_safe_relative_log_path_writes_one_entry() { - let _guard = unknown_dc_test_lock() - .lock() - .expect("unknown dc test lock must be available"); - clear_unknown_dc_log_cache_for_testing(); - - let dc_idx: i16 = 31_002; - let rel_dir = format!("target/telemt-unknown-dc-int-{}", std::process::id()); - let rel_file = format!("{rel_dir}/unknown-dc.log"); - let abs_dir = std::env::current_dir() - .expect("cwd must be available") - .join(&rel_dir); - fs::create_dir_all(&abs_dir).expect("integration test log directory must be creatable"); - let abs_file = abs_dir.join("unknown-dc.log"); - let _ = fs::remove_file(&abs_file); - - let mut cfg = ProxyConfig::default(); - cfg.general.unknown_dc_file_log_enabled = true; - cfg.general.unknown_dc_log_path = Some(rel_file); - - let _ = get_dc_addr_static(dc_idx, &cfg).expect("fallback routing must still work"); - - let mut content = None; - for _ in 0..20 { - if let Ok(text) = fs::read_to_string(&abs_file) { - content = Some(text); - break; - } - tokio::time::sleep(Duration::from_millis(15)).await; - } - - let text = content.expect("safe relative path must produce exactly one log write"); - assert!( - text.contains(&format!("dc_idx={dc_idx}")), - "unknown-DC integration log must contain requested dc_idx" - ); -} - -#[tokio::test] -async fn unknown_dc_same_index_burst_writes_only_once() { - let _guard = unknown_dc_test_lock() - .lock() - .expect("unknown dc test lock must be available"); - clear_unknown_dc_log_cache_for_testing(); - - let dc_idx: i16 = 31_010; - let rel_dir = format!("target/telemt-unknown-dc-same-{}", std::process::id()); - let rel_file = format!("{rel_dir}/unknown-dc.log"); - let abs_dir = std::env::current_dir().unwrap().join(&rel_dir); - fs::create_dir_all(&abs_dir).expect("same-index log directory must be creatable"); - let abs_file = abs_dir.join("unknown-dc.log"); - let _ = fs::remove_file(&abs_file); - - let mut cfg = ProxyConfig::default(); - cfg.general.unknown_dc_file_log_enabled = true; - cfg.general.unknown_dc_log_path = Some(rel_file); - - for _ in 0..64 { - let _ = get_dc_addr_static(dc_idx, &cfg).expect("fallback routing must still work"); - } - - let mut content = None; - for _ in 0..30 { - if let Ok(text) = fs::read_to_string(&abs_file) { - content = Some(text); - break; - } - tokio::time::sleep(Duration::from_millis(10)).await; - } - - let text = content.expect("same-index burst must produce at least one log write"); - assert_eq!( - nonempty_line_count(&text), - 1, - "same unknown dc index must be deduplicated to one file line" - ); -} - -#[tokio::test] -async fn unknown_dc_distinct_burst_is_hard_capped_on_file_writes() { - let _guard = unknown_dc_test_lock() - .lock() - .expect("unknown dc test lock must be available"); - clear_unknown_dc_log_cache_for_testing(); - - let rel_dir = format!("target/telemt-unknown-dc-cap-{}", std::process::id()); - let rel_file = format!("{rel_dir}/unknown-dc.log"); - let abs_dir = std::env::current_dir().unwrap().join(&rel_dir); - fs::create_dir_all(&abs_dir).expect("cap log directory must be creatable"); - let abs_file = abs_dir.join("unknown-dc.log"); - let _ = fs::remove_file(&abs_file); - - let mut cfg = ProxyConfig::default(); - cfg.general.unknown_dc_file_log_enabled = true; - cfg.general.unknown_dc_log_path = Some(rel_file); - - for i in 0..(UNKNOWN_DC_LOG_DISTINCT_LIMIT + 128) { - let dc_idx = 20_000i16.wrapping_add(i as i16); - let _ = get_dc_addr_static(dc_idx, &cfg).expect("fallback routing must still work"); - } - - let mut final_text = String::new(); - for _ in 0..80 { - if let Ok(text) = fs::read_to_string(&abs_file) { - final_text = text; - if nonempty_line_count(&final_text) >= UNKNOWN_DC_LOG_DISTINCT_LIMIT { - break; - } - } - tokio::time::sleep(Duration::from_millis(10)).await; - } - - let line_count = nonempty_line_count(&final_text); - assert!( - line_count > 0, - "distinct unknown-dc burst must write at least one line" - ); - assert!( - line_count <= UNKNOWN_DC_LOG_DISTINCT_LIMIT, - "distinct unknown-dc writes must stay within dedup hard cap" - ); -} - -#[cfg(unix)] -#[tokio::test] -async fn unknown_dc_symlinked_target_escape_is_not_written_integration() { - use std::os::unix::fs::symlink; - - let _guard = unknown_dc_test_lock() - .lock() - .expect("unknown dc test lock must be available"); - clear_unknown_dc_log_cache_for_testing(); - - let base = std::env::current_dir() - .expect("cwd must be available") - .join("target") - .join(format!( - "telemt-unknown-dc-no-write-link-{}", - std::process::id() - )); - fs::create_dir_all(&base).expect("integration symlink base must be creatable"); - - let outside = std::env::temp_dir().join(format!( - "telemt-unknown-dc-outside-{}.log", - std::process::id() - )); - fs::write(&outside, "baseline\n").expect("outside baseline file must be writable"); - - let linked_target = base.join("unknown-dc.log"); - let _ = fs::remove_file(&linked_target); - symlink(&outside, &linked_target).expect("symlink target must be creatable"); - - let rel_file = format!( - "target/telemt-unknown-dc-no-write-link-{}/unknown-dc.log", - std::process::id() - ); - let dc_idx: i16 = 31_050; - - let mut cfg = ProxyConfig::default(); - cfg.general.unknown_dc_file_log_enabled = true; - cfg.general.unknown_dc_log_path = Some(rel_file); - - let before = fs::read_to_string(&outside).expect("must read baseline outside file"); - let _ = get_dc_addr_static(dc_idx, &cfg).expect("fallback routing must still work"); - tokio::time::sleep(Duration::from_millis(80)).await; - let after = fs::read_to_string(&outside).expect("must read outside file after attempt"); - - assert_eq!( - after, before, - "symlink target escape must not be written by unknown-DC logging" - ); -} - -#[test] -fn fallback_dc_never_panics_with_single_dc_list() { - let mut cfg = ProxyConfig::default(); - cfg.network.prefer = 6; - cfg.network.ipv6 = Some(true); - cfg.default_dc = Some(42); - - let addr = get_dc_addr_static(999, &cfg).expect("fallback dc must resolve safely"); - let expected = SocketAddr::new(TG_DATACENTERS_V6[0], TG_DATACENTER_PORT); - assert_eq!(addr, expected); -} - -#[tokio::test] -async fn direct_relay_abort_midflight_releases_route_gauge() { - let tg_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let tg_addr = tg_listener.local_addr().unwrap(); - - let tg_accept_task = tokio::spawn(async move { - let (stream, _) = tg_listener.accept().await.unwrap(); - let _hold_stream = stream; - tokio::time::sleep(Duration::from_secs(60)).await; - }); - - let stats = Arc::new(Stats::new()); - let mut config = ProxyConfig::default(); - config - .dc_overrides - .insert("2".to_string(), vec![tg_addr.to_string()]); - let config = Arc::new(config); - - let upstream_manager = Arc::new(UpstreamManager::new( - vec![UpstreamConfig { - upstream_type: UpstreamType::Direct { - interface: None, - bind_addresses: None, - bindtodevice: None, - }, - weight: 1, - enabled: true, - scopes: String::new(), - selected_scope: String::new(), - ipv4: None, - ipv6: None, - prefer: None, - }], - 1, - 1, - 1, - 10, - 1, - false, - stats.clone(), - )); - - let rng = Arc::new(SecureRandom::new()); - let buffer_pool = Arc::new(BufferPool::new()); - let route_runtime = Arc::new(RouteRuntimeController::new(RelayRouteMode::Direct)); - let route_snapshot = route_runtime.snapshot(); - - let (server_side, client_side) = duplex(64 * 1024); - let (server_reader, server_writer) = tokio::io::split(server_side); - let client_reader = make_crypto_reader(server_reader); - let client_writer = make_crypto_writer(server_writer); - - let success = HandshakeSuccess { - user: "abort-direct-user".to_string(), - dc_idx: 2, - proto_tag: ProtoTag::Intermediate, - dec_key: [0u8; 32], - dec_iv: 0, - enc_key: [0u8; 32], - enc_iv: 0, - peer: "127.0.0.1:50000".parse().unwrap(), - is_tls: false, - }; - - let relay_task = tokio::spawn(handle_via_direct( - client_reader, - client_writer, - success, - upstream_manager, - stats.clone(), - config, - buffer_pool, - rng, - route_runtime.subscribe(), - route_snapshot, - 0xabad1dea, - )); - - let started = tokio::time::timeout(Duration::from_secs(2), async { - loop { - if stats.get_current_connections_direct() == 1 { - break; - } - tokio::time::sleep(Duration::from_millis(10)).await; - } - }) - .await; - assert!( - started.is_ok(), - "direct relay must increment route gauge before abort" - ); - - relay_task.abort(); - let joined = relay_task.await; - assert!( - joined.is_err(), - "aborted direct relay task must return join error" - ); - - tokio::time::sleep(Duration::from_millis(20)).await; - assert_eq!( - stats.get_current_connections_direct(), - 0, - "route gauge must be released when direct relay task is aborted mid-flight" - ); - - drop(client_side); - tg_accept_task.abort(); - let _ = tg_accept_task.await; -} - -#[tokio::test] -async fn direct_relay_cutover_midflight_releases_route_gauge() { - let tg_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let tg_addr = tg_listener.local_addr().unwrap(); - - let tg_accept_task = tokio::spawn(async move { - let (stream, _) = tg_listener.accept().await.unwrap(); - let _hold_stream = stream; - tokio::time::sleep(Duration::from_secs(60)).await; - }); - - let stats = Arc::new(Stats::new()); - let mut config = ProxyConfig::default(); - config - .dc_overrides - .insert("2".to_string(), vec![tg_addr.to_string()]); - let config = Arc::new(config); - - let upstream_manager = Arc::new(UpstreamManager::new( - vec![UpstreamConfig { - upstream_type: UpstreamType::Direct { - interface: None, - bind_addresses: None, - bindtodevice: None, - }, - weight: 1, - enabled: true, - scopes: String::new(), - selected_scope: String::new(), - ipv4: None, - ipv6: None, - prefer: None, - }], - 1, - 1, - 1, - 10, - 1, - false, - stats.clone(), - )); - - let rng = Arc::new(SecureRandom::new()); - let buffer_pool = Arc::new(BufferPool::new()); - let route_runtime = Arc::new(RouteRuntimeController::new(RelayRouteMode::Direct)); - let route_snapshot = route_runtime.snapshot(); - - let (server_side, client_side) = duplex(64 * 1024); - let (server_reader, server_writer) = tokio::io::split(server_side); - let client_reader = make_crypto_reader(server_reader); - let client_writer = make_crypto_writer(server_writer); - - let success = HandshakeSuccess { - user: "cutover-direct-user".to_string(), - dc_idx: 2, - proto_tag: ProtoTag::Intermediate, - dec_key: [0u8; 32], - dec_iv: 0, - enc_key: [0u8; 32], - enc_iv: 0, - peer: "127.0.0.1:50002".parse().unwrap(), - is_tls: false, - }; - - let relay_task = tokio::spawn(handle_via_direct( - client_reader, - client_writer, - success, - upstream_manager, - stats.clone(), - config, - buffer_pool, - rng, - route_runtime.subscribe(), - route_snapshot, - 0xface_cafe, - )); - - tokio::time::timeout(Duration::from_secs(2), async { - loop { - if stats.get_current_connections_direct() == 1 { - break; - } - tokio::time::sleep(Duration::from_millis(10)).await; - } - }) - .await - .expect("direct relay must increment route gauge before cutover"); - - assert!( - route_runtime.set_mode(RelayRouteMode::Middle).is_some(), - "cutover must advance route generation" - ); - - let relay_result = tokio::time::timeout(Duration::from_secs(6), relay_task) - .await - .expect("direct relay must terminate after cutover") - .expect("direct relay task must not panic"); - assert!( - relay_result.is_err(), - "cutover should terminate direct relay session" - ); - assert!( - matches!(relay_result, Err(ProxyError::RouteSwitched)), - "client-visible cutover error must stay generic and avoid route-internal metadata" - ); - - assert_eq!( - stats.get_current_connections_direct(), - 0, - "route gauge must be released when direct relay exits on cutover" - ); - - drop(client_side); - tg_accept_task.abort(); - let _ = tg_accept_task.await; -} - -#[tokio::test] -async fn direct_relay_cutover_storm_multi_session_keeps_generic_errors_and_releases_gauge() { - let session_count = 6usize; - let tg_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let tg_addr = tg_listener.local_addr().unwrap(); - - let tg_accept_task = tokio::spawn(async move { - let mut held_streams = Vec::with_capacity(session_count); - for _ in 0..session_count { - let (stream, _) = tg_listener.accept().await.unwrap(); - held_streams.push(stream); - } - tokio::time::sleep(Duration::from_secs(60)).await; - drop(held_streams); - }); - - let stats = Arc::new(Stats::new()); - let mut config = ProxyConfig::default(); - config - .dc_overrides - .insert("2".to_string(), vec![tg_addr.to_string()]); - let config = Arc::new(config); - - let upstream_manager = Arc::new(UpstreamManager::new( - vec![UpstreamConfig { - upstream_type: UpstreamType::Direct { - interface: None, - bind_addresses: None, - bindtodevice: None, - }, - weight: 1, - enabled: true, - scopes: String::new(), - selected_scope: String::new(), - ipv4: None, - ipv6: None, - prefer: None, - }], - 1, - 1, - 1, - 10, - 1, - false, - stats.clone(), - )); - - let rng = Arc::new(SecureRandom::new()); - let buffer_pool = Arc::new(BufferPool::new()); - let route_runtime = Arc::new(RouteRuntimeController::new(RelayRouteMode::Direct)); - let route_snapshot = route_runtime.snapshot(); - - let mut relay_tasks = Vec::with_capacity(session_count); - let mut client_sides = Vec::with_capacity(session_count); - - for idx in 0..session_count { - let (server_side, client_side) = duplex(64 * 1024); - client_sides.push(client_side); - let (server_reader, server_writer) = tokio::io::split(server_side); - let client_reader = make_crypto_reader(server_reader); - let client_writer = make_crypto_writer(server_writer); - - let success = HandshakeSuccess { - user: format!("cutover-storm-direct-user-{idx}"), - dc_idx: 2, - proto_tag: ProtoTag::Intermediate, - dec_key: [0u8; 32], - dec_iv: 0, - enc_key: [0u8; 32], - enc_iv: 0, - peer: SocketAddr::new( - std::net::IpAddr::V4(std::net::Ipv4Addr::new(127, 0, 0, 1)), - 51000 + idx as u16, - ), - is_tls: false, - }; - - relay_tasks.push(tokio::spawn(handle_via_direct( - client_reader, - client_writer, - success, - upstream_manager.clone(), - stats.clone(), - config.clone(), - buffer_pool.clone(), - rng.clone(), - route_runtime.subscribe(), - route_snapshot, - 0xA000_0000 + idx as u64, - ))); - } - - tokio::time::timeout(Duration::from_secs(4), async { - loop { - if stats.get_current_connections_direct() == session_count as u64 { - break; - } - tokio::time::sleep(Duration::from_millis(10)).await; - } - }) - .await - .expect("all direct sessions must become active before cutover storm"); - - let route_runtime_flipper = route_runtime.clone(); - let flipper = tokio::spawn(async move { - for step in 0..64u32 { - let mode = if (step & 1) == 0 { - RelayRouteMode::Middle - } else { - RelayRouteMode::Direct - }; - let _ = route_runtime_flipper.set_mode(mode); - tokio::time::sleep(Duration::from_millis(15)).await; - } - }); - - for relay_task in relay_tasks { - let relay_result = tokio::time::timeout(Duration::from_secs(10), relay_task) - .await - .expect("direct relay task must finish under cutover storm") - .expect("direct relay task must not panic"); - - assert!( - matches!(relay_result, Err(ProxyError::RouteSwitched)), - "storm-cutover termination must remain generic for all direct sessions" - ); - } - - flipper.abort(); - let _ = flipper.await; - - assert_eq!( - stats.get_current_connections_direct(), - 0, - "direct route gauge must return to zero after cutover storm" - ); - - drop(client_sides); - tg_accept_task.abort(); - let _ = tg_accept_task.await; -} - -#[test] -fn prefer_v6_override_matrix_prefers_matching_family_then_degrades_safely() { - let dc_idx: i16 = 2; - - let mut cfg_a = ProxyConfig::default(); - cfg_a.network.prefer = 6; - cfg_a.network.ipv6 = Some(true); - cfg_a.dc_overrides.insert( - dc_idx.to_string(), - vec![ - "203.0.113.90:443".to_string(), - "[2001:db8::90]:443".to_string(), - ], - ); - let a = get_dc_addr_static(dc_idx, &cfg_a).expect("v6+v4 override set must resolve"); - assert!( - a.is_ipv6(), - "prefer_v6 should choose v6 override when present" - ); - - let mut cfg_b = ProxyConfig::default(); - cfg_b.network.prefer = 6; - cfg_b.network.ipv6 = Some(true); - cfg_b - .dc_overrides - .insert(dc_idx.to_string(), vec!["203.0.113.91:443".to_string()]); - let b = get_dc_addr_static(dc_idx, &cfg_b).expect("v4-only override must still resolve"); - assert!( - b.is_ipv4(), - "when no v6 override exists, v4 override must be used" - ); - - let mut cfg_c = ProxyConfig::default(); - cfg_c.network.prefer = 6; - cfg_c.network.ipv6 = Some(true); - let c = get_dc_addr_static(dc_idx, &cfg_c).expect("table fallback must resolve"); - assert_eq!( - c, - SocketAddr::new(TG_DATACENTERS_V6[(dc_idx as usize) - 1], TG_DATACENTER_PORT), - "without overrides, prefer_v6 path must resolve from static v6 datacenter table" - ); -} - -#[test] -fn prefer_v6_override_matrix_ignores_invalid_entries_and_keeps_fail_closed_fallback() { - let dc_idx: i16 = 3; - - let mut cfg = ProxyConfig::default(); - cfg.network.prefer = 6; - cfg.network.ipv6 = Some(true); - cfg.dc_overrides.insert( - dc_idx.to_string(), - vec![ - "not-an-addr".to_string(), - "also:bad".to_string(), - "203.0.113.55:443".to_string(), - ], - ); - - let addr = get_dc_addr_static(dc_idx, &cfg) - .expect("at least one valid override must keep resolution alive"); - assert_eq!(addr, "203.0.113.55:443".parse::().unwrap()); -} - -#[test] -fn stress_prefer_v6_override_matrix_is_deterministic_under_mixed_inputs() { - for idx in 1..=5i16 { - let mut cfg = ProxyConfig::default(); - cfg.network.prefer = 6; - cfg.network.ipv6 = Some(true); - cfg.dc_overrides.insert( - idx.to_string(), - vec![ - format!("203.0.113.{}:443", 100 + idx), - format!("[2001:db8::{}]:443", 100 + idx), - ], - ); - - let first = get_dc_addr_static(idx, &cfg).expect("first lookup must resolve"); - let second = get_dc_addr_static(idx, &cfg).expect("second lookup must resolve"); - assert_eq!( - first, second, - "override resolution must stay deterministic for dc {idx}" - ); - assert!(first.is_ipv6(), "dc {idx}: v6 override should be preferred"); - } -} - -#[tokio::test] -async fn negative_direct_relay_dc_connection_refused_fails_fast() { - let (client_reader_side, _client_writer_side) = duplex(1024); - let (_client_reader_relay, client_writer_side) = duplex(1024); - - let key = [0u8; 32]; - let iv = 0u128; - let client_reader = CryptoReader::new(client_reader_side, AesCtr::new(&key, iv)); - let client_writer = CryptoWriter::new(client_writer_side, AesCtr::new(&key, iv), 1024); - - let stats = Arc::new(Stats::new()); - let buffer_pool = Arc::new(BufferPool::with_config(1024, 1)); - let rng = Arc::new(SecureRandom::new()); - let route_runtime = RouteRuntimeController::new(RelayRouteMode::Direct); - - // Reserve an ephemeral port and immediately release it to deterministically - // exercise the direct-connect failure path without long-lived hangs. - let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let dc_addr = listener.local_addr().unwrap(); - drop(listener); - - let mut config_with_override = ProxyConfig::default(); - config_with_override - .dc_overrides - .insert("1".to_string(), vec![dc_addr.to_string()]); - let config = Arc::new(config_with_override); - - let upstream_manager = Arc::new(UpstreamManager::new( - vec![UpstreamConfig { - enabled: true, - weight: 1, - scopes: String::new(), - upstream_type: UpstreamType::Direct { - interface: None, - bind_addresses: None, - bindtodevice: None, - }, - selected_scope: String::new(), - ipv4: None, - ipv6: None, - prefer: None, - }], - 1, - 100, - 5000, - 10, - 3, - false, - stats.clone(), - )); - - let success = HandshakeSuccess { - user: "test-user".to_string(), - peer: "127.0.0.1:12345".parse().unwrap(), - dc_idx: 1, - proto_tag: ProtoTag::Intermediate, - enc_key: key, - enc_iv: iv, - dec_key: key, - dec_iv: iv, - is_tls: false, - }; - - let result = timeout( - TokioDuration::from_secs(2), - handle_via_direct( - client_reader, - client_writer, - success, - upstream_manager, - stats, - config, - buffer_pool, - rng, - route_runtime.subscribe(), - route_runtime.snapshot(), - 0xABCD_1234, - ), - ) - .await - .expect("direct relay must fail fast on connection-refused upstream"); - - assert!( - result.is_err(), - "connection-refused upstream must fail closed" - ); -} - -#[tokio::test] -async fn adversarial_direct_relay_cutover_integrity() { - let (client_reader_side, _client_writer_side) = duplex(1024); - let (_client_reader_relay, client_writer_side) = duplex(1024); - - let key = [0u8; 32]; - let iv = 0u128; - let client_reader = CryptoReader::new(client_reader_side, AesCtr::new(&key, iv)); - let client_writer = CryptoWriter::new(client_writer_side, AesCtr::new(&key, iv), 1024); - - let stats = Arc::new(Stats::new()); - let buffer_pool = Arc::new(BufferPool::with_config(1024, 1)); - let rng = Arc::new(SecureRandom::new()); - let route_runtime = RouteRuntimeController::new(RelayRouteMode::Direct); - - // Mock upstream server. - let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let dc_addr = listener.local_addr().unwrap(); - - tokio::spawn(async move { - let (mut stream, _) = listener.accept().await.unwrap(); - // Read handshake nonce. - let mut nonce = [0u8; 64]; - let _ = stream.read_exact(&mut nonce).await; - // Keep connection open. - tokio::time::sleep(TokioDuration::from_secs(5)).await; - }); - - let mut config_with_override = ProxyConfig::default(); - config_with_override - .dc_overrides - .insert("1".to_string(), vec![dc_addr.to_string()]); - let config = Arc::new(config_with_override); - - let upstream_manager = Arc::new(UpstreamManager::new( - vec![UpstreamConfig { - enabled: true, - weight: 1, - scopes: String::new(), - upstream_type: UpstreamType::Direct { - interface: None, - bind_addresses: None, - bindtodevice: None, - }, - selected_scope: String::new(), - ipv4: None, - ipv6: None, - prefer: None, - }], - 1, - 100, - 5000, - 10, - 3, - false, - stats.clone(), - )); - - let success = HandshakeSuccess { - user: "test-user".to_string(), - peer: "127.0.0.1:12345".parse().unwrap(), - dc_idx: 1, - proto_tag: ProtoTag::Intermediate, - enc_key: key, - enc_iv: iv, - dec_key: key, - dec_iv: iv, - is_tls: false, - }; - - let stats_for_task = stats.clone(); - let runtime_clone = route_runtime.clone(); - let session_task = tokio::spawn(async move { - handle_via_direct( - client_reader, - client_writer, - success, - upstream_manager, - stats_for_task, - config, - buffer_pool, - rng, - runtime_clone.subscribe(), - runtime_clone.snapshot(), - 0xABCD_1234, - ) - .await - }); - - timeout(TokioDuration::from_secs(2), async { - loop { - if stats.get_current_connections_direct() == 1 { - break; - } - tokio::time::sleep(TokioDuration::from_millis(10)).await; - } - }) - .await - .expect("direct relay session must start before cutover"); - - // Trigger cutover. - route_runtime.set_mode(RelayRouteMode::Middle).unwrap(); - - // The session should terminate after the staggered delay (1000-2000ms). - let result = timeout(TokioDuration::from_secs(5), session_task) - .await - .expect("Session must terminate after cutover") - .expect("Session must not panic"); - - assert!( - matches!(result, Err(ProxyError::RouteSwitched)), - "Session must terminate with route switch error on cutover" - ); -} +// Unknown-DC deduplication and path validation. +#[path = "direct_relay_security_tests/unknown_paths.rs"] +mod unknown_paths; +// No-follow file opening and target-swap defenses. +#[path = "direct_relay_security_tests/nofollow.rs"] +mod nofollow; +// Directory-anchored append and descriptor integrity. +#[path = "direct_relay_security_tests/anchored.rs"] +mod anchored; +// Asynchronous unknown-DC logging integration. +#[path = "direct_relay_security_tests/logging_integration.rs"] +mod logging_integration; +// Direct relay cancellation and cutover lifecycle. +#[path = "direct_relay_security_tests/relay_lifecycle.rs"] +mod relay_lifecycle; +// DC override routing and negative connection paths. +#[path = "direct_relay_security_tests/routing.rs"] +mod routing; diff --git a/src/proxy/tests/direct_relay_security_tests/anchored.rs b/src/proxy/tests/direct_relay_security_tests/anchored.rs new file mode 100644 index 0000000..6f0b242 --- /dev/null +++ b/src/proxy/tests/direct_relay_security_tests/anchored.rs @@ -0,0 +1,358 @@ +use super::*; + +#[cfg(unix)] +#[test] +fn adversarial_parent_swap_after_check_is_blocked_by_anchored_open() { + use std::os::unix::fs::symlink; + + let base = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!( + "telemt-unknown-dc-parent-swap-openat-{}", + std::process::id() + )); + if let Ok(meta) = fs::symlink_metadata(&base) { + if meta.file_type().is_symlink() || meta.is_file() { + fs::remove_file(&base).expect("stale parent-swap-openat path must be removable"); + } else { + fs::remove_dir_all(&base) + .expect("stale parent-swap-openat directory must be removable"); + } + } + let moved = base.with_extension("bak"); + if let Ok(meta) = fs::symlink_metadata(&moved) { + if meta.file_type().is_symlink() || meta.is_file() { + fs::remove_file(&moved) + .expect("stale parent-swap-openat backup path must be removable"); + } else { + fs::remove_dir_all(&moved) + .expect("stale parent-swap-openat backup directory must be removable"); + } + } + fs::create_dir_all(&base).expect("parent-swap-openat base must be creatable"); + + let rel_candidate = format!( + "target/telemt-unknown-dc-parent-swap-openat-{}/unknown-dc.log", + std::process::id() + ); + let sanitized = sanitize_unknown_dc_log_path(&rel_candidate) + .expect("candidate must sanitize before parent swap"); + fs::write(&sanitized.resolved_path, "seed\n").expect("seed target file must be writable"); + + assert!( + unknown_dc_log_path_is_still_safe(&sanitized), + "precondition: target should initially pass revalidation" + ); + + let outside_parent = std::env::temp_dir().join(format!( + "telemt-unknown-dc-parent-swap-openat-outside-{}", + std::process::id() + )); + fs::create_dir_all(&outside_parent).expect("outside parent directory must be creatable"); + let outside_target = outside_parent.join("unknown-dc.log"); + let _ = fs::remove_file(&outside_target); + + fs::rename(&base, &moved).expect("base parent must be movable for swap simulation"); + symlink(&outside_parent, &base).expect("base parent symlink replacement must be creatable"); + + let err = open_unknown_dc_log_append_anchored(&sanitized) + .expect_err("anchored open must fail when parent is swapped to symlink"); + let raw = err.raw_os_error(); + assert!( + matches!( + raw, + Some(libc::ELOOP) | Some(libc::ENOTDIR) | Some(libc::ENOENT) + ), + "anchored open must fail closed on parent swap race, got raw_os_error={raw:?}" + ); + assert!( + !outside_target.exists(), + "anchored open must never create a log file in swapped outside parent" + ); +} + +#[cfg(unix)] +#[test] +fn anchored_open_nix_path_writes_expected_lines() { + let base = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!( + "telemt-unknown-dc-anchored-open-ok-{}", + std::process::id() + )); + fs::create_dir_all(&base).expect("anchored-open-ok base must be creatable"); + + let rel_candidate = format!( + "target/telemt-unknown-dc-anchored-open-ok-{}/unknown-dc.log", + std::process::id() + ); + let sanitized = sanitize_unknown_dc_log_path(&rel_candidate).expect("candidate must sanitize"); + let _ = fs::remove_file(&sanitized.resolved_path); + + let mut first = open_unknown_dc_log_append_anchored(&sanitized) + .expect("anchored open must create log file in allowed parent"); + append_unknown_dc_line(&mut first, 31_200).expect("first append must succeed"); + + let mut second = open_unknown_dc_log_append_anchored(&sanitized) + .expect("anchored reopen must succeed for existing regular file"); + append_unknown_dc_line(&mut second, 31_201).expect("second append must succeed"); + + let content = + fs::read_to_string(&sanitized.resolved_path).expect("anchored log file must be readable"); + let lines: Vec<&str> = content + .lines() + .filter(|line| !line.trim().is_empty()) + .collect(); + assert_eq!(lines.len(), 2, "expected one line per anchored append call"); + assert!( + lines.contains(&"dc_idx=31200") && lines.contains(&"dc_idx=31201"), + "anchored append output must contain both expected dc_idx lines" + ); +} + +#[cfg(unix)] +#[test] +fn anchored_open_parallel_appends_preserve_line_integrity() { + let base = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!( + "telemt-unknown-dc-anchored-open-parallel-{}", + std::process::id() + )); + fs::create_dir_all(&base).expect("anchored-open-parallel base must be creatable"); + + let rel_candidate = format!( + "target/telemt-unknown-dc-anchored-open-parallel-{}/unknown-dc.log", + std::process::id() + ); + let sanitized = sanitize_unknown_dc_log_path(&rel_candidate).expect("candidate must sanitize"); + let _ = fs::remove_file(&sanitized.resolved_path); + + let mut workers = Vec::new(); + for idx in 0..64i16 { + let sanitized = sanitized.clone(); + workers.push(std::thread::spawn(move || { + let mut file = open_unknown_dc_log_append_anchored(&sanitized) + .expect("anchored open must succeed in worker"); + append_unknown_dc_line(&mut file, 32_000 + idx).expect("worker append must succeed"); + })); + } + + for worker in workers { + worker.join().expect("worker must not panic"); + } + + let content = + fs::read_to_string(&sanitized.resolved_path).expect("parallel log file must be readable"); + let lines: Vec<&str> = content + .lines() + .filter(|line| !line.trim().is_empty()) + .collect(); + assert_eq!( + lines.len(), + 64, + "expected one complete line per worker append" + ); + for line in lines { + assert!( + line.starts_with("dc_idx="), + "line must keep dc_idx prefix and not be interleaved: {line}" + ); + let value = line + .strip_prefix("dc_idx=") + .expect("prefix checked above") + .parse::(); + assert!( + value.is_ok(), + "line payload must remain parseable i16 and not be corrupted: {line}" + ); + } +} + +#[cfg(unix)] +#[test] +fn anchored_open_creates_private_0600_file_permissions() { + use std::os::unix::fs::PermissionsExt; + + let base = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!( + "telemt-unknown-dc-anchored-perms-{}", + std::process::id() + )); + fs::create_dir_all(&base).expect("anchored-perms base must be creatable"); + + let rel_candidate = format!( + "target/telemt-unknown-dc-anchored-perms-{}/unknown-dc.log", + std::process::id() + ); + let sanitized = sanitize_unknown_dc_log_path(&rel_candidate).expect("candidate must sanitize"); + let _ = fs::remove_file(&sanitized.resolved_path); + + let mut file = open_unknown_dc_log_append_anchored(&sanitized) + .expect("anchored open must create file with restricted mode"); + append_unknown_dc_line(&mut file, 31_210).expect("initial append must succeed"); + drop(file); + + let mode = fs::metadata(&sanitized.resolved_path) + .expect("created log file metadata must be readable") + .permissions() + .mode() + & 0o777; + assert_eq!( + mode, 0o600, + "anchored open must create unknown-dc log file with owner-only rw permissions" + ); +} + +#[cfg(unix)] +#[test] +fn anchored_open_rejects_existing_symlink_target() { + use std::os::unix::fs::symlink; + + let base = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!( + "telemt-unknown-dc-anchored-symlink-target-{}", + std::process::id() + )); + fs::create_dir_all(&base).expect("anchored-symlink-target base must be creatable"); + + let rel_candidate = format!( + "target/telemt-unknown-dc-anchored-symlink-target-{}/unknown-dc.log", + std::process::id() + ); + let sanitized = sanitize_unknown_dc_log_path(&rel_candidate).expect("candidate must sanitize"); + + let outside = std::env::temp_dir().join(format!( + "telemt-unknown-dc-anchored-symlink-outside-{}.log", + std::process::id() + )); + fs::write(&outside, "outside\n").expect("outside baseline file must be writable"); + + let _ = fs::remove_file(&sanitized.resolved_path); + symlink(&outside, &sanitized.resolved_path) + .expect("target symlink for anchored-open rejection test must be creatable"); + + let err = open_unknown_dc_log_append_anchored(&sanitized) + .expect_err("anchored open must reject symlinked filename target"); + assert_eq!( + err.raw_os_error(), + Some(libc::ELOOP), + "anchored open should fail closed with ELOOP on symlinked target" + ); +} + +#[cfg(unix)] +#[test] +fn anchored_open_high_contention_multi_write_preserves_complete_lines() { + let base = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!( + "telemt-unknown-dc-anchored-contention-{}", + std::process::id() + )); + fs::create_dir_all(&base).expect("anchored-contention base must be creatable"); + + let rel_candidate = format!( + "target/telemt-unknown-dc-anchored-contention-{}/unknown-dc.log", + std::process::id() + ); + let sanitized = sanitize_unknown_dc_log_path(&rel_candidate).expect("candidate must sanitize"); + let _ = fs::remove_file(&sanitized.resolved_path); + + let workers = 24usize; + let rounds = 40usize; + let mut threads = Vec::new(); + + for worker in 0..workers { + let sanitized = sanitized.clone(); + threads.push(std::thread::spawn(move || { + for round in 0..rounds { + let mut file = open_unknown_dc_log_append_anchored(&sanitized) + .expect("anchored open must succeed under contention"); + let dc_idx = 20_000i16.wrapping_add((worker * rounds + round) as i16); + append_unknown_dc_line(&mut file, dc_idx) + .expect("each contention append must complete"); + } + })); + } + + for thread in threads { + thread.join().expect("contention worker must not panic"); + } + + let content = fs::read_to_string(&sanitized.resolved_path) + .expect("contention output file must be readable"); + let lines: Vec<&str> = content + .lines() + .filter(|line| !line.trim().is_empty()) + .collect(); + assert_eq!( + lines.len(), + workers * rounds, + "every contention append must produce exactly one line" + ); + + let mut unique = std::collections::HashSet::new(); + for line in lines { + assert!( + line.starts_with("dc_idx="), + "line must preserve expected prefix under heavy contention: {line}" + ); + let value = line + .strip_prefix("dc_idx=") + .expect("prefix validated") + .parse::() + .expect("line payload must remain parseable i16 under contention"); + unique.insert(value); + } + + assert_eq!( + unique.len(), + workers * rounds, + "contention output must not lose or duplicate logical writes" + ); +} + +#[cfg(unix)] +#[test] +fn append_unknown_dc_line_returns_error_for_read_only_descriptor() { + let base = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!( + "telemt-unknown-dc-append-ro-{}", + std::process::id() + )); + fs::create_dir_all(&base).expect("append-ro base must be creatable"); + + let rel_candidate = format!( + "target/telemt-unknown-dc-append-ro-{}/unknown-dc.log", + std::process::id() + ); + let sanitized = sanitize_unknown_dc_log_path(&rel_candidate).expect("candidate must sanitize"); + fs::write(&sanitized.resolved_path, "seed\n").expect("seed file must be writable"); + + let mut readonly = std::fs::OpenOptions::new() + .read(true) + .open(&sanitized.resolved_path) + .expect("readonly file open must succeed"); + + append_unknown_dc_line(&mut readonly, 31_222) + .expect_err("append on readonly descriptor must fail closed"); + + let content_after = + fs::read_to_string(&sanitized.resolved_path).expect("seed file must remain readable"); + assert_eq!( + nonempty_line_count(&content_after), + 1, + "failed readonly append must not modify persisted unknown-dc log content" + ); +} diff --git a/src/proxy/tests/direct_relay_security_tests/logging_integration.rs b/src/proxy/tests/direct_relay_security_tests/logging_integration.rs new file mode 100644 index 0000000..6178805 --- /dev/null +++ b/src/proxy/tests/direct_relay_security_tests/logging_integration.rs @@ -0,0 +1,207 @@ +use super::*; + +#[tokio::test] +async fn unknown_dc_absolute_log_path_writes_one_entry() { + let _guard = unknown_dc_test_lock().lock().await; + clear_unknown_dc_log_cache_for_testing(); + + let dc_idx: i16 = 31_001; + let file_path = std::env::temp_dir().join(format!( + "telemt-unknown-dc-abs-{}-{}.log", + std::process::id(), + dc_idx + )); + let _ = fs::remove_file(&file_path); + + let mut cfg = ProxyConfig::default(); + cfg.general.unknown_dc_file_log_enabled = true; + cfg.general.unknown_dc_log_path = Some( + file_path + .to_str() + .expect("temp file path must be valid UTF-8") + .to_string(), + ); + + let _ = get_dc_addr_static(dc_idx, &cfg).expect("fallback routing must still work"); + + let mut content = None; + for _ in 0..20 { + if let Ok(text) = fs::read_to_string(&file_path) { + content = Some(text); + break; + } + tokio::time::sleep(Duration::from_millis(15)).await; + } + + let text = content.expect("absolute unknown-DC log path must produce exactly one log write"); + assert!( + text.contains(&format!("dc_idx={dc_idx}")), + "absolute unknown-DC integration log must contain requested dc_idx" + ); +} + +#[tokio::test] +async fn unknown_dc_safe_relative_log_path_writes_one_entry() { + let _guard = unknown_dc_test_lock().lock().await; + clear_unknown_dc_log_cache_for_testing(); + + let dc_idx: i16 = 31_002; + let rel_dir = format!("target/telemt-unknown-dc-int-{}", std::process::id()); + let rel_file = format!("{rel_dir}/unknown-dc.log"); + let abs_dir = std::env::current_dir() + .expect("cwd must be available") + .join(&rel_dir); + fs::create_dir_all(&abs_dir).expect("integration test log directory must be creatable"); + let abs_file = abs_dir.join("unknown-dc.log"); + let _ = fs::remove_file(&abs_file); + + let mut cfg = ProxyConfig::default(); + cfg.general.unknown_dc_file_log_enabled = true; + cfg.general.unknown_dc_log_path = Some(rel_file); + + let _ = get_dc_addr_static(dc_idx, &cfg).expect("fallback routing must still work"); + + let mut content = None; + for _ in 0..20 { + if let Ok(text) = fs::read_to_string(&abs_file) { + content = Some(text); + break; + } + tokio::time::sleep(Duration::from_millis(15)).await; + } + + let text = content.expect("safe relative path must produce exactly one log write"); + assert!( + text.contains(&format!("dc_idx={dc_idx}")), + "unknown-DC integration log must contain requested dc_idx" + ); +} + +#[tokio::test] +async fn unknown_dc_same_index_burst_writes_only_once() { + let _guard = unknown_dc_test_lock().lock().await; + clear_unknown_dc_log_cache_for_testing(); + + let dc_idx: i16 = 31_010; + let rel_dir = format!("target/telemt-unknown-dc-same-{}", std::process::id()); + let rel_file = format!("{rel_dir}/unknown-dc.log"); + let abs_dir = std::env::current_dir().unwrap().join(&rel_dir); + fs::create_dir_all(&abs_dir).expect("same-index log directory must be creatable"); + let abs_file = abs_dir.join("unknown-dc.log"); + let _ = fs::remove_file(&abs_file); + + let mut cfg = ProxyConfig::default(); + cfg.general.unknown_dc_file_log_enabled = true; + cfg.general.unknown_dc_log_path = Some(rel_file); + + for _ in 0..64 { + let _ = get_dc_addr_static(dc_idx, &cfg).expect("fallback routing must still work"); + } + + let mut content = None; + for _ in 0..30 { + if let Ok(text) = fs::read_to_string(&abs_file) { + content = Some(text); + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + + let text = content.expect("same-index burst must produce at least one log write"); + assert_eq!( + nonempty_line_count(&text), + 1, + "same unknown dc index must be deduplicated to one file line" + ); +} + +#[tokio::test] +async fn unknown_dc_distinct_burst_is_hard_capped_on_file_writes() { + let _guard = unknown_dc_test_lock().lock().await; + clear_unknown_dc_log_cache_for_testing(); + + let rel_dir = format!("target/telemt-unknown-dc-cap-{}", std::process::id()); + let rel_file = format!("{rel_dir}/unknown-dc.log"); + let abs_dir = std::env::current_dir().unwrap().join(&rel_dir); + fs::create_dir_all(&abs_dir).expect("cap log directory must be creatable"); + let abs_file = abs_dir.join("unknown-dc.log"); + let _ = fs::remove_file(&abs_file); + + let mut cfg = ProxyConfig::default(); + cfg.general.unknown_dc_file_log_enabled = true; + cfg.general.unknown_dc_log_path = Some(rel_file); + + for i in 0..(UNKNOWN_DC_LOG_DISTINCT_LIMIT + 128) { + let dc_idx = 20_000i16.wrapping_add(i as i16); + let _ = get_dc_addr_static(dc_idx, &cfg).expect("fallback routing must still work"); + } + + let mut final_text = String::new(); + for _ in 0..80 { + if let Ok(text) = fs::read_to_string(&abs_file) { + final_text = text; + if nonempty_line_count(&final_text) >= UNKNOWN_DC_LOG_DISTINCT_LIMIT { + break; + } + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + + let line_count = nonempty_line_count(&final_text); + assert!( + line_count > 0, + "distinct unknown-dc burst must write at least one line" + ); + assert!( + line_count <= UNKNOWN_DC_LOG_DISTINCT_LIMIT, + "distinct unknown-dc writes must stay within dedup hard cap" + ); +} + +#[cfg(unix)] +#[tokio::test] +async fn unknown_dc_symlinked_target_escape_is_not_written_integration() { + use std::os::unix::fs::symlink; + + let _guard = unknown_dc_test_lock().lock().await; + clear_unknown_dc_log_cache_for_testing(); + + let base = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!( + "telemt-unknown-dc-no-write-link-{}", + std::process::id() + )); + fs::create_dir_all(&base).expect("integration symlink base must be creatable"); + + let outside = std::env::temp_dir().join(format!( + "telemt-unknown-dc-outside-{}.log", + std::process::id() + )); + fs::write(&outside, "baseline\n").expect("outside baseline file must be writable"); + + let linked_target = base.join("unknown-dc.log"); + let _ = fs::remove_file(&linked_target); + symlink(&outside, &linked_target).expect("symlink target must be creatable"); + + let rel_file = format!( + "target/telemt-unknown-dc-no-write-link-{}/unknown-dc.log", + std::process::id() + ); + let dc_idx: i16 = 31_050; + + let mut cfg = ProxyConfig::default(); + cfg.general.unknown_dc_file_log_enabled = true; + cfg.general.unknown_dc_log_path = Some(rel_file); + + let before = fs::read_to_string(&outside).expect("must read baseline outside file"); + let _ = get_dc_addr_static(dc_idx, &cfg).expect("fallback routing must still work"); + tokio::time::sleep(Duration::from_millis(80)).await; + let after = fs::read_to_string(&outside).expect("must read outside file after attempt"); + + assert_eq!( + after, before, + "symlink target escape must not be written by unknown-DC logging" + ); +} diff --git a/src/proxy/tests/direct_relay_security_tests/nofollow.rs b/src/proxy/tests/direct_relay_security_tests/nofollow.rs new file mode 100644 index 0000000..cced77d --- /dev/null +++ b/src/proxy/tests/direct_relay_security_tests/nofollow.rs @@ -0,0 +1,303 @@ +use super::*; + +#[cfg(unix)] +#[test] +fn unknown_dc_open_append_rejects_symlink_target_with_nofollow() { + use std::os::unix::fs::symlink; + + let base = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!("telemt-unknown-dc-nofollow-{}", std::process::id())); + fs::create_dir_all(&base).expect("nofollow base must be creatable"); + + let outside = std::env::temp_dir().join(format!( + "telemt-unknown-dc-nofollow-outside-{}.log", + std::process::id() + )); + let _ = fs::remove_file(&outside); + fs::write(&outside, "outside\n").expect("outside file must be writable"); + + let linked_target = base.join("unknown-dc.log"); + let _ = fs::remove_file(&linked_target); + symlink(&outside, &linked_target).expect("symlink target must be creatable"); + + let err = open_unknown_dc_log_append(&linked_target) + .expect_err("O_NOFOLLOW open must fail for symlink target"); + assert_eq!( + err.raw_os_error(), + Some(libc::ELOOP), + "symlink target must be rejected with ELOOP when O_NOFOLLOW is applied" + ); +} + +#[cfg(unix)] +#[test] +fn unknown_dc_open_append_rejects_broken_symlink_target_with_nofollow() { + use std::os::unix::fs::symlink; + + let base = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!( + "telemt-unknown-dc-broken-link-{}", + std::process::id() + )); + fs::create_dir_all(&base).expect("broken-link base must be creatable"); + + let linked_target = base.join("unknown-dc.log"); + let _ = fs::remove_file(&linked_target); + symlink(base.join("missing-target.log"), &linked_target) + .expect("broken symlink target must be creatable"); + + let err = open_unknown_dc_log_append(&linked_target) + .expect_err("O_NOFOLLOW open must fail for broken symlink target"); + assert_eq!( + err.raw_os_error(), + Some(libc::ELOOP), + "broken symlink target must be rejected with ELOOP when O_NOFOLLOW is applied" + ); +} + +#[cfg(unix)] +#[test] +fn adversarial_unknown_dc_open_append_symlink_flip_never_writes_outside_file() { + use std::os::unix::fs::symlink; + + let base = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!( + "telemt-unknown-dc-symlink-flip-{}", + std::process::id() + )); + fs::create_dir_all(&base).expect("symlink-flip base must be creatable"); + + let outside = std::env::temp_dir().join(format!( + "telemt-unknown-dc-symlink-flip-outside-{}.log", + std::process::id() + )); + fs::write(&outside, "outside-baseline\n").expect("outside baseline file must be writable"); + let outside_before = fs::read_to_string(&outside).expect("outside baseline must be readable"); + + let target = base.join("unknown-dc.log"); + let _ = fs::remove_file(&target); + + for step in 0..1024usize { + let _ = fs::remove_file(&target); + if step % 2 == 0 { + symlink(&outside, &target).expect("symlink creation in flip loop must succeed"); + } + if let Ok(mut file) = open_unknown_dc_log_append(&target) { + writeln!(file, "dc_idx={step}").expect("append on regular file must succeed"); + } + } + + let outside_after = fs::read_to_string(&outside).expect("outside file must remain readable"); + assert_eq!( + outside_after, outside_before, + "outside file must never be modified under symlink-flip adversarial churn" + ); +} + +#[test] +fn unknown_dc_open_append_creates_regular_file() { + let base = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!("telemt-unknown-dc-open-{}", std::process::id())); + fs::create_dir_all(&base).expect("open test base must be creatable"); + + let target = base.join("unknown-dc.log"); + let _ = fs::remove_file(&target); + + { + let mut file = open_unknown_dc_log_append(&target) + .expect("regular target must be creatable with append open"); + writeln!(file, "dc_idx=1234").expect("append write must succeed"); + } + + let meta = fs::symlink_metadata(&target).expect("created target metadata must be readable"); + assert!(meta.file_type().is_file(), "target must be a regular file"); + assert!( + !meta.file_type().is_symlink(), + "regular target open path must not produce symlink artifacts" + ); +} + +#[test] +fn stress_unknown_dc_open_append_regular_file_preserves_line_integrity() { + let base = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!( + "telemt-unknown-dc-open-stress-{}", + std::process::id() + )); + fs::create_dir_all(&base).expect("stress open base must be creatable"); + + let target = base.join("unknown-dc.log"); + let _ = fs::remove_file(&target); + + let writes = 2048usize; + for idx in 0..writes { + let mut file = open_unknown_dc_log_append(&target) + .expect("stress append open on regular file must succeed"); + writeln!(file, "dc_idx={idx}").expect("stress append write must succeed"); + } + + let content = fs::read_to_string(&target).expect("stress output file must be readable"); + assert_eq!( + nonempty_line_count(&content), + writes, + "regular-file append stress must preserve one logical line per write" + ); +} + +#[test] +fn unknown_dc_log_path_revalidation_accepts_regular_existing_target() { + let base = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!( + "telemt-unknown-dc-safe-target-{}", + std::process::id() + )); + fs::create_dir_all(&base).expect("safe target base must be creatable"); + + let target = base.join("unknown-dc.log"); + fs::write(&target, "seed\n").expect("safe target seed write must succeed"); + + let rel_candidate = format!( + "target/telemt-unknown-dc-safe-target-{}/unknown-dc.log", + std::process::id() + ); + let sanitized = + sanitize_unknown_dc_log_path(&rel_candidate).expect("safe candidate must sanitize"); + assert!( + unknown_dc_log_path_is_still_safe(&sanitized), + "revalidation must allow safe existing regular files" + ); +} + +#[test] +fn unknown_dc_log_path_revalidation_rejects_deleted_parent_after_sanitize() { + let base = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!( + "telemt-unknown-dc-vanish-parent-{}", + std::process::id() + )); + fs::create_dir_all(&base).expect("vanish-parent base must be creatable"); + + let rel_candidate = format!( + "target/telemt-unknown-dc-vanish-parent-{}/unknown-dc.log", + std::process::id() + ); + let sanitized = sanitize_unknown_dc_log_path(&rel_candidate) + .expect("candidate must sanitize before parent deletion"); + + fs::remove_dir_all(&base).expect("test parent directory must be removable"); + assert!( + !unknown_dc_log_path_is_still_safe(&sanitized), + "revalidation must fail when sanitized parent disappears before write" + ); +} + +#[cfg(unix)] +#[test] +fn unknown_dc_log_path_revalidation_rejects_parent_swapped_to_symlink() { + use std::os::unix::fs::symlink; + + let parent = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!( + "telemt-unknown-dc-parent-swap-{}", + std::process::id() + )); + if let Ok(meta) = fs::symlink_metadata(&parent) { + if meta.file_type().is_symlink() || meta.is_file() { + fs::remove_file(&parent).expect("stale parent-swap path must be removable"); + } else { + fs::remove_dir_all(&parent).expect("stale parent-swap directory must be removable"); + } + } + let moved = parent.with_extension("bak"); + if let Ok(meta) = fs::symlink_metadata(&moved) { + if meta.file_type().is_symlink() || meta.is_file() { + fs::remove_file(&moved).expect("stale parent-swap backup path must be removable"); + } else { + fs::remove_dir_all(&moved) + .expect("stale parent-swap backup directory must be removable"); + } + } + fs::create_dir_all(&parent).expect("parent-swap test parent must be creatable"); + + let rel_candidate = format!( + "target/telemt-unknown-dc-parent-swap-{}/unknown-dc.log", + std::process::id() + ); + let sanitized = sanitize_unknown_dc_log_path(&rel_candidate) + .expect("candidate must sanitize before parent swap"); + + fs::rename(&parent, &moved).expect("parent must be movable for swap simulation"); + symlink("/tmp", &parent).expect("symlink replacement for parent must be creatable"); + + assert!( + !unknown_dc_log_path_is_still_safe(&sanitized), + "revalidation must fail when canonical parent is swapped to a symlinked target" + ); +} + +#[cfg(unix)] +#[test] +fn adversarial_check_then_symlink_flip_is_blocked_by_nofollow_open() { + use std::os::unix::fs::symlink; + + let parent = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!( + "telemt-unknown-dc-check-open-race-{}", + std::process::id() + )); + if let Ok(meta) = fs::symlink_metadata(&parent) { + if meta.file_type().is_symlink() || meta.is_file() { + fs::remove_file(&parent).expect("stale check-open-race path must be removable"); + } else { + fs::remove_dir_all(&parent).expect("stale check-open-race parent must be removable"); + } + } + fs::create_dir_all(&parent).expect("check-open-race parent must be creatable"); + + let target = parent.join("unknown-dc.log"); + fs::write(&target, "seed\n").expect("seed target file must be writable"); + let rel_candidate = format!( + "target/telemt-unknown-dc-check-open-race-{}/unknown-dc.log", + std::process::id() + ); + let sanitized = sanitize_unknown_dc_log_path(&rel_candidate).expect("candidate must sanitize"); + + assert!( + unknown_dc_log_path_is_still_safe(&sanitized), + "precondition: target should initially pass revalidation" + ); + + let outside = std::env::temp_dir().join(format!( + "telemt-unknown-dc-check-open-race-outside-{}.log", + std::process::id() + )); + fs::write(&outside, "outside\n").expect("outside file must be writable"); + fs::remove_file(&target).expect("target removal before flip must succeed"); + symlink(&outside, &target).expect("target symlink flip must be creatable"); + + let err = open_unknown_dc_log_append(&sanitized.resolved_path) + .expect_err("nofollow open must fail after symlink flip between check and open"); + assert_eq!( + err.raw_os_error(), + Some(libc::ELOOP), + "symlink flip in check/open window must be neutralized by O_NOFOLLOW" + ); +} diff --git a/src/proxy/tests/direct_relay_security_tests/relay_lifecycle.rs b/src/proxy/tests/direct_relay_security_tests/relay_lifecycle.rs new file mode 100644 index 0000000..a513ba9 --- /dev/null +++ b/src/proxy/tests/direct_relay_security_tests/relay_lifecycle.rs @@ -0,0 +1,384 @@ +use super::*; + +#[test] +fn fallback_dc_never_panics_with_single_dc_list() { + let mut cfg = ProxyConfig::default(); + cfg.network.prefer = 6; + cfg.network.ipv6 = Some(true); + cfg.default_dc = Some(42); + + let addr = get_dc_addr_static(999, &cfg).expect("fallback dc must resolve safely"); + let expected = SocketAddr::new(TG_DATACENTERS_V6[0], TG_DATACENTER_PORT); + assert_eq!(addr, expected); +} + +#[tokio::test] +async fn direct_relay_abort_midflight_releases_route_gauge() { + let tg_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let tg_addr = tg_listener.local_addr().unwrap(); + + let tg_accept_task = tokio::spawn(async move { + let (stream, _) = tg_listener.accept().await.unwrap(); + let _hold_stream = stream; + tokio::time::sleep(Duration::from_secs(60)).await; + }); + + let stats = Arc::new(Stats::new()); + let mut config = ProxyConfig::default(); + config + .dc_overrides + .insert("2".to_string(), vec![tg_addr.to_string()]); + let config = Arc::new(config); + + let upstream_manager = Arc::new(UpstreamManager::new( + vec![UpstreamConfig { + upstream_type: UpstreamType::Direct { + interface: None, + bind_addresses: None, + bindtodevice: None, + }, + weight: 1, + enabled: true, + scopes: String::new(), + selected_scope: String::new(), + ipv4: None, + ipv6: None, + prefer: None, + }], + 1, + 1, + 1, + 10, + 1, + false, + stats.clone(), + )); + + let rng = Arc::new(SecureRandom::new()); + let buffer_pool = Arc::new(BufferPool::new()); + let route_runtime = Arc::new(RouteRuntimeController::new(RelayRouteMode::Direct)); + let route_snapshot = route_runtime.snapshot(); + + let (server_side, client_side) = duplex(64 * 1024); + let (server_reader, server_writer) = tokio::io::split(server_side); + let client_reader = make_crypto_reader(server_reader); + let client_writer = make_crypto_writer(server_writer); + + let success = HandshakeSuccess { + user: "abort-direct-user".to_string(), + dc_idx: 2, + proto_tag: ProtoTag::Intermediate, + dec_key: [0u8; 32], + dec_iv: 0, + enc_key: [0u8; 32], + enc_iv: 0, + peer: "127.0.0.1:50000".parse().unwrap(), + is_tls: false, + }; + + let relay_task = tokio::spawn(handle_via_direct( + client_reader, + client_writer, + success, + upstream_manager, + stats.clone(), + config, + buffer_pool, + rng, + route_runtime.subscribe(), + route_snapshot, + 0xabad1dea, + )); + + let started = tokio::time::timeout(Duration::from_secs(2), async { + loop { + if stats.get_current_connections_direct() == 1 { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await; + assert!( + started.is_ok(), + "direct relay must increment route gauge before abort" + ); + + relay_task.abort(); + let joined = relay_task.await; + assert!( + joined.is_err(), + "aborted direct relay task must return join error" + ); + + tokio::time::sleep(Duration::from_millis(20)).await; + assert_eq!( + stats.get_current_connections_direct(), + 0, + "route gauge must be released when direct relay task is aborted mid-flight" + ); + + drop(client_side); + tg_accept_task.abort(); + let _ = tg_accept_task.await; +} + +#[tokio::test] +async fn direct_relay_cutover_midflight_releases_route_gauge() { + let tg_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let tg_addr = tg_listener.local_addr().unwrap(); + + let tg_accept_task = tokio::spawn(async move { + let (stream, _) = tg_listener.accept().await.unwrap(); + let _hold_stream = stream; + tokio::time::sleep(Duration::from_secs(60)).await; + }); + + let stats = Arc::new(Stats::new()); + let mut config = ProxyConfig::default(); + config + .dc_overrides + .insert("2".to_string(), vec![tg_addr.to_string()]); + let config = Arc::new(config); + + let upstream_manager = Arc::new(UpstreamManager::new( + vec![UpstreamConfig { + upstream_type: UpstreamType::Direct { + interface: None, + bind_addresses: None, + bindtodevice: None, + }, + weight: 1, + enabled: true, + scopes: String::new(), + selected_scope: String::new(), + ipv4: None, + ipv6: None, + prefer: None, + }], + 1, + 1, + 1, + 10, + 1, + false, + stats.clone(), + )); + + let rng = Arc::new(SecureRandom::new()); + let buffer_pool = Arc::new(BufferPool::new()); + let route_runtime = Arc::new(RouteRuntimeController::new(RelayRouteMode::Direct)); + let route_snapshot = route_runtime.snapshot(); + + let (server_side, client_side) = duplex(64 * 1024); + let (server_reader, server_writer) = tokio::io::split(server_side); + let client_reader = make_crypto_reader(server_reader); + let client_writer = make_crypto_writer(server_writer); + + let success = HandshakeSuccess { + user: "cutover-direct-user".to_string(), + dc_idx: 2, + proto_tag: ProtoTag::Intermediate, + dec_key: [0u8; 32], + dec_iv: 0, + enc_key: [0u8; 32], + enc_iv: 0, + peer: "127.0.0.1:50002".parse().unwrap(), + is_tls: false, + }; + + let relay_task = tokio::spawn(handle_via_direct( + client_reader, + client_writer, + success, + upstream_manager, + stats.clone(), + config, + buffer_pool, + rng, + route_runtime.subscribe(), + route_snapshot, + 0xface_cafe, + )); + + tokio::time::timeout(Duration::from_secs(2), async { + loop { + if stats.get_current_connections_direct() == 1 { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("direct relay must increment route gauge before cutover"); + + assert!( + route_runtime.set_mode(RelayRouteMode::Middle).is_some(), + "cutover must advance route generation" + ); + + let relay_result = tokio::time::timeout(Duration::from_secs(6), relay_task) + .await + .expect("direct relay must terminate after cutover") + .expect("direct relay task must not panic"); + assert!( + relay_result.is_err(), + "cutover should terminate direct relay session" + ); + assert!( + matches!(relay_result, Err(ProxyError::RouteSwitched)), + "client-visible cutover error must stay generic and avoid route-internal metadata" + ); + + assert_eq!( + stats.get_current_connections_direct(), + 0, + "route gauge must be released when direct relay exits on cutover" + ); + + drop(client_side); + tg_accept_task.abort(); + let _ = tg_accept_task.await; +} + +#[tokio::test] +async fn direct_relay_cutover_storm_multi_session_keeps_generic_errors_and_releases_gauge() { + let session_count = 6usize; + let tg_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let tg_addr = tg_listener.local_addr().unwrap(); + + let tg_accept_task = tokio::spawn(async move { + let mut held_streams = Vec::with_capacity(session_count); + for _ in 0..session_count { + let (stream, _) = tg_listener.accept().await.unwrap(); + held_streams.push(stream); + } + tokio::time::sleep(Duration::from_secs(60)).await; + drop(held_streams); + }); + + let stats = Arc::new(Stats::new()); + let mut config = ProxyConfig::default(); + config + .dc_overrides + .insert("2".to_string(), vec![tg_addr.to_string()]); + let config = Arc::new(config); + + let upstream_manager = Arc::new(UpstreamManager::new( + vec![UpstreamConfig { + upstream_type: UpstreamType::Direct { + interface: None, + bind_addresses: None, + bindtodevice: None, + }, + weight: 1, + enabled: true, + scopes: String::new(), + selected_scope: String::new(), + ipv4: None, + ipv6: None, + prefer: None, + }], + 1, + 1, + 1, + 10, + 1, + false, + stats.clone(), + )); + + let rng = Arc::new(SecureRandom::new()); + let buffer_pool = Arc::new(BufferPool::new()); + let route_runtime = Arc::new(RouteRuntimeController::new(RelayRouteMode::Direct)); + let route_snapshot = route_runtime.snapshot(); + + let mut relay_tasks = Vec::with_capacity(session_count); + let mut client_sides = Vec::with_capacity(session_count); + + for idx in 0..session_count { + let (server_side, client_side) = duplex(64 * 1024); + client_sides.push(client_side); + let (server_reader, server_writer) = tokio::io::split(server_side); + let client_reader = make_crypto_reader(server_reader); + let client_writer = make_crypto_writer(server_writer); + + let success = HandshakeSuccess { + user: format!("cutover-storm-direct-user-{idx}"), + dc_idx: 2, + proto_tag: ProtoTag::Intermediate, + dec_key: [0u8; 32], + dec_iv: 0, + enc_key: [0u8; 32], + enc_iv: 0, + peer: SocketAddr::new( + std::net::IpAddr::V4(std::net::Ipv4Addr::new(127, 0, 0, 1)), + 51000 + idx as u16, + ), + is_tls: false, + }; + + relay_tasks.push(tokio::spawn(handle_via_direct( + client_reader, + client_writer, + success, + upstream_manager.clone(), + stats.clone(), + config.clone(), + buffer_pool.clone(), + rng.clone(), + route_runtime.subscribe(), + route_snapshot, + 0xA000_0000 + idx as u64, + ))); + } + + tokio::time::timeout(Duration::from_secs(4), async { + loop { + if stats.get_current_connections_direct() == session_count as u64 { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("all direct sessions must become active before cutover storm"); + + let route_runtime_flipper = route_runtime.clone(); + let flipper = tokio::spawn(async move { + for step in 0..64u32 { + let mode = if (step & 1) == 0 { + RelayRouteMode::Middle + } else { + RelayRouteMode::Direct + }; + let _ = route_runtime_flipper.set_mode(mode); + tokio::time::sleep(Duration::from_millis(15)).await; + } + }); + + for relay_task in relay_tasks { + let relay_result = tokio::time::timeout(Duration::from_secs(10), relay_task) + .await + .expect("direct relay task must finish under cutover storm") + .expect("direct relay task must not panic"); + + assert!( + matches!(relay_result, Err(ProxyError::RouteSwitched)), + "storm-cutover termination must remain generic for all direct sessions" + ); + } + + flipper.abort(); + let _ = flipper.await; + + assert_eq!( + stats.get_current_connections_direct(), + 0, + "direct route gauge must return to zero after cutover storm" + ); + + drop(client_sides); + tg_accept_task.abort(); + let _ = tg_accept_task.await; +} diff --git a/src/proxy/tests/direct_relay_security_tests/routing.rs b/src/proxy/tests/direct_relay_security_tests/routing.rs new file mode 100644 index 0000000..d629e6d --- /dev/null +++ b/src/proxy/tests/direct_relay_security_tests/routing.rs @@ -0,0 +1,292 @@ +use super::*; + +#[test] +fn prefer_v6_override_matrix_prefers_matching_family_then_degrades_safely() { + let dc_idx: i16 = 2; + + let mut cfg_a = ProxyConfig::default(); + cfg_a.network.prefer = 6; + cfg_a.network.ipv6 = Some(true); + cfg_a.dc_overrides.insert( + dc_idx.to_string(), + vec![ + "203.0.113.90:443".to_string(), + "[2001:db8::90]:443".to_string(), + ], + ); + let a = get_dc_addr_static(dc_idx, &cfg_a).expect("v6+v4 override set must resolve"); + assert!( + a.is_ipv6(), + "prefer_v6 should choose v6 override when present" + ); + + let mut cfg_b = ProxyConfig::default(); + cfg_b.network.prefer = 6; + cfg_b.network.ipv6 = Some(true); + cfg_b + .dc_overrides + .insert(dc_idx.to_string(), vec!["203.0.113.91:443".to_string()]); + let b = get_dc_addr_static(dc_idx, &cfg_b).expect("v4-only override must still resolve"); + assert!( + b.is_ipv4(), + "when no v6 override exists, v4 override must be used" + ); + + let mut cfg_c = ProxyConfig::default(); + cfg_c.network.prefer = 6; + cfg_c.network.ipv6 = Some(true); + let c = get_dc_addr_static(dc_idx, &cfg_c).expect("table fallback must resolve"); + assert_eq!( + c, + SocketAddr::new(TG_DATACENTERS_V6[(dc_idx as usize) - 1], TG_DATACENTER_PORT), + "without overrides, prefer_v6 path must resolve from static v6 datacenter table" + ); +} + +#[test] +fn prefer_v6_override_matrix_ignores_invalid_entries_and_keeps_fail_closed_fallback() { + let dc_idx: i16 = 3; + + let mut cfg = ProxyConfig::default(); + cfg.network.prefer = 6; + cfg.network.ipv6 = Some(true); + cfg.dc_overrides.insert( + dc_idx.to_string(), + vec![ + "not-an-addr".to_string(), + "also:bad".to_string(), + "203.0.113.55:443".to_string(), + ], + ); + + let addr = get_dc_addr_static(dc_idx, &cfg) + .expect("at least one valid override must keep resolution alive"); + assert_eq!(addr, "203.0.113.55:443".parse::().unwrap()); +} + +#[test] +fn stress_prefer_v6_override_matrix_is_deterministic_under_mixed_inputs() { + for idx in 1..=5i16 { + let mut cfg = ProxyConfig::default(); + cfg.network.prefer = 6; + cfg.network.ipv6 = Some(true); + cfg.dc_overrides.insert( + idx.to_string(), + vec![ + format!("203.0.113.{}:443", 100 + idx), + format!("[2001:db8::{}]:443", 100 + idx), + ], + ); + + let first = get_dc_addr_static(idx, &cfg).expect("first lookup must resolve"); + let second = get_dc_addr_static(idx, &cfg).expect("second lookup must resolve"); + assert_eq!( + first, second, + "override resolution must stay deterministic for dc {idx}" + ); + assert!(first.is_ipv6(), "dc {idx}: v6 override should be preferred"); + } +} + +#[tokio::test] +async fn negative_direct_relay_dc_connection_refused_fails_fast() { + let (client_reader_side, _client_writer_side) = duplex(1024); + let (_client_reader_relay, client_writer_side) = duplex(1024); + + let key = [0u8; 32]; + let iv = 0u128; + let client_reader = CryptoReader::new(client_reader_side, AesCtr::new(&key, iv)); + let client_writer = CryptoWriter::new(client_writer_side, AesCtr::new(&key, iv), 1024); + + let stats = Arc::new(Stats::new()); + let buffer_pool = Arc::new(BufferPool::with_config(1024, 1)); + let rng = Arc::new(SecureRandom::new()); + let route_runtime = RouteRuntimeController::new(RelayRouteMode::Direct); + + // Reserve an ephemeral port and immediately release it to deterministically + // exercise the direct-connect failure path without long-lived hangs. + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let dc_addr = listener.local_addr().unwrap(); + drop(listener); + + let mut config_with_override = ProxyConfig::default(); + config_with_override + .dc_overrides + .insert("1".to_string(), vec![dc_addr.to_string()]); + let config = Arc::new(config_with_override); + + let upstream_manager = Arc::new(UpstreamManager::new( + vec![UpstreamConfig { + enabled: true, + weight: 1, + scopes: String::new(), + upstream_type: UpstreamType::Direct { + interface: None, + bind_addresses: None, + bindtodevice: None, + }, + selected_scope: String::new(), + ipv4: None, + ipv6: None, + prefer: None, + }], + 1, + 100, + 5000, + 10, + 3, + false, + stats.clone(), + )); + + let success = HandshakeSuccess { + user: "test-user".to_string(), + peer: "127.0.0.1:12345".parse().unwrap(), + dc_idx: 1, + proto_tag: ProtoTag::Intermediate, + enc_key: key, + enc_iv: iv, + dec_key: key, + dec_iv: iv, + is_tls: false, + }; + + let result = timeout( + TokioDuration::from_secs(2), + handle_via_direct( + client_reader, + client_writer, + success, + upstream_manager, + stats, + config, + buffer_pool, + rng, + route_runtime.subscribe(), + route_runtime.snapshot(), + 0xABCD_1234, + ), + ) + .await + .expect("direct relay must fail fast on connection-refused upstream"); + + assert!( + result.is_err(), + "connection-refused upstream must fail closed" + ); +} + +#[tokio::test] +async fn adversarial_direct_relay_cutover_integrity() { + let (client_reader_side, _client_writer_side) = duplex(1024); + let (_client_reader_relay, client_writer_side) = duplex(1024); + + let key = [0u8; 32]; + let iv = 0u128; + let client_reader = CryptoReader::new(client_reader_side, AesCtr::new(&key, iv)); + let client_writer = CryptoWriter::new(client_writer_side, AesCtr::new(&key, iv), 1024); + + let stats = Arc::new(Stats::new()); + let buffer_pool = Arc::new(BufferPool::with_config(1024, 1)); + let rng = Arc::new(SecureRandom::new()); + let route_runtime = RouteRuntimeController::new(RelayRouteMode::Direct); + + // Mock upstream server. + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let dc_addr = listener.local_addr().unwrap(); + + tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.unwrap(); + // Read handshake nonce. + let mut nonce = [0u8; 64]; + let _ = stream.read_exact(&mut nonce).await; + // Keep connection open. + tokio::time::sleep(TokioDuration::from_secs(5)).await; + }); + + let mut config_with_override = ProxyConfig::default(); + config_with_override + .dc_overrides + .insert("1".to_string(), vec![dc_addr.to_string()]); + let config = Arc::new(config_with_override); + + let upstream_manager = Arc::new(UpstreamManager::new( + vec![UpstreamConfig { + enabled: true, + weight: 1, + scopes: String::new(), + upstream_type: UpstreamType::Direct { + interface: None, + bind_addresses: None, + bindtodevice: None, + }, + selected_scope: String::new(), + ipv4: None, + ipv6: None, + prefer: None, + }], + 1, + 100, + 5000, + 10, + 3, + false, + stats.clone(), + )); + + let success = HandshakeSuccess { + user: "test-user".to_string(), + peer: "127.0.0.1:12345".parse().unwrap(), + dc_idx: 1, + proto_tag: ProtoTag::Intermediate, + enc_key: key, + enc_iv: iv, + dec_key: key, + dec_iv: iv, + is_tls: false, + }; + + let stats_for_task = stats.clone(); + let runtime_clone = route_runtime.clone(); + let session_task = tokio::spawn(async move { + handle_via_direct( + client_reader, + client_writer, + success, + upstream_manager, + stats_for_task, + config, + buffer_pool, + rng, + runtime_clone.subscribe(), + runtime_clone.snapshot(), + 0xABCD_1234, + ) + .await + }); + + timeout(TokioDuration::from_secs(2), async { + loop { + if stats.get_current_connections_direct() == 1 { + break; + } + tokio::time::sleep(TokioDuration::from_millis(10)).await; + } + }) + .await + .expect("direct relay session must start before cutover"); + + // Trigger cutover. + route_runtime.set_mode(RelayRouteMode::Middle).unwrap(); + + // The session should terminate after the staggered delay (1000-2000ms). + let result = timeout(TokioDuration::from_secs(5), session_task) + .await + .expect("Session must terminate after cutover") + .expect("Session must not panic"); + + assert!( + matches!(result, Err(ProxyError::RouteSwitched)), + "Session must terminate with route switch error on cutover" + ); +} diff --git a/src/proxy/tests/direct_relay_security_tests/unknown_paths.rs b/src/proxy/tests/direct_relay_security_tests/unknown_paths.rs new file mode 100644 index 0000000..36020a1 --- /dev/null +++ b/src/proxy/tests/direct_relay_security_tests/unknown_paths.rs @@ -0,0 +1,372 @@ +use super::*; + +#[test] +fn unknown_dc_log_is_deduplicated_per_dc_idx() { + let _guard = unknown_dc_test_lock().blocking_lock(); + clear_unknown_dc_log_cache_for_testing(); + + assert!(should_log_unknown_dc(777)); + assert!( + !should_log_unknown_dc(777), + "same unknown dc_idx must not be logged repeatedly" + ); + assert!( + should_log_unknown_dc(778), + "different unknown dc_idx must still be loggable" + ); +} + +#[test] +fn unknown_dc_log_respects_distinct_limit() { + let _guard = unknown_dc_test_lock().blocking_lock(); + clear_unknown_dc_log_cache_for_testing(); + + for dc in 1..=UNKNOWN_DC_LOG_DISTINCT_LIMIT { + assert!( + should_log_unknown_dc(dc as i16), + "expected first-time unknown dc_idx to be loggable" + ); + } + + assert!( + !should_log_unknown_dc(i16::MAX), + "distinct unknown dc_idx entries above limit must not be logged" + ); +} + +#[test] +fn unknown_dc_log_fails_closed_when_dedup_lock_is_poisoned() { + let poisoned = Arc::new(std::sync::Mutex::new( + std::collections::HashSet::::new(), + )); + let poisoned_for_thread = poisoned.clone(); + + let _ = std::thread::spawn(move || { + let _guard = poisoned_for_thread + .lock() + .expect("poison setup lock must be available"); + panic!("intentional poison for fail-closed regression"); + }) + .join(); + + assert!( + !should_log_unknown_dc_with_set(poisoned.as_ref(), 4242), + "poisoned unknown-DC dedup lock must fail closed" + ); +} + +#[test] +fn unsafe_unknown_dc_log_path_does_not_consume_dedup_slot() { + let _guard = unknown_dc_test_lock().blocking_lock(); + clear_unknown_dc_log_cache_for_testing(); + + let dc_idx: i16 = 31_123; + let mut cfg = ProxyConfig::default(); + cfg.general.unknown_dc_file_log_enabled = true; + cfg.general.unknown_dc_log_path = Some("../telemt-unknown-dc-unsafe.log".to_string()); + + let _ = get_dc_addr_static(dc_idx, &cfg).expect("fallback routing must still work"); + + assert!( + should_log_unknown_dc(dc_idx), + "rejected unsafe log path must not consume unknown-dc dedup entry" + ); +} + +#[test] +fn stress_unknown_dc_log_concurrent_unique_churn_respects_cap() { + let _guard = unknown_dc_test_lock().blocking_lock(); + clear_unknown_dc_log_cache_for_testing(); + + let accepted = Arc::new(AtomicUsize::new(0)); + let mut workers = Vec::new(); + + // Adversarial model: many concurrent peers rotate dc_idx values rapidly. + for worker in 0..16usize { + let accepted = Arc::clone(&accepted); + workers.push(std::thread::spawn(move || { + let base = (worker * 2048) as i32; + for offset in 0..512i32 { + let raw = base + offset; + let dc = (raw % i16::MAX as i32) as i16; + if should_log_unknown_dc(dc) { + accepted.fetch_add(1, Ordering::Relaxed); + } + } + })); + } + + for worker in workers { + worker.join().expect("worker thread must not panic"); + } + + assert_eq!( + accepted.load(Ordering::Relaxed), + UNKNOWN_DC_LOG_DISTINCT_LIMIT, + "concurrent unique churn must never admit more than the configured distinct cap" + ); +} + +#[test] +fn light_fuzz_unknown_dc_log_mixed_duplicates_never_exceeds_cap() { + let _guard = unknown_dc_test_lock().blocking_lock(); + clear_unknown_dc_log_cache_for_testing(); + + // Deterministic xorshift sequence for reproducible mixed duplicate fuzzing. + let mut s: u64 = 0xA5A5_5A5A_C3C3_3C3C; + let mut admitted = 0usize; + + for _ in 0..20_000 { + s ^= s << 7; + s ^= s >> 9; + s ^= s << 8; + + let dc = (s as i16).wrapping_sub(i16::MAX / 2); + if should_log_unknown_dc(dc) { + admitted += 1; + } + } + + assert!( + admitted <= UNKNOWN_DC_LOG_DISTINCT_LIMIT, + "mixed-duplicate fuzzed inputs must not admit more than cap" + ); +} + +#[test] +fn scope_hint_accepts_ascii_alnum_and_dash_within_limit() { + assert_eq!(validated_scope_hint("scope_alpha-1"), Some("alpha-1")); + assert_eq!(validated_scope_hint("scope_AZ09"), Some("AZ09")); +} + +#[test] +fn scope_hint_rejects_invalid_or_oversized_values() { + assert_eq!(validated_scope_hint("plain_user"), None); + assert_eq!(validated_scope_hint("scope_"), None); + assert_eq!(validated_scope_hint("scope_a/b"), None); + assert_eq!(validated_scope_hint("scope_bad space"), None); + assert_eq!(validated_scope_hint("scope_bad.dot"), None); + + let oversized = format!("scope_{}", "a".repeat(MAX_SCOPE_HINT_LEN + 1)); + assert_eq!(validated_scope_hint(&oversized), None); +} + +#[test] +fn unknown_dc_log_path_sanitizer_rejects_parent_traversal_inputs() { + assert!( + sanitize_unknown_dc_log_path("../unknown-dc.txt").is_none(), + "parent traversal paths must be rejected" + ); + assert!( + sanitize_unknown_dc_log_path("logs/../unknown-dc.txt").is_none(), + "embedded parent traversal must be rejected" + ); + assert!( + sanitize_unknown_dc_log_path("./../unknown-dc.txt").is_none(), + "relative parent traversal must be rejected" + ); +} + +#[test] +fn unknown_dc_log_path_sanitizer_accepts_absolute_paths_with_existing_parent() { + let absolute = std::env::temp_dir().join("unknown-dc.txt"); + let absolute_str = absolute + .to_str() + .expect("temp absolute path must be valid UTF-8"); + + let sanitized = sanitize_unknown_dc_log_path(absolute_str) + .expect("absolute paths with existing parent must be accepted"); + assert_eq!(sanitized.resolved_path, absolute); +} + +#[test] +fn unknown_dc_log_path_sanitizer_rejects_absolute_parent_traversal() { + assert!( + sanitize_unknown_dc_log_path("/tmp/../etc/passwd").is_none(), + "absolute parent traversal must be rejected" + ); +} + +#[test] +fn unknown_dc_log_path_sanitizer_accepts_safe_relative_path() { + let base = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!("telemt-unknown-dc-log-{}", std::process::id())); + fs::create_dir_all(&base).expect("temp test directory must be creatable"); + + let candidate = base.join("unknown-dc.txt"); + let candidate_relative = format!( + "target/telemt-unknown-dc-log-{}/unknown-dc.txt", + std::process::id() + ); + + let sanitized = sanitize_unknown_dc_log_path(&candidate_relative) + .expect("safe relative path with existing parent must be accepted"); + assert_eq!(sanitized.resolved_path, candidate); +} + +#[test] +fn unknown_dc_log_path_sanitizer_rejects_empty_or_dot_only_inputs() { + assert!( + sanitize_unknown_dc_log_path("").is_none(), + "empty path must be rejected" + ); + assert!( + sanitize_unknown_dc_log_path(".").is_none(), + "dot-only path without filename must be rejected" + ); +} + +#[test] +fn unknown_dc_log_path_sanitizer_accepts_directory_only_as_filename_projection() { + let sanitized = sanitize_unknown_dc_log_path("target/") + .expect("directory-only input is interpreted as filename projection in current sanitizer"); + assert!( + sanitized.resolved_path.ends_with("target"), + "directory-only input should resolve to canonical parent plus filename projection" + ); +} + +#[test] +fn unknown_dc_log_path_sanitizer_accepts_dot_prefixed_relative_path() { + let rel_dir = format!("target/telemt-unknown-dc-dot-{}", std::process::id()); + let abs_dir = std::env::current_dir() + .expect("cwd must be available") + .join(&rel_dir); + fs::create_dir_all(&abs_dir).expect("dot-prefixed test directory must be creatable"); + + let rel_candidate = format!("./{rel_dir}/unknown-dc.log"); + let expected = abs_dir.join("unknown-dc.log"); + let sanitized = sanitize_unknown_dc_log_path(&rel_candidate) + .expect("dot-prefixed safe path must be accepted"); + assert_eq!(sanitized.resolved_path, expected); +} + +#[test] +fn light_fuzz_unknown_dc_path_parentdir_inputs_always_rejected() { + let mut s: u64 = 0xD00D_BAAD_1234_5678; + for _ in 0..4096 { + s ^= s << 7; + s ^= s >> 9; + s ^= s << 8; + let a = (s as usize) % 32; + let b = ((s >> 8) as usize) % 32; + let candidate = format!("target/{a}/../{b}/unknown-dc.log"); + assert!( + sanitize_unknown_dc_log_path(&candidate).is_none(), + "parent-dir candidate must be rejected: {candidate}" + ); + } +} + +#[test] +fn unknown_dc_log_path_sanitizer_rejects_nonexistent_parent_directory() { + let rel_candidate = format!( + "target/telemt-unknown-dc-missing-{}/nested/unknown-dc.txt", + std::process::id() + ); + + assert!( + sanitize_unknown_dc_log_path(&rel_candidate).is_none(), + "path with missing parent must be rejected to avoid implicit directory creation" + ); +} + +#[cfg(unix)] +#[test] +fn unknown_dc_log_path_sanitizer_accepts_symlinked_parent_inside_workspace() { + use std::os::unix::fs::symlink; + + let base = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!( + "telemt-unknown-dc-log-symlink-internal-{}", + std::process::id() + )); + let real_parent = base.join("real_parent"); + fs::create_dir_all(&real_parent).expect("real parent dir must be creatable"); + + let symlink_parent = base.join("internal_link"); + let _ = fs::remove_file(&symlink_parent); + symlink(&real_parent, &symlink_parent).expect("internal symlink must be creatable"); + + let rel_candidate = format!( + "target/telemt-unknown-dc-log-symlink-internal-{}/internal_link/unknown-dc.txt", + std::process::id() + ); + + let sanitized = sanitize_unknown_dc_log_path(&rel_candidate) + .expect("symlinked parent that resolves inside workspace must be accepted"); + assert!( + sanitized.resolved_path.starts_with(&real_parent), + "sanitized path must resolve to canonical internal parent" + ); +} + +#[cfg(unix)] +#[test] +fn unknown_dc_log_path_sanitizer_accepts_symlink_parent_escape_as_canonical_path() { + use std::os::unix::fs::symlink; + + let base = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!( + "telemt-unknown-dc-log-symlink-{}", + std::process::id() + )); + fs::create_dir_all(&base).expect("symlink test directory must be creatable"); + + let symlink_parent = base.join("escape_link"); + let _ = fs::remove_file(&symlink_parent); + symlink("/tmp", &symlink_parent).expect("symlink parent must be creatable"); + + let rel_candidate = format!( + "target/telemt-unknown-dc-log-symlink-{}/escape_link/unknown-dc.txt", + std::process::id() + ); + + let sanitized = sanitize_unknown_dc_log_path(&rel_candidate) + .expect("symlinked parent must canonicalize to target path"); + assert!( + sanitized.resolved_path.starts_with(Path::new("/tmp")), + "sanitized path must resolve to canonical symlink target" + ); +} + +#[cfg(unix)] +#[test] +fn unknown_dc_log_path_revalidation_rejects_symlinked_target_escape() { + use std::os::unix::fs::symlink; + + let base = std::env::current_dir() + .expect("cwd must be available") + .join("target") + .join(format!( + "telemt-unknown-dc-target-link-{}", + std::process::id() + )); + fs::create_dir_all(&base).expect("target-link base must be creatable"); + + let outside = std::env::temp_dir().join(format!("telemt-outside-{}", std::process::id())); + let _ = fs::remove_file(&outside); + fs::write(&outside, "outside").expect("outside file must be writable"); + + let linked_target = base.join("unknown-dc.log"); + let _ = fs::remove_file(&linked_target); + symlink(&outside, &linked_target).expect("target symlink must be creatable"); + + let rel_candidate = format!( + "target/telemt-unknown-dc-target-link-{}/unknown-dc.log", + std::process::id() + ); + let sanitized = sanitize_unknown_dc_log_path(&rel_candidate) + .expect("candidate should sanitize before final revalidation"); + + assert!( + !unknown_dc_log_path_is_still_safe(&sanitized), + "final revalidation must reject symlinked target escape" + ); +} diff --git a/src/proxy/tests/direct_relay_subtle_adversarial_tests.rs b/src/proxy/tests/direct_relay_subtle_adversarial_tests.rs index 325cffd..b99a106 100644 --- a/src/proxy/tests/direct_relay_subtle_adversarial_tests.rs +++ b/src/proxy/tests/direct_relay_subtle_adversarial_tests.rs @@ -8,9 +8,7 @@ fn nonempty_line_count(text: &str) -> usize { #[test] fn subtle_stress_single_unknown_dc_under_concurrency_logs_once() { - let _guard = unknown_dc_test_lock() - .lock() - .expect("unknown dc test lock must be available"); + let _guard = unknown_dc_test_lock().blocking_lock(); clear_unknown_dc_log_cache_for_testing(); let winners = Arc::new(AtomicUsize::new(0)); @@ -103,9 +101,7 @@ fn subtle_light_fuzz_dc_resolution_never_panics_and_preserves_port() { #[tokio::test(flavor = "multi_thread", worker_threads = 4)] async fn subtle_integration_parallel_same_dc_logs_one_line() { - let _guard = unknown_dc_test_lock() - .lock() - .expect("unknown dc test lock must be available"); + let _guard = unknown_dc_test_lock().lock().await; clear_unknown_dc_log_cache_for_testing(); let rel_dir = format!("target/telemt-direct-relay-same-{}", std::process::id()); @@ -148,9 +144,7 @@ async fn subtle_integration_parallel_same_dc_logs_one_line() { #[tokio::test(flavor = "multi_thread", worker_threads = 4)] async fn subtle_integration_parallel_unique_dcs_log_unique_lines() { - let _guard = unknown_dc_test_lock() - .lock() - .expect("unknown dc test lock must be available"); + let _guard = unknown_dc_test_lock().lock().await; clear_unknown_dc_log_cache_for_testing(); let rel_dir = format!("target/telemt-direct-relay-unique-{}", std::process::id()); diff --git a/src/proxy/tests/masking_adversarial_tests.rs b/src/proxy/tests/masking_adversarial_tests.rs index 6d930b6..996dd35 100644 --- a/src/proxy/tests/masking_adversarial_tests.rs +++ b/src/proxy/tests/masking_adversarial_tests.rs @@ -13,795 +13,9 @@ use tokio::time::{Duration, Instant}; // Probing Indistinguishability (OWASP ASVS 5.1.7) // ------------------------------------------------------------------ -#[tokio::test] -async fn masking_probes_indistinguishable_timing() { - let mut config = ProxyConfig::default(); - config.censorship.mask = true; - config.censorship.mask_host = Some("127.0.0.1".to_string()); - config.censorship.mask_port = 80; // Should timeout/refuse - - let peer: SocketAddr = "192.0.2.10:443".parse().unwrap(); - let local_addr: SocketAddr = "127.0.0.1:443".parse().unwrap(); - let beobachten = BeobachtenStore::new(); - - // Test different probe types - let probes = vec![ - (b"GET / HTTP/1.1\r\nHost: x\r\n\r\n".to_vec(), "HTTP"), - (b"SSH-2.0-probe".to_vec(), "SSH"), - ( - vec![0x16, 0x03, 0x03, 0x00, 0x05, 0x01, 0x00, 0x00, 0x01, 0x00], - "TLS-scanner", - ), - (vec![0x42; 5], "port-scanner"), - ]; - - for (probe, type_name) in probes { - let (client_reader, _client_writer) = duplex(256); - let (_client_visible_reader, client_visible_writer) = duplex(256); - - let start = Instant::now(); - handle_bad_client( - client_reader, - client_visible_writer, - &probe, - peer, - local_addr, - &config, - &beobachten, - ) - .await; - - let elapsed = start.elapsed(); - - // We expect any outcome to take roughly MASK_TIMEOUT (50ms in tests) - // to mask whether the backend was reachable or refused. - assert!( - elapsed >= Duration::from_millis(30), - "Probe {type_name} finished too fast: {elapsed:?}" - ); - } -} - -// ------------------------------------------------------------------ -// Masking Budget Stress Tests (OWASP ASVS 5.1.6) -// ------------------------------------------------------------------ - -#[tokio::test] -async fn masking_budget_stress_under_load() { - let mut config = ProxyConfig::default(); - config.censorship.mask = true; - config.censorship.mask_host = Some("127.0.0.1".to_string()); - config.censorship.mask_port = 1; // Unlikely port - - let peer: SocketAddr = "192.0.2.20:443".parse().unwrap(); - let local_addr: SocketAddr = "127.0.0.1:443".parse().unwrap(); - let beobachten = Arc::new(BeobachtenStore::new()); - - let mut tasks = Vec::new(); - for _ in 0..50 { - let (client_reader, _client_writer) = duplex(256); - let (_client_visible_reader, client_visible_writer) = duplex(256); - let config = config.clone(); - let beobachten = Arc::clone(&beobachten); - - tasks.push(tokio::spawn(async move { - let start = Instant::now(); - handle_bad_client( - client_reader, - client_visible_writer, - b"probe", - peer, - local_addr, - &config, - &beobachten, - ) - .await; - start.elapsed() - })); - } - - for task in tasks { - let elapsed = task.await.unwrap(); - assert!( - elapsed >= Duration::from_millis(30), - "Stress probe finished too fast: {elapsed:?}" - ); - } -} - -// ------------------------------------------------------------------ -// detect_client_type Fingerprint Check -// ------------------------------------------------------------------ - -#[test] -fn test_detect_client_type_boundary_cases() { - // 9 bytes = port-scanner - assert_eq!(detect_client_type(&[0x42; 9]), "port-scanner"); - // 10 bytes = unknown - assert_eq!(detect_client_type(&[0x42; 10]), "unknown"); - - // HTTP verbs without trailing space - assert_eq!(detect_client_type(b"GET/"), "port-scanner"); // because len < 10 - assert_eq!(detect_client_type(b"GET /path"), "HTTP"); -} - -// ------------------------------------------------------------------ -// Priority 2: Slowloris and Slow Read Attacks (OWASP ASVS 5.1.5) -// ------------------------------------------------------------------ - -#[tokio::test] -async fn masking_slowloris_client_idle_timeout_rejected() { - let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let backend_addr = listener.local_addr().unwrap(); - let initial = b"GET / HTTP/1.1\r\nHost: front.example\r\n\r\n".to_vec(); - - let accept_task = tokio::spawn({ - let initial = initial.clone(); - async move { - let (mut stream, _) = listener.accept().await.unwrap(); - let mut observed = vec![0u8; initial.len()]; - stream.read_exact(&mut observed).await.unwrap(); - assert_eq!(observed, initial); - - let mut drip = [0u8; 1]; - let drip_read = - tokio::time::timeout(Duration::from_millis(220), stream.read_exact(&mut drip)) - .await; - assert!( - drip_read.is_err() || drip_read.unwrap().is_err(), - "backend must not receive post-timeout slowloris drip bytes" - ); - } - }); - - let mut config = ProxyConfig::default(); - config.censorship.mask = true; - config.censorship.mask_host = Some("127.0.0.1".to_string()); - config.censorship.mask_port = backend_addr.port(); - - let beobachten = BeobachtenStore::new(); - let peer: SocketAddr = "192.0.2.10:12345".parse().unwrap(); - let local: SocketAddr = "192.0.2.1:443".parse().unwrap(); - - let (mut client_writer, client_reader) = duplex(1024); - let (_client_visible_reader, client_visible_writer) = duplex(1024); - - let handle = tokio::spawn(async move { - handle_bad_client( - client_reader, - client_visible_writer, - &initial, - peer, - local, - &config, - &beobachten, - ) - .await; - }); - - tokio::time::sleep(Duration::from_millis(160)).await; - let _ = client_writer.write_all(b"X").await; - - handle.await.unwrap(); - accept_task.await.unwrap(); -} - -// ------------------------------------------------------------------ -// Priority 2: Fallback Server Down / Fingerprinting (OWASP ASVS 5.1.7) -// ------------------------------------------------------------------ - -#[tokio::test] -async fn masking_fallback_down_mimics_timeout() { - let mut config = ProxyConfig::default(); - config.censorship.mask = true; - config.censorship.mask_host = Some("127.0.0.1".to_string()); - config.censorship.mask_port = 1; // Unlikely port - - let (server_reader, server_writer) = duplex(1024); - let beobachten = BeobachtenStore::new(); - let peer: SocketAddr = "192.0.2.12:12345".parse().unwrap(); - let local: SocketAddr = "192.0.2.1:443".parse().unwrap(); - - let start = Instant::now(); - handle_bad_client( - server_reader, - server_writer, - b"GET / HTTP/1.1\r\n", - peer, - local, - &config, - &beobachten, - ) - .await; - - let elapsed = start.elapsed(); - // It should wait for MASK_TIMEOUT (50ms in tests) even if connection was refused immediately - assert!( - elapsed >= Duration::from_millis(40), - "Must respect connect budget even on failure: {:?}", - elapsed - ); -} - -// ------------------------------------------------------------------ -// Priority 2: SSRF Prevention (OWASP ASVS 5.1.2) -// ------------------------------------------------------------------ - -#[tokio::test] -async fn masking_ssrf_resolve_internal_ranges_blocked() { - use crate::network::dns_overrides::DnsOverrides; - - let blocked_ips = [ - "127.0.0.1", - "169.254.169.254", - "10.0.0.1", - "192.168.1.1", - "0.0.0.0", - ]; - let resolver = DnsOverrides::default(); - - for ip in blocked_ips { - assert!( - resolver.resolve_socket_addr(ip, 80).is_none(), - "runtime DNS overrides must not resolve unconfigured literal host targets" - ); - } -} - -#[tokio::test] -async fn masking_unknown_proxy_protocol_version_falls_back_to_v1_unknown_header() { - let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let backend_addr = listener.local_addr().unwrap(); - - let accept_task = tokio::spawn(async move { - let (mut stream, _) = listener.accept().await.unwrap(); - - let mut header = [0u8; 15]; - stream.read_exact(&mut header).await.unwrap(); - assert_eq!(&header, b"PROXY UNKNOWN\r\n"); - - let mut payload = [0u8; 5]; - stream.read_exact(&mut payload).await.unwrap(); - assert_eq!(&payload, b"probe"); - }); - - let mut config = ProxyConfig::default(); - config.censorship.mask = true; - config.censorship.mask_host = Some("127.0.0.1".to_string()); - config.censorship.mask_port = backend_addr.port(); - config.censorship.mask_proxy_protocol = 255; - - let peer: SocketAddr = "198.51.100.77:50001".parse().unwrap(); - let local_addr: SocketAddr = "[2001:db8::10]:443".parse().unwrap(); - let beobachten = BeobachtenStore::new(); - let (client_reader, _client_writer) = duplex(128); - let (_client_visible_reader, client_visible_writer) = duplex(128); - - handle_bad_client( - client_reader, - client_visible_writer, - b"probe", - peer, - local_addr, - &config, - &beobachten, - ) - .await; - - accept_task.await.unwrap(); -} - -#[tokio::test] -async fn masking_zero_length_initial_data_does_not_hang_or_panic() { - let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let backend_addr = listener.local_addr().unwrap(); - - let accept_task = tokio::spawn(async move { - let (mut stream, _) = listener.accept().await.unwrap(); - let mut one = [0u8; 1]; - let n = tokio::time::timeout(Duration::from_millis(150), stream.read(&mut one)) - .await - .unwrap() - .unwrap(); - assert_eq!( - n, 0, - "backend must observe clean EOF for empty initial payload" - ); - }); - - let mut config = ProxyConfig::default(); - config.censorship.mask = true; - config.censorship.mask_host = Some("127.0.0.1".to_string()); - config.censorship.mask_port = backend_addr.port(); - - let peer: SocketAddr = "203.0.113.70:50002".parse().unwrap(); - let local: SocketAddr = "127.0.0.1:443".parse().unwrap(); - let beobachten = BeobachtenStore::new(); - - let (client_reader, client_writer) = duplex(64); - drop(client_writer); - let (_client_visible_reader, client_visible_writer) = duplex(64); - - handle_bad_client( - client_reader, - client_visible_writer, - b"", - peer, - local, - &config, - &beobachten, - ) - .await; - - accept_task.await.unwrap(); -} - -#[tokio::test] -async fn masking_oversized_initial_payload_is_forwarded_verbatim() { - let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let backend_addr = listener.local_addr().unwrap(); - let payload = vec![0xA5u8; 32 * 1024]; - - let accept_task = tokio::spawn({ - let payload = payload.clone(); - async move { - let (mut stream, _) = listener.accept().await.unwrap(); - let mut observed = vec![0u8; payload.len()]; - stream.read_exact(&mut observed).await.unwrap(); - assert_eq!( - observed, payload, - "large initial payload must stay byte-for-byte" - ); - } - }); - - let mut config = ProxyConfig::default(); - config.censorship.mask = true; - config.censorship.mask_host = Some("127.0.0.1".to_string()); - config.censorship.mask_port = backend_addr.port(); - - let peer: SocketAddr = "203.0.113.71:50003".parse().unwrap(); - let local: SocketAddr = "127.0.0.1:443".parse().unwrap(); - let beobachten = BeobachtenStore::new(); - let (client_reader, _client_writer) = duplex(64); - let (_client_visible_reader, client_visible_writer) = duplex(64); - - handle_bad_client( - client_reader, - client_visible_writer, - &payload, - peer, - local, - &config, - &beobachten, - ) - .await; - - accept_task.await.unwrap(); -} - -#[tokio::test] -async fn masking_refused_backend_keeps_constantish_timing_floor_under_burst() { - let mut config = ProxyConfig::default(); - config.censorship.mask = true; - config.censorship.mask_host = Some("127.0.0.1".to_string()); - config.censorship.mask_port = 1; - - let peer: SocketAddr = "203.0.113.72:50004".parse().unwrap(); - let local: SocketAddr = "127.0.0.1:443".parse().unwrap(); - let beobachten = BeobachtenStore::new(); - - for _ in 0..16 { - let (client_reader, _client_writer) = duplex(128); - let (_client_visible_reader, client_visible_writer) = duplex(128); - let started = Instant::now(); - handle_bad_client( - client_reader, - client_visible_writer, - b"GET / HTTP/1.1\r\n", - peer, - local, - &config, - &beobachten, - ) - .await; - assert!( - started.elapsed() >= Duration::from_millis(30), - "refused-backend path must keep timing floor to reduce fingerprinting" - ); - } -} - -#[tokio::test] -async fn masking_backend_half_close_then_client_half_close_completes_without_hang() { - let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let backend_addr = listener.local_addr().unwrap(); - - let accept_task = tokio::spawn(async move { - let (mut stream, _) = listener.accept().await.unwrap(); - let mut pre = [0u8; 4]; - stream.read_exact(&mut pre).await.unwrap(); - assert_eq!(&pre, b"PING"); - stream.write_all(b"PONG").await.unwrap(); - stream.shutdown().await.unwrap(); - }); - - let mut config = ProxyConfig::default(); - config.censorship.mask = true; - config.censorship.mask_host = Some("127.0.0.1".to_string()); - config.censorship.mask_port = backend_addr.port(); - - let peer: SocketAddr = "203.0.113.73:50005".parse().unwrap(); - let local: SocketAddr = "127.0.0.1:443".parse().unwrap(); - let beobachten = BeobachtenStore::new(); - - let (mut client_writer, client_reader) = duplex(256); - let (mut client_visible_reader, client_visible_writer) = duplex(256); - - let handle = tokio::spawn(async move { - handle_bad_client( - client_reader, - client_visible_writer, - b"PING", - peer, - local, - &config, - &beobachten, - ) - .await; - }); - - client_writer.shutdown().await.unwrap(); - - let mut got = [0u8; 4]; - client_visible_reader.read_exact(&mut got).await.unwrap(); - assert_eq!(&got, b"PONG"); - - timeout(Duration::from_secs(2), handle) - .await - .expect("masking task must terminate after bilateral half-close") - .unwrap(); - accept_task.await.unwrap(); -} - -#[tokio::test] -async fn chaos_burst_reconnect_storm_for_masking_and_relay_concurrently() { - const MASKING_SESSIONS: usize = 48; - const RELAY_SESSIONS: usize = 48; - - let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let backend_addr = listener.local_addr().unwrap(); - let backend_reply = b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nOK".to_vec(); - - let backend_task = tokio::spawn({ - let backend_reply = backend_reply.clone(); - async move { - for _ in 0..MASKING_SESSIONS { - let (mut stream, _) = listener.accept().await.unwrap(); - let mut req = [0u8; 32]; - stream.read_exact(&mut req).await.unwrap(); - assert!( - req.starts_with(b"GET /storm/"), - "masking backend must receive storm reconnect probes" - ); - stream.write_all(&backend_reply).await.unwrap(); - stream.shutdown().await.unwrap(); - } - } - }); - - let mut config = ProxyConfig::default(); - config.censorship.mask = true; - config.censorship.mask_host = Some("127.0.0.1".to_string()); - config.censorship.mask_port = backend_addr.port(); - config.censorship.mask_proxy_protocol = 0; - - let config = Arc::new(config); - let beobachten = Arc::new(BeobachtenStore::new()); - let peer: SocketAddr = "198.51.100.200:55555".parse().unwrap(); - let local: SocketAddr = "127.0.0.1:443".parse().unwrap(); - - let mut masking_tasks = Vec::with_capacity(MASKING_SESSIONS); - for i in 0..MASKING_SESSIONS { - let config = Arc::clone(&config); - let beobachten = Arc::clone(&beobachten); - let expected_reply = backend_reply.clone(); - masking_tasks.push(tokio::spawn(async move { - let mut probe = [0u8; 32]; - let template = format!("GET /storm/{i:04} HTTP/1.1\r\n\r\n"); - let bytes = template.as_bytes(); - probe[..bytes.len()].copy_from_slice(bytes); - - let (client_reader, client_writer) = duplex(256); - drop(client_writer); - let (mut client_visible_reader, client_visible_writer) = duplex(1024); - - let handle = tokio::spawn(async move { - handle_bad_client( - client_reader, - client_visible_writer, - &probe, - peer, - local, - &config, - &beobachten, - ) - .await; - }); - - let mut observed = vec![0u8; expected_reply.len()]; - client_visible_reader - .read_exact(&mut observed) - .await - .unwrap(); - assert_eq!(observed, expected_reply); - - timeout(Duration::from_secs(2), handle) - .await - .expect("masking reconnect task must complete") - .unwrap(); - })); - } - - let mut relay_tasks = Vec::with_capacity(RELAY_SESSIONS); - for i in 0..RELAY_SESSIONS { - relay_tasks.push(tokio::spawn(async move { - let stats = Arc::new(Stats::new()); - let (mut client_peer, relay_client) = duplex(4096); - let (relay_server, mut server_peer) = duplex(4096); - - let (client_reader, client_writer) = tokio::io::split(relay_client); - let (server_reader, server_writer) = tokio::io::split(relay_server); - - let relay_task = tokio::spawn(relay_bidirectional( - client_reader, - client_writer, - server_reader, - server_writer, - 1024, - 1024, - "chaos-storm-relay", - stats, - None, - Arc::new(BufferPool::new()), - )); - - let c2s = vec![(i as u8).wrapping_add(1); 64]; - client_peer.write_all(&c2s).await.unwrap(); - let mut c2s_seen = vec![0u8; c2s.len()]; - server_peer.read_exact(&mut c2s_seen).await.unwrap(); - assert_eq!(c2s_seen, c2s); - - let s2c = vec![(i as u8).wrapping_add(17); 96]; - server_peer.write_all(&s2c).await.unwrap(); - let mut s2c_seen = vec![0u8; s2c.len()]; - client_peer.read_exact(&mut s2c_seen).await.unwrap(); - assert_eq!(s2c_seen, s2c); - - drop(client_peer); - drop(server_peer); - timeout(Duration::from_secs(2), relay_task) - .await - .expect("relay reconnect task must complete") - .unwrap() - .unwrap(); - })); - } - - for task in masking_tasks { - timeout(Duration::from_secs(3), task) - .await - .expect("masking storm join must complete") - .unwrap(); - } - - for task in relay_tasks { - timeout(Duration::from_secs(3), task) - .await - .expect("relay storm join must complete") - .unwrap(); - } - - timeout(Duration::from_secs(3), backend_task) - .await - .expect("masking backend accept loop must complete") - .unwrap(); -} - -fn read_env_usize_or_default(name: &str, default: usize) -> usize { - match std::env::var(name) { - Ok(raw) => match raw.parse::() { - Ok(parsed) if parsed > 0 => parsed, - _ => default, - }, - Err(_) => default, - } -} - -#[tokio::test] -#[ignore = "heavy soak; run manually"] -async fn chaos_burst_reconnect_storm_for_masking_and_relay_multiwave_soak() { - let waves = read_env_usize_or_default("CHAOS_WAVES", 4); - let masking_per_wave = read_env_usize_or_default("CHAOS_MASKING_PER_WAVE", 160); - let relay_per_wave = read_env_usize_or_default("CHAOS_RELAY_PER_WAVE", 160); - let total_masking = waves * masking_per_wave; - - let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let backend_addr = listener.local_addr().unwrap(); - let backend_reply = b"HTTP/1.1 204 No Content\r\nContent-Length: 0\r\n\r\n".to_vec(); - - let backend_task = tokio::spawn({ - let backend_reply = backend_reply.clone(); - async move { - for _ in 0..total_masking { - let (mut stream, _) = listener.accept().await.unwrap(); - let mut req = [0u8; 32]; - stream.read_exact(&mut req).await.unwrap(); - assert!( - req.starts_with(b"GET /storm/"), - "mask backend must only receive storm probes" - ); - stream.write_all(&backend_reply).await.unwrap(); - stream.shutdown().await.unwrap(); - } - } - }); - - let mut config = ProxyConfig::default(); - config.censorship.mask = true; - config.censorship.mask_host = Some("127.0.0.1".to_string()); - config.censorship.mask_port = backend_addr.port(); - config.censorship.mask_proxy_protocol = 0; - - let config = Arc::new(config); - let beobachten = Arc::new(BeobachtenStore::new()); - let peer: SocketAddr = "198.51.100.201:56565".parse().unwrap(); - let local: SocketAddr = "127.0.0.1:443".parse().unwrap(); - - for wave in 0..waves { - let mut masking_tasks = Vec::with_capacity(masking_per_wave); - for i in 0..masking_per_wave { - let config = Arc::clone(&config); - let beobachten = Arc::clone(&beobachten); - let expected_reply = backend_reply.clone(); - masking_tasks.push(tokio::spawn(async move { - let mut probe = [0u8; 32]; - let template = format!("GET /storm/{wave:02}-{i:03}\r\n\r\n"); - let bytes = template.as_bytes(); - probe[..bytes.len()].copy_from_slice(bytes); - - let (client_reader, client_writer) = duplex(256); - drop(client_writer); - let (mut client_visible_reader, client_visible_writer) = duplex(1024); - - let handle = tokio::spawn(async move { - handle_bad_client( - client_reader, - client_visible_writer, - &probe, - peer, - local, - &config, - &beobachten, - ) - .await; - }); - - let mut observed = vec![0u8; expected_reply.len()]; - client_visible_reader - .read_exact(&mut observed) - .await - .unwrap(); - assert_eq!(observed, expected_reply); - - timeout(Duration::from_secs(3), handle) - .await - .expect("masking storm task must complete") - .unwrap(); - })); - } - - let mut relay_tasks = Vec::with_capacity(relay_per_wave); - for i in 0..relay_per_wave { - relay_tasks.push(tokio::spawn(async move { - let stats = Arc::new(Stats::new()); - let (mut client_peer, relay_client) = duplex(4096); - let (relay_server, mut server_peer) = duplex(4096); - - let (client_reader, client_writer) = tokio::io::split(relay_client); - let (server_reader, server_writer) = tokio::io::split(relay_server); - - let relay_task = tokio::spawn(relay_bidirectional( - client_reader, - client_writer, - server_reader, - server_writer, - 1024, - 1024, - "chaos-multiwave-relay", - stats, - None, - Arc::new(BufferPool::new()), - )); - - let c2s = vec![(wave as u8).wrapping_add(i as u8).wrapping_add(1); 32]; - client_peer.write_all(&c2s).await.unwrap(); - let mut c2s_seen = vec![0u8; c2s.len()]; - server_peer.read_exact(&mut c2s_seen).await.unwrap(); - assert_eq!(c2s_seen, c2s); - - let s2c = vec![(wave as u8).wrapping_add(i as u8).wrapping_add(17); 48]; - server_peer.write_all(&s2c).await.unwrap(); - let mut s2c_seen = vec![0u8; s2c.len()]; - client_peer.read_exact(&mut s2c_seen).await.unwrap(); - assert_eq!(s2c_seen, s2c); - - drop(client_peer); - drop(server_peer); - timeout(Duration::from_secs(3), relay_task) - .await - .expect("relay storm task must complete") - .unwrap() - .unwrap(); - })); - } - - for task in masking_tasks { - timeout(Duration::from_secs(6), task) - .await - .expect("masking wave task join must complete") - .unwrap(); - } - - for task in relay_tasks { - timeout(Duration::from_secs(6), task) - .await - .expect("relay wave task join must complete") - .unwrap(); - } - } - - timeout(Duration::from_secs(8), backend_task) - .await - .expect("mask backend must complete all accepted storm sessions") - .unwrap(); -} - -#[tokio::test] -#[ignore = "heavy soak; run manually"] -async fn masking_timing_bucket_soak_refused_backend_stays_within_narrow_band() { - let mut config = ProxyConfig::default(); - config.censorship.mask = true; - config.censorship.mask_host = Some("127.0.0.1".to_string()); - config.censorship.mask_port = 1; - - let peer: SocketAddr = "203.0.113.74:50006".parse().unwrap(); - let local: SocketAddr = "127.0.0.1:443".parse().unwrap(); - let beobachten = BeobachtenStore::new(); - - let mut samples = Vec::with_capacity(128); - for _ in 0..128 { - let (client_reader, _client_writer) = duplex(128); - let (_client_visible_reader, client_visible_writer) = duplex(128); - let started = Instant::now(); - handle_bad_client( - client_reader, - client_visible_writer, - b"GET / HTTP/1.1\r\n", - peer, - local, - &config, - &beobachten, - ) - .await; - samples.push(started.elapsed().as_millis()); - } - - samples.sort_unstable(); - let p10 = samples[samples.len() / 10]; - let p90 = samples[(samples.len() * 9) / 10]; - assert!( - p90.saturating_sub(p10) <= 40, - "timing spread too wide for refused-backend masking path: p10={p10}ms p90={p90}ms" - ); -} +// Masking timing, fallback, and relay boundary cases. +#[path = "masking_adversarial_tests/boundaries.rs"] +mod boundaries; +// Concurrent reconnect storms and manual soak cases. +#[path = "masking_adversarial_tests/chaos.rs"] +mod chaos; diff --git a/src/proxy/tests/masking_adversarial_tests/boundaries.rs b/src/proxy/tests/masking_adversarial_tests/boundaries.rs new file mode 100644 index 0000000..1226476 --- /dev/null +++ b/src/proxy/tests/masking_adversarial_tests/boundaries.rs @@ -0,0 +1,452 @@ +use super::*; + +#[tokio::test] +async fn masking_probes_indistinguishable_timing() { + let mut config = ProxyConfig::default(); + config.censorship.mask = true; + config.censorship.mask_host = Some("127.0.0.1".to_string()); + config.censorship.mask_port = 80; // Should timeout/refuse + + let peer: SocketAddr = "192.0.2.10:443".parse().unwrap(); + let local_addr: SocketAddr = "127.0.0.1:443".parse().unwrap(); + let beobachten = BeobachtenStore::new(); + + // Test different probe types + let probes = vec![ + (b"GET / HTTP/1.1\r\nHost: x\r\n\r\n".to_vec(), "HTTP"), + (b"SSH-2.0-probe".to_vec(), "SSH"), + ( + vec![0x16, 0x03, 0x03, 0x00, 0x05, 0x01, 0x00, 0x00, 0x01, 0x00], + "TLS-scanner", + ), + (vec![0x42; 5], "port-scanner"), + ]; + + for (probe, type_name) in probes { + let (client_reader, _client_writer) = duplex(256); + let (_client_visible_reader, client_visible_writer) = duplex(256); + + let start = Instant::now(); + handle_bad_client( + client_reader, + client_visible_writer, + &probe, + peer, + local_addr, + &config, + &beobachten, + ) + .await; + + let elapsed = start.elapsed(); + + // We expect any outcome to take roughly MASK_TIMEOUT (50ms in tests) + // to mask whether the backend was reachable or refused. + assert!( + elapsed >= Duration::from_millis(30), + "Probe {type_name} finished too fast: {elapsed:?}" + ); + } +} + +// ------------------------------------------------------------------ +// Masking Budget Stress Tests (OWASP ASVS 5.1.6) +// ------------------------------------------------------------------ + +#[tokio::test] +async fn masking_budget_stress_under_load() { + let mut config = ProxyConfig::default(); + config.censorship.mask = true; + config.censorship.mask_host = Some("127.0.0.1".to_string()); + config.censorship.mask_port = 1; // Unlikely port + + let peer: SocketAddr = "192.0.2.20:443".parse().unwrap(); + let local_addr: SocketAddr = "127.0.0.1:443".parse().unwrap(); + let beobachten = Arc::new(BeobachtenStore::new()); + + let mut tasks = Vec::new(); + for _ in 0..50 { + let (client_reader, _client_writer) = duplex(256); + let (_client_visible_reader, client_visible_writer) = duplex(256); + let config = config.clone(); + let beobachten = Arc::clone(&beobachten); + + tasks.push(tokio::spawn(async move { + let start = Instant::now(); + handle_bad_client( + client_reader, + client_visible_writer, + b"probe", + peer, + local_addr, + &config, + &beobachten, + ) + .await; + start.elapsed() + })); + } + + for task in tasks { + let elapsed = task.await.unwrap(); + assert!( + elapsed >= Duration::from_millis(30), + "Stress probe finished too fast: {elapsed:?}" + ); + } +} + +// ------------------------------------------------------------------ +// detect_client_type Fingerprint Check +// ------------------------------------------------------------------ + +#[test] +fn test_detect_client_type_boundary_cases() { + // 9 bytes = port-scanner + assert_eq!(detect_client_type(&[0x42; 9]), "port-scanner"); + // 10 bytes = unknown + assert_eq!(detect_client_type(&[0x42; 10]), "unknown"); + + // HTTP verbs without trailing space + assert_eq!(detect_client_type(b"GET/"), "port-scanner"); // because len < 10 + assert_eq!(detect_client_type(b"GET /path"), "HTTP"); +} + +// ------------------------------------------------------------------ +// Priority 2: Slowloris and Slow Read Attacks (OWASP ASVS 5.1.5) +// ------------------------------------------------------------------ + +#[tokio::test] +async fn masking_slowloris_client_idle_timeout_rejected() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let backend_addr = listener.local_addr().unwrap(); + let initial = b"GET / HTTP/1.1\r\nHost: front.example\r\n\r\n".to_vec(); + + let accept_task = tokio::spawn({ + let initial = initial.clone(); + async move { + let (mut stream, _) = listener.accept().await.unwrap(); + let mut observed = vec![0u8; initial.len()]; + stream.read_exact(&mut observed).await.unwrap(); + assert_eq!(observed, initial); + + let mut drip = [0u8; 1]; + let drip_read = + tokio::time::timeout(Duration::from_millis(220), stream.read_exact(&mut drip)) + .await; + assert!( + drip_read.is_err() || drip_read.unwrap().is_err(), + "backend must not receive post-timeout slowloris drip bytes" + ); + } + }); + + let mut config = ProxyConfig::default(); + config.censorship.mask = true; + config.censorship.mask_host = Some("127.0.0.1".to_string()); + config.censorship.mask_port = backend_addr.port(); + + let beobachten = BeobachtenStore::new(); + let peer: SocketAddr = "192.0.2.10:12345".parse().unwrap(); + let local: SocketAddr = "192.0.2.1:443".parse().unwrap(); + + let (mut client_writer, client_reader) = duplex(1024); + let (_client_visible_reader, client_visible_writer) = duplex(1024); + + let handle = tokio::spawn(async move { + handle_bad_client( + client_reader, + client_visible_writer, + &initial, + peer, + local, + &config, + &beobachten, + ) + .await; + }); + + tokio::time::sleep(Duration::from_millis(160)).await; + let _ = client_writer.write_all(b"X").await; + + handle.await.unwrap(); + accept_task.await.unwrap(); +} + +// ------------------------------------------------------------------ +// Priority 2: Fallback Server Down / Fingerprinting (OWASP ASVS 5.1.7) +// ------------------------------------------------------------------ + +#[tokio::test] +async fn masking_fallback_down_mimics_timeout() { + let mut config = ProxyConfig::default(); + config.censorship.mask = true; + config.censorship.mask_host = Some("127.0.0.1".to_string()); + config.censorship.mask_port = 1; // Unlikely port + + let (server_reader, server_writer) = duplex(1024); + let beobachten = BeobachtenStore::new(); + let peer: SocketAddr = "192.0.2.12:12345".parse().unwrap(); + let local: SocketAddr = "192.0.2.1:443".parse().unwrap(); + + let start = Instant::now(); + handle_bad_client( + server_reader, + server_writer, + b"GET / HTTP/1.1\r\n", + peer, + local, + &config, + &beobachten, + ) + .await; + + let elapsed = start.elapsed(); + // It should wait for MASK_TIMEOUT (50ms in tests) even if connection was refused immediately + assert!( + elapsed >= Duration::from_millis(40), + "Must respect connect budget even on failure: {:?}", + elapsed + ); +} + +// ------------------------------------------------------------------ +// Priority 2: SSRF Prevention (OWASP ASVS 5.1.2) +// ------------------------------------------------------------------ + +#[tokio::test] +async fn masking_ssrf_resolve_internal_ranges_blocked() { + use crate::network::dns_overrides::DnsOverrides; + + let blocked_ips = [ + "127.0.0.1", + "169.254.169.254", + "10.0.0.1", + "192.168.1.1", + "0.0.0.0", + ]; + let resolver = DnsOverrides::default(); + + for ip in blocked_ips { + assert!( + resolver.resolve_socket_addr(ip, 80).is_none(), + "runtime DNS overrides must not resolve unconfigured literal host targets" + ); + } +} + +#[tokio::test] +async fn masking_unknown_proxy_protocol_version_falls_back_to_v1_unknown_header() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let backend_addr = listener.local_addr().unwrap(); + + let accept_task = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.unwrap(); + + let mut header = [0u8; 15]; + stream.read_exact(&mut header).await.unwrap(); + assert_eq!(&header, b"PROXY UNKNOWN\r\n"); + + let mut payload = [0u8; 5]; + stream.read_exact(&mut payload).await.unwrap(); + assert_eq!(&payload, b"probe"); + }); + + let mut config = ProxyConfig::default(); + config.censorship.mask = true; + config.censorship.mask_host = Some("127.0.0.1".to_string()); + config.censorship.mask_port = backend_addr.port(); + config.censorship.mask_proxy_protocol = 255; + + let peer: SocketAddr = "198.51.100.77:50001".parse().unwrap(); + let local_addr: SocketAddr = "[2001:db8::10]:443".parse().unwrap(); + let beobachten = BeobachtenStore::new(); + let (client_reader, _client_writer) = duplex(128); + let (_client_visible_reader, client_visible_writer) = duplex(128); + + handle_bad_client( + client_reader, + client_visible_writer, + b"probe", + peer, + local_addr, + &config, + &beobachten, + ) + .await; + + accept_task.await.unwrap(); +} + +#[tokio::test] +async fn masking_zero_length_initial_data_does_not_hang_or_panic() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let backend_addr = listener.local_addr().unwrap(); + + let accept_task = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.unwrap(); + let mut one = [0u8; 1]; + let n = tokio::time::timeout(Duration::from_millis(150), stream.read(&mut one)) + .await + .unwrap() + .unwrap(); + assert_eq!( + n, 0, + "backend must observe clean EOF for empty initial payload" + ); + }); + + let mut config = ProxyConfig::default(); + config.censorship.mask = true; + config.censorship.mask_host = Some("127.0.0.1".to_string()); + config.censorship.mask_port = backend_addr.port(); + + let peer: SocketAddr = "203.0.113.70:50002".parse().unwrap(); + let local: SocketAddr = "127.0.0.1:443".parse().unwrap(); + let beobachten = BeobachtenStore::new(); + + let (client_reader, client_writer) = duplex(64); + drop(client_writer); + let (_client_visible_reader, client_visible_writer) = duplex(64); + + handle_bad_client( + client_reader, + client_visible_writer, + b"", + peer, + local, + &config, + &beobachten, + ) + .await; + + accept_task.await.unwrap(); +} + +#[tokio::test] +async fn masking_oversized_initial_payload_is_forwarded_verbatim() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let backend_addr = listener.local_addr().unwrap(); + let payload = vec![0xA5u8; 32 * 1024]; + + let accept_task = tokio::spawn({ + let payload = payload.clone(); + async move { + let (mut stream, _) = listener.accept().await.unwrap(); + let mut observed = vec![0u8; payload.len()]; + stream.read_exact(&mut observed).await.unwrap(); + assert_eq!( + observed, payload, + "large initial payload must stay byte-for-byte" + ); + } + }); + + let mut config = ProxyConfig::default(); + config.censorship.mask = true; + config.censorship.mask_host = Some("127.0.0.1".to_string()); + config.censorship.mask_port = backend_addr.port(); + + let peer: SocketAddr = "203.0.113.71:50003".parse().unwrap(); + let local: SocketAddr = "127.0.0.1:443".parse().unwrap(); + let beobachten = BeobachtenStore::new(); + let (client_reader, _client_writer) = duplex(64); + let (_client_visible_reader, client_visible_writer) = duplex(64); + + handle_bad_client( + client_reader, + client_visible_writer, + &payload, + peer, + local, + &config, + &beobachten, + ) + .await; + + accept_task.await.unwrap(); +} + +#[tokio::test] +async fn masking_refused_backend_keeps_constantish_timing_floor_under_burst() { + let mut config = ProxyConfig::default(); + config.censorship.mask = true; + config.censorship.mask_host = Some("127.0.0.1".to_string()); + config.censorship.mask_port = 1; + + let peer: SocketAddr = "203.0.113.72:50004".parse().unwrap(); + let local: SocketAddr = "127.0.0.1:443".parse().unwrap(); + let beobachten = BeobachtenStore::new(); + + for _ in 0..16 { + let (client_reader, _client_writer) = duplex(128); + let (_client_visible_reader, client_visible_writer) = duplex(128); + let started = Instant::now(); + handle_bad_client( + client_reader, + client_visible_writer, + b"GET / HTTP/1.1\r\n", + peer, + local, + &config, + &beobachten, + ) + .await; + assert!( + started.elapsed() >= Duration::from_millis(30), + "refused-backend path must keep timing floor to reduce fingerprinting" + ); + } +} + +#[tokio::test] +async fn masking_backend_half_close_then_client_half_close_completes_without_hang() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let backend_addr = listener.local_addr().unwrap(); + + let accept_task = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.unwrap(); + let mut pre = [0u8; 4]; + stream.read_exact(&mut pre).await.unwrap(); + assert_eq!(&pre, b"PING"); + stream.write_all(b"PONG").await.unwrap(); + stream.shutdown().await.unwrap(); + }); + + let mut config = ProxyConfig::default(); + config.censorship.mask = true; + config.censorship.mask_host = Some("127.0.0.1".to_string()); + config.censorship.mask_port = backend_addr.port(); + + let peer: SocketAddr = "203.0.113.73:50005".parse().unwrap(); + let local: SocketAddr = "127.0.0.1:443".parse().unwrap(); + let beobachten = BeobachtenStore::new(); + + let (mut client_writer, client_reader) = duplex(256); + let (mut client_visible_reader, client_visible_writer) = duplex(256); + + let handle = tokio::spawn(async move { + handle_bad_client( + client_reader, + client_visible_writer, + b"PING", + peer, + local, + &config, + &beobachten, + ) + .await; + }); + + client_writer.shutdown().await.unwrap(); + + let mut got = [0u8; 4]; + client_visible_reader.read_exact(&mut got).await.unwrap(); + assert_eq!(&got, b"PONG"); + + timeout(Duration::from_secs(2), handle) + .await + .expect("masking task must terminate after bilateral half-close") + .unwrap(); + accept_task.await.unwrap(); +} diff --git a/src/proxy/tests/masking_adversarial_tests/chaos.rs b/src/proxy/tests/masking_adversarial_tests/chaos.rs new file mode 100644 index 0000000..ae89136 --- /dev/null +++ b/src/proxy/tests/masking_adversarial_tests/chaos.rs @@ -0,0 +1,343 @@ +use super::*; + +#[tokio::test] +async fn chaos_burst_reconnect_storm_for_masking_and_relay_concurrently() { + const MASKING_SESSIONS: usize = 48; + const RELAY_SESSIONS: usize = 48; + + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let backend_addr = listener.local_addr().unwrap(); + let backend_reply = b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nOK".to_vec(); + + let backend_task = tokio::spawn({ + let backend_reply = backend_reply.clone(); + async move { + for _ in 0..MASKING_SESSIONS { + let (mut stream, _) = listener.accept().await.unwrap(); + let mut req = [0u8; 32]; + stream.read_exact(&mut req).await.unwrap(); + assert!( + req.starts_with(b"GET /storm/"), + "masking backend must receive storm reconnect probes" + ); + stream.write_all(&backend_reply).await.unwrap(); + stream.shutdown().await.unwrap(); + } + } + }); + + let mut config = ProxyConfig::default(); + config.censorship.mask = true; + config.censorship.mask_host = Some("127.0.0.1".to_string()); + config.censorship.mask_port = backend_addr.port(); + config.censorship.mask_proxy_protocol = 0; + + let config = Arc::new(config); + let beobachten = Arc::new(BeobachtenStore::new()); + let peer: SocketAddr = "198.51.100.200:55555".parse().unwrap(); + let local: SocketAddr = "127.0.0.1:443".parse().unwrap(); + + let mut masking_tasks = Vec::with_capacity(MASKING_SESSIONS); + for i in 0..MASKING_SESSIONS { + let config = Arc::clone(&config); + let beobachten = Arc::clone(&beobachten); + let expected_reply = backend_reply.clone(); + masking_tasks.push(tokio::spawn(async move { + let mut probe = [0u8; 32]; + let template = format!("GET /storm/{i:04} HTTP/1.1\r\n\r\n"); + let bytes = template.as_bytes(); + probe[..bytes.len()].copy_from_slice(bytes); + + let (client_reader, client_writer) = duplex(256); + drop(client_writer); + let (mut client_visible_reader, client_visible_writer) = duplex(1024); + + let handle = tokio::spawn(async move { + handle_bad_client( + client_reader, + client_visible_writer, + &probe, + peer, + local, + &config, + &beobachten, + ) + .await; + }); + + let mut observed = vec![0u8; expected_reply.len()]; + client_visible_reader + .read_exact(&mut observed) + .await + .unwrap(); + assert_eq!(observed, expected_reply); + + timeout(Duration::from_secs(2), handle) + .await + .expect("masking reconnect task must complete") + .unwrap(); + })); + } + + let mut relay_tasks = Vec::with_capacity(RELAY_SESSIONS); + for i in 0..RELAY_SESSIONS { + relay_tasks.push(tokio::spawn(async move { + let stats = Arc::new(Stats::new()); + let (mut client_peer, relay_client) = duplex(4096); + let (relay_server, mut server_peer) = duplex(4096); + + let (client_reader, client_writer) = tokio::io::split(relay_client); + let (server_reader, server_writer) = tokio::io::split(relay_server); + + let relay_task = tokio::spawn(relay_bidirectional( + client_reader, + client_writer, + server_reader, + server_writer, + 1024, + 1024, + "chaos-storm-relay", + stats, + None, + Arc::new(BufferPool::new()), + )); + + let c2s = vec![(i as u8).wrapping_add(1); 64]; + client_peer.write_all(&c2s).await.unwrap(); + let mut c2s_seen = vec![0u8; c2s.len()]; + server_peer.read_exact(&mut c2s_seen).await.unwrap(); + assert_eq!(c2s_seen, c2s); + + let s2c = vec![(i as u8).wrapping_add(17); 96]; + server_peer.write_all(&s2c).await.unwrap(); + let mut s2c_seen = vec![0u8; s2c.len()]; + client_peer.read_exact(&mut s2c_seen).await.unwrap(); + assert_eq!(s2c_seen, s2c); + + drop(client_peer); + drop(server_peer); + timeout(Duration::from_secs(2), relay_task) + .await + .expect("relay reconnect task must complete") + .unwrap() + .unwrap(); + })); + } + + for task in masking_tasks { + timeout(Duration::from_secs(3), task) + .await + .expect("masking storm join must complete") + .unwrap(); + } + + for task in relay_tasks { + timeout(Duration::from_secs(3), task) + .await + .expect("relay storm join must complete") + .unwrap(); + } + + timeout(Duration::from_secs(3), backend_task) + .await + .expect("masking backend accept loop must complete") + .unwrap(); +} + +fn read_env_usize_or_default(name: &str, default: usize) -> usize { + match std::env::var(name) { + Ok(raw) => match raw.parse::() { + Ok(parsed) if parsed > 0 => parsed, + _ => default, + }, + Err(_) => default, + } +} + +#[tokio::test] +#[ignore = "heavy soak; run manually"] +async fn chaos_burst_reconnect_storm_for_masking_and_relay_multiwave_soak() { + let waves = read_env_usize_or_default("CHAOS_WAVES", 4); + let masking_per_wave = read_env_usize_or_default("CHAOS_MASKING_PER_WAVE", 160); + let relay_per_wave = read_env_usize_or_default("CHAOS_RELAY_PER_WAVE", 160); + let total_masking = waves * masking_per_wave; + + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let backend_addr = listener.local_addr().unwrap(); + let backend_reply = b"HTTP/1.1 204 No Content\r\nContent-Length: 0\r\n\r\n".to_vec(); + + let backend_task = tokio::spawn({ + let backend_reply = backend_reply.clone(); + async move { + for _ in 0..total_masking { + let (mut stream, _) = listener.accept().await.unwrap(); + let mut req = [0u8; 32]; + stream.read_exact(&mut req).await.unwrap(); + assert!( + req.starts_with(b"GET /storm/"), + "mask backend must only receive storm probes" + ); + stream.write_all(&backend_reply).await.unwrap(); + stream.shutdown().await.unwrap(); + } + } + }); + + let mut config = ProxyConfig::default(); + config.censorship.mask = true; + config.censorship.mask_host = Some("127.0.0.1".to_string()); + config.censorship.mask_port = backend_addr.port(); + config.censorship.mask_proxy_protocol = 0; + + let config = Arc::new(config); + let beobachten = Arc::new(BeobachtenStore::new()); + let peer: SocketAddr = "198.51.100.201:56565".parse().unwrap(); + let local: SocketAddr = "127.0.0.1:443".parse().unwrap(); + + for wave in 0..waves { + let mut masking_tasks = Vec::with_capacity(masking_per_wave); + for i in 0..masking_per_wave { + let config = Arc::clone(&config); + let beobachten = Arc::clone(&beobachten); + let expected_reply = backend_reply.clone(); + masking_tasks.push(tokio::spawn(async move { + let mut probe = [0u8; 32]; + let template = format!("GET /storm/{wave:02}-{i:03}\r\n\r\n"); + let bytes = template.as_bytes(); + probe[..bytes.len()].copy_from_slice(bytes); + + let (client_reader, client_writer) = duplex(256); + drop(client_writer); + let (mut client_visible_reader, client_visible_writer) = duplex(1024); + + let handle = tokio::spawn(async move { + handle_bad_client( + client_reader, + client_visible_writer, + &probe, + peer, + local, + &config, + &beobachten, + ) + .await; + }); + + let mut observed = vec![0u8; expected_reply.len()]; + client_visible_reader + .read_exact(&mut observed) + .await + .unwrap(); + assert_eq!(observed, expected_reply); + + timeout(Duration::from_secs(3), handle) + .await + .expect("masking storm task must complete") + .unwrap(); + })); + } + + let mut relay_tasks = Vec::with_capacity(relay_per_wave); + for i in 0..relay_per_wave { + relay_tasks.push(tokio::spawn(async move { + let stats = Arc::new(Stats::new()); + let (mut client_peer, relay_client) = duplex(4096); + let (relay_server, mut server_peer) = duplex(4096); + + let (client_reader, client_writer) = tokio::io::split(relay_client); + let (server_reader, server_writer) = tokio::io::split(relay_server); + + let relay_task = tokio::spawn(relay_bidirectional( + client_reader, + client_writer, + server_reader, + server_writer, + 1024, + 1024, + "chaos-multiwave-relay", + stats, + None, + Arc::new(BufferPool::new()), + )); + + let c2s = vec![(wave as u8).wrapping_add(i as u8).wrapping_add(1); 32]; + client_peer.write_all(&c2s).await.unwrap(); + let mut c2s_seen = vec![0u8; c2s.len()]; + server_peer.read_exact(&mut c2s_seen).await.unwrap(); + assert_eq!(c2s_seen, c2s); + + let s2c = vec![(wave as u8).wrapping_add(i as u8).wrapping_add(17); 48]; + server_peer.write_all(&s2c).await.unwrap(); + let mut s2c_seen = vec![0u8; s2c.len()]; + client_peer.read_exact(&mut s2c_seen).await.unwrap(); + assert_eq!(s2c_seen, s2c); + + drop(client_peer); + drop(server_peer); + timeout(Duration::from_secs(3), relay_task) + .await + .expect("relay storm task must complete") + .unwrap() + .unwrap(); + })); + } + + for task in masking_tasks { + timeout(Duration::from_secs(6), task) + .await + .expect("masking wave task join must complete") + .unwrap(); + } + + for task in relay_tasks { + timeout(Duration::from_secs(6), task) + .await + .expect("relay wave task join must complete") + .unwrap(); + } + } + + timeout(Duration::from_secs(8), backend_task) + .await + .expect("mask backend must complete all accepted storm sessions") + .unwrap(); +} + +#[tokio::test] +#[ignore = "heavy soak; run manually"] +async fn masking_timing_bucket_soak_refused_backend_stays_within_narrow_band() { + let mut config = ProxyConfig::default(); + config.censorship.mask = true; + config.censorship.mask_host = Some("127.0.0.1".to_string()); + config.censorship.mask_port = 1; + + let peer: SocketAddr = "203.0.113.74:50006".parse().unwrap(); + let local: SocketAddr = "127.0.0.1:443".parse().unwrap(); + let beobachten = BeobachtenStore::new(); + + let mut samples = Vec::with_capacity(128); + for _ in 0..128 { + let (client_reader, _client_writer) = duplex(128); + let (_client_visible_reader, client_visible_writer) = duplex(128); + let started = Instant::now(); + handle_bad_client( + client_reader, + client_visible_writer, + b"GET / HTTP/1.1\r\n", + peer, + local, + &config, + &beobachten, + ) + .await; + samples.push(started.elapsed().as_millis()); + } + + samples.sort_unstable(); + let p10 = samples[samples.len() / 10]; + let p90 = samples[(samples.len() * 9) / 10]; + assert!( + p90.saturating_sub(p10) <= 40, + "timing spread too wide for refused-backend masking path: p10={p10}ms p90={p90}ms" + ); +} diff --git a/src/proxy/tests/masking_connect_failure_close_matrix_security_tests.rs b/src/proxy/tests/masking_connect_failure_close_matrix_security_tests.rs index af3f118..f6e5b40 100644 --- a/src/proxy/tests/masking_connect_failure_close_matrix_security_tests.rs +++ b/src/proxy/tests/masking_connect_failure_close_matrix_security_tests.rs @@ -86,15 +86,14 @@ async fn connect_failure_refusal_close_behavior_matrix() { let peer: SocketAddr = format!("203.0.113.210:{}", 54100 + idx as u16) .parse() .unwrap(); - let elapsed = - run_connect_failure_case( - "127.0.0.1", - unused_port, - timing_normalization_enabled, - peer, - Vec::new(), - ) - .await; + let elapsed = run_connect_failure_case( + "127.0.0.1", + unused_port, + timing_normalization_enabled, + peer, + Vec::new(), + ) + .await; if timing_normalization_enabled { assert!( diff --git a/src/proxy/tests/masking_interface_cache_concurrency_security_tests.rs b/src/proxy/tests/masking_interface_cache_concurrency_security_tests.rs index a1584fc..2444b8c 100644 --- a/src/proxy/tests/masking_interface_cache_concurrency_security_tests.rs +++ b/src/proxy/tests/masking_interface_cache_concurrency_security_tests.rs @@ -1,19 +1,11 @@ #![cfg(unix)] use super::*; -use std::sync::{Mutex, OnceLock}; use tokio::sync::Barrier; -fn interface_cache_test_lock() -> &'static Mutex<()> { - static LOCK: OnceLock> = OnceLock::new(); - LOCK.get_or_init(|| Mutex::new(())) -} - #[tokio::test(flavor = "multi_thread", worker_threads = 4)] async fn adversarial_parallel_cold_miss_performs_single_interface_refresh() { - let _guard = interface_cache_test_lock() - .lock() - .unwrap_or_else(|poison| poison.into_inner()); + let _guard = interface_cache_test_lock().lock().await; reset_local_interface_enumerations_for_tests(); let local_addr: SocketAddr = "0.0.0.0:443".parse().expect("valid local addr"); diff --git a/src/proxy/tests/masking_interface_cache_security_tests.rs b/src/proxy/tests/masking_interface_cache_security_tests.rs index 4be2857..6d39912 100644 --- a/src/proxy/tests/masking_interface_cache_security_tests.rs +++ b/src/proxy/tests/masking_interface_cache_security_tests.rs @@ -1,18 +1,10 @@ #![cfg(unix)] use super::*; -use std::sync::{Mutex, OnceLock}; - -fn interface_cache_test_lock() -> &'static Mutex<()> { - static LOCK: OnceLock> = OnceLock::new(); - LOCK.get_or_init(|| Mutex::new(())) -} #[tokio::test] async fn tdd_repeated_local_listener_checks_do_not_repeat_interface_enumeration_within_window() { - let _guard = interface_cache_test_lock() - .lock() - .unwrap_or_else(|poison| poison.into_inner()); + let _guard = interface_cache_test_lock().lock().await; reset_local_interface_enumerations_for_tests(); let local_addr: SocketAddr = "0.0.0.0:443".parse().expect("valid local addr"); @@ -29,9 +21,7 @@ async fn tdd_repeated_local_listener_checks_do_not_repeat_interface_enumeration_ #[tokio::test] async fn tdd_non_local_port_short_circuit_does_not_enumerate_interfaces() { - let _guard = interface_cache_test_lock() - .lock() - .unwrap_or_else(|poison| poison.into_inner()); + let _guard = interface_cache_test_lock().lock().await; reset_local_interface_enumerations_for_tests(); let local_addr: SocketAddr = "0.0.0.0:443".parse().expect("valid local addr"); diff --git a/src/proxy/tests/middle_relay_tiny_frame_debt_security_tests.rs b/src/proxy/tests/middle_relay_tiny_frame_debt_security_tests.rs index fad87d0..b5285fc 100644 --- a/src/proxy/tests/middle_relay_tiny_frame_debt_security_tests.rs +++ b/src/proxy/tests/middle_relay_tiny_frame_debt_security_tests.rs @@ -95,710 +95,12 @@ fn simulate_tiny_debt_pattern(pattern: &[bool], max_steps: usize) -> (Option> 9; - seed ^= seed << 8; - - let len = 512 + ((seed as usize) & 0x3ff); - let mut pattern = Vec::with_capacity(len); - let mut local_seed = seed; - for _ in 0..len { - local_seed ^= local_seed << 7; - local_seed ^= local_seed >> 9; - local_seed ^= local_seed << 8; - pattern.push((local_seed & 1) == 0); - } - - let (closed_at, debt, _) = simulate_tiny_debt_pattern(&pattern, pattern.len()); - if closed_at.is_none() { - assert!(debt < TINY_FRAME_DEBT_LIMIT); - } - assert!(debt <= u32::MAX); - } -} - -#[test] -fn stress_many_independent_simulations_keep_isolated_debt_state() { - for idx in 0..2048usize { - let mut pattern = Vec::with_capacity(64); - for j in 0..64usize { - pattern.push(((idx ^ j) & 3) == 0); - } - let (_closed_at, debt, _reals) = simulate_tiny_debt_pattern(&pattern, pattern.len()); - assert!(debt <= TINY_FRAME_DEBT_LIMIT.saturating_add(TINY_FRAME_DEBT_PER_TINY)); - } -} - -#[tokio::test] -async fn idle_policy_enabled_intermediate_zero_length_flood_is_fail_closed() { - let (reader, mut writer) = duplex(4096); - let mut crypto_reader = make_crypto_reader(reader); - let buffer_pool = Arc::new(BufferPool::new()); - let stats = Stats::new(); - let session_started_at = Instant::now(); - let forensics = make_forensics(11, session_started_at); - let mut frame_counter = 0u64; - let mut idle_state = RelayClientIdleState::new(session_started_at); - let idle_policy = make_enabled_idle_policy(); - let last_downstream_activity_ms = AtomicU64::new(0); - - let flood_plaintext = vec![0u8; 4 * 256]; - let flood_encrypted = encrypt_for_reader(&flood_plaintext); - writer.write_all(&flood_encrypted).await.unwrap(); - drop(writer); - - let result = read_bounded( - &mut crypto_reader, - ProtoTag::Intermediate, - &buffer_pool, - &forensics, - &mut frame_counter, - &stats, - &idle_policy, - &mut idle_state, - &last_downstream_activity_ms, - session_started_at, - ) - .await; - - assert!(matches!(result, Err(ProxyError::Proxy(_)))); -} - -#[tokio::test] -async fn idle_policy_enabled_secure_zero_length_flood_is_fail_closed() { - let (reader, mut writer) = duplex(4096); - let mut crypto_reader = make_crypto_reader(reader); - let buffer_pool = Arc::new(BufferPool::new()); - let stats = Stats::new(); - let session_started_at = Instant::now(); - let forensics = make_forensics(12, session_started_at); - let mut frame_counter = 0u64; - let mut idle_state = RelayClientIdleState::new(session_started_at); - let idle_policy = make_enabled_idle_policy(); - let last_downstream_activity_ms = AtomicU64::new(0); - - let flood_plaintext = vec![0u8; 4 * 256]; - let flood_encrypted = encrypt_for_reader(&flood_plaintext); - writer.write_all(&flood_encrypted).await.unwrap(); - drop(writer); - - let result = read_bounded( - &mut crypto_reader, - ProtoTag::Secure, - &buffer_pool, - &forensics, - &mut frame_counter, - &stats, - &idle_policy, - &mut idle_state, - &last_downstream_activity_ms, - session_started_at, - ) - .await; - - assert!(matches!(result, Err(ProxyError::Proxy(_)))); -} - -#[tokio::test] -async fn intermediate_alternating_zero_and_real_eventually_closes() { - let (reader, mut writer) = duplex(8192); - let mut crypto_reader = make_crypto_reader(reader); - let buffer_pool = Arc::new(BufferPool::new()); - let stats = Stats::new(); - let session_started_at = Instant::now(); - let forensics = make_forensics(13, session_started_at); - let mut frame_counter = 0u64; - let mut idle_state = RelayClientIdleState::new(session_started_at); - let idle_policy = make_enabled_idle_policy(); - let last_downstream_activity_ms = AtomicU64::new(0); - - let mut plaintext = Vec::with_capacity(3000); - for idx in 0..160u8 { - plaintext.extend_from_slice(&0u32.to_le_bytes()); - plaintext.extend_from_slice(&4u32.to_le_bytes()); - plaintext.extend_from_slice(&[idx, idx ^ 0x11, idx ^ 0x22, idx ^ 0x33]); - } - let encrypted = encrypt_for_reader(&plaintext); - writer.write_all(&encrypted).await.unwrap(); - drop(writer); - - let mut closed = false; - for _ in 0..220 { - let result = read_bounded( - &mut crypto_reader, - ProtoTag::Intermediate, - &buffer_pool, - &forensics, - &mut frame_counter, - &stats, - &idle_policy, - &mut idle_state, - &last_downstream_activity_ms, - session_started_at, - ) - .await; - - match result { - Ok(Some(_)) => {} - Err(ProxyError::Proxy(_)) => { - closed = true; - break; - } - Ok(None) => break, - Err(other) => panic!("unexpected error while probing alternating close: {other}"), - } - } - - assert!(closed, "intermediate alternating attack must fail closed"); -} - -#[tokio::test] -async fn small_tiny_burst_followed_by_real_frame_does_not_spuriously_close() { - let (reader, mut writer) = duplex(1024); - let mut crypto_reader = make_crypto_reader(reader); - let buffer_pool = Arc::new(BufferPool::new()); - let stats = Stats::new(); - let session_started_at = Instant::now(); - let forensics = make_forensics(14, session_started_at); - let mut frame_counter = 0u64; - let mut idle_state = RelayClientIdleState::new(session_started_at); - let idle_policy = make_enabled_idle_policy(); - let last_downstream_activity_ms = AtomicU64::new(0); - - let mut plaintext = Vec::with_capacity(64); - for _ in 0..8 { - plaintext.push(0x00); - } - plaintext.push(0x01); - plaintext.extend_from_slice(&[1, 2, 3, 4]); - - let encrypted = encrypt_for_reader(&plaintext); - writer.write_all(&encrypted).await.unwrap(); - - let first = read_bounded( - &mut crypto_reader, - ProtoTag::Abridged, - &buffer_pool, - &forensics, - &mut frame_counter, - &stats, - &idle_policy, - &mut idle_state, - &last_downstream_activity_ms, - session_started_at, - ) - .await; - - match first { - Ok(Some((payload, _))) => assert_eq!(payload.as_ref(), &[1, 2, 3, 4]), - Err(e) => panic!("unexpected close after small tiny burst: {e}"), - Ok(None) => panic!("unexpected EOF before real frame"), - } -} - -#[tokio::test] -async fn idle_policy_enabled_zero_length_flood_is_fail_closed() { - let (reader, mut writer) = duplex(4096); - let mut crypto_reader = make_crypto_reader(reader); - let buffer_pool = Arc::new(BufferPool::new()); - let stats = Stats::new(); - let session_started_at = Instant::now(); - let forensics = make_forensics(1, session_started_at); - let mut frame_counter = 0u64; - let mut idle_state = RelayClientIdleState::new(session_started_at); - let idle_policy = make_enabled_idle_policy(); - let last_downstream_activity_ms = AtomicU64::new(0); - - let flood_plaintext = vec![0u8; 1024]; - let flood_encrypted = encrypt_for_reader(&flood_plaintext); - writer - .write_all(&flood_encrypted) - .await - .expect("zero-length flood bytes must be writable"); - drop(writer); - - let result = read_bounded( - &mut crypto_reader, - ProtoTag::Abridged, - &buffer_pool, - &forensics, - &mut frame_counter, - &stats, - &idle_policy, - &mut idle_state, - &last_downstream_activity_ms, - session_started_at, - ) - .await; - - assert!( - matches!(result, Err(ProxyError::Proxy(_))), - "idle policy enabled must fail closed for pure zero-length flood" - ); -} - -#[tokio::test] -async fn idle_policy_enabled_alternating_tiny_real_eventually_closes() { - let (reader, mut writer) = duplex(8192); - let mut crypto_reader = make_crypto_reader(reader); - let buffer_pool = Arc::new(BufferPool::new()); - let stats = Stats::new(); - let session_started_at = Instant::now(); - let forensics = make_forensics(2, session_started_at); - let mut frame_counter = 0u64; - let mut idle_state = RelayClientIdleState::new(session_started_at); - let idle_policy = make_enabled_idle_policy(); - let last_downstream_activity_ms = AtomicU64::new(0); - - let mut plaintext = Vec::with_capacity(256 * 6); - for idx in 0..=255u8 { - plaintext.push(0x00); - plaintext.push(0x01); - plaintext.extend_from_slice(&[idx, idx ^ 0x55, idx ^ 0xAA, 0x11]); - } - - let encrypted = encrypt_for_reader(&plaintext); - writer - .write_all(&encrypted) - .await - .expect("alternating flood bytes must be writable"); - drop(writer); - - let mut saw_proxy_close = false; - for _ in 0..300 { - let result = read_bounded( - &mut crypto_reader, - ProtoTag::Abridged, - &buffer_pool, - &forensics, - &mut frame_counter, - &stats, - &idle_policy, - &mut idle_state, - &last_downstream_activity_ms, - session_started_at, - ) - .await; - - match result { - Ok(Some((_payload, _quickack))) => {} - Err(ProxyError::Proxy(_)) => { - saw_proxy_close = true; - break; - } - Err(ProxyError::Io(e)) => panic!("unexpected IO error before close: {e}"), - Ok(None) => panic!("unexpected EOF before debt-based closure"), - Err(other) => panic!("unexpected error before close: {other}"), - } - } - - assert!( - saw_proxy_close, - "alternating tiny/real sequence must eventually fail closed" - ); -} - -#[tokio::test] -async fn enabled_idle_policy_valid_nonzero_frame_still_passes() { - let (reader, mut writer) = duplex(1024); - let mut crypto_reader = make_crypto_reader(reader); - let buffer_pool = Arc::new(BufferPool::new()); - let stats = Stats::new(); - let session_started_at = Instant::now(); - let forensics = make_forensics(3, session_started_at); - let mut frame_counter = 0u64; - let mut idle_state = RelayClientIdleState::new(session_started_at); - let idle_policy = make_enabled_idle_policy(); - let last_downstream_activity_ms = AtomicU64::new(0); - - let payload = [7u8, 8, 9, 10]; - let mut plaintext = Vec::with_capacity(1 + payload.len()); - plaintext.push(0x01); - plaintext.extend_from_slice(&payload); - - let encrypted = encrypt_for_reader(&plaintext); - writer - .write_all(&encrypted) - .await - .expect("nonzero frame must be writable"); - - let result = read_bounded( - &mut crypto_reader, - ProtoTag::Abridged, - &buffer_pool, - &forensics, - &mut frame_counter, - &stats, - &idle_policy, - &mut idle_state, - &last_downstream_activity_ms, - session_started_at, - ) - .await - .expect("valid frame should decode") - .expect("valid frame should return payload"); - - assert_eq!(result.0.as_ref(), &payload); - assert!(!result.1); - assert_eq!(frame_counter, 1); -} - -#[tokio::test] -async fn abridged_quickack_tiny_flood_is_fail_closed() { - let (reader, mut writer) = duplex(4096); - let mut crypto_reader = make_crypto_reader(reader); - let buffer_pool = Arc::new(BufferPool::new()); - let stats = Stats::new(); - let session_started_at = Instant::now(); - let forensics = make_forensics(21, session_started_at); - let mut frame_counter = 0u64; - let mut idle_state = RelayClientIdleState::new(session_started_at); - let idle_policy = make_enabled_idle_policy(); - let last_downstream_activity_ms = AtomicU64::new(0); - - let flood_plaintext = vec![0x80u8; 256]; - let flood_encrypted = encrypt_for_reader(&flood_plaintext); - writer.write_all(&flood_encrypted).await.unwrap(); - drop(writer); - - let result = read_bounded( - &mut crypto_reader, - ProtoTag::Abridged, - &buffer_pool, - &forensics, - &mut frame_counter, - &stats, - &idle_policy, - &mut idle_state, - &last_downstream_activity_ms, - session_started_at, - ) - .await; - - assert!( - matches!(result, Err(ProxyError::Proxy(_))), - "quickack-marked zero-length flood must fail closed" - ); -} - -#[tokio::test] -async fn abridged_extended_zero_len_flood_is_fail_closed() { - let (reader, mut writer) = duplex(4096); - let mut crypto_reader = make_crypto_reader(reader); - let buffer_pool = Arc::new(BufferPool::new()); - let stats = Stats::new(); - let session_started_at = Instant::now(); - let forensics = make_forensics(22, session_started_at); - let mut frame_counter = 0u64; - let mut idle_state = RelayClientIdleState::new(session_started_at); - let idle_policy = make_enabled_idle_policy(); - let last_downstream_activity_ms = AtomicU64::new(0); - - let mut flood_plaintext = Vec::with_capacity(4 * 256); - for _ in 0..256 { - flood_plaintext.extend_from_slice(&[0x7f, 0x00, 0x00, 0x00]); - } - let flood_encrypted = encrypt_for_reader(&flood_plaintext); - writer.write_all(&flood_encrypted).await.unwrap(); - drop(writer); - - let result = read_bounded( - &mut crypto_reader, - ProtoTag::Abridged, - &buffer_pool, - &forensics, - &mut frame_counter, - &stats, - &idle_policy, - &mut idle_state, - &last_downstream_activity_ms, - session_started_at, - ) - .await; - - assert!( - matches!(result, Err(ProxyError::Proxy(_))), - "extended zero-length abridged flood must fail closed" - ); -} - -#[tokio::test] -async fn one_to_eight_abridged_wire_pattern_survives_without_false_positive_close() { - let mut plaintext = Vec::with_capacity(9 * 300); - for idx in 0..300usize { - plaintext.push(0x00); - for _ in 0..8 { - let b = idx as u8; - plaintext.push(0x01); - plaintext.extend_from_slice(&[b, b ^ 0x11, b ^ 0x22, b ^ 0x33]); - } - } - - // Keep the test single-task and deterministic: make duplex capacity larger than the - // generated ciphertext so write_all cannot block waiting for a concurrent reader. - let duplex_capacity = plaintext.len().saturating_add(1024); - let (reader, mut writer) = duplex(duplex_capacity); - let mut crypto_reader = make_crypto_reader(reader); - let buffer_pool = Arc::new(BufferPool::new()); - let stats = Stats::new(); - let session_started_at = Instant::now(); - let forensics = make_forensics(23, session_started_at); - let mut frame_counter = 0u64; - let mut idle_state = RelayClientIdleState::new(session_started_at); - let idle_policy = make_enabled_idle_policy(); - let last_downstream_activity_ms = AtomicU64::new(0); - - let encrypted = encrypt_for_reader(&plaintext); - writer.write_all(&encrypted).await.unwrap(); - drop(writer); - - let mut closed = false; - for _ in 0..3000 { - match read_bounded( - &mut crypto_reader, - ProtoTag::Abridged, - &buffer_pool, - &forensics, - &mut frame_counter, - &stats, - &idle_policy, - &mut idle_state, - &last_downstream_activity_ms, - session_started_at, - ) - .await - { - Ok(Some(_)) => {} - Ok(None) => break, - Err(ProxyError::Proxy(_)) => { - closed = true; - break; - } - Err(other) => panic!("unexpected error in 1:8 wire test: {other}"), - } - } - - assert!( - !closed, - "wire-level 1:8 tiny-to-real pattern should not trigger debt close" - ); -} - -#[tokio::test] -async fn deterministic_light_fuzz_abridged_wire_behavior_matches_model() { - let mut seed = 0xD1CE_BAAD_2026_0322u64; - - for case_idx in 0..32u64 { - seed ^= seed << 7; - seed ^= seed >> 9; - seed ^= seed << 8; - - let events = 300 + ((seed as usize) & 0xff); - let mut pattern = Vec::with_capacity(events); - let mut local = seed; - for _ in 0..events { - local ^= local << 7; - local ^= local >> 9; - local ^= local << 8; - pattern.push((local & 0x03) == 0); - } - - let mut plaintext = Vec::with_capacity(events * 6); - for (idx, tiny) in pattern.iter().copied().enumerate() { - if tiny { - plaintext.push(0x00); - } else { - let b = (idx as u8) ^ (case_idx as u8); - plaintext.push(0x01); - plaintext.extend_from_slice(&[b, b ^ 0x1F, b ^ 0x7A, b ^ 0xC3]); - } - } - - let (reader, mut writer) = duplex(16 * 1024); - let mut crypto_reader = make_crypto_reader(reader); - let buffer_pool = Arc::new(BufferPool::new()); - let stats = Stats::new(); - let session_started_at = Instant::now(); - let forensics = make_forensics(500 + case_idx, session_started_at); - let mut frame_counter = 0u64; - let mut idle_state = RelayClientIdleState::new(session_started_at); - let idle_policy = make_enabled_idle_policy(); - let last_downstream_activity_ms = AtomicU64::new(0); - - writer - .write_all(&encrypt_for_reader(&plaintext)) - .await - .unwrap(); - drop(writer); - - let (expected_close, _, _) = simulate_tiny_debt_pattern(&pattern, pattern.len()); - let mut observed_close = false; - - for _ in 0..(events + 8) { - match read_bounded( - &mut crypto_reader, - ProtoTag::Abridged, - &buffer_pool, - &forensics, - &mut frame_counter, - &stats, - &idle_policy, - &mut idle_state, - &last_downstream_activity_ms, - session_started_at, - ) - .await - { - Ok(Some(_)) => {} - Ok(None) => break, - Err(ProxyError::Proxy(_)) => { - observed_close = true; - break; - } - Err(other) => panic!("unexpected fuzz error: {other}"), - } - } - - assert_eq!( - observed_close, - expected_close.is_some(), - "wire parser behavior must match debt model for case {case_idx}" - ); - } -} +// Pure tiny-frame debt model invariants. +#[path = "middle_relay_tiny_frame_debt_security_tests/model.rs"] +mod model; +// Intermediate and secure transport debt behavior. +#[path = "middle_relay_tiny_frame_debt_security_tests/transport.rs"] +mod transport; +// Abridged framing debt behavior. +#[path = "middle_relay_tiny_frame_debt_security_tests/abridged.rs"] +mod abridged; diff --git a/src/proxy/tests/middle_relay_tiny_frame_debt_security_tests/abridged.rs b/src/proxy/tests/middle_relay_tiny_frame_debt_security_tests/abridged.rs new file mode 100644 index 0000000..dd34ac6 --- /dev/null +++ b/src/proxy/tests/middle_relay_tiny_frame_debt_security_tests/abridged.rs @@ -0,0 +1,225 @@ +use super::*; + +#[tokio::test] +async fn abridged_quickack_tiny_flood_is_fail_closed() { + let (reader, mut writer) = duplex(4096); + let mut crypto_reader = make_crypto_reader(reader); + let buffer_pool = Arc::new(BufferPool::new()); + let stats = Stats::new(); + let session_started_at = Instant::now(); + let forensics = make_forensics(21, session_started_at); + let mut frame_counter = 0u64; + let mut idle_state = RelayClientIdleState::new(session_started_at); + let idle_policy = make_enabled_idle_policy(); + let last_downstream_activity_ms = AtomicU64::new(0); + + let flood_plaintext = vec![0x80u8; 256]; + let flood_encrypted = encrypt_for_reader(&flood_plaintext); + writer.write_all(&flood_encrypted).await.unwrap(); + drop(writer); + + let result = read_bounded( + &mut crypto_reader, + ProtoTag::Abridged, + &buffer_pool, + &forensics, + &mut frame_counter, + &stats, + &idle_policy, + &mut idle_state, + &last_downstream_activity_ms, + session_started_at, + ) + .await; + + assert!( + matches!(result, Err(ProxyError::Proxy(_))), + "quickack-marked zero-length flood must fail closed" + ); +} + +#[tokio::test] +async fn abridged_extended_zero_len_flood_is_fail_closed() { + let (reader, mut writer) = duplex(4096); + let mut crypto_reader = make_crypto_reader(reader); + let buffer_pool = Arc::new(BufferPool::new()); + let stats = Stats::new(); + let session_started_at = Instant::now(); + let forensics = make_forensics(22, session_started_at); + let mut frame_counter = 0u64; + let mut idle_state = RelayClientIdleState::new(session_started_at); + let idle_policy = make_enabled_idle_policy(); + let last_downstream_activity_ms = AtomicU64::new(0); + + let mut flood_plaintext = Vec::with_capacity(4 * 256); + for _ in 0..256 { + flood_plaintext.extend_from_slice(&[0x7f, 0x00, 0x00, 0x00]); + } + let flood_encrypted = encrypt_for_reader(&flood_plaintext); + writer.write_all(&flood_encrypted).await.unwrap(); + drop(writer); + + let result = read_bounded( + &mut crypto_reader, + ProtoTag::Abridged, + &buffer_pool, + &forensics, + &mut frame_counter, + &stats, + &idle_policy, + &mut idle_state, + &last_downstream_activity_ms, + session_started_at, + ) + .await; + + assert!( + matches!(result, Err(ProxyError::Proxy(_))), + "extended zero-length abridged flood must fail closed" + ); +} + +#[tokio::test] +async fn one_to_eight_abridged_wire_pattern_survives_without_false_positive_close() { + let mut plaintext = Vec::with_capacity(9 * 300); + for idx in 0..300usize { + plaintext.push(0x00); + for _ in 0..8 { + let b = idx as u8; + plaintext.push(0x01); + plaintext.extend_from_slice(&[b, b ^ 0x11, b ^ 0x22, b ^ 0x33]); + } + } + + // Keep the test single-task and deterministic: make duplex capacity larger than the + // generated ciphertext so write_all cannot block waiting for a concurrent reader. + let duplex_capacity = plaintext.len().saturating_add(1024); + let (reader, mut writer) = duplex(duplex_capacity); + let mut crypto_reader = make_crypto_reader(reader); + let buffer_pool = Arc::new(BufferPool::new()); + let stats = Stats::new(); + let session_started_at = Instant::now(); + let forensics = make_forensics(23, session_started_at); + let mut frame_counter = 0u64; + let mut idle_state = RelayClientIdleState::new(session_started_at); + let idle_policy = make_enabled_idle_policy(); + let last_downstream_activity_ms = AtomicU64::new(0); + + let encrypted = encrypt_for_reader(&plaintext); + writer.write_all(&encrypted).await.unwrap(); + drop(writer); + + let mut closed = false; + for _ in 0..3000 { + match read_bounded( + &mut crypto_reader, + ProtoTag::Abridged, + &buffer_pool, + &forensics, + &mut frame_counter, + &stats, + &idle_policy, + &mut idle_state, + &last_downstream_activity_ms, + session_started_at, + ) + .await + { + Ok(Some(_)) => {} + Ok(None) => break, + Err(ProxyError::Proxy(_)) => { + closed = true; + break; + } + Err(other) => panic!("unexpected error in 1:8 wire test: {other}"), + } + } + + assert!( + !closed, + "wire-level 1:8 tiny-to-real pattern should not trigger debt close" + ); +} + +#[tokio::test] +async fn deterministic_light_fuzz_abridged_wire_behavior_matches_model() { + let mut seed = 0xD1CE_BAAD_2026_0322u64; + + for case_idx in 0..32u64 { + seed ^= seed << 7; + seed ^= seed >> 9; + seed ^= seed << 8; + + let events = 300 + ((seed as usize) & 0xff); + let mut pattern = Vec::with_capacity(events); + let mut local = seed; + for _ in 0..events { + local ^= local << 7; + local ^= local >> 9; + local ^= local << 8; + pattern.push((local & 0x03) == 0); + } + + let mut plaintext = Vec::with_capacity(events * 6); + for (idx, tiny) in pattern.iter().copied().enumerate() { + if tiny { + plaintext.push(0x00); + } else { + let b = (idx as u8) ^ (case_idx as u8); + plaintext.push(0x01); + plaintext.extend_from_slice(&[b, b ^ 0x1F, b ^ 0x7A, b ^ 0xC3]); + } + } + + let (reader, mut writer) = duplex(16 * 1024); + let mut crypto_reader = make_crypto_reader(reader); + let buffer_pool = Arc::new(BufferPool::new()); + let stats = Stats::new(); + let session_started_at = Instant::now(); + let forensics = make_forensics(500 + case_idx, session_started_at); + let mut frame_counter = 0u64; + let mut idle_state = RelayClientIdleState::new(session_started_at); + let idle_policy = make_enabled_idle_policy(); + let last_downstream_activity_ms = AtomicU64::new(0); + + writer + .write_all(&encrypt_for_reader(&plaintext)) + .await + .unwrap(); + drop(writer); + + let (expected_close, _, _) = simulate_tiny_debt_pattern(&pattern, pattern.len()); + let mut observed_close = false; + + for _ in 0..(events + 8) { + match read_bounded( + &mut crypto_reader, + ProtoTag::Abridged, + &buffer_pool, + &forensics, + &mut frame_counter, + &stats, + &idle_policy, + &mut idle_state, + &last_downstream_activity_ms, + session_started_at, + ) + .await + { + Ok(Some(_)) => {} + Ok(None) => break, + Err(ProxyError::Proxy(_)) => { + observed_close = true; + break; + } + Err(other) => panic!("unexpected fuzz error: {other}"), + } + } + + assert_eq!( + observed_close, + expected_close.is_some(), + "wire parser behavior must match debt model for case {case_idx}" + ); + } +} diff --git a/src/proxy/tests/middle_relay_tiny_frame_debt_security_tests/model.rs b/src/proxy/tests/middle_relay_tiny_frame_debt_security_tests/model.rs new file mode 100644 index 0000000..2854d43 --- /dev/null +++ b/src/proxy/tests/middle_relay_tiny_frame_debt_security_tests/model.rs @@ -0,0 +1,171 @@ +use super::*; + +#[test] +fn tiny_frame_debt_constants_match_security_budget_expectations() { + assert_eq!(TINY_FRAME_DEBT_PER_TINY, 8); + assert_eq!(TINY_FRAME_DEBT_LIMIT, 512); +} + +#[test] +fn relay_client_idle_state_initial_debt_is_zero() { + let state = RelayClientIdleState::new(Instant::now()); + assert_eq!(state.tiny_frame_debt, 0); +} + +#[test] +fn on_client_frame_does_not_reset_tiny_frame_debt() { + let now = Instant::now(); + let mut state = RelayClientIdleState::new(now); + state.tiny_frame_debt = 77; + state.on_client_frame(now); + assert_eq!(state.tiny_frame_debt, 77); +} + +#[test] +fn tiny_frame_debt_increment_is_saturating() { + let mut debt = u32::MAX - 1; + debt = debt.saturating_add(TINY_FRAME_DEBT_PER_TINY); + assert_eq!(debt, u32::MAX); +} + +#[test] +fn tiny_frame_debt_decrement_is_saturating() { + let mut debt = 0u32; + debt = debt.saturating_sub(1); + assert_eq!(debt, 0); +} + +#[test] +fn consecutive_tiny_frames_close_exactly_at_threshold() { + let max_tiny_without_close = (TINY_FRAME_DEBT_LIMIT / TINY_FRAME_DEBT_PER_TINY) as usize; + let pattern = vec![true; max_tiny_without_close]; + let (closed_at, _, _) = simulate_tiny_debt_pattern(&pattern, pattern.len()); + assert_eq!(closed_at, Some(max_tiny_without_close)); +} + +#[test] +fn one_less_than_threshold_tiny_frames_do_not_close() { + let tiny_count = (TINY_FRAME_DEBT_LIMIT / TINY_FRAME_DEBT_PER_TINY) as usize - 1; + let pattern = vec![true; tiny_count]; + let (closed_at, debt, _) = simulate_tiny_debt_pattern(&pattern, pattern.len()); + assert_eq!(closed_at, None); + assert!(debt < TINY_FRAME_DEBT_LIMIT); +} + +#[test] +fn alternating_one_to_one_closes_with_bounded_real_frame_count() { + let mut pattern = Vec::with_capacity(512); + for _ in 0..256 { + pattern.push(true); + pattern.push(false); + } + let (closed_at, _, reals) = simulate_tiny_debt_pattern(&pattern, pattern.len()); + assert!(closed_at.is_some()); + assert!( + reals <= 80, + "expected bounded real frames before close, got {reals}" + ); +} + +#[test] +fn alternating_one_to_eight_is_stable_for_long_runs() { + let mut pattern = Vec::with_capacity(9 * 5000); + for _ in 0..5000 { + pattern.push(true); + for _ in 0..8 { + pattern.push(false); + } + } + let (closed_at, debt, _) = simulate_tiny_debt_pattern(&pattern, pattern.len()); + assert_eq!(closed_at, None); + assert!(debt <= TINY_FRAME_DEBT_PER_TINY); +} + +#[test] +fn alternating_one_to_seven_eventually_closes() { + let mut pattern = Vec::with_capacity(8 * 2000); + for _ in 0..2000 { + pattern.push(true); + for _ in 0..7 { + pattern.push(false); + } + } + let (closed_at, _, _) = simulate_tiny_debt_pattern(&pattern, pattern.len()); + assert!( + closed_at.is_some(), + "1:7 tiny-to-real must eventually close" + ); +} + +#[test] +fn two_tiny_one_real_closes_faster_than_one_to_one() { + let mut one_to_one = Vec::with_capacity(512); + for _ in 0..256 { + one_to_one.push(true); + one_to_one.push(false); + } + + let mut two_to_one = Vec::with_capacity(768); + for _ in 0..256 { + two_to_one.push(true); + two_to_one.push(true); + two_to_one.push(false); + } + + let (a_close, _, _) = simulate_tiny_debt_pattern(&one_to_one, one_to_one.len()); + let (b_close, _, _) = simulate_tiny_debt_pattern(&two_to_one, two_to_one.len()); + assert!(a_close.is_some() && b_close.is_some()); + assert!(b_close.unwrap_or(usize::MAX) < a_close.unwrap_or(0)); +} + +#[test] +fn burst_then_drain_can_recover_without_close() { + let burst_tiny = ((TINY_FRAME_DEBT_LIMIT / TINY_FRAME_DEBT_PER_TINY) / 2) as usize; + let mut pattern = Vec::with_capacity(burst_tiny + 600); + for _ in 0..burst_tiny { + pattern.push(true); + } + pattern.extend(std::iter::repeat_n(false, 600)); + + let (closed_at, debt, _) = simulate_tiny_debt_pattern(&pattern, pattern.len()); + assert_eq!(closed_at, None); + assert_eq!(debt, 0); +} + +#[test] +fn light_fuzz_tiny_frame_debt_model_stays_within_bounds() { + let mut seed = 0xA5A5_91C3_2026_0322u64; + for _case in 0..128 { + seed ^= seed << 7; + seed ^= seed >> 9; + seed ^= seed << 8; + + let len = 512 + ((seed as usize) & 0x3ff); + let mut pattern = Vec::with_capacity(len); + let mut local_seed = seed; + for _ in 0..len { + local_seed ^= local_seed << 7; + local_seed ^= local_seed >> 9; + local_seed ^= local_seed << 8; + pattern.push((local_seed & 1) == 0); + } + + let (closed_at, debt, _) = simulate_tiny_debt_pattern(&pattern, pattern.len()); + if closed_at.is_none() { + assert!(debt < TINY_FRAME_DEBT_LIMIT); + } + assert!(debt <= TINY_FRAME_DEBT_LIMIT.saturating_add(TINY_FRAME_DEBT_PER_TINY)); + } +} + +#[test] +fn stress_many_independent_simulations_keep_isolated_debt_state() { + for idx in 0..2048usize { + let mut pattern = Vec::with_capacity(64); + for j in 0..64usize { + pattern.push(((idx ^ j) & 3) == 0); + } + let (_closed_at, debt, _reals) = simulate_tiny_debt_pattern(&pattern, pattern.len()); + assert!(debt <= TINY_FRAME_DEBT_LIMIT.saturating_add(TINY_FRAME_DEBT_PER_TINY)); + } +} diff --git a/src/proxy/tests/middle_relay_tiny_frame_debt_security_tests/transport.rs b/src/proxy/tests/middle_relay_tiny_frame_debt_security_tests/transport.rs new file mode 100644 index 0000000..4b4d9c4 --- /dev/null +++ b/src/proxy/tests/middle_relay_tiny_frame_debt_security_tests/transport.rs @@ -0,0 +1,315 @@ +use super::*; + +#[tokio::test] +async fn idle_policy_enabled_intermediate_zero_length_flood_is_fail_closed() { + let (reader, mut writer) = duplex(4096); + let mut crypto_reader = make_crypto_reader(reader); + let buffer_pool = Arc::new(BufferPool::new()); + let stats = Stats::new(); + let session_started_at = Instant::now(); + let forensics = make_forensics(11, session_started_at); + let mut frame_counter = 0u64; + let mut idle_state = RelayClientIdleState::new(session_started_at); + let idle_policy = make_enabled_idle_policy(); + let last_downstream_activity_ms = AtomicU64::new(0); + + let flood_plaintext = vec![0u8; 4 * 256]; + let flood_encrypted = encrypt_for_reader(&flood_plaintext); + writer.write_all(&flood_encrypted).await.unwrap(); + drop(writer); + + let result = read_bounded( + &mut crypto_reader, + ProtoTag::Intermediate, + &buffer_pool, + &forensics, + &mut frame_counter, + &stats, + &idle_policy, + &mut idle_state, + &last_downstream_activity_ms, + session_started_at, + ) + .await; + + assert!(matches!(result, Err(ProxyError::Proxy(_)))); +} + +#[tokio::test] +async fn idle_policy_enabled_secure_zero_length_flood_is_fail_closed() { + let (reader, mut writer) = duplex(4096); + let mut crypto_reader = make_crypto_reader(reader); + let buffer_pool = Arc::new(BufferPool::new()); + let stats = Stats::new(); + let session_started_at = Instant::now(); + let forensics = make_forensics(12, session_started_at); + let mut frame_counter = 0u64; + let mut idle_state = RelayClientIdleState::new(session_started_at); + let idle_policy = make_enabled_idle_policy(); + let last_downstream_activity_ms = AtomicU64::new(0); + + let flood_plaintext = vec![0u8; 4 * 256]; + let flood_encrypted = encrypt_for_reader(&flood_plaintext); + writer.write_all(&flood_encrypted).await.unwrap(); + drop(writer); + + let result = read_bounded( + &mut crypto_reader, + ProtoTag::Secure, + &buffer_pool, + &forensics, + &mut frame_counter, + &stats, + &idle_policy, + &mut idle_state, + &last_downstream_activity_ms, + session_started_at, + ) + .await; + + assert!(matches!(result, Err(ProxyError::Proxy(_)))); +} + +#[tokio::test] +async fn intermediate_alternating_zero_and_real_eventually_closes() { + let (reader, mut writer) = duplex(8192); + let mut crypto_reader = make_crypto_reader(reader); + let buffer_pool = Arc::new(BufferPool::new()); + let stats = Stats::new(); + let session_started_at = Instant::now(); + let forensics = make_forensics(13, session_started_at); + let mut frame_counter = 0u64; + let mut idle_state = RelayClientIdleState::new(session_started_at); + let idle_policy = make_enabled_idle_policy(); + let last_downstream_activity_ms = AtomicU64::new(0); + + let mut plaintext = Vec::with_capacity(3000); + for idx in 0..160u8 { + plaintext.extend_from_slice(&0u32.to_le_bytes()); + plaintext.extend_from_slice(&4u32.to_le_bytes()); + plaintext.extend_from_slice(&[idx, idx ^ 0x11, idx ^ 0x22, idx ^ 0x33]); + } + let encrypted = encrypt_for_reader(&plaintext); + writer.write_all(&encrypted).await.unwrap(); + drop(writer); + + let mut closed = false; + for _ in 0..220 { + let result = read_bounded( + &mut crypto_reader, + ProtoTag::Intermediate, + &buffer_pool, + &forensics, + &mut frame_counter, + &stats, + &idle_policy, + &mut idle_state, + &last_downstream_activity_ms, + session_started_at, + ) + .await; + + match result { + Ok(Some(_)) => {} + Err(ProxyError::Proxy(_)) => { + closed = true; + break; + } + Ok(None) => break, + Err(other) => panic!("unexpected error while probing alternating close: {other}"), + } + } + + assert!(closed, "intermediate alternating attack must fail closed"); +} + +#[tokio::test] +async fn small_tiny_burst_followed_by_real_frame_does_not_spuriously_close() { + let (reader, mut writer) = duplex(1024); + let mut crypto_reader = make_crypto_reader(reader); + let buffer_pool = Arc::new(BufferPool::new()); + let stats = Stats::new(); + let session_started_at = Instant::now(); + let forensics = make_forensics(14, session_started_at); + let mut frame_counter = 0u64; + let mut idle_state = RelayClientIdleState::new(session_started_at); + let idle_policy = make_enabled_idle_policy(); + let last_downstream_activity_ms = AtomicU64::new(0); + + let mut plaintext = Vec::with_capacity(64); + for _ in 0..8 { + plaintext.push(0x00); + } + plaintext.push(0x01); + plaintext.extend_from_slice(&[1, 2, 3, 4]); + + let encrypted = encrypt_for_reader(&plaintext); + writer.write_all(&encrypted).await.unwrap(); + + let first = read_bounded( + &mut crypto_reader, + ProtoTag::Abridged, + &buffer_pool, + &forensics, + &mut frame_counter, + &stats, + &idle_policy, + &mut idle_state, + &last_downstream_activity_ms, + session_started_at, + ) + .await; + + match first { + Ok(Some((payload, _))) => assert_eq!(payload.as_ref(), &[1, 2, 3, 4]), + Err(e) => panic!("unexpected close after small tiny burst: {e}"), + Ok(None) => panic!("unexpected EOF before real frame"), + } +} + +#[tokio::test] +async fn idle_policy_enabled_zero_length_flood_is_fail_closed() { + let (reader, mut writer) = duplex(4096); + let mut crypto_reader = make_crypto_reader(reader); + let buffer_pool = Arc::new(BufferPool::new()); + let stats = Stats::new(); + let session_started_at = Instant::now(); + let forensics = make_forensics(1, session_started_at); + let mut frame_counter = 0u64; + let mut idle_state = RelayClientIdleState::new(session_started_at); + let idle_policy = make_enabled_idle_policy(); + let last_downstream_activity_ms = AtomicU64::new(0); + + let flood_plaintext = vec![0u8; 1024]; + let flood_encrypted = encrypt_for_reader(&flood_plaintext); + writer + .write_all(&flood_encrypted) + .await + .expect("zero-length flood bytes must be writable"); + drop(writer); + + let result = read_bounded( + &mut crypto_reader, + ProtoTag::Abridged, + &buffer_pool, + &forensics, + &mut frame_counter, + &stats, + &idle_policy, + &mut idle_state, + &last_downstream_activity_ms, + session_started_at, + ) + .await; + + assert!( + matches!(result, Err(ProxyError::Proxy(_))), + "idle policy enabled must fail closed for pure zero-length flood" + ); +} + +#[tokio::test] +async fn idle_policy_enabled_alternating_tiny_real_eventually_closes() { + let (reader, mut writer) = duplex(8192); + let mut crypto_reader = make_crypto_reader(reader); + let buffer_pool = Arc::new(BufferPool::new()); + let stats = Stats::new(); + let session_started_at = Instant::now(); + let forensics = make_forensics(2, session_started_at); + let mut frame_counter = 0u64; + let mut idle_state = RelayClientIdleState::new(session_started_at); + let idle_policy = make_enabled_idle_policy(); + let last_downstream_activity_ms = AtomicU64::new(0); + + let mut plaintext = Vec::with_capacity(256 * 6); + for idx in 0..=255u8 { + plaintext.push(0x00); + plaintext.push(0x01); + plaintext.extend_from_slice(&[idx, idx ^ 0x55, idx ^ 0xAA, 0x11]); + } + + let encrypted = encrypt_for_reader(&plaintext); + writer + .write_all(&encrypted) + .await + .expect("alternating flood bytes must be writable"); + drop(writer); + + let mut saw_proxy_close = false; + for _ in 0..300 { + let result = read_bounded( + &mut crypto_reader, + ProtoTag::Abridged, + &buffer_pool, + &forensics, + &mut frame_counter, + &stats, + &idle_policy, + &mut idle_state, + &last_downstream_activity_ms, + session_started_at, + ) + .await; + + match result { + Ok(Some((_payload, _quickack))) => {} + Err(ProxyError::Proxy(_)) => { + saw_proxy_close = true; + break; + } + Err(ProxyError::Io(e)) => panic!("unexpected IO error before close: {e}"), + Ok(None) => panic!("unexpected EOF before debt-based closure"), + Err(other) => panic!("unexpected error before close: {other}"), + } + } + + assert!( + saw_proxy_close, + "alternating tiny/real sequence must eventually fail closed" + ); +} + +#[tokio::test] +async fn enabled_idle_policy_valid_nonzero_frame_still_passes() { + let (reader, mut writer) = duplex(1024); + let mut crypto_reader = make_crypto_reader(reader); + let buffer_pool = Arc::new(BufferPool::new()); + let stats = Stats::new(); + let session_started_at = Instant::now(); + let forensics = make_forensics(3, session_started_at); + let mut frame_counter = 0u64; + let mut idle_state = RelayClientIdleState::new(session_started_at); + let idle_policy = make_enabled_idle_policy(); + let last_downstream_activity_ms = AtomicU64::new(0); + + let payload = [7u8, 8, 9, 10]; + let mut plaintext = Vec::with_capacity(1 + payload.len()); + plaintext.push(0x01); + plaintext.extend_from_slice(&payload); + + let encrypted = encrypt_for_reader(&plaintext); + writer + .write_all(&encrypted) + .await + .expect("nonzero frame must be writable"); + + let result = read_bounded( + &mut crypto_reader, + ProtoTag::Abridged, + &buffer_pool, + &forensics, + &mut frame_counter, + &stats, + &idle_policy, + &mut idle_state, + &last_downstream_activity_ms, + session_started_at, + ) + .await + .expect("valid frame should decode") + .expect("valid frame should return payload"); + + assert_eq!(result.0.as_ref(), &payload); + assert!(!result.1); + assert_eq!(frame_counter, 1); +} diff --git a/src/proxy/traffic_limiter.rs b/src/proxy/traffic_limiter.rs index 4785075..a7a6068 100644 --- a/src/proxy/traffic_limiter.rs +++ b/src/proxy/traffic_limiter.rs @@ -1,17 +1,33 @@ use std::collections::{HashMap, HashSet}; -use std::hash::{Hash, Hasher}; use std::net::IpAddr; use std::sync::Arc; -use std::sync::OnceLock; use std::sync::atomic::{AtomicU64, Ordering}; -use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH}; use arc_swap::ArcSwap; use dashmap::DashMap; use ipnetwork::IpNetwork; -use crate::config::{CidrRateLimitKey, RateLimitBps}; +use crate::config::RateLimitBps; +// Atomic per-user and per-CIDR accounting. +mod buckets; +// Immutable policy matching and sharded registries. +mod policy; +// Traffic lease accounting and cleanup. +mod lease; +// Runtime policy application and admission. +mod limiter; +// Epoch and arithmetic helpers. +mod helpers; + +pub use helpers::next_refill_delay; +use helpers::{ + auto_cidr_bucket_key, bytes_per_epoch, current_epoch, decrement_atomic_saturating, + now_epoch_secs, +}; + +#[cfg(test)] +mod tests; const REGISTRY_SHARDS: usize = 64; const FAIR_EPOCH_MS: u64 = 20; const MAX_BORROW_CHUNK_BYTES: u64 = 32 * 1024; @@ -56,114 +72,18 @@ struct ScopeMetrics { policy_entries: AtomicU64, } -impl ScopeMetrics { - fn throttle(&self, direction: RateDirection) { - match direction { - RateDirection::Up => { - self.throttle_up_total.fetch_add(1, Ordering::Relaxed); - } - RateDirection::Down => { - self.throttle_down_total.fetch_add(1, Ordering::Relaxed); - } - } - } - - fn wait_ms(&self, direction: RateDirection, wait_ms: u64) { - match direction { - RateDirection::Up => { - self.wait_up_ms_total.fetch_add(wait_ms, Ordering::Relaxed); - } - RateDirection::Down => { - self.wait_down_ms_total - .fetch_add(wait_ms, Ordering::Relaxed); - } - } - } -} - #[derive(Default)] struct AtomicRatePair { up_bps: AtomicU64, down_bps: AtomicU64, } -impl AtomicRatePair { - fn set(&self, limits: RateLimitBps) { - self.up_bps.store(limits.up_bps, Ordering::Relaxed); - self.down_bps.store(limits.down_bps, Ordering::Relaxed); - } - - fn get(&self, direction: RateDirection) -> u64 { - match direction { - RateDirection::Up => self.up_bps.load(Ordering::Relaxed), - RateDirection::Down => self.down_bps.load(Ordering::Relaxed), - } - } -} - #[derive(Default)] struct DirectionBucket { epoch: AtomicU64, used: AtomicU64, } -impl DirectionBucket { - fn sync_epoch(&self, epoch: u64) { - let current = self.epoch.load(Ordering::Relaxed); - if current == epoch { - return; - } - if current < epoch - && self - .epoch - .compare_exchange(current, epoch, Ordering::Relaxed, Ordering::Relaxed) - .is_ok() - { - self.used.store(0, Ordering::Relaxed); - } - } - - fn try_consume(&self, cap_bps: u64, requested: u64) -> u64 { - if requested == 0 { - return 0; - } - if cap_bps == 0 { - return requested; - } - - let epoch = current_epoch(); - self.sync_epoch(epoch); - let cap_epoch = bytes_per_epoch(cap_bps); - - loop { - let used = self.used.load(Ordering::Relaxed); - if used >= cap_epoch { - return 0; - } - let remaining = cap_epoch.saturating_sub(used); - let grant = requested.min(remaining); - if grant == 0 { - return 0; - } - let next = used.saturating_add(grant); - if self - .used - .compare_exchange_weak(used, next, Ordering::Relaxed, Ordering::Relaxed) - .is_ok() - { - return grant; - } - } - } - - fn refund(&self, bytes: u64) { - if bytes == 0 { - return; - } - decrement_atomic_saturating(&self.used, bytes); - } -} - struct UserBucket { rates: AtomicRatePair, up: DirectionBucket, @@ -171,38 +91,6 @@ struct UserBucket { active_leases: AtomicU64, } -impl UserBucket { - fn new(limits: RateLimitBps) -> Self { - let rates = AtomicRatePair::default(); - rates.set(limits); - Self { - rates, - up: DirectionBucket::default(), - down: DirectionBucket::default(), - active_leases: AtomicU64::new(0), - } - } - - fn set_rates(&self, limits: RateLimitBps) { - self.rates.set(limits); - } - - fn try_consume(&self, direction: RateDirection, requested: u64) -> u64 { - let cap_bps = self.rates.get(direction); - match direction { - RateDirection::Up => self.up.try_consume(cap_bps, requested), - RateDirection::Down => self.down.try_consume(cap_bps, requested), - } - } - - fn refund(&self, direction: RateDirection, bytes: u64) { - match direction { - RateDirection::Up => self.up.refund(bytes), - RateDirection::Down => self.down.refund(bytes), - } - } -} - #[derive(Default)] struct CidrDirectionBucket { epoch: AtomicU64, @@ -210,125 +98,18 @@ struct CidrDirectionBucket { active_users: AtomicU64, } -impl CidrDirectionBucket { - fn sync_epoch(&self, epoch: u64) { - let current = self.epoch.load(Ordering::Relaxed); - if current == epoch { - return; - } - if current < epoch - && self - .epoch - .compare_exchange(current, epoch, Ordering::Relaxed, Ordering::Relaxed) - .is_ok() - { - self.used.store(0, Ordering::Relaxed); - self.active_users.store(0, Ordering::Relaxed); - } - } - - fn try_consume( - &self, - user_state: &CidrUserDirectionState, - cap_epoch: u64, - requested: u64, - ) -> u64 { - if requested == 0 || cap_epoch == 0 { - return 0; - } - - let epoch = current_epoch(); - self.sync_epoch(epoch); - user_state.sync_epoch_and_mark_active(epoch, &self.active_users); - let active_users = self.active_users.load(Ordering::Relaxed).max(1); - let fair_share = cap_epoch.saturating_div(active_users).max(1); - - loop { - let total_used = self.used.load(Ordering::Relaxed); - if total_used >= cap_epoch { - return 0; - } - let total_remaining = cap_epoch.saturating_sub(total_used); - let user_used = user_state.used.load(Ordering::Relaxed); - let guaranteed_remaining = fair_share.saturating_sub(user_used); - - let grant = if guaranteed_remaining > 0 { - requested.min(guaranteed_remaining).min(total_remaining) - } else { - requested.min(total_remaining).min(MAX_BORROW_CHUNK_BYTES) - }; - - if grant == 0 { - return 0; - } - - let next_total = total_used.saturating_add(grant); - if self - .used - .compare_exchange_weak(total_used, next_total, Ordering::Relaxed, Ordering::Relaxed) - .is_ok() - { - user_state.used.fetch_add(grant, Ordering::Relaxed); - return grant; - } - } - } - - fn refund(&self, bytes: u64) { - if bytes == 0 { - return; - } - decrement_atomic_saturating(&self.used, bytes); - } -} - #[derive(Default)] struct CidrUserDirectionState { epoch: AtomicU64, used: AtomicU64, } -impl CidrUserDirectionState { - fn sync_epoch_and_mark_active(&self, epoch: u64, active_users: &AtomicU64) { - let current = self.epoch.load(Ordering::Relaxed); - if current == epoch { - return; - } - if current < epoch - && self - .epoch - .compare_exchange(current, epoch, Ordering::Relaxed, Ordering::Relaxed) - .is_ok() - { - self.used.store(0, Ordering::Relaxed); - active_users.fetch_add(1, Ordering::Relaxed); - } - } - - fn refund(&self, bytes: u64) { - if bytes == 0 { - return; - } - decrement_atomic_saturating(&self.used, bytes); - } -} - struct CidrUserShare { active_conns: AtomicU64, up: CidrUserDirectionState, down: CidrUserDirectionState, } -impl CidrUserShare { - fn new() -> Self { - Self { - active_conns: AtomicU64::new(0), - up: CidrUserDirectionState::default(), - down: CidrUserDirectionState::default(), - } - } -} - struct CidrBucket { rates: AtomicRatePair, up: CidrDirectionBucket, @@ -337,75 +118,6 @@ struct CidrBucket { active_leases: AtomicU64, } -impl CidrBucket { - fn new(limits: RateLimitBps) -> Self { - let rates = AtomicRatePair::default(); - rates.set(limits); - Self { - rates, - up: CidrDirectionBucket::default(), - down: CidrDirectionBucket::default(), - users: ShardedRegistry::new(REGISTRY_SHARDS), - active_leases: AtomicU64::new(0), - } - } - - fn set_rates(&self, limits: RateLimitBps) { - self.rates.set(limits); - } - - fn acquire_user_share(&self, user: &str) -> Arc { - self.users - .get_or_insert_with(user, CidrUserShare::new, |share| { - share.active_conns.fetch_add(1, Ordering::Relaxed); - }) - } - - fn release_user_share(&self, user: &str, share: &Arc) { - decrement_atomic_saturating(&share.active_conns, 1); - let share_for_remove = Arc::clone(share); - let _ = self.users.remove_if(user, |candidate| { - Arc::ptr_eq(candidate, &share_for_remove) - && candidate.active_conns.load(Ordering::Relaxed) == 0 - }); - } - - fn try_consume_for_user( - &self, - direction: RateDirection, - share: &CidrUserShare, - requested: u64, - ) -> u64 { - let cap_bps = self.rates.get(direction); - if cap_bps == 0 { - return requested; - } - let cap_epoch = bytes_per_epoch(cap_bps); - match direction { - RateDirection::Up => self.up.try_consume(&share.up, cap_epoch, requested), - RateDirection::Down => self.down.try_consume(&share.down, cap_epoch, requested), - } - } - - fn refund_for_user(&self, direction: RateDirection, share: &CidrUserShare, bytes: u64) { - match direction { - RateDirection::Up => { - self.up.refund(bytes); - share.up.refund(bytes); - } - RateDirection::Down => { - self.down.refund(bytes); - share.down.refund(bytes); - } - } - } - - fn cleanup_idle_users(&self) { - self.users - .retain(|_, share| share.active_conns.load(Ordering::Relaxed) > 0); - } -} - #[derive(Clone)] struct CidrRule { key: String, @@ -435,97 +147,11 @@ struct PolicySnapshot { cidr_rule_keys: HashSet, } -impl PolicySnapshot { - fn match_cidr(&self, ip: IpAddr) -> Option> { - match ip { - IpAddr::V4(_) => self - .cidr_rules_v4 - .iter() - .find(|rule| rule.cidr.contains(ip)), - IpAddr::V6(_) => self - .cidr_rules_v6 - .iter() - .find(|rule| rule.cidr.contains(ip)), - } - .map(CidrPolicyMatch::Explicit) - .or_else(|| self.match_auto_cidr(ip)) - } - - fn match_auto_cidr(&self, ip: IpAddr) -> Option> { - let rule = match ip { - IpAddr::V4(_) => self.cidr_auto_rules_v4.first()?, - IpAddr::V6(_) => self.cidr_auto_rules_v6.first()?, - }; - let key = auto_cidr_bucket_key(ip, rule.prefix_len)?; - Some(CidrPolicyMatch::Auto { - key, - limits: rule.limits, - }) - } -} - struct ShardedRegistry { shards: Box<[DashMap>]>, mask: usize, } -impl ShardedRegistry { - fn new(shards: usize) -> Self { - let shard_count = shards.max(1).next_power_of_two(); - let mut items = Vec::with_capacity(shard_count); - for _ in 0..shard_count { - items.push(DashMap::>::new()); - } - Self { - shards: items.into_boxed_slice(), - mask: shard_count.saturating_sub(1), - } - } - - fn shard_index(&self, key: &str) -> usize { - let mut hasher = std::collections::hash_map::DefaultHasher::new(); - key.hash(&mut hasher); - (hasher.finish() as usize) & self.mask - } - - fn get_or_insert_with(&self, key: &str, make: F, activate: A) -> Arc - where - F: FnOnce() -> T, - A: FnOnce(&Arc), - { - let shard = &self.shards[self.shard_index(key)]; - match shard.entry(key.to_string()) { - dashmap::mapref::entry::Entry::Occupied(entry) => { - activate(entry.get()); - Arc::clone(entry.get()) - } - dashmap::mapref::entry::Entry::Vacant(slot) => { - let value = Arc::new(make()); - activate(&value); - slot.insert(Arc::clone(&value)); - value - } - } - } - - fn retain(&self, predicate: F) - where - F: Fn(&String, &Arc) -> bool + Copy, - { - for shard in &*self.shards { - shard.retain(|key, value| predicate(key, value)); - } - } - - fn remove_if(&self, key: &str, predicate: F) -> bool - where - F: Fn(&Arc) -> bool, - { - let shard = &self.shards[self.shard_index(key)]; - shard.remove_if(key, |_, value| predicate(value)).is_some() - } -} - pub struct TrafficLease { limiter: Arc, user_bucket: Option>, @@ -534,107 +160,6 @@ pub struct TrafficLease { cidr_user_share: Option>, } -impl TrafficLease { - pub fn try_consume(&self, direction: RateDirection, requested: u64) -> TrafficConsumeResult { - if requested == 0 { - return TrafficConsumeResult { - granted: 0, - blocked_user: false, - blocked_cidr: false, - }; - } - - let mut granted = requested; - if let Some(user_bucket) = self.user_bucket.as_ref() { - let user_granted = user_bucket.try_consume(direction, granted); - if user_granted == 0 { - self.limiter.observe_throttle(direction, true, false); - return TrafficConsumeResult { - granted: 0, - blocked_user: true, - blocked_cidr: false, - }; - } - granted = user_granted; - } - - if let (Some(cidr_bucket), Some(cidr_user_share)) = - (self.cidr_bucket.as_ref(), self.cidr_user_share.as_ref()) - { - let cidr_granted = - cidr_bucket.try_consume_for_user(direction, cidr_user_share, granted); - if cidr_granted < granted - && let Some(user_bucket) = self.user_bucket.as_ref() - { - user_bucket.refund(direction, granted.saturating_sub(cidr_granted)); - } - if cidr_granted == 0 { - self.limiter.observe_throttle(direction, false, true); - return TrafficConsumeResult { - granted: 0, - blocked_user: false, - blocked_cidr: true, - }; - } - granted = cidr_granted; - } - - TrafficConsumeResult { - granted, - blocked_user: false, - blocked_cidr: false, - } - } - - pub fn refund(&self, direction: RateDirection, bytes: u64) { - if bytes == 0 { - return; - } - - if let Some(user_bucket) = self.user_bucket.as_ref() { - user_bucket.refund(direction, bytes); - } - if let (Some(cidr_bucket), Some(cidr_user_share)) = - (self.cidr_bucket.as_ref(), self.cidr_user_share.as_ref()) - { - cidr_bucket.refund_for_user(direction, cidr_user_share, bytes); - } - } - - pub fn observe_wait_ms( - &self, - direction: RateDirection, - blocked_user: bool, - blocked_cidr: bool, - wait_ms: u64, - ) { - if wait_ms == 0 { - return; - } - self.limiter - .observe_wait(direction, blocked_user, blocked_cidr, wait_ms); - } -} - -impl Drop for TrafficLease { - fn drop(&mut self) { - if let Some(bucket) = self.user_bucket.as_ref() { - decrement_atomic_saturating(&bucket.active_leases, 1); - decrement_atomic_saturating(&self.limiter.user_scope.active_leases, 1); - } - - if let Some(bucket) = self.cidr_bucket.as_ref() { - if let (Some(user_key), Some(share)) = - (self.cidr_user_key.as_ref(), self.cidr_user_share.as_ref()) - { - bucket.release_user_share(user_key, share); - } - decrement_atomic_saturating(&bucket.active_leases, 1); - decrement_atomic_saturating(&self.limiter.cidr_scope.active_leases, 1); - } - } -} - pub struct TrafficLimiter { policy: ArcSwap, user_buckets: ShardedRegistry, @@ -643,357 +168,3 @@ pub struct TrafficLimiter { cidr_scope: ScopeMetrics, last_cleanup_epoch_secs: AtomicU64, } - -impl TrafficLimiter { - pub fn new() -> Arc { - Arc::new(Self { - policy: ArcSwap::from_pointee(PolicySnapshot::default()), - user_buckets: ShardedRegistry::new(REGISTRY_SHARDS), - cidr_buckets: ShardedRegistry::new(REGISTRY_SHARDS), - user_scope: ScopeMetrics::default(), - cidr_scope: ScopeMetrics::default(), - last_cleanup_epoch_secs: AtomicU64::new(0), - }) - } - - pub fn apply_policy( - &self, - user_limits: HashMap, - cidr_limits: HashMap, - ) { - let filtered_users = user_limits - .into_iter() - .filter(|(_, limit)| limit.up_bps > 0 || limit.down_bps > 0) - .collect::>(); - - let mut cidr_rules_v4 = Vec::new(); - let mut cidr_rules_v6 = Vec::new(); - let mut cidr_auto_rules_v4 = Vec::new(); - let mut cidr_auto_rules_v6 = Vec::new(); - let mut cidr_rule_keys = HashSet::new(); - for (key, limits) in cidr_limits { - if limits.up_bps == 0 && limits.down_bps == 0 { - continue; - } - match key { - CidrRateLimitKey::Network(cidr) => { - let key = cidr.to_string(); - let rule = CidrRule { - key: key.clone(), - cidr, - limits, - prefix_len: cidr.prefix(), - }; - cidr_rule_keys.insert(key); - match rule.cidr { - IpNetwork::V4(_) => cidr_rules_v4.push(rule), - IpNetwork::V6(_) => cidr_rules_v6.push(rule), - } - } - CidrRateLimitKey::AutoV4(prefix_len) => { - cidr_auto_rules_v4.push(CidrAutoRule { prefix_len, limits }); - } - CidrRateLimitKey::AutoV6(prefix_len) => { - cidr_auto_rules_v6.push(CidrAutoRule { prefix_len, limits }); - } - CidrRateLimitKey::AutoDual(prefix_len) => { - cidr_auto_rules_v4.push(CidrAutoRule { prefix_len, limits }); - cidr_auto_rules_v6.push(CidrAutoRule { - prefix_len: prefix_len.saturating_mul(4), - limits, - }); - } - } - } - - cidr_rules_v4.sort_by(|a, b| b.prefix_len.cmp(&a.prefix_len)); - cidr_rules_v6.sort_by(|a, b| b.prefix_len.cmp(&a.prefix_len)); - cidr_auto_rules_v4.sort_by(|a, b| b.prefix_len.cmp(&a.prefix_len)); - cidr_auto_rules_v6.sort_by(|a, b| b.prefix_len.cmp(&a.prefix_len)); - let cidr_policy_entries = - cidr_rule_keys.len() + cidr_auto_rules_v4.len() + cidr_auto_rules_v6.len(); - - self.user_scope - .policy_entries - .store(filtered_users.len() as u64, Ordering::Relaxed); - self.cidr_scope - .policy_entries - .store(cidr_policy_entries as u64, Ordering::Relaxed); - - self.policy.store(Arc::new(PolicySnapshot { - user_limits: filtered_users, - cidr_rules_v4, - cidr_rules_v6, - cidr_auto_rules_v4, - cidr_auto_rules_v6, - cidr_rule_keys, - })); - - self.maybe_cleanup(); - } - - pub fn acquire_lease( - self: &Arc, - user: &str, - client_ip: IpAddr, - ) -> Option> { - let policy = self.policy.load_full(); - let mut user_bucket = None; - if let Some(limit) = policy.user_limits.get(user).copied() { - let bucket = self - .user_buckets - .get_or_insert_with(user, || UserBucket::new(limit), |bucket| { - bucket.active_leases.fetch_add(1, Ordering::Relaxed); - }); - bucket.set_rates(limit); - self.user_scope - .active_leases - .fetch_add(1, Ordering::Relaxed); - user_bucket = Some(bucket); - } - - let mut cidr_bucket = None; - let mut cidr_user_key = None; - let mut cidr_user_share = None; - if let Some(rule_match) = policy.match_cidr(client_ip) { - let (key, limits) = match &rule_match { - CidrPolicyMatch::Explicit(rule) => (rule.key.as_str(), rule.limits), - CidrPolicyMatch::Auto { key, limits } => (key.as_str(), *limits), - }; - let bucket = self - .cidr_buckets - .get_or_insert_with(key, || CidrBucket::new(limits), |bucket| { - bucket.active_leases.fetch_add(1, Ordering::Relaxed); - }); - bucket.set_rates(limits); - self.cidr_scope - .active_leases - .fetch_add(1, Ordering::Relaxed); - let share = bucket.acquire_user_share(user); - cidr_user_key = Some(user.to_string()); - cidr_user_share = Some(share); - cidr_bucket = Some(bucket); - } - - if user_bucket.is_none() && cidr_bucket.is_none() { - return None; - } - - self.maybe_cleanup(); - Some(Arc::new(TrafficLease { - limiter: Arc::clone(self), - user_bucket, - cidr_bucket, - cidr_user_key, - cidr_user_share, - })) - } - - pub fn metrics_snapshot(&self) -> TrafficLimiterMetricsSnapshot { - TrafficLimiterMetricsSnapshot { - user_throttle_up_total: self.user_scope.throttle_up_total.load(Ordering::Relaxed), - user_throttle_down_total: self.user_scope.throttle_down_total.load(Ordering::Relaxed), - cidr_throttle_up_total: self.cidr_scope.throttle_up_total.load(Ordering::Relaxed), - cidr_throttle_down_total: self.cidr_scope.throttle_down_total.load(Ordering::Relaxed), - user_wait_up_ms_total: self.user_scope.wait_up_ms_total.load(Ordering::Relaxed), - user_wait_down_ms_total: self.user_scope.wait_down_ms_total.load(Ordering::Relaxed), - cidr_wait_up_ms_total: self.cidr_scope.wait_up_ms_total.load(Ordering::Relaxed), - cidr_wait_down_ms_total: self.cidr_scope.wait_down_ms_total.load(Ordering::Relaxed), - user_active_leases: self.user_scope.active_leases.load(Ordering::Relaxed), - cidr_active_leases: self.cidr_scope.active_leases.load(Ordering::Relaxed), - user_policy_entries: self.user_scope.policy_entries.load(Ordering::Relaxed), - cidr_policy_entries: self.cidr_scope.policy_entries.load(Ordering::Relaxed), - } - } - - fn observe_throttle(&self, direction: RateDirection, blocked_user: bool, blocked_cidr: bool) { - if blocked_user { - self.user_scope.throttle(direction); - } - if blocked_cidr { - self.cidr_scope.throttle(direction); - } - } - - fn observe_wait( - &self, - direction: RateDirection, - blocked_user: bool, - blocked_cidr: bool, - wait_ms: u64, - ) { - if blocked_user { - self.user_scope.wait_ms(direction, wait_ms); - } - if blocked_cidr { - self.cidr_scope.wait_ms(direction, wait_ms); - } - } - - fn maybe_cleanup(&self) { - let now_epoch_secs = now_epoch_secs(); - let last = self.last_cleanup_epoch_secs.load(Ordering::Relaxed); - if now_epoch_secs.saturating_sub(last) < CLEANUP_INTERVAL_SECS { - return; - } - if self - .last_cleanup_epoch_secs - .compare_exchange(last, now_epoch_secs, Ordering::Relaxed, Ordering::Relaxed) - .is_err() - { - return; - } - - let policy = self.policy.load_full(); - self.user_buckets.retain(|user, bucket| { - bucket.active_leases.load(Ordering::Relaxed) > 0 - || policy.user_limits.contains_key(user) - }); - self.cidr_buckets.retain(|cidr_key, bucket| { - bucket.cleanup_idle_users(); - bucket.active_leases.load(Ordering::Relaxed) > 0 - || policy.cidr_rule_keys.contains(cidr_key) - }); - } -} - -pub fn next_refill_delay() -> Duration { - let start = limiter_epoch_start(); - let elapsed_ms = start.elapsed().as_millis() as u64; - let epoch_pos = elapsed_ms % FAIR_EPOCH_MS; - let wait_ms = FAIR_EPOCH_MS.saturating_sub(epoch_pos).max(1); - Duration::from_millis(wait_ms) -} - -fn decrement_atomic_saturating(counter: &AtomicU64, by: u64) { - if by == 0 { - return; - } - let mut current = counter.load(Ordering::Relaxed); - loop { - if current == 0 { - return; - } - let next = current.saturating_sub(by); - match counter.compare_exchange_weak(current, next, Ordering::Relaxed, Ordering::Relaxed) { - Ok(_) => return, - Err(actual) => current = actual, - } - } -} - -fn now_epoch_secs() -> u64 { - SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap_or_default() - .as_secs() -} - -fn bytes_per_epoch(bps: u64) -> u64 { - if bps == 0 { - return 0; - } - let numerator = bps.saturating_mul(FAIR_EPOCH_MS); - let bytes = numerator.saturating_div(8_000); - bytes.max(1) -} - -fn auto_cidr_bucket_key(ip: IpAddr, prefix_len: u8) -> Option { - let cidr = IpNetwork::new(ip, prefix_len).ok()?; - let network = IpNetwork::new(cidr.network(), prefix_len).ok()?; - let family = match network { - IpNetwork::V4(_) => "4", - IpNetwork::V6(_) => "6", - }; - Some(format!("auto:{family}:{network}")) -} - -fn current_epoch() -> u64 { - let start = limiter_epoch_start(); - let elapsed_ms = start.elapsed().as_millis() as u64; - elapsed_ms / FAIR_EPOCH_MS -} - -fn limiter_epoch_start() -> &'static Instant { - static START: OnceLock = OnceLock::new(); - START.get_or_init(Instant::now) -} - -#[cfg(test)] -mod tests { - use super::*; - - fn rate(up_bps: u64, down_bps: u64) -> RateLimitBps { - RateLimitBps { up_bps, down_bps } - } - - #[test] - fn explicit_cidr_rule_wins_over_auto_template() { - let limiter = TrafficLimiter::new(); - let mut cidr_limits = HashMap::new(); - cidr_limits.insert(CidrRateLimitKey::AutoV4(24), rate(1_000, 0)); - cidr_limits.insert( - CidrRateLimitKey::Network("203.0.113.7/32".parse().unwrap()), - rate(2_000, 0), - ); - - limiter.apply_policy(HashMap::new(), cidr_limits); - let policy = limiter.policy.load_full(); - let matched = policy.match_cidr("203.0.113.7".parse().unwrap()).unwrap(); - - match matched { - CidrPolicyMatch::Explicit(rule) => assert_eq!(rule.key.as_str(), "203.0.113.7/32"), - CidrPolicyMatch::Auto { .. } => panic!("explicit CIDR must have priority"), - } - } - - #[test] - fn auto_template_uses_longest_prefix() { - let limiter = TrafficLimiter::new(); - let mut cidr_limits = HashMap::new(); - cidr_limits.insert(CidrRateLimitKey::AutoV4(24), rate(1_000, 0)); - cidr_limits.insert(CidrRateLimitKey::AutoV4(32), rate(2_000, 0)); - - limiter.apply_policy(HashMap::new(), cidr_limits); - let policy = limiter.policy.load_full(); - let matched = policy.match_cidr("203.0.113.129".parse().unwrap()).unwrap(); - - match matched { - CidrPolicyMatch::Auto { key, limits } => { - assert_eq!(key, "auto:4:203.0.113.129/32"); - assert_eq!(limits.up_bps, 2_000); - } - CidrPolicyMatch::Explicit(_) => panic!("auto-template match expected"), - } - } - - #[test] - fn dual_auto_template_maps_v6_prefix_by_four() { - let limiter = TrafficLimiter::new(); - let mut cidr_limits = HashMap::new(); - cidr_limits.insert(CidrRateLimitKey::AutoDual(32), rate(1_000, 0)); - - limiter.apply_policy(HashMap::new(), cidr_limits); - let policy = limiter.policy.load_full(); - let matched = policy.match_cidr("2001:db8::1".parse().unwrap()).unwrap(); - - match matched { - CidrPolicyMatch::Auto { key, .. } => { - assert_eq!(key, "auto:6:2001:db8::1/128"); - } - CidrPolicyMatch::Explicit(_) => panic!("auto-template match expected"), - } - } - - #[test] - fn auto_cidr_bucket_key_canonicalizes_network_address() { - assert_eq!( - auto_cidr_bucket_key("203.0.113.129".parse().unwrap(), 24).unwrap(), - "auto:4:203.0.113.0/24" - ); - assert_eq!( - auto_cidr_bucket_key("2001:db8::abcd".parse().unwrap(), 64).unwrap(), - "auto:6:2001:db8::/64" - ); - } -} diff --git a/src/proxy/traffic_limiter/buckets.rs b/src/proxy/traffic_limiter/buckets.rs new file mode 100644 index 0000000..f6fe592 --- /dev/null +++ b/src/proxy/traffic_limiter/buckets.rs @@ -0,0 +1,310 @@ +use super::*; + +impl ScopeMetrics { + pub(super) fn throttle(&self, direction: RateDirection) { + match direction { + RateDirection::Up => { + self.throttle_up_total.fetch_add(1, Ordering::Relaxed); + } + RateDirection::Down => { + self.throttle_down_total.fetch_add(1, Ordering::Relaxed); + } + } + } + + pub(super) fn wait_ms(&self, direction: RateDirection, wait_ms: u64) { + match direction { + RateDirection::Up => { + self.wait_up_ms_total.fetch_add(wait_ms, Ordering::Relaxed); + } + RateDirection::Down => { + self.wait_down_ms_total + .fetch_add(wait_ms, Ordering::Relaxed); + } + } + } +} + +impl AtomicRatePair { + pub(super) fn set(&self, limits: RateLimitBps) { + self.up_bps.store(limits.up_bps, Ordering::Relaxed); + self.down_bps.store(limits.down_bps, Ordering::Relaxed); + } + + pub(super) fn get(&self, direction: RateDirection) -> u64 { + match direction { + RateDirection::Up => self.up_bps.load(Ordering::Relaxed), + RateDirection::Down => self.down_bps.load(Ordering::Relaxed), + } + } +} + +impl DirectionBucket { + pub(super) fn sync_epoch(&self, epoch: u64) { + let current = self.epoch.load(Ordering::Relaxed); + if current == epoch { + return; + } + if current < epoch + && self + .epoch + .compare_exchange(current, epoch, Ordering::Relaxed, Ordering::Relaxed) + .is_ok() + { + self.used.store(0, Ordering::Relaxed); + } + } + + pub(super) fn try_consume(&self, cap_bps: u64, requested: u64) -> u64 { + if requested == 0 { + return 0; + } + if cap_bps == 0 { + return requested; + } + + let epoch = current_epoch(); + self.sync_epoch(epoch); + let cap_epoch = bytes_per_epoch(cap_bps); + + loop { + let used = self.used.load(Ordering::Relaxed); + if used >= cap_epoch { + return 0; + } + let remaining = cap_epoch.saturating_sub(used); + let grant = requested.min(remaining); + if grant == 0 { + return 0; + } + let next = used.saturating_add(grant); + if self + .used + .compare_exchange_weak(used, next, Ordering::Relaxed, Ordering::Relaxed) + .is_ok() + { + return grant; + } + } + } + + pub(super) fn refund(&self, bytes: u64) { + if bytes == 0 { + return; + } + decrement_atomic_saturating(&self.used, bytes); + } +} + +impl UserBucket { + pub(super) fn new(limits: RateLimitBps) -> Self { + let rates = AtomicRatePair::default(); + rates.set(limits); + Self { + rates, + up: DirectionBucket::default(), + down: DirectionBucket::default(), + active_leases: AtomicU64::new(0), + } + } + + pub(super) fn set_rates(&self, limits: RateLimitBps) { + self.rates.set(limits); + } + + pub(super) fn try_consume(&self, direction: RateDirection, requested: u64) -> u64 { + let cap_bps = self.rates.get(direction); + match direction { + RateDirection::Up => self.up.try_consume(cap_bps, requested), + RateDirection::Down => self.down.try_consume(cap_bps, requested), + } + } + + pub(super) fn refund(&self, direction: RateDirection, bytes: u64) { + match direction { + RateDirection::Up => self.up.refund(bytes), + RateDirection::Down => self.down.refund(bytes), + } + } +} + +impl CidrDirectionBucket { + pub(super) fn sync_epoch(&self, epoch: u64) { + let current = self.epoch.load(Ordering::Relaxed); + if current == epoch { + return; + } + if current < epoch + && self + .epoch + .compare_exchange(current, epoch, Ordering::Relaxed, Ordering::Relaxed) + .is_ok() + { + self.used.store(0, Ordering::Relaxed); + self.active_users.store(0, Ordering::Relaxed); + } + } + + pub(super) fn try_consume( + &self, + user_state: &CidrUserDirectionState, + cap_epoch: u64, + requested: u64, + ) -> u64 { + if requested == 0 || cap_epoch == 0 { + return 0; + } + + let epoch = current_epoch(); + self.sync_epoch(epoch); + user_state.sync_epoch_and_mark_active(epoch, &self.active_users); + let active_users = self.active_users.load(Ordering::Relaxed).max(1); + let fair_share = cap_epoch.saturating_div(active_users).max(1); + + loop { + let total_used = self.used.load(Ordering::Relaxed); + if total_used >= cap_epoch { + return 0; + } + let total_remaining = cap_epoch.saturating_sub(total_used); + let user_used = user_state.used.load(Ordering::Relaxed); + let guaranteed_remaining = fair_share.saturating_sub(user_used); + + let grant = if guaranteed_remaining > 0 { + requested.min(guaranteed_remaining).min(total_remaining) + } else { + requested.min(total_remaining).min(MAX_BORROW_CHUNK_BYTES) + }; + + if grant == 0 { + return 0; + } + + let next_total = total_used.saturating_add(grant); + if self + .used + .compare_exchange_weak(total_used, next_total, Ordering::Relaxed, Ordering::Relaxed) + .is_ok() + { + user_state.used.fetch_add(grant, Ordering::Relaxed); + return grant; + } + } + } + + pub(super) fn refund(&self, bytes: u64) { + if bytes == 0 { + return; + } + decrement_atomic_saturating(&self.used, bytes); + } +} + +impl CidrUserDirectionState { + pub(super) fn sync_epoch_and_mark_active(&self, epoch: u64, active_users: &AtomicU64) { + let current = self.epoch.load(Ordering::Relaxed); + if current == epoch { + return; + } + if current < epoch + && self + .epoch + .compare_exchange(current, epoch, Ordering::Relaxed, Ordering::Relaxed) + .is_ok() + { + self.used.store(0, Ordering::Relaxed); + active_users.fetch_add(1, Ordering::Relaxed); + } + } + + pub(super) fn refund(&self, bytes: u64) { + if bytes == 0 { + return; + } + decrement_atomic_saturating(&self.used, bytes); + } +} + +impl CidrUserShare { + pub(super) fn new() -> Self { + Self { + active_conns: AtomicU64::new(0), + up: CidrUserDirectionState::default(), + down: CidrUserDirectionState::default(), + } + } +} + +impl CidrBucket { + pub(super) fn new(limits: RateLimitBps) -> Self { + let rates = AtomicRatePair::default(); + rates.set(limits); + Self { + rates, + up: CidrDirectionBucket::default(), + down: CidrDirectionBucket::default(), + users: ShardedRegistry::new(REGISTRY_SHARDS), + active_leases: AtomicU64::new(0), + } + } + + pub(super) fn set_rates(&self, limits: RateLimitBps) { + self.rates.set(limits); + } + + pub(super) fn acquire_user_share(&self, user: &str) -> Arc { + self.users + .get_or_insert_with(user, CidrUserShare::new, |share| { + share.active_conns.fetch_add(1, Ordering::Relaxed); + }) + } + + pub(super) fn release_user_share(&self, user: &str, share: &Arc) { + decrement_atomic_saturating(&share.active_conns, 1); + let share_for_remove = Arc::clone(share); + let _ = self.users.remove_if(user, |candidate| { + Arc::ptr_eq(candidate, &share_for_remove) + && candidate.active_conns.load(Ordering::Relaxed) == 0 + }); + } + + pub(super) fn try_consume_for_user( + &self, + direction: RateDirection, + share: &CidrUserShare, + requested: u64, + ) -> u64 { + let cap_bps = self.rates.get(direction); + if cap_bps == 0 { + return requested; + } + let cap_epoch = bytes_per_epoch(cap_bps); + match direction { + RateDirection::Up => self.up.try_consume(&share.up, cap_epoch, requested), + RateDirection::Down => self.down.try_consume(&share.down, cap_epoch, requested), + } + } + + pub(super) fn refund_for_user( + &self, + direction: RateDirection, + share: &CidrUserShare, + bytes: u64, + ) { + match direction { + RateDirection::Up => { + self.up.refund(bytes); + share.up.refund(bytes); + } + RateDirection::Down => { + self.down.refund(bytes); + share.down.refund(bytes); + } + } + } + + pub(super) fn cleanup_idle_users(&self) { + self.users + .retain(|_, share| share.active_conns.load(Ordering::Relaxed) > 0); + } +} diff --git a/src/proxy/traffic_limiter/helpers.rs b/src/proxy/traffic_limiter/helpers.rs new file mode 100644 index 0000000..272fe03 --- /dev/null +++ b/src/proxy/traffic_limiter/helpers.rs @@ -0,0 +1,65 @@ +use std::sync::OnceLock; +use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH}; + +use super::*; +pub fn next_refill_delay() -> Duration { + let start = limiter_epoch_start(); + let elapsed_ms = start.elapsed().as_millis() as u64; + let epoch_pos = elapsed_ms % FAIR_EPOCH_MS; + let wait_ms = FAIR_EPOCH_MS.saturating_sub(epoch_pos).max(1); + Duration::from_millis(wait_ms) +} + +pub(super) fn decrement_atomic_saturating(counter: &AtomicU64, by: u64) { + if by == 0 { + return; + } + let mut current = counter.load(Ordering::Relaxed); + loop { + if current == 0 { + return; + } + let next = current.saturating_sub(by); + match counter.compare_exchange_weak(current, next, Ordering::Relaxed, Ordering::Relaxed) { + Ok(_) => return, + Err(actual) => current = actual, + } + } +} + +pub(super) fn now_epoch_secs() -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs() +} + +pub(super) fn bytes_per_epoch(bps: u64) -> u64 { + if bps == 0 { + return 0; + } + let numerator = bps.saturating_mul(FAIR_EPOCH_MS); + let bytes = numerator.saturating_div(8_000); + bytes.max(1) +} + +pub(super) fn auto_cidr_bucket_key(ip: IpAddr, prefix_len: u8) -> Option { + let cidr = IpNetwork::new(ip, prefix_len).ok()?; + let network = IpNetwork::new(cidr.network(), prefix_len).ok()?; + let family = match network { + IpNetwork::V4(_) => "4", + IpNetwork::V6(_) => "6", + }; + Some(format!("auto:{family}:{network}")) +} + +pub(super) fn current_epoch() -> u64 { + let start = limiter_epoch_start(); + let elapsed_ms = start.elapsed().as_millis() as u64; + elapsed_ms / FAIR_EPOCH_MS +} + +pub(super) fn limiter_epoch_start() -> &'static Instant { + static START: OnceLock = OnceLock::new(); + START.get_or_init(Instant::now) +} diff --git a/src/proxy/traffic_limiter/lease.rs b/src/proxy/traffic_limiter/lease.rs new file mode 100644 index 0000000..ad136af --- /dev/null +++ b/src/proxy/traffic_limiter/lease.rs @@ -0,0 +1,102 @@ +use super::*; + +impl TrafficLease { + pub fn try_consume(&self, direction: RateDirection, requested: u64) -> TrafficConsumeResult { + if requested == 0 { + return TrafficConsumeResult { + granted: 0, + blocked_user: false, + blocked_cidr: false, + }; + } + + let mut granted = requested; + if let Some(user_bucket) = self.user_bucket.as_ref() { + let user_granted = user_bucket.try_consume(direction, granted); + if user_granted == 0 { + self.limiter.observe_throttle(direction, true, false); + return TrafficConsumeResult { + granted: 0, + blocked_user: true, + blocked_cidr: false, + }; + } + granted = user_granted; + } + + if let (Some(cidr_bucket), Some(cidr_user_share)) = + (self.cidr_bucket.as_ref(), self.cidr_user_share.as_ref()) + { + let cidr_granted = + cidr_bucket.try_consume_for_user(direction, cidr_user_share, granted); + if cidr_granted < granted + && let Some(user_bucket) = self.user_bucket.as_ref() + { + user_bucket.refund(direction, granted.saturating_sub(cidr_granted)); + } + if cidr_granted == 0 { + self.limiter.observe_throttle(direction, false, true); + return TrafficConsumeResult { + granted: 0, + blocked_user: false, + blocked_cidr: true, + }; + } + granted = cidr_granted; + } + + TrafficConsumeResult { + granted, + blocked_user: false, + blocked_cidr: false, + } + } + + pub fn refund(&self, direction: RateDirection, bytes: u64) { + if bytes == 0 { + return; + } + + if let Some(user_bucket) = self.user_bucket.as_ref() { + user_bucket.refund(direction, bytes); + } + if let (Some(cidr_bucket), Some(cidr_user_share)) = + (self.cidr_bucket.as_ref(), self.cidr_user_share.as_ref()) + { + cidr_bucket.refund_for_user(direction, cidr_user_share, bytes); + } + } + + pub fn observe_wait_ms( + &self, + direction: RateDirection, + blocked_user: bool, + blocked_cidr: bool, + wait_ms: u64, + ) { + if wait_ms == 0 { + return; + } + self.limiter + .observe_wait(direction, blocked_user, blocked_cidr, wait_ms); + } +} + +impl Drop for TrafficLease { + fn drop(&mut self) { + if let Some(bucket) = self.user_bucket.as_ref() { + decrement_atomic_saturating(&bucket.active_leases, 1); + decrement_atomic_saturating(&self.limiter.user_scope.active_leases, 1); + } + + if let Some(bucket) = self.cidr_bucket.as_ref() { + if let (Some(user_key), Some(share)) = + (self.cidr_user_key.as_ref(), self.cidr_user_share.as_ref()) + { + bucket.release_user_share(user_key, share); + } + decrement_atomic_saturating(&bucket.active_leases, 1); + decrement_atomic_saturating(&self.limiter.cidr_scope.active_leases, 1); + } + } +} diff --git a/src/proxy/traffic_limiter/limiter.rs b/src/proxy/traffic_limiter/limiter.rs new file mode 100644 index 0000000..4aaf411 --- /dev/null +++ b/src/proxy/traffic_limiter/limiter.rs @@ -0,0 +1,224 @@ +use crate::config::CidrRateLimitKey; + +use super::*; +impl TrafficLimiter { + pub fn new() -> Arc { + Arc::new(Self { + policy: ArcSwap::from_pointee(PolicySnapshot::default()), + user_buckets: ShardedRegistry::new(REGISTRY_SHARDS), + cidr_buckets: ShardedRegistry::new(REGISTRY_SHARDS), + user_scope: ScopeMetrics::default(), + cidr_scope: ScopeMetrics::default(), + last_cleanup_epoch_secs: AtomicU64::new(0), + }) + } + + pub fn apply_policy( + &self, + user_limits: HashMap, + cidr_limits: HashMap, + ) { + let filtered_users = user_limits + .into_iter() + .filter(|(_, limit)| limit.up_bps > 0 || limit.down_bps > 0) + .collect::>(); + + let mut cidr_rules_v4 = Vec::new(); + let mut cidr_rules_v6 = Vec::new(); + let mut cidr_auto_rules_v4 = Vec::new(); + let mut cidr_auto_rules_v6 = Vec::new(); + let mut cidr_rule_keys = HashSet::new(); + for (key, limits) in cidr_limits { + if limits.up_bps == 0 && limits.down_bps == 0 { + continue; + } + match key { + CidrRateLimitKey::Network(cidr) => { + let key = cidr.to_string(); + let rule = CidrRule { + key: key.clone(), + cidr, + limits, + prefix_len: cidr.prefix(), + }; + cidr_rule_keys.insert(key); + match rule.cidr { + IpNetwork::V4(_) => cidr_rules_v4.push(rule), + IpNetwork::V6(_) => cidr_rules_v6.push(rule), + } + } + CidrRateLimitKey::AutoV4(prefix_len) => { + cidr_auto_rules_v4.push(CidrAutoRule { prefix_len, limits }); + } + CidrRateLimitKey::AutoV6(prefix_len) => { + cidr_auto_rules_v6.push(CidrAutoRule { prefix_len, limits }); + } + CidrRateLimitKey::AutoDual(prefix_len) => { + cidr_auto_rules_v4.push(CidrAutoRule { prefix_len, limits }); + cidr_auto_rules_v6.push(CidrAutoRule { + prefix_len: prefix_len.saturating_mul(4), + limits, + }); + } + } + } + + cidr_rules_v4.sort_by(|a, b| b.prefix_len.cmp(&a.prefix_len)); + cidr_rules_v6.sort_by(|a, b| b.prefix_len.cmp(&a.prefix_len)); + cidr_auto_rules_v4.sort_by(|a, b| b.prefix_len.cmp(&a.prefix_len)); + cidr_auto_rules_v6.sort_by(|a, b| b.prefix_len.cmp(&a.prefix_len)); + let cidr_policy_entries = + cidr_rule_keys.len() + cidr_auto_rules_v4.len() + cidr_auto_rules_v6.len(); + + self.user_scope + .policy_entries + .store(filtered_users.len() as u64, Ordering::Relaxed); + self.cidr_scope + .policy_entries + .store(cidr_policy_entries as u64, Ordering::Relaxed); + + self.policy.store(Arc::new(PolicySnapshot { + user_limits: filtered_users, + cidr_rules_v4, + cidr_rules_v6, + cidr_auto_rules_v4, + cidr_auto_rules_v6, + cidr_rule_keys, + })); + + self.maybe_cleanup(); + } + + pub fn acquire_lease( + self: &Arc, + user: &str, + client_ip: IpAddr, + ) -> Option> { + let policy = self.policy.load_full(); + let mut user_bucket = None; + if let Some(limit) = policy.user_limits.get(user).copied() { + let bucket = self.user_buckets.get_or_insert_with( + user, + || UserBucket::new(limit), + |bucket| { + bucket.active_leases.fetch_add(1, Ordering::Relaxed); + }, + ); + bucket.set_rates(limit); + self.user_scope + .active_leases + .fetch_add(1, Ordering::Relaxed); + user_bucket = Some(bucket); + } + + let mut cidr_bucket = None; + let mut cidr_user_key = None; + let mut cidr_user_share = None; + if let Some(rule_match) = policy.match_cidr(client_ip) { + let (key, limits) = match &rule_match { + CidrPolicyMatch::Explicit(rule) => (rule.key.as_str(), rule.limits), + CidrPolicyMatch::Auto { key, limits } => (key.as_str(), *limits), + }; + let bucket = self.cidr_buckets.get_or_insert_with( + key, + || CidrBucket::new(limits), + |bucket| { + bucket.active_leases.fetch_add(1, Ordering::Relaxed); + }, + ); + bucket.set_rates(limits); + self.cidr_scope + .active_leases + .fetch_add(1, Ordering::Relaxed); + let share = bucket.acquire_user_share(user); + cidr_user_key = Some(user.to_string()); + cidr_user_share = Some(share); + cidr_bucket = Some(bucket); + } + + if user_bucket.is_none() && cidr_bucket.is_none() { + return None; + } + + self.maybe_cleanup(); + Some(Arc::new(TrafficLease { + limiter: Arc::clone(self), + user_bucket, + cidr_bucket, + cidr_user_key, + cidr_user_share, + })) + } + + pub fn metrics_snapshot(&self) -> TrafficLimiterMetricsSnapshot { + TrafficLimiterMetricsSnapshot { + user_throttle_up_total: self.user_scope.throttle_up_total.load(Ordering::Relaxed), + user_throttle_down_total: self.user_scope.throttle_down_total.load(Ordering::Relaxed), + cidr_throttle_up_total: self.cidr_scope.throttle_up_total.load(Ordering::Relaxed), + cidr_throttle_down_total: self.cidr_scope.throttle_down_total.load(Ordering::Relaxed), + user_wait_up_ms_total: self.user_scope.wait_up_ms_total.load(Ordering::Relaxed), + user_wait_down_ms_total: self.user_scope.wait_down_ms_total.load(Ordering::Relaxed), + cidr_wait_up_ms_total: self.cidr_scope.wait_up_ms_total.load(Ordering::Relaxed), + cidr_wait_down_ms_total: self.cidr_scope.wait_down_ms_total.load(Ordering::Relaxed), + user_active_leases: self.user_scope.active_leases.load(Ordering::Relaxed), + cidr_active_leases: self.cidr_scope.active_leases.load(Ordering::Relaxed), + user_policy_entries: self.user_scope.policy_entries.load(Ordering::Relaxed), + cidr_policy_entries: self.cidr_scope.policy_entries.load(Ordering::Relaxed), + } + } + + pub(super) fn observe_throttle( + &self, + direction: RateDirection, + blocked_user: bool, + blocked_cidr: bool, + ) { + if blocked_user { + self.user_scope.throttle(direction); + } + if blocked_cidr { + self.cidr_scope.throttle(direction); + } + } + + pub(super) fn observe_wait( + &self, + direction: RateDirection, + blocked_user: bool, + blocked_cidr: bool, + wait_ms: u64, + ) { + if blocked_user { + self.user_scope.wait_ms(direction, wait_ms); + } + if blocked_cidr { + self.cidr_scope.wait_ms(direction, wait_ms); + } + } + + pub(super) fn maybe_cleanup(&self) { + let now_epoch_secs = now_epoch_secs(); + let last = self.last_cleanup_epoch_secs.load(Ordering::Relaxed); + if now_epoch_secs.saturating_sub(last) < CLEANUP_INTERVAL_SECS { + return; + } + if self + .last_cleanup_epoch_secs + .compare_exchange(last, now_epoch_secs, Ordering::Relaxed, Ordering::Relaxed) + .is_err() + { + return; + } + + let policy = self.policy.load_full(); + self.user_buckets.retain(|user, bucket| { + bucket.active_leases.load(Ordering::Relaxed) > 0 + || policy.user_limits.contains_key(user) + }); + self.cidr_buckets.retain(|cidr_key, bucket| { + bucket.cleanup_idle_users(); + bucket.active_leases.load(Ordering::Relaxed) > 0 + || policy.cidr_rule_keys.contains(cidr_key) + }); + } +} diff --git a/src/proxy/traffic_limiter/policy.rs b/src/proxy/traffic_limiter/policy.rs new file mode 100644 index 0000000..21f0e67 --- /dev/null +++ b/src/proxy/traffic_limiter/policy.rs @@ -0,0 +1,88 @@ +use std::hash::{Hash, Hasher}; + +use super::*; +impl PolicySnapshot { + pub(super) fn match_cidr(&self, ip: IpAddr) -> Option> { + match ip { + IpAddr::V4(_) => self + .cidr_rules_v4 + .iter() + .find(|rule| rule.cidr.contains(ip)), + IpAddr::V6(_) => self + .cidr_rules_v6 + .iter() + .find(|rule| rule.cidr.contains(ip)), + } + .map(CidrPolicyMatch::Explicit) + .or_else(|| self.match_auto_cidr(ip)) + } + + pub(super) fn match_auto_cidr(&self, ip: IpAddr) -> Option> { + let rule = match ip { + IpAddr::V4(_) => self.cidr_auto_rules_v4.first()?, + IpAddr::V6(_) => self.cidr_auto_rules_v6.first()?, + }; + let key = auto_cidr_bucket_key(ip, rule.prefix_len)?; + Some(CidrPolicyMatch::Auto { + key, + limits: rule.limits, + }) + } +} + +impl ShardedRegistry { + pub(super) fn new(shards: usize) -> Self { + let shard_count = shards.max(1).next_power_of_two(); + let mut items = Vec::with_capacity(shard_count); + for _ in 0..shard_count { + items.push(DashMap::>::new()); + } + Self { + shards: items.into_boxed_slice(), + mask: shard_count.saturating_sub(1), + } + } + + pub(super) fn shard_index(&self, key: &str) -> usize { + let mut hasher = std::collections::hash_map::DefaultHasher::new(); + key.hash(&mut hasher); + (hasher.finish() as usize) & self.mask + } + + pub(super) fn get_or_insert_with(&self, key: &str, make: F, activate: A) -> Arc + where + F: FnOnce() -> T, + A: FnOnce(&Arc), + { + let shard = &self.shards[self.shard_index(key)]; + match shard.entry(key.to_string()) { + dashmap::mapref::entry::Entry::Occupied(entry) => { + activate(entry.get()); + Arc::clone(entry.get()) + } + dashmap::mapref::entry::Entry::Vacant(slot) => { + let value = Arc::new(make()); + activate(&value); + slot.insert(Arc::clone(&value)); + value + } + } + } + + pub(super) fn retain(&self, predicate: F) + where + F: Fn(&String, &Arc) -> bool + Copy, + { + for shard in &*self.shards { + shard.retain(|key, value| predicate(key, value)); + } + } + + pub(super) fn remove_if(&self, key: &str, predicate: F) -> bool + where + F: Fn(&Arc) -> bool, + { + let shard = &self.shards[self.shard_index(key)]; + shard.remove_if(key, |_, value| predicate(value)).is_some() + } +} diff --git a/src/proxy/traffic_limiter/tests.rs b/src/proxy/traffic_limiter/tests.rs new file mode 100644 index 0000000..b9147da --- /dev/null +++ b/src/proxy/traffic_limiter/tests.rs @@ -0,0 +1,76 @@ +use super::*; +use crate::config::CidrRateLimitKey; + +fn rate(up_bps: u64, down_bps: u64) -> RateLimitBps { + RateLimitBps { up_bps, down_bps } +} + +#[test] +fn explicit_cidr_rule_wins_over_auto_template() { + let limiter = TrafficLimiter::new(); + let mut cidr_limits = HashMap::new(); + cidr_limits.insert(CidrRateLimitKey::AutoV4(24), rate(1_000, 0)); + cidr_limits.insert( + CidrRateLimitKey::Network("203.0.113.7/32".parse().unwrap()), + rate(2_000, 0), + ); + + limiter.apply_policy(HashMap::new(), cidr_limits); + let policy = limiter.policy.load_full(); + let matched = policy.match_cidr("203.0.113.7".parse().unwrap()).unwrap(); + + match matched { + CidrPolicyMatch::Explicit(rule) => assert_eq!(rule.key.as_str(), "203.0.113.7/32"), + CidrPolicyMatch::Auto { .. } => panic!("explicit CIDR must have priority"), + } +} + +#[test] +fn auto_template_uses_longest_prefix() { + let limiter = TrafficLimiter::new(); + let mut cidr_limits = HashMap::new(); + cidr_limits.insert(CidrRateLimitKey::AutoV4(24), rate(1_000, 0)); + cidr_limits.insert(CidrRateLimitKey::AutoV4(32), rate(2_000, 0)); + + limiter.apply_policy(HashMap::new(), cidr_limits); + let policy = limiter.policy.load_full(); + let matched = policy.match_cidr("203.0.113.129".parse().unwrap()).unwrap(); + + match matched { + CidrPolicyMatch::Auto { key, limits } => { + assert_eq!(key, "auto:4:203.0.113.129/32"); + assert_eq!(limits.up_bps, 2_000); + } + CidrPolicyMatch::Explicit(_) => panic!("auto-template match expected"), + } +} + +#[test] +fn dual_auto_template_maps_v6_prefix_by_four() { + let limiter = TrafficLimiter::new(); + let mut cidr_limits = HashMap::new(); + cidr_limits.insert(CidrRateLimitKey::AutoDual(32), rate(1_000, 0)); + + limiter.apply_policy(HashMap::new(), cidr_limits); + let policy = limiter.policy.load_full(); + let matched = policy.match_cidr("2001:db8::1".parse().unwrap()).unwrap(); + + match matched { + CidrPolicyMatch::Auto { key, .. } => { + assert_eq!(key, "auto:6:2001:db8::1/128"); + } + CidrPolicyMatch::Explicit(_) => panic!("auto-template match expected"), + } +} + +#[test] +fn auto_cidr_bucket_key_canonicalizes_network_address() { + assert_eq!( + auto_cidr_bucket_key("203.0.113.129".parse().unwrap(), 24).unwrap(), + "auto:4:203.0.113.0/24" + ); + assert_eq!( + auto_cidr_bucket_key("2001:db8::abcd".parse().unwrap(), 64).unwrap(), + "auto:6:2001:db8::/64" + ); +} diff --git a/src/quota_state.rs b/src/quota_state.rs index 96727d7..bf3149b 100644 --- a/src/quota_state.rs +++ b/src/quota_state.rs @@ -83,10 +83,7 @@ impl QuotaStateOwner { } /// Persists a checkpoint filtered to the active configured user set. - pub(crate) async fn save( - &self, - configured_users: &BTreeSet, - ) -> std::io::Result<()> { + pub(crate) async fn save(&self, configured_users: &BTreeSet) -> std::io::Result<()> { let guard = Arc::clone(&self.mutation).lock_owned().await; let state = self.state_for_users(configured_users, None); let path = self.path.clone(); @@ -313,9 +310,8 @@ mod tests { let alice_owner = owner.clone(); let alice_users = configured.clone(); - let alice = tokio::spawn(async move { - alice_owner.reset_user(&alice_users, "alice").await - }); + let alice = + tokio::spawn(async move { alice_owner.reset_user(&alice_users, "alice").await }); let bob_owner = owner.clone(); let bob_users = configured.clone(); let bob = tokio::spawn(async move { bob_owner.reset_user(&bob_users, "bob").await }); diff --git a/src/stats/me_counters.rs b/src/stats/me_counters.rs index e4c588c..a2dd1f1 100644 --- a/src/stats/me_counters.rs +++ b/src/stats/me_counters.rs @@ -82,9 +82,7 @@ impl Stats { return; } - let mut slots = self - .me_handshake_error_code_slots - .load(Ordering::Acquire); + let mut slots = self.me_handshake_error_code_slots.load(Ordering::Acquire); loop { if slots >= ME_HANDSHAKE_ERROR_CODE_MAX { if let Some(entry) = self.me_handshake_error_codes.get(&code) { @@ -508,9 +506,7 @@ mod tests { (WORKERS * CODES_PER_WORKER) as u64 ); assert_eq!( - stats - .me_handshake_error_code_slots - .load(Ordering::Acquire), + stats.me_handshake_error_code_slots.load(Ordering::Acquire), counts.len() ); } diff --git a/src/tls_front/cache.rs b/src/tls_front/cache.rs index a5516e6..a177a91 100644 --- a/src/tls_front/cache.rs +++ b/src/tls_front/cache.rs @@ -1,22 +1,28 @@ -use std::collections::{HashMap, HashSet}; use std::collections::hash_map::RandomState; +use std::collections::{HashMap, HashSet}; use std::hash::{BuildHasher, Hash, Hasher}; use std::net::IpAddr; -use std::path::{Path, PathBuf}; +use std::path::PathBuf; use std::sync::Arc; use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering}; use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH}; use tokio::sync::RwLock; -use tokio::io::AsyncReadExt; -use tokio::time::sleep; -use tracing::{debug, info, warn}; use crate::tls_front::types::{ CachedTlsData, ParsedServerHello, TlsBehaviorProfile, TlsFetchResult, TlsProfileQuality, TlsProfileSource, }; +// Runtime TLS profile cache operations. +mod runtime; +// Bounded disk-cache reads and domain matching. +mod disk; + +use disk::*; + +#[cfg(test)] +mod tests; const FULL_CERT_SENT_SWEEP_INTERVAL_SECS: u64 = 1; const FULL_CERT_SENT_MAX_ENTRIES: usize = 65_536; const FULL_CERT_SENT_SHARDS: usize = 64; @@ -97,8 +103,7 @@ impl TlsFullCertBudget { } async fn sweep_one_shard(&self, now: Instant) { - let shard_index = self.sweep_cursor.fetch_add(1, Ordering::Relaxed) - % FULL_CERT_SENT_SHARDS; + let shard_index = self.sweep_cursor.fetch_add(1, Ordering::Relaxed) % FULL_CERT_SENT_SHARDS; let mut guard = self.shards[shard_index].write().await; let before = guard.len(); guard.retain(|_, entry| entry.expires_at.map_or(true, |expires_at| expires_at > now)); @@ -113,9 +118,7 @@ impl TlsFullCertBudget { let should_sweep = self .last_sweep_epoch_secs .fetch_update(Ordering::AcqRel, Ordering::Relaxed, |last_sweep| { - if now_epoch_secs.saturating_sub(last_sweep) - >= FULL_CERT_SENT_SWEEP_INTERVAL_SECS - { + if now_epoch_secs.saturating_sub(last_sweep) >= FULL_CERT_SENT_SWEEP_INTERVAL_SECS { Some(now_epoch_secs) } else { None @@ -149,9 +152,7 @@ impl TlsFullCertBudget { if guard.len() >= FULL_CERT_SENT_MAX_ENTRIES_PER_SHARD { let before = guard.len(); - guard.retain(|_, entry| { - entry.expires_at.map_or(true, |expires_at| expires_at > now) - }); + guard.retain(|_, entry| entry.expires_at.map_or(true, |expires_at| expires_at > now)); self.decrement_entries(before.saturating_sub(guard.len())); } if guard.len() >= FULL_CERT_SENT_MAX_ENTRIES_PER_SHARD || !self.try_reserve_entry() { @@ -228,735 +229,3 @@ fn key_share_group_label(group: Option) -> &'static str { None => "none", } } - -#[allow(dead_code)] -impl TlsFrontCache { - pub fn new(domains: &[String], default_len: usize, disk_path: impl AsRef) -> Self { - Self::new_with_full_cert_budget( - domains, - default_len, - disk_path, - Arc::new(TlsFullCertBudget::new()), - ) - } - - /// Creates a generation-local cache backed by the process-owned full-cert budget. - pub(crate) fn new_with_full_cert_budget( - domains: &[String], - default_len: usize, - disk_path: impl AsRef, - full_cert_budget: Arc, - ) -> Self { - let default_template = ParsedServerHello { - version: [0x03, 0x03], - random: [0u8; 32], - session_id: Vec::new(), - cipher_suite: [0x13, 0x01], - compression: 0, - extensions: Vec::new(), - }; - - let default = Arc::new(CachedTlsData { - server_hello_template: default_template, - cert_info: None, - cert_payload: None, - app_data_records_sizes: vec![default_len], - total_app_data_len: default_len, - behavior_profile: TlsBehaviorProfile::default(), - fetched_at: SystemTime::now(), - domain: "default".to_string(), - }); - - let mut map = HashMap::new(); - let mut full_cert_domain_keys = HashMap::new(); - let mut disk_entry_names = HashSet::new(); - for d in domains { - map.insert(d.clone(), default.clone()); - disk_entry_names.insert(format!("{}.json", d.replace(['/', '\\'], "_"))); - let canonical: Arc = Arc::from(normalize_dns_name(d)); - full_cert_domain_keys.insert(d.clone(), canonical.clone()); - full_cert_domain_keys - .entry(canonical.to_string()) - .or_insert(canonical); - } - - Self { - memory: RwLock::new(map), - default, - full_cert_budget, - full_cert_domain_keys, - disk_entry_names, - disk_path: disk_path.as_ref().to_path_buf(), - } - } - - pub async fn get(&self, sni: &str) -> Arc { - let guard = self.memory.read().await; - guard - .get(sni) - .cloned() - .unwrap_or_else(|| self.default.clone()) - } - - pub async fn contains_domain(&self, domain: &str) -> bool { - self.memory.read().await.contains_key(domain) - } - - pub(crate) async fn profile_health_snapshot( - &self, - domains: &[String], - max_domains: usize, - ) -> (Vec, usize) { - let guard = self.memory.read().await; - let now = SystemTime::now(); - let mut snapshot = Vec::with_capacity(domains.len().min(max_domains)); - let mut suppressed = 0usize; - - for domain in domains { - if snapshot.len() >= max_domains { - suppressed = suppressed.saturating_add(1); - continue; - } - - let cached = guard - .get(domain) - .cloned() - .unwrap_or_else(|| self.default.clone()); - let mut behavior = cached.behavior_profile.clone(); - behavior.refresh_server_hello_summary(&cached.server_hello_template); - let age_seconds = now - .duration_since(cached.fetched_at) - .map(|duration| duration.as_secs()) - .unwrap_or(0); - - snapshot.push(TlsFrontProfileHealth { - domain: domain.clone(), - source: profile_source_label(behavior.source), - quality: profile_quality_label(behavior.quality), - key_share_group: key_share_group_label(behavior.server_hello_key_share_group), - age_seconds, - is_default: cached.domain == "default", - has_cert_info: cached.cert_info.is_some(), - has_cert_payload: cached.cert_payload.is_some(), - server_hello_record_len: behavior.server_hello_record_len, - server_hello_extensions: behavior.server_hello_extension_types.len(), - app_data_records: cached - .app_data_records_sizes - .len() - .max(behavior.app_data_record_sizes.len()), - ticket_records: behavior.ticket_record_sizes.len(), - change_cipher_spec_count: behavior.change_cipher_spec_count, - total_app_data_len: cached.total_app_data_len, - }); - } - - (snapshot, suppressed) - } - - /// Returns configured domains that still resolve to the synthetic default profile. - pub(crate) async fn default_profile_domains(&self, domains: &[String]) -> Vec { - let guard = self.memory.read().await; - domains - .iter() - .filter(|domain| { - guard.get(domain.as_str()).unwrap_or(&self.default).domain == "default" - }) - .cloned() - .collect() - } - - fn full_cert_domain_key(&self, domain: &str) -> Arc { - self.full_cert_domain_keys - .get(domain) - .cloned() - .unwrap_or_else(|| Arc::from(normalize_dns_name(domain))) - } - - /// Returns true when the selected domain and client IP may receive a full cert payload. - pub async fn take_full_cert_budget_for_ip( - &self, - domain: &str, - client_ip: IpAddr, - ttl: Duration, - ) -> bool { - self.full_cert_budget - .take(self.full_cert_domain_key(domain), client_ip, ttl) - .await - } - - /// Returns the current process-owned full-cert budget entry count. - pub(crate) fn full_cert_budget_entries_for_metrics(&self) -> u64 { - self.full_cert_budget.entries_for_metrics() - } - - /// Returns the cumulative process-owned full-cert budget cap drops. - pub(crate) fn full_cert_budget_cap_drops_for_metrics(&self) -> u64 { - self.full_cert_budget.cap_drops_for_metrics() - } - - #[cfg(test)] - async fn insert_full_cert_sent_for_tests( - &self, - domain: &str, - client_ip: IpAddr, - expires_at: Instant, - ) { - let key = FullCertBudgetKey { - domain: self.full_cert_domain_key(domain), - client_ip, - }; - let shard_index = self.full_cert_budget.shard_index(&key); - let mut guard = self.full_cert_budget.shards[shard_index].write().await; - if guard - .insert( - key, - FullCertBudgetEntry { - expires_at: Some(expires_at), - }, - ) - .is_none() - { - self.full_cert_budget - .entries - .fetch_add(1, Ordering::Relaxed); - } - } - - #[cfg(test)] - async fn full_cert_sent_is_empty_for_tests(&self) -> bool { - for shard in &self.full_cert_budget.shards { - if !shard.read().await.is_empty() { - return false; - } - } - true - } - - #[cfg(test)] - async fn full_cert_sent_contains_for_tests(&self, domain: &str, client_ip: IpAddr) -> bool { - let key = FullCertBudgetKey { - domain: self.full_cert_domain_key(domain), - client_ip, - }; - let shard_index = self.full_cert_budget.shard_index(&key); - self.full_cert_budget.shards[shard_index] - .read() - .await - .contains_key(&key) - } - - pub async fn set(&self, domain: &str, data: CachedTlsData) { - let mut guard = self.memory.write().await; - guard.insert(domain.to_string(), Arc::new(data)); - } - - pub async fn load_from_disk(&self) { - let path = self.disk_path.clone(); - if tokio::fs::create_dir_all(&path).await.is_err() { - return; - } - let mut loaded = 0usize; - for name in &self.disk_entry_names { - let entry_path = path.join(name); - let Ok(metadata) = tokio::fs::symlink_metadata(&entry_path).await else { - continue; - }; - if !metadata.file_type().is_file() { - continue; - } - if let Ok(data) = read_disk_entry_bounded(&entry_path).await - && let Ok(mut cached) = serde_json::from_slice::(&data) - { - if cached.domain.is_empty() - || cached.domain.len() > 255 - || !cached - .domain - .chars() - .all(|c| c.is_ascii_alphanumeric() || c == '.' || c == '-') - { - warn!(file = %name, "Skipping TLS cache entry with invalid domain"); - continue; - } - if !self.full_cert_domain_keys.contains_key(&cached.domain) { - warn!( - file = %name, - domain = %cached.domain, - "Skipping TLS cache entry outside configured domains" - ); - continue; - } - if !cert_info_matches_domain(&cached) { - warn!( - file = %name, - domain = %cached.domain, - "Skipping TLS cache entry with mismatched certificate metadata" - ); - continue; - } - // fetched_at is skipped during deserialization; approximate with file mtime if available. - if let Ok(modified) = metadata.modified() { - cached.fetched_at = modified; - } - // Drop entries older than 72h - if let Ok(age) = cached.fetched_at.elapsed() - && age > Duration::from_secs(72 * 3600) - { - warn!(domain = %cached.domain, "Skipping stale TLS cache entry (>72h)"); - continue; - } - cached - .behavior_profile - .refresh_server_hello_summary(&cached.server_hello_template); - let domain = cached.domain.clone(); - self.set(&domain, cached).await; - loaded += 1; - } - } - if loaded > 0 { - info!(count = loaded, "Loaded TLS cache entries from disk"); - } - } - - async fn persist(&self, domain: &str, data: &CachedTlsData) { - if tokio::fs::create_dir_all(&self.disk_path).await.is_err() { - return; - } - let fname = format!("{}.json", domain.replace(['/', '\\'], "_")); - let path = self.disk_path.join(fname); - if let Ok(json) = serde_json::to_vec_pretty(data) { - if json.len() as u64 > TLS_FRONT_DISK_ENTRY_MAX_BYTES { - warn!( - domain, - bytes = json.len(), - "Skipping oversized TLS cache persistence" - ); - return; - } - // best-effort write - let _ = tokio::fs::write(path, json).await; - } - } - - /// Spawn background updater that periodically refreshes cached domains using provided fetcher. - pub fn spawn_updater(self: Arc, domains: Vec, interval: Duration, fetcher: F) - where - F: Fn(String) -> tokio::task::JoinHandle<()> + Send + Sync + 'static, - { - tokio::spawn(async move { - loop { - for domain in &domains { - let _ = fetcher(domain.clone()).await; - } - sleep(interval).await; - } - }); - } - - /// Replace cached entry from a fetch result. - pub async fn update_from_fetch(&self, domain: &str, fetched: TlsFetchResult) { - let TlsFetchResult { - server_hello_parsed, - app_data_records_sizes, - total_app_data_len, - mut behavior_profile, - cert_info, - cert_payload, - } = fetched; - behavior_profile.refresh_server_hello_summary(&server_hello_parsed); - let quality = behavior_profile.quality; - let data = CachedTlsData { - server_hello_template: server_hello_parsed, - cert_info, - cert_payload, - app_data_records_sizes: app_data_records_sizes.clone(), - total_app_data_len, - behavior_profile, - fetched_at: SystemTime::now(), - domain: domain.to_string(), - }; - - self.set(domain, data.clone()).await; - self.persist(domain, &data).await; - if quality == TlsProfileQuality::RawStrict { - debug!(domain = %domain, len = total_app_data_len, "TLS cache updated"); - } else { - warn!( - domain = %domain, - quality = profile_quality_label(quality), - len = total_app_data_len, - "TLS cache updated with non-strict front profile" - ); - } - } - - pub fn default_entry(&self) -> Arc { - self.default.clone() - } - - pub fn disk_path(&self) -> &Path { - &self.disk_path - } -} - -async fn read_disk_entry_bounded(path: &Path) -> std::io::Result> { - let file = tokio::fs::File::open(path).await?; - if file.metadata().await?.len() > TLS_FRONT_DISK_ENTRY_MAX_BYTES { - return Err(std::io::Error::new( - std::io::ErrorKind::InvalidData, - "TLS cache entry exceeds the 1 MiB limit", - )); - } - let mut bytes = Vec::new(); - file.take(TLS_FRONT_DISK_ENTRY_MAX_BYTES.saturating_add(1)) - .read_to_end(&mut bytes) - .await?; - if bytes.len() as u64 > TLS_FRONT_DISK_ENTRY_MAX_BYTES { - return Err(std::io::Error::new( - std::io::ErrorKind::InvalidData, - "TLS cache entry grew beyond the 1 MiB limit while reading", - )); - } - Ok(bytes) -} - -fn cert_info_matches_domain(cached: &CachedTlsData) -> bool { - let Some(cert_info) = cached.cert_info.as_ref() else { - return true; - }; - if !cert_info.san_names.is_empty() { - return cert_info - .san_names - .iter() - .any(|name| dns_name_matches_domain(name, &cached.domain)); - } - cert_info - .subject_cn - .as_deref() - .map_or(true, |name| dns_name_matches_domain(name, &cached.domain)) -} - -fn dns_name_matches_domain(pattern: &str, domain: &str) -> bool { - let pattern = normalize_dns_name(pattern); - let domain = normalize_dns_name(domain); - if pattern == domain { - return true; - } - - let Some(suffix) = pattern.strip_prefix("*.") else { - return false; - }; - let Some(prefix) = domain.strip_suffix(suffix) else { - return false; - }; - prefix.ends_with('.') && !prefix[..prefix.len() - 1].contains('.') -} - -fn normalize_dns_name(value: &str) -> String { - value.trim().trim_end_matches('.').to_ascii_lowercase() -} - -#[cfg(test)] -mod tests { - use super::*; - - fn cached_with_cert_info( - domain: &str, - subject_cn: Option<&str>, - san_names: Vec<&str>, - ) -> CachedTlsData { - CachedTlsData { - server_hello_template: ParsedServerHello { - version: [0x03, 0x03], - random: [0u8; 32], - session_id: Vec::new(), - cipher_suite: [0x13, 0x01], - compression: 0, - extensions: Vec::new(), - }, - cert_info: Some(crate::tls_front::types::ParsedCertificateInfo { - not_after_unix: None, - not_before_unix: None, - issuer_cn: None, - subject_cn: subject_cn.map(str::to_string), - san_names: san_names.into_iter().map(str::to_string).collect(), - }), - cert_payload: None, - app_data_records_sizes: vec![1024], - total_app_data_len: 1024, - behavior_profile: TlsBehaviorProfile::default(), - fetched_at: SystemTime::now(), - domain: domain.to_string(), - } - } - - #[test] - fn cert_info_domain_match_accepts_exact_san() { - let cached = cached_with_cert_info("b.com", Some("a.com"), vec!["b.com"]); - assert!(cert_info_matches_domain(&cached)); - } - - #[test] - fn cert_info_domain_match_rejects_wrong_san() { - let cached = cached_with_cert_info("b.com", Some("b.com"), vec!["a.com"]); - assert!(!cert_info_matches_domain(&cached)); - } - - #[test] - fn cert_info_domain_match_accepts_single_label_wildcard_san() { - let cached = cached_with_cert_info("api.b.com", None, vec!["*.b.com"]); - assert!(cert_info_matches_domain(&cached)); - } - - #[test] - fn cert_info_domain_match_rejects_multi_label_wildcard_san() { - let cached = cached_with_cert_info("deep.api.b.com", None, vec!["*.b.com"]); - assert!(!cert_info_matches_domain(&cached)); - } - - #[tokio::test] - async fn default_profile_domains_reports_only_unprepared_entries() { - let domains = vec!["ready.example".to_string(), "pending.example".to_string()]; - let cache = TlsFrontCache::new(&domains, 1024, "tlsfront-test-cache"); - cache - .set( - "ready.example", - cached_with_cert_info("ready.example", None, Vec::new()), - ) - .await; - - assert_eq!( - cache.default_profile_domains(&domains).await, - vec!["pending.example".to_string()] - ); - } - - #[tokio::test] - async fn test_take_full_cert_budget_for_ip_uses_ttl() { - let cache = TlsFrontCache::new(&["example.com".to_string()], 1024, "tlsfront-test-cache"); - let ip: IpAddr = "127.0.0.1".parse().expect("ip"); - let ttl = Duration::from_millis(80); - - assert!( - cache - .take_full_cert_budget_for_ip("example.com", ip, ttl) - .await - ); - assert!( - !cache - .take_full_cert_budget_for_ip("example.com", ip, ttl) - .await - ); - - tokio::time::sleep(Duration::from_millis(90)).await; - - assert!( - cache - .take_full_cert_budget_for_ip("example.com", ip, ttl) - .await - ); - } - - #[tokio::test] - async fn test_take_full_cert_budget_for_ip_zero_ttl_always_allows_full_payload() { - let cache = TlsFrontCache::new(&["example.com".to_string()], 1024, "tlsfront-test-cache"); - let ttl = Duration::ZERO; - - for idx in 0..100_000u32 { - let ip = IpAddr::V4(std::net::Ipv4Addr::new( - 10, - ((idx >> 16) & 0xff) as u8, - ((idx >> 8) & 0xff) as u8, - (idx & 0xff) as u8, - )); - assert!( - cache - .take_full_cert_budget_for_ip("example.com", ip, ttl) - .await - ); - } - - assert!(cache.full_cert_sent_is_empty_for_tests().await); - } - - #[tokio::test] - async fn test_take_full_cert_budget_for_ip_sweeps_expired_entries_when_due() { - let cache = TlsFrontCache::new(&["example.com".to_string()], 1024, "tlsfront-test-cache"); - let stale_ip: IpAddr = "127.0.0.1".parse().expect("ip"); - let new_ip: IpAddr = "127.0.0.2".parse().expect("ip"); - let ttl = Duration::from_secs(1); - let stale_expires_at = Instant::now() - .checked_sub(Duration::from_secs(1)) - .unwrap_or_else(Instant::now); - - cache - .insert_full_cert_sent_for_tests("example.com", stale_ip, stale_expires_at) - .await; - let stale_key = FullCertBudgetKey { - domain: cache.full_cert_domain_key("example.com"), - client_ip: stale_ip, - }; - cache.full_cert_budget.sweep_cursor.store( - cache.full_cert_budget.shard_index(&stale_key), - Ordering::Relaxed, - ); - cache - .full_cert_budget - .last_sweep_epoch_secs - .store(0, Ordering::Relaxed); - - assert!( - cache - .take_full_cert_budget_for_ip("example.com", new_ip, ttl) - .await - ); - - assert!( - !cache - .full_cert_sent_contains_for_tests("example.com", stale_ip) - .await - ); - assert!( - cache - .full_cert_sent_contains_for_tests("example.com", new_ip) - .await - ); - } - - #[tokio::test] - async fn test_take_full_cert_budget_for_ip_does_not_sweep_every_call() { - let cache = TlsFrontCache::new(&["example.com".to_string()], 1024, "tlsfront-test-cache"); - let stale_ip: IpAddr = "127.0.0.1".parse().expect("ip"); - let new_ip: IpAddr = "127.0.0.2".parse().expect("ip"); - let ttl = Duration::from_secs(1); - let stale_expires_at = Instant::now() - .checked_sub(Duration::from_secs(1)) - .unwrap_or_else(Instant::now); - let now_epoch_secs = SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap_or_default() - .as_secs(); - - cache - .insert_full_cert_sent_for_tests("example.com", stale_ip, stale_expires_at) - .await; - cache - .full_cert_budget - .last_sweep_epoch_secs - .store(now_epoch_secs, Ordering::Relaxed); - - assert!( - cache - .take_full_cert_budget_for_ip("example.com", new_ip, ttl) - .await - ); - - assert!( - cache - .full_cert_sent_contains_for_tests("example.com", stale_ip) - .await - ); - assert!( - cache - .full_cert_sent_contains_for_tests("example.com", new_ip) - .await - ); - } - - #[tokio::test] - async fn full_cert_budget_is_shared_across_cache_generations_and_scoped_by_domain() { - let budget = Arc::new(TlsFullCertBudget::new()); - let domains = ["one.example".to_string(), "two.example".to_string()]; - let first = TlsFrontCache::new_with_full_cert_budget( - &domains, - 1024, - "tlsfront-test-cache", - budget.clone(), - ); - let second = TlsFrontCache::new_with_full_cert_budget( - &domains, - 1024, - "tlsfront-test-cache", - budget, - ); - let ip: IpAddr = "127.0.0.1".parse().expect("ip"); - let ttl = Duration::from_secs(60); - - assert!( - first - .take_full_cert_budget_for_ip("one.example", ip, ttl) - .await - ); - assert!( - !second - .take_full_cert_budget_for_ip("one.example", ip, ttl) - .await - ); - assert!( - second - .take_full_cert_budget_for_ip("two.example", ip, ttl) - .await - ); - assert_eq!(second.full_cert_budget_entries_for_metrics(), 2); - } - - #[tokio::test] - async fn existing_full_cert_entry_keeps_its_own_expiry_after_ttl_change() { - let cache = TlsFrontCache::new(&["example.com".to_string()], 1024, "tlsfront-test-cache"); - let ip: IpAddr = "127.0.0.1".parse().expect("ip"); - - assert!( - cache - .take_full_cert_budget_for_ip("example.com", ip, Duration::from_millis(80)) - .await - ); - tokio::time::sleep(Duration::from_millis(20)).await; - assert!( - !cache - .take_full_cert_budget_for_ip("example.com", ip, Duration::from_millis(1)) - .await - ); - tokio::time::sleep(Duration::from_millis(70)).await; - assert!( - cache - .take_full_cert_budget_for_ip("example.com", ip, Duration::from_millis(1)) - .await - ); - } - - #[tokio::test] - async fn disk_reader_rejects_an_entry_above_the_hard_limit() { - let directory = tempfile::tempdir().unwrap(); - let path = directory.path().join("oversized.json"); - tokio::fs::write( - &path, - vec![0u8; TLS_FRONT_DISK_ENTRY_MAX_BYTES as usize + 1], - ) - .await - .unwrap(); - - let error = read_disk_entry_bounded(&path).await.unwrap_err(); - - assert_eq!(error.kind(), std::io::ErrorKind::InvalidData); - } - - #[cfg(unix)] - #[tokio::test] - async fn disk_loader_does_not_follow_a_configured_name_symlink() { - let directory = tempfile::tempdir().unwrap(); - let target = directory.path().join("outside.json"); - let cached = cached_with_cert_info("example.com", None, Vec::new()); - tokio::fs::write(&target, serde_json::to_vec(&cached).unwrap()) - .await - .unwrap(); - std::os::unix::fs::symlink(&target, directory.path().join("example.com.json")).unwrap(); - let cache = TlsFrontCache::new( - &["example.com".to_string()], - 1024, - directory.path(), - ); - - cache.load_from_disk().await; - - assert_eq!(cache.get("example.com").await.domain, "default"); - } -} diff --git a/src/tls_front/cache/disk.rs b/src/tls_front/cache/disk.rs new file mode 100644 index 0000000..3a73263 --- /dev/null +++ b/src/tls_front/cache/disk.rs @@ -0,0 +1,61 @@ +use std::path::Path; + +use tokio::io::AsyncReadExt; + +use super::*; +pub(super) async fn read_disk_entry_bounded(path: &Path) -> std::io::Result> { + let file = tokio::fs::File::open(path).await?; + if file.metadata().await?.len() > TLS_FRONT_DISK_ENTRY_MAX_BYTES { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + "TLS cache entry exceeds the 1 MiB limit", + )); + } + let mut bytes = Vec::new(); + file.take(TLS_FRONT_DISK_ENTRY_MAX_BYTES.saturating_add(1)) + .read_to_end(&mut bytes) + .await?; + if bytes.len() as u64 > TLS_FRONT_DISK_ENTRY_MAX_BYTES { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + "TLS cache entry grew beyond the 1 MiB limit while reading", + )); + } + Ok(bytes) +} + +pub(super) fn cert_info_matches_domain(cached: &CachedTlsData) -> bool { + let Some(cert_info) = cached.cert_info.as_ref() else { + return true; + }; + if !cert_info.san_names.is_empty() { + return cert_info + .san_names + .iter() + .any(|name| dns_name_matches_domain(name, &cached.domain)); + } + cert_info + .subject_cn + .as_deref() + .map_or(true, |name| dns_name_matches_domain(name, &cached.domain)) +} + +pub(super) fn dns_name_matches_domain(pattern: &str, domain: &str) -> bool { + let pattern = normalize_dns_name(pattern); + let domain = normalize_dns_name(domain); + if pattern == domain { + return true; + } + + let Some(suffix) = pattern.strip_prefix("*.") else { + return false; + }; + let Some(prefix) = domain.strip_suffix(suffix) else { + return false; + }; + prefix.ends_with('.') && !prefix[..prefix.len() - 1].contains('.') +} + +pub(super) fn normalize_dns_name(value: &str) -> String { + value.trim().trim_end_matches('.').to_ascii_lowercase() +} diff --git a/src/tls_front/cache/runtime.rs b/src/tls_front/cache/runtime.rs new file mode 100644 index 0000000..91132e4 --- /dev/null +++ b/src/tls_front/cache/runtime.rs @@ -0,0 +1,378 @@ +use std::path::Path; + +use tokio::time::sleep; +use tracing::{debug, info, warn}; + +use super::*; +#[allow(dead_code)] +impl TlsFrontCache { + pub fn new(domains: &[String], default_len: usize, disk_path: impl AsRef) -> Self { + Self::new_with_full_cert_budget( + domains, + default_len, + disk_path, + Arc::new(TlsFullCertBudget::new()), + ) + } + + /// Creates a generation-local cache backed by the process-owned full-cert budget. + pub(crate) fn new_with_full_cert_budget( + domains: &[String], + default_len: usize, + disk_path: impl AsRef, + full_cert_budget: Arc, + ) -> Self { + let default_template = ParsedServerHello { + version: [0x03, 0x03], + random: [0u8; 32], + session_id: Vec::new(), + cipher_suite: [0x13, 0x01], + compression: 0, + extensions: Vec::new(), + }; + + let default = Arc::new(CachedTlsData { + server_hello_template: default_template, + cert_info: None, + cert_payload: None, + app_data_records_sizes: vec![default_len], + total_app_data_len: default_len, + behavior_profile: TlsBehaviorProfile::default(), + fetched_at: SystemTime::now(), + domain: "default".to_string(), + }); + + let mut map = HashMap::new(); + let mut full_cert_domain_keys = HashMap::new(); + let mut disk_entry_names = HashSet::new(); + for d in domains { + map.insert(d.clone(), default.clone()); + disk_entry_names.insert(format!("{}.json", d.replace(['/', '\\'], "_"))); + let canonical: Arc = Arc::from(normalize_dns_name(d)); + full_cert_domain_keys.insert(d.clone(), canonical.clone()); + full_cert_domain_keys + .entry(canonical.to_string()) + .or_insert(canonical); + } + + Self { + memory: RwLock::new(map), + default, + full_cert_budget, + full_cert_domain_keys, + disk_entry_names, + disk_path: disk_path.as_ref().to_path_buf(), + } + } + + pub async fn get(&self, sni: &str) -> Arc { + let guard = self.memory.read().await; + guard + .get(sni) + .cloned() + .unwrap_or_else(|| self.default.clone()) + } + + pub async fn contains_domain(&self, domain: &str) -> bool { + self.memory.read().await.contains_key(domain) + } + + pub(crate) async fn profile_health_snapshot( + &self, + domains: &[String], + max_domains: usize, + ) -> (Vec, usize) { + let guard = self.memory.read().await; + let now = SystemTime::now(); + let mut snapshot = Vec::with_capacity(domains.len().min(max_domains)); + let mut suppressed = 0usize; + + for domain in domains { + if snapshot.len() >= max_domains { + suppressed = suppressed.saturating_add(1); + continue; + } + + let cached = guard + .get(domain) + .cloned() + .unwrap_or_else(|| self.default.clone()); + let mut behavior = cached.behavior_profile.clone(); + behavior.refresh_server_hello_summary(&cached.server_hello_template); + let age_seconds = now + .duration_since(cached.fetched_at) + .map(|duration| duration.as_secs()) + .unwrap_or(0); + + snapshot.push(TlsFrontProfileHealth { + domain: domain.clone(), + source: profile_source_label(behavior.source), + quality: profile_quality_label(behavior.quality), + key_share_group: key_share_group_label(behavior.server_hello_key_share_group), + age_seconds, + is_default: cached.domain == "default", + has_cert_info: cached.cert_info.is_some(), + has_cert_payload: cached.cert_payload.is_some(), + server_hello_record_len: behavior.server_hello_record_len, + server_hello_extensions: behavior.server_hello_extension_types.len(), + app_data_records: cached + .app_data_records_sizes + .len() + .max(behavior.app_data_record_sizes.len()), + ticket_records: behavior.ticket_record_sizes.len(), + change_cipher_spec_count: behavior.change_cipher_spec_count, + total_app_data_len: cached.total_app_data_len, + }); + } + + (snapshot, suppressed) + } + + /// Returns configured domains that still resolve to the synthetic default profile. + pub(crate) async fn default_profile_domains(&self, domains: &[String]) -> Vec { + let guard = self.memory.read().await; + domains + .iter() + .filter(|domain| { + guard.get(domain.as_str()).unwrap_or(&self.default).domain == "default" + }) + .cloned() + .collect() + } + + pub(super) fn full_cert_domain_key(&self, domain: &str) -> Arc { + self.full_cert_domain_keys + .get(domain) + .cloned() + .unwrap_or_else(|| Arc::from(normalize_dns_name(domain))) + } + + /// Returns true when the selected domain and client IP may receive a full cert payload. + pub async fn take_full_cert_budget_for_ip( + &self, + domain: &str, + client_ip: IpAddr, + ttl: Duration, + ) -> bool { + self.full_cert_budget + .take(self.full_cert_domain_key(domain), client_ip, ttl) + .await + } + + /// Returns the current process-owned full-cert budget entry count. + pub(crate) fn full_cert_budget_entries_for_metrics(&self) -> u64 { + self.full_cert_budget.entries_for_metrics() + } + + /// Returns the cumulative process-owned full-cert budget cap drops. + pub(crate) fn full_cert_budget_cap_drops_for_metrics(&self) -> u64 { + self.full_cert_budget.cap_drops_for_metrics() + } + + #[cfg(test)] + pub(super) async fn insert_full_cert_sent_for_tests( + &self, + domain: &str, + client_ip: IpAddr, + expires_at: Instant, + ) { + let key = FullCertBudgetKey { + domain: self.full_cert_domain_key(domain), + client_ip, + }; + let shard_index = self.full_cert_budget.shard_index(&key); + let mut guard = self.full_cert_budget.shards[shard_index].write().await; + if guard + .insert( + key, + FullCertBudgetEntry { + expires_at: Some(expires_at), + }, + ) + .is_none() + { + self.full_cert_budget + .entries + .fetch_add(1, Ordering::Relaxed); + } + } + + #[cfg(test)] + pub(super) async fn full_cert_sent_is_empty_for_tests(&self) -> bool { + for shard in &self.full_cert_budget.shards { + if !shard.read().await.is_empty() { + return false; + } + } + true + } + + #[cfg(test)] + pub(super) async fn full_cert_sent_contains_for_tests( + &self, + domain: &str, + client_ip: IpAddr, + ) -> bool { + let key = FullCertBudgetKey { + domain: self.full_cert_domain_key(domain), + client_ip, + }; + let shard_index = self.full_cert_budget.shard_index(&key); + self.full_cert_budget.shards[shard_index] + .read() + .await + .contains_key(&key) + } + + pub async fn set(&self, domain: &str, data: CachedTlsData) { + let mut guard = self.memory.write().await; + guard.insert(domain.to_string(), Arc::new(data)); + } + + pub async fn load_from_disk(&self) { + let path = self.disk_path.clone(); + if tokio::fs::create_dir_all(&path).await.is_err() { + return; + } + let mut loaded = 0usize; + for name in &self.disk_entry_names { + let entry_path = path.join(name); + let Ok(metadata) = tokio::fs::symlink_metadata(&entry_path).await else { + continue; + }; + if !metadata.file_type().is_file() { + continue; + } + if let Ok(data) = read_disk_entry_bounded(&entry_path).await + && let Ok(mut cached) = serde_json::from_slice::(&data) + { + if cached.domain.is_empty() + || cached.domain.len() > 255 + || !cached + .domain + .chars() + .all(|c| c.is_ascii_alphanumeric() || c == '.' || c == '-') + { + warn!(file = %name, "Skipping TLS cache entry with invalid domain"); + continue; + } + if !self.full_cert_domain_keys.contains_key(&cached.domain) { + warn!( + file = %name, + domain = %cached.domain, + "Skipping TLS cache entry outside configured domains" + ); + continue; + } + if !cert_info_matches_domain(&cached) { + warn!( + file = %name, + domain = %cached.domain, + "Skipping TLS cache entry with mismatched certificate metadata" + ); + continue; + } + // fetched_at is skipped during deserialization; approximate with file mtime if available. + if let Ok(modified) = metadata.modified() { + cached.fetched_at = modified; + } + // Drop entries older than 72h + if let Ok(age) = cached.fetched_at.elapsed() + && age > Duration::from_secs(72 * 3600) + { + warn!(domain = %cached.domain, "Skipping stale TLS cache entry (>72h)"); + continue; + } + cached + .behavior_profile + .refresh_server_hello_summary(&cached.server_hello_template); + let domain = cached.domain.clone(); + self.set(&domain, cached).await; + loaded += 1; + } + } + if loaded > 0 { + info!(count = loaded, "Loaded TLS cache entries from disk"); + } + } + + async fn persist(&self, domain: &str, data: &CachedTlsData) { + if tokio::fs::create_dir_all(&self.disk_path).await.is_err() { + return; + } + let fname = format!("{}.json", domain.replace(['/', '\\'], "_")); + let path = self.disk_path.join(fname); + if let Ok(json) = serde_json::to_vec_pretty(data) { + if json.len() as u64 > TLS_FRONT_DISK_ENTRY_MAX_BYTES { + warn!( + domain, + bytes = json.len(), + "Skipping oversized TLS cache persistence" + ); + return; + } + // best-effort write + let _ = tokio::fs::write(path, json).await; + } + } + + /// Spawn background updater that periodically refreshes cached domains using provided fetcher. + pub fn spawn_updater(self: Arc, domains: Vec, interval: Duration, fetcher: F) + where + F: Fn(String) -> tokio::task::JoinHandle<()> + Send + Sync + 'static, + { + tokio::spawn(async move { + loop { + for domain in &domains { + let _ = fetcher(domain.clone()).await; + } + sleep(interval).await; + } + }); + } + + /// Replace cached entry from a fetch result. + pub async fn update_from_fetch(&self, domain: &str, fetched: TlsFetchResult) { + let TlsFetchResult { + server_hello_parsed, + app_data_records_sizes, + total_app_data_len, + mut behavior_profile, + cert_info, + cert_payload, + } = fetched; + behavior_profile.refresh_server_hello_summary(&server_hello_parsed); + let quality = behavior_profile.quality; + let data = CachedTlsData { + server_hello_template: server_hello_parsed, + cert_info, + cert_payload, + app_data_records_sizes: app_data_records_sizes.clone(), + total_app_data_len, + behavior_profile, + fetched_at: SystemTime::now(), + domain: domain.to_string(), + }; + + self.set(domain, data.clone()).await; + self.persist(domain, &data).await; + if quality == TlsProfileQuality::RawStrict { + debug!(domain = %domain, len = total_app_data_len, "TLS cache updated"); + } else { + warn!( + domain = %domain, + quality = profile_quality_label(quality), + len = total_app_data_len, + "TLS cache updated with non-strict front profile" + ); + } + } + + pub fn default_entry(&self) -> Arc { + self.default.clone() + } + + pub fn disk_path(&self) -> &Path { + &self.disk_path + } +} diff --git a/src/tls_front/cache/tests.rs b/src/tls_front/cache/tests.rs new file mode 100644 index 0000000..9bebc64 --- /dev/null +++ b/src/tls_front/cache/tests.rs @@ -0,0 +1,294 @@ +use super::*; + +fn cached_with_cert_info( + domain: &str, + subject_cn: Option<&str>, + san_names: Vec<&str>, +) -> CachedTlsData { + CachedTlsData { + server_hello_template: ParsedServerHello { + version: [0x03, 0x03], + random: [0u8; 32], + session_id: Vec::new(), + cipher_suite: [0x13, 0x01], + compression: 0, + extensions: Vec::new(), + }, + cert_info: Some(crate::tls_front::types::ParsedCertificateInfo { + not_after_unix: None, + not_before_unix: None, + issuer_cn: None, + subject_cn: subject_cn.map(str::to_string), + san_names: san_names.into_iter().map(str::to_string).collect(), + }), + cert_payload: None, + app_data_records_sizes: vec![1024], + total_app_data_len: 1024, + behavior_profile: TlsBehaviorProfile::default(), + fetched_at: SystemTime::now(), + domain: domain.to_string(), + } +} + +#[test] +fn cert_info_domain_match_accepts_exact_san() { + let cached = cached_with_cert_info("b.com", Some("a.com"), vec!["b.com"]); + assert!(cert_info_matches_domain(&cached)); +} + +#[test] +fn cert_info_domain_match_rejects_wrong_san() { + let cached = cached_with_cert_info("b.com", Some("b.com"), vec!["a.com"]); + assert!(!cert_info_matches_domain(&cached)); +} + +#[test] +fn cert_info_domain_match_accepts_single_label_wildcard_san() { + let cached = cached_with_cert_info("api.b.com", None, vec!["*.b.com"]); + assert!(cert_info_matches_domain(&cached)); +} + +#[test] +fn cert_info_domain_match_rejects_multi_label_wildcard_san() { + let cached = cached_with_cert_info("deep.api.b.com", None, vec!["*.b.com"]); + assert!(!cert_info_matches_domain(&cached)); +} + +#[tokio::test] +async fn default_profile_domains_reports_only_unprepared_entries() { + let domains = vec!["ready.example".to_string(), "pending.example".to_string()]; + let cache = TlsFrontCache::new(&domains, 1024, "tlsfront-test-cache"); + cache + .set( + "ready.example", + cached_with_cert_info("ready.example", None, Vec::new()), + ) + .await; + + assert_eq!( + cache.default_profile_domains(&domains).await, + vec!["pending.example".to_string()] + ); +} + +#[tokio::test] +async fn test_take_full_cert_budget_for_ip_uses_ttl() { + let cache = TlsFrontCache::new(&["example.com".to_string()], 1024, "tlsfront-test-cache"); + let ip: IpAddr = "127.0.0.1".parse().expect("ip"); + let ttl = Duration::from_millis(80); + + assert!( + cache + .take_full_cert_budget_for_ip("example.com", ip, ttl) + .await + ); + assert!( + !cache + .take_full_cert_budget_for_ip("example.com", ip, ttl) + .await + ); + + tokio::time::sleep(Duration::from_millis(90)).await; + + assert!( + cache + .take_full_cert_budget_for_ip("example.com", ip, ttl) + .await + ); +} + +#[tokio::test] +async fn test_take_full_cert_budget_for_ip_zero_ttl_always_allows_full_payload() { + let cache = TlsFrontCache::new(&["example.com".to_string()], 1024, "tlsfront-test-cache"); + let ttl = Duration::ZERO; + + for idx in 0..100_000u32 { + let ip = IpAddr::V4(std::net::Ipv4Addr::new( + 10, + ((idx >> 16) & 0xff) as u8, + ((idx >> 8) & 0xff) as u8, + (idx & 0xff) as u8, + )); + assert!( + cache + .take_full_cert_budget_for_ip("example.com", ip, ttl) + .await + ); + } + + assert!(cache.full_cert_sent_is_empty_for_tests().await); +} + +#[tokio::test] +async fn test_take_full_cert_budget_for_ip_sweeps_expired_entries_when_due() { + let cache = TlsFrontCache::new(&["example.com".to_string()], 1024, "tlsfront-test-cache"); + let stale_ip: IpAddr = "127.0.0.1".parse().expect("ip"); + let new_ip: IpAddr = "127.0.0.2".parse().expect("ip"); + let ttl = Duration::from_secs(1); + let stale_expires_at = Instant::now() + .checked_sub(Duration::from_secs(1)) + .unwrap_or_else(Instant::now); + + cache + .insert_full_cert_sent_for_tests("example.com", stale_ip, stale_expires_at) + .await; + let stale_key = FullCertBudgetKey { + domain: cache.full_cert_domain_key("example.com"), + client_ip: stale_ip, + }; + cache.full_cert_budget.sweep_cursor.store( + cache.full_cert_budget.shard_index(&stale_key), + Ordering::Relaxed, + ); + cache + .full_cert_budget + .last_sweep_epoch_secs + .store(0, Ordering::Relaxed); + + assert!( + cache + .take_full_cert_budget_for_ip("example.com", new_ip, ttl) + .await + ); + + assert!( + !cache + .full_cert_sent_contains_for_tests("example.com", stale_ip) + .await + ); + assert!( + cache + .full_cert_sent_contains_for_tests("example.com", new_ip) + .await + ); +} + +#[tokio::test] +async fn test_take_full_cert_budget_for_ip_does_not_sweep_every_call() { + let cache = TlsFrontCache::new(&["example.com".to_string()], 1024, "tlsfront-test-cache"); + let stale_ip: IpAddr = "127.0.0.1".parse().expect("ip"); + let new_ip: IpAddr = "127.0.0.2".parse().expect("ip"); + let ttl = Duration::from_secs(1); + let stale_expires_at = Instant::now() + .checked_sub(Duration::from_secs(1)) + .unwrap_or_else(Instant::now); + let now_epoch_secs = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs(); + + cache + .insert_full_cert_sent_for_tests("example.com", stale_ip, stale_expires_at) + .await; + cache + .full_cert_budget + .last_sweep_epoch_secs + .store(now_epoch_secs, Ordering::Relaxed); + + assert!( + cache + .take_full_cert_budget_for_ip("example.com", new_ip, ttl) + .await + ); + + assert!( + cache + .full_cert_sent_contains_for_tests("example.com", stale_ip) + .await + ); + assert!( + cache + .full_cert_sent_contains_for_tests("example.com", new_ip) + .await + ); +} + +#[tokio::test] +async fn full_cert_budget_is_shared_across_cache_generations_and_scoped_by_domain() { + let budget = Arc::new(TlsFullCertBudget::new()); + let domains = ["one.example".to_string(), "two.example".to_string()]; + let first = TlsFrontCache::new_with_full_cert_budget( + &domains, + 1024, + "tlsfront-test-cache", + budget.clone(), + ); + let second = + TlsFrontCache::new_with_full_cert_budget(&domains, 1024, "tlsfront-test-cache", budget); + let ip: IpAddr = "127.0.0.1".parse().expect("ip"); + let ttl = Duration::from_secs(60); + + assert!( + first + .take_full_cert_budget_for_ip("one.example", ip, ttl) + .await + ); + assert!( + !second + .take_full_cert_budget_for_ip("one.example", ip, ttl) + .await + ); + assert!( + second + .take_full_cert_budget_for_ip("two.example", ip, ttl) + .await + ); + assert_eq!(second.full_cert_budget_entries_for_metrics(), 2); +} + +#[tokio::test] +async fn existing_full_cert_entry_keeps_its_own_expiry_after_ttl_change() { + let cache = TlsFrontCache::new(&["example.com".to_string()], 1024, "tlsfront-test-cache"); + let ip: IpAddr = "127.0.0.1".parse().expect("ip"); + + assert!( + cache + .take_full_cert_budget_for_ip("example.com", ip, Duration::from_millis(80)) + .await + ); + tokio::time::sleep(Duration::from_millis(20)).await; + assert!( + !cache + .take_full_cert_budget_for_ip("example.com", ip, Duration::from_millis(1)) + .await + ); + tokio::time::sleep(Duration::from_millis(70)).await; + assert!( + cache + .take_full_cert_budget_for_ip("example.com", ip, Duration::from_millis(1)) + .await + ); +} + +#[tokio::test] +async fn disk_reader_rejects_an_entry_above_the_hard_limit() { + let directory = tempfile::tempdir().unwrap(); + let path = directory.path().join("oversized.json"); + tokio::fs::write( + &path, + vec![0u8; TLS_FRONT_DISK_ENTRY_MAX_BYTES as usize + 1], + ) + .await + .unwrap(); + + let error = read_disk_entry_bounded(&path).await.unwrap_err(); + + assert_eq!(error.kind(), std::io::ErrorKind::InvalidData); +} + +#[cfg(unix)] +#[tokio::test] +async fn disk_loader_does_not_follow_a_configured_name_symlink() { + let directory = tempfile::tempdir().unwrap(); + let target = directory.path().join("outside.json"); + let cached = cached_with_cert_info("example.com", None, Vec::new()); + tokio::fs::write(&target, serde_json::to_vec(&cached).unwrap()) + .await + .unwrap(); + std::os::unix::fs::symlink(&target, directory.path().join("example.com.json")).unwrap(); + let cache = TlsFrontCache::new(&["example.com".to_string()], 1024, directory.path()); + + cache.load_from_disk().await; + + assert_eq!(cache.get("example.com").await.domain, "default"); +} diff --git a/src/tls_front/fetcher.rs b/src/tls_front/fetcher.rs index e6bf81c..e49ae86 100644 --- a/src/tls_front/fetcher.rs +++ b/src/tls_front/fetcher.rs @@ -345,1150 +345,25 @@ fn remember_profile_success_with_cap( ); } -fn build_client_config(alpn_protocols: &[&[u8]]) -> Arc { - let root = rustls::RootCertStore::empty(); - - let provider = rustls::crypto::ring::default_provider(); - let mut config = ClientConfig::builder_with_provider(Arc::new(provider)) - .with_protocol_versions(&[&rustls::version::TLS13, &rustls::version::TLS12]) - .expect("protocol versions") - .with_root_certificates(root) - .with_no_client_auth(); - - config - .dangerous() - .set_certificate_verifier(Arc::new(NoVerify)); - config.alpn_protocols = alpn_protocols.iter().map(|proto| proto.to_vec()).collect(); - - Arc::new(config) -} - -fn deterministic_bytes(seed: &str, len: usize) -> Vec { - let mut out = Vec::with_capacity(len); - let mut counter: u32 = 0; - while out.len() < len { - let mut chunk_seed = Vec::with_capacity(seed.len() + std::mem::size_of::()); - chunk_seed.extend_from_slice(seed.as_bytes()); - chunk_seed.extend_from_slice(&counter.to_le_bytes()); - out.extend_from_slice(&sha256(&chunk_seed)); - counter = counter.wrapping_add(1); - } - out.truncate(len); - out -} - -fn profile_cipher_suites(profile: TlsFetchProfile) -> &'static [u16] { - const MODERN_CHROME: &[u16] = &[ - 0x1301, 0x1302, 0x1303, 0xc02b, 0xc02c, 0xcca9, 0xc02f, 0xc030, 0xcca8, 0x009e, 0x00ff, - ]; - const MODERN_FIREFOX: &[u16] = &[ - 0x1301, 0x1303, 0x1302, 0xc02b, 0xcca9, 0xc02c, 0xc02f, 0xcca8, 0xc030, 0x009e, 0x00ff, - ]; - const COMPAT_TLS12: &[u16] = &[ - 0xc02b, 0xc02c, 0xc02f, 0xc030, 0xcca9, 0xcca8, 0x1301, 0x1302, 0x1303, 0x009e, 0x00ff, - ]; - const LEGACY_MINIMAL: &[u16] = &[0xc02b, 0xc02f, 0x1301, 0x1302, 0x00ff]; - - match profile { - TlsFetchProfile::ModernChromeLike => MODERN_CHROME, - TlsFetchProfile::ModernFirefoxLike => MODERN_FIREFOX, - TlsFetchProfile::CompatTls12 => COMPAT_TLS12, - TlsFetchProfile::LegacyMinimal => LEGACY_MINIMAL, - } -} - -fn profile_groups(profile: TlsFetchProfile) -> &'static [u16] { - const MODERN: &[u16] = &[ - TLS_NAMED_GROUP_X25519MLKEM768, - TLS_NAMED_GROUP_X25519, - 0x0017, - 0x0018, - ]; - const COMPAT: &[u16] = &[TLS_NAMED_GROUP_X25519, 0x0017]; - const LEGACY: &[u16] = &[0x0017]; - - match profile { - TlsFetchProfile::ModernChromeLike | TlsFetchProfile::ModernFirefoxLike => MODERN, - TlsFetchProfile::CompatTls12 => COMPAT, - TlsFetchProfile::LegacyMinimal => LEGACY, - } -} - -fn profile_sig_algs(profile: TlsFetchProfile) -> &'static [u16] { - const MODERN: &[u16] = &[0x0804, 0x0805, 0x0403, 0x0503, 0x0806]; - const COMPAT: &[u16] = &[0x0403, 0x0503, 0x0804, 0x0805]; - const LEGACY: &[u16] = &[0x0403, 0x0804]; - - match profile { - TlsFetchProfile::ModernChromeLike | TlsFetchProfile::ModernFirefoxLike => MODERN, - TlsFetchProfile::CompatTls12 => COMPAT, - TlsFetchProfile::LegacyMinimal => LEGACY, - } -} - -fn profile_alpn(profile: TlsFetchProfile) -> &'static [&'static [u8]] { - const H2_HTTP11: &[&[u8]] = &[b"h2", b"http/1.1"]; - const HTTP11: &[&[u8]] = &[b"http/1.1"]; - match profile { - TlsFetchProfile::ModernChromeLike | TlsFetchProfile::ModernFirefoxLike => H2_HTTP11, - TlsFetchProfile::CompatTls12 | TlsFetchProfile::LegacyMinimal => HTTP11, - } -} - -fn profile_alpn_labels(profile: TlsFetchProfile) -> &'static [&'static str] { - const H2_HTTP11: &[&str] = &["h2", "http/1.1"]; - const HTTP11: &[&str] = &["http/1.1"]; - match profile { - TlsFetchProfile::ModernChromeLike | TlsFetchProfile::ModernFirefoxLike => H2_HTTP11, - TlsFetchProfile::CompatTls12 | TlsFetchProfile::LegacyMinimal => HTTP11, - } -} - -fn profile_session_id_len(profile: TlsFetchProfile) -> usize { - match profile { - TlsFetchProfile::ModernChromeLike | TlsFetchProfile::ModernFirefoxLike => 32, - TlsFetchProfile::CompatTls12 | TlsFetchProfile::LegacyMinimal => 0, - } -} - -fn profile_supported_versions(profile: TlsFetchProfile) -> &'static [u16] { - const MODERN: &[u16] = &[0x0304, 0x0303]; - const COMPAT: &[u16] = &[0x0303, 0x0304]; - const LEGACY: &[u16] = &[0x0303]; - match profile { - TlsFetchProfile::ModernChromeLike | TlsFetchProfile::ModernFirefoxLike => MODERN, - TlsFetchProfile::CompatTls12 => COMPAT, - TlsFetchProfile::LegacyMinimal => LEGACY, - } -} - -fn profile_padding_target(profile: TlsFetchProfile) -> usize { - match profile { - // X25519MLKEM768 makes the Chrome-like ClientHello much larger than - // legacy pre-hybrid profiles; keep enough headroom for padding. - TlsFetchProfile::ModernChromeLike => 1450, - TlsFetchProfile::ModernFirefoxLike => 200, - TlsFetchProfile::CompatTls12 => 180, - TlsFetchProfile::LegacyMinimal => 64, - } -} - -fn grease_value(rng: &SecureRandom, deterministic: bool, seed: &str) -> u16 { - const GREASE_VALUES: [u16; 16] = [ - 0x0a0a, 0x1a1a, 0x2a2a, 0x3a3a, 0x4a4a, 0x5a5a, 0x6a6a, 0x7a7a, 0x8a8a, 0x9a9a, 0xaaaa, - 0xbaba, 0xcaca, 0xdada, 0xeaea, 0xfafa, - ]; - if deterministic { - let idx = deterministic_bytes(seed, 1)[0] as usize % GREASE_VALUES.len(); - GREASE_VALUES[idx] - } else { - let idx = (rng.bytes(1)[0] as usize) % GREASE_VALUES.len(); - GREASE_VALUES[idx] - } -} - -fn gen_mlkem768_client_encapsulation_key( - rng: &SecureRandom, - deterministic: bool, - seed: &str, -) -> Option> { - let seed_bytes = if deterministic { - deterministic_bytes(seed, 64) - } else { - rng.bytes(64) - }; - let seed = MlKemSeed::try_from(seed_bytes.as_slice()).ok()?; - let decapsulation_key = MlKemDecapsulationKey::::from_seed(seed); - let encapsulation_key = decapsulation_key.encapsulation_key().to_bytes(); - let bytes = encapsulation_key.as_slice(); - if bytes.len() == MLKEM768_CLIENT_ENCAPSULATION_KEY_LEN { - Some(bytes.to_vec()) - } else { - None - } -} - -fn gen_x25519mlkem768_client_key_share( - rng: &SecureRandom, - deterministic: bool, - seed: &str, -) -> Option> { - let mlkem_key = - gen_mlkem768_client_encapsulation_key(rng, deterministic, &format!("{seed}:mlkem768"))?; - let x25519_key = gen_key_share(rng, deterministic, &format!("{seed}:x25519")); - let mut key_share = - Vec::with_capacity(MLKEM768_CLIENT_ENCAPSULATION_KEY_LEN + x25519_key.len()); - key_share.extend_from_slice(&mlkem_key); - key_share.extend_from_slice(&x25519_key); - Some(key_share) -} - -fn push_client_key_share_entry(keyshare: &mut Vec, group: u16, key: &[u8]) { - keyshare.extend_from_slice(&group.to_be_bytes()); - keyshare.extend_from_slice(&(key.len() as u16).to_be_bytes()); - keyshare.extend_from_slice(key); -} - -fn build_client_hello( - sni: &str, - rng: &SecureRandom, - profile: TlsFetchProfile, - grease_enabled: bool, - deterministic: bool, -) -> Vec { - // === ClientHello body === - let mut body = Vec::new(); - - // Legacy version (TLS 1.0) as in real ClientHello headers - body.extend_from_slice(&[0x03, 0x03]); - - // Random - if deterministic { - body.extend_from_slice(&deterministic_bytes(&format!("tls-fetch-random:{sni}"), 32)); - } else { - body.extend_from_slice(&rng.bytes(32)); - } - - // Use non-empty Session ID for modern TLS 1.3-like profiles to reduce middlebox friction. - let session_id_len = profile_session_id_len(profile); - let session_id = if session_id_len == 0 { - Vec::new() - } else if deterministic { - deterministic_bytes( - &format!("tls-fetch-session:{sni}:{}", profile.as_str()), - session_id_len, - ) - } else { - rng.bytes(session_id_len) - }; - body.push(session_id.len() as u8); - body.extend_from_slice(&session_id); - - let mut cipher_suites = profile_cipher_suites(profile).to_vec(); - if grease_enabled { - let grease = grease_value(rng, deterministic, &format!("cipher:{sni}")); - cipher_suites.insert(0, grease); - } - body.extend_from_slice(&((cipher_suites.len() * 2) as u16).to_be_bytes()); - for suite in cipher_suites { - body.extend_from_slice(&suite.to_be_bytes()); - } - - // Compression methods: null only - body.push(1); - body.push(0); - - // === Extensions === - let mut exts = Vec::new(); - - let mut push_extension = |ext_type: u16, data: &[u8]| { - exts.extend_from_slice(&ext_type.to_be_bytes()); - exts.extend_from_slice(&(data.len() as u16).to_be_bytes()); - exts.extend_from_slice(data); - }; - - // server_name (SNI) - let sni_bytes = sni.as_bytes(); - let mut sni_ext = Vec::with_capacity(5 + sni_bytes.len()); - sni_ext.extend_from_slice(&(sni_bytes.len() as u16 + 3).to_be_bytes()); - sni_ext.push(0); - sni_ext.extend_from_slice(&(sni_bytes.len() as u16).to_be_bytes()); - sni_ext.extend_from_slice(sni_bytes); - push_extension(0x0000, &sni_ext); - - // Chrome-like profile keeps browser-like ordering and extension set. - if matches!(profile, TlsFetchProfile::ModernChromeLike) { - // ec_point_formats: uncompressed only. - push_extension(0x000b, &[0x01, 0x00]); - } - - // supported_groups - let mut groups = profile_groups(profile).to_vec(); - if grease_enabled { - let grease = grease_value(rng, deterministic, &format!("group:{sni}")); - groups.insert(0, grease); - } - let mut groups_ext = Vec::with_capacity(2 + groups.len() * 2); - groups_ext.extend_from_slice(&(groups.len() as u16 * 2).to_be_bytes()); - for g in groups { - groups_ext.extend_from_slice(&g.to_be_bytes()); - } - push_extension(0x000a, &groups_ext); - - if matches!(profile, TlsFetchProfile::ModernChromeLike) { - // session_ticket - push_extension(0x0023, &[]); - } - - // signature_algorithms - let mut sig_algs = profile_sig_algs(profile).to_vec(); - if grease_enabled { - let grease = grease_value(rng, deterministic, &format!("sigalg:{sni}")); - sig_algs.insert(0, grease); - } - let mut sig_algs_ext = Vec::with_capacity(2 + sig_algs.len() * 2); - sig_algs_ext.extend_from_slice(&(sig_algs.len() as u16 * 2).to_be_bytes()); - for a in sig_algs { - sig_algs_ext.extend_from_slice(&a.to_be_bytes()); - } - push_extension(0x000d, &sig_algs_ext); - - // supported_versions - let mut versions = profile_supported_versions(profile).to_vec(); - if grease_enabled { - let grease = grease_value(rng, deterministic, &format!("version:{sni}")); - versions.insert(0, grease); - } - let mut versions_ext = Vec::with_capacity(1 + versions.len() * 2); - versions_ext.push((versions.len() * 2) as u8); - for v in versions { - versions_ext.extend_from_slice(&v.to_be_bytes()); - } - push_extension(0x002b, &versions_ext); - - if matches!(profile, TlsFetchProfile::ModernChromeLike) { - // psk_key_exchange_modes: psk_dhe_ke - push_extension(0x002d, &[0x01, 0x01]); - } - - // key_share - let key_share_seed = format!("tls-fetch-keyshare:{sni}:{}", profile.as_str()); - let mut keyshare = Vec::new(); - if matches!( - profile, - TlsFetchProfile::ModernChromeLike | TlsFetchProfile::ModernFirefoxLike - ) { - if let Some(key) = gen_x25519mlkem768_client_key_share(rng, deterministic, &key_share_seed) - { - push_client_key_share_entry(&mut keyshare, TLS_NAMED_GROUP_X25519MLKEM768, &key); - } - } - let key = gen_key_share(rng, deterministic, &key_share_seed); - push_client_key_share_entry(&mut keyshare, TLS_NAMED_GROUP_X25519, &key); - let mut keyshare_ext = Vec::with_capacity(2 + keyshare.len()); - keyshare_ext.extend_from_slice(&(keyshare.len() as u16).to_be_bytes()); - keyshare_ext.extend_from_slice(&keyshare); - push_extension(0x0033, &keyshare_ext); - - // ALPN - let mut alpn_list = Vec::new(); - for proto in profile_alpn(profile) { - alpn_list.push(proto.len() as u8); - alpn_list.extend_from_slice(proto); - } - if !alpn_list.is_empty() { - let mut alpn_ext = Vec::with_capacity(2 + alpn_list.len()); - alpn_ext.extend_from_slice(&(alpn_list.len() as u16).to_be_bytes()); - alpn_ext.extend_from_slice(&alpn_list); - push_extension(0x0010, &alpn_ext); - } - - if grease_enabled { - let grease = grease_value(rng, deterministic, &format!("ext:{sni}")); - push_extension(grease, &[]); - } - - // padding to reduce recognizability and keep length ~500 bytes - let target_ext_len = profile_padding_target(profile); - if exts.len() < target_ext_len { - let remaining = target_ext_len - exts.len(); - if remaining > 4 { - let pad_len = remaining - 4; // minus type+len - exts.extend_from_slice(&0x0015u16.to_be_bytes()); // padding extension - exts.extend_from_slice(&(pad_len as u16).to_be_bytes()); - exts.resize(exts.len() + pad_len, 0); - } - } - - // Extensions length prefix - body.extend_from_slice(&(exts.len() as u16).to_be_bytes()); - body.extend_from_slice(&exts); - - // === Handshake wrapper === - let mut handshake = Vec::new(); - handshake.push(0x01); // ClientHello - let len_bytes = (body.len() as u32).to_be_bytes(); - handshake.extend_from_slice(&len_bytes[1..4]); - handshake.extend_from_slice(&body); - - // === Record === - let mut record = Vec::new(); - record.push(TLS_RECORD_HANDSHAKE); - record.extend_from_slice(&[0x03, 0x01]); // legacy record version - record.extend_from_slice(&(handshake.len() as u16).to_be_bytes()); - record.extend_from_slice(&handshake); - - record -} - -fn gen_key_share(rng: &SecureRandom, deterministic: bool, seed: &str) -> [u8; 32] { - let mut scalar = [0u8; 32]; - if deterministic { - scalar.copy_from_slice(&deterministic_bytes(seed, 32)); - } else { - scalar.copy_from_slice(&rng.bytes(32)); - } - x25519(scalar, X25519_BASEPOINT_BYTES) -} - -async fn read_tls_record(stream: &mut S) -> Result<(u8, Vec)> -where - S: AsyncRead + Unpin, -{ - let mut header = [0u8; 5]; - stream.read_exact(&mut header).await?; - let len = u16::from_be_bytes([header[3], header[4]]) as usize; - let mut body = vec![0u8; len]; - stream.read_exact(&mut body).await?; - Ok((header[0], body)) -} - -fn parse_server_hello(body: &[u8]) -> Option { - if body.len() < 4 || body[0] != 0x02 { - return None; - } - - let msg_len = u32::from_be_bytes([0, body[1], body[2], body[3]]) as usize; - if msg_len + 4 > body.len() { - return None; - } - - let mut pos = 4; - let version = [*body.get(pos)?, *body.get(pos + 1)?]; - pos += 2; - - let mut random = [0u8; 32]; - random.copy_from_slice(body.get(pos..pos + 32)?); - pos += 32; - - let session_len = *body.get(pos)? as usize; - pos += 1; - let session_id = body.get(pos..pos + session_len)?.to_vec(); - pos += session_len; - - let cipher_suite = [*body.get(pos)?, *body.get(pos + 1)?]; - pos += 2; - - let compression = *body.get(pos)?; - pos += 1; - - let ext_len = u16::from_be_bytes([*body.get(pos)?, *body.get(pos + 1)?]) as usize; - pos += 2; - let ext_end = pos.checked_add(ext_len)?; - if ext_end > body.len() { - return None; - } - - let mut extensions = Vec::new(); - while pos + 4 <= ext_end { - let etype = u16::from_be_bytes([body[pos], body[pos + 1]]); - let elen = u16::from_be_bytes([body[pos + 2], body[pos + 3]]) as usize; - pos += 4; - let data = body.get(pos..pos + elen)?.to_vec(); - pos += elen; - extensions.push(TlsExtension { - ext_type: etype, - data, - }); - } - - Some(ParsedServerHello { - version, - random, - session_id, - cipher_suite, - compression, - extensions, - }) -} - -fn derive_behavior_profile(records: &[(u8, Vec)]) -> TlsBehaviorProfile { - let mut change_cipher_spec_count = 0u8; - let mut app_data_record_sizes = Vec::new(); - - for (record_type, body) in records { - match *record_type { - TLS_RECORD_CHANGE_CIPHER => { - change_cipher_spec_count = change_cipher_spec_count.saturating_add(1); - } - TLS_RECORD_APPLICATION => { - app_data_record_sizes.push(body.len()); - } - _ => {} - } - } - - let mut ticket_record_sizes = Vec::new(); - while app_data_record_sizes - .last() - .is_some_and(|size| *size <= 256 && ticket_record_sizes.len() < 2) - { - if let Some(size) = app_data_record_sizes.pop() { - ticket_record_sizes.push(size); - } - } - ticket_record_sizes.reverse(); - - TlsBehaviorProfile { - change_cipher_spec_count: change_cipher_spec_count.max(1), - app_data_record_sizes, - ticket_record_sizes, - source: TlsProfileSource::Raw, - ..TlsBehaviorProfile::default() - } -} - -fn parse_cert_info(certs: &[CertificateDer<'static>]) -> Option { - let first = certs.first()?; - let (_rem, cert) = X509Certificate::from_der(first.as_ref()).ok()?; - - let not_before = Some(cert.validity().not_before.to_datetime().unix_timestamp()); - let not_after = Some(cert.validity().not_after.to_datetime().unix_timestamp()); - - let issuer_cn = cert - .issuer() - .iter_common_name() - .next() - .and_then(|cn| cn.as_str().ok()) - .map(|s| s.to_string()); - - let subject_cn = cert - .subject() - .iter_common_name() - .next() - .and_then(|cn| cn.as_str().ok()) - .map(|s| s.to_string()); - - let san_names = cert - .subject_alternative_name() - .ok() - .flatten() - .map(|san| { - san.value - .general_names - .iter() - .filter_map(|gn| match gn { - x509_parser::extensions::GeneralName::DNSName(n) => Some(n.to_string()), - _ => None, - }) - .collect::>() - }) - .unwrap_or_default(); - - Some(ParsedCertificateInfo { - not_after_unix: not_after, - not_before_unix: not_before, - issuer_cn, - subject_cn, - san_names, - }) -} - -fn u24_bytes(value: usize) -> Option<[u8; 3]> { - if value > 0x00ff_ffff { - return None; - } - Some([ - ((value >> 16) & 0xff) as u8, - ((value >> 8) & 0xff) as u8, - (value & 0xff) as u8, - ]) -} - -async fn connect_with_dns_override( - host: &str, - port: u16, - connect_timeout: Duration, -) -> Result { - Ok(timeout(connect_timeout, TcpStream::connect((host, port))).await??) -} - -async fn connect_tcp_with_upstream( - host: &str, - port: u16, - connect_timeout: Duration, - upstream: Option>, - scope: Option<&str>, - strict_route: bool, -) -> Result { - if let Some(manager) = upstream { - let resolved = match manager.resolve_hostname(host, port).await { - Ok(addr) => Some(addr), - Err(e) => { - if strict_route { - return Err(anyhow!( - "upstream route DNS resolution failed for {host}:{port}: {e}" - )); - } - warn!( - host = %host, - port = port, - scope = ?scope, - error = %e, - "Upstream DNS resolution failed, using direct connect" - ); - None - } - }; - - if let Some(addr) = resolved { - match manager.connect(addr, None, scope).await { - Ok(stream) => return Ok(stream), - Err(e) => { - if strict_route { - return Err(anyhow!( - "upstream route connect failed for {host}:{port}: {e}" - )); - } - warn!( - host = %host, - port = port, - scope = ?scope, - error = %e, - "Upstream connect failed, using direct connect" - ); - return Ok(UpstreamStream::Tcp( - timeout(connect_timeout, TcpStream::connect(addr)).await??, - )); - } - } - } else if strict_route { - return Err(anyhow!( - "upstream route resolution produced no usable address for {host}:{port}" - )); - } - } - Ok(UpstreamStream::Tcp( - connect_with_dns_override(host, port, connect_timeout).await?, - )) -} - -fn socket_addrs_from_upstream_stream( - stream: &UpstreamStream, -) -> (Option, Option) { - match stream { - UpstreamStream::Tcp(tcp) => (tcp.local_addr().ok(), tcp.peer_addr().ok()), - UpstreamStream::Shadowsocks(_) => (None, None), - } -} - -fn build_tls_fetch_proxy_header( - proxy_protocol: u8, - src_addr: Option, - dst_addr: Option, -) -> Option> { - match proxy_protocol { - 0 => None, - 2 => { - let header = match (src_addr, dst_addr) { - (Some(src @ SocketAddr::V4(_)), Some(dst @ SocketAddr::V4(_))) - | (Some(src @ SocketAddr::V6(_)), Some(dst @ SocketAddr::V6(_))) => { - ProxyProtocolV2Builder::new().with_addrs(src, dst).build() - } - _ => ProxyProtocolV2Builder::new().build(), - }; - Some(header) - } - _ => { - let header = match (src_addr, dst_addr) { - (Some(SocketAddr::V4(src)), Some(SocketAddr::V4(dst))) => { - ProxyProtocolV1Builder::new() - .tcp4(src.into(), dst.into()) - .build() - } - (Some(SocketAddr::V6(src)), Some(SocketAddr::V6(dst))) => { - ProxyProtocolV1Builder::new() - .tcp6(src.into(), dst.into()) - .build() - } - _ => ProxyProtocolV1Builder::new().build(), - }; - Some(header) - } - } -} - -fn encode_tls13_certificate_message(cert_chain_der: &[Vec]) -> Option> { - if cert_chain_der.is_empty() { - return None; - } - - let mut certificate_list = Vec::new(); - for cert in cert_chain_der { - if cert.is_empty() { - return None; - } - certificate_list.extend_from_slice(&u24_bytes(cert.len())?); - certificate_list.extend_from_slice(cert); - certificate_list.extend_from_slice(&0u16.to_be_bytes()); // cert_entry extensions - } - - // Certificate = context_len(1) + certificate_list_len(3) + entries - let body_len = 1usize.checked_add(3)?.checked_add(certificate_list.len())?; - - let mut message = Vec::with_capacity(4 + body_len); - message.push(0x0b); // HandshakeType::certificate - message.extend_from_slice(&u24_bytes(body_len)?); - message.push(0x00); // certificate_request_context length - message.extend_from_slice(&u24_bytes(certificate_list.len())?); - message.extend_from_slice(&certificate_list); - Some(message) -} - -async fn fetch_via_raw_tls_stream( - mut stream: S, - sni: &str, - connect_timeout: Duration, - proxy_header: Option>, - profile: TlsFetchProfile, - grease_enabled: bool, - deterministic: bool, -) -> Result -where - S: AsyncRead + AsyncWrite + Unpin, -{ - let rng = SecureRandom::new(); - let client_hello = build_client_hello(sni, &rng, profile, grease_enabled, deterministic); - timeout(connect_timeout, async { - if let Some(header) = proxy_header.as_ref() { - stream.write_all(&header).await?; - } - stream.write_all(&client_hello).await?; - stream.flush().await?; - Ok::<(), std::io::Error>(()) - }) - .await??; - - let mut records = Vec::new(); - let mut app_records_seen = 0usize; - // Read a bounded encrypted flight: ServerHello, CCS, certificate-like data, - // and a small number of ticket-like tail records. - for _ in 0..8 { - match timeout(connect_timeout, read_tls_record(&mut stream)).await { - Ok(Ok(rec)) => { - if rec.0 == TLS_RECORD_APPLICATION { - app_records_seen += 1; - } - records.push(rec); - } - Ok(Err(e)) => return Err(e), - Err(_) => break, - } - if app_records_seen >= 4 { - break; - } - } - - let mut server_hello = None; - let mut server_hello_record_len = 0usize; - for (t, body) in &records { - if *t == TLS_RECORD_HANDSHAKE && server_hello.is_none() { - server_hello = parse_server_hello(body); - server_hello_record_len = body.len(); - } - } - - let parsed = server_hello.ok_or_else(|| anyhow!("ServerHello not received"))?; - let mut behavior_profile = derive_behavior_profile(&records); - behavior_profile.server_hello_record_len = server_hello_record_len; - behavior_profile.refresh_server_hello_summary(&parsed); - let mut app_sizes = behavior_profile.app_data_record_sizes.clone(); - app_sizes.extend_from_slice(&behavior_profile.ticket_record_sizes); - let total_app_data_len = app_sizes.iter().sum::().max(1024); - let app_data_records_sizes = if app_sizes.is_empty() { - vec![total_app_data_len] - } else { - app_sizes - }; - - Ok(TlsFetchResult { - server_hello_parsed: parsed, - app_data_records_sizes, - total_app_data_len, - behavior_profile, - cert_info: None, - cert_payload: None, - }) -} - -async fn fetch_via_raw_tls( - host: &str, - port: u16, - sni: &str, - connect_timeout: Duration, - upstream: Option>, - scope: Option<&str>, - proxy_protocol: u8, - unix_sock: Option<&str>, - strict_route: bool, - profile: TlsFetchProfile, - grease_enabled: bool, - deterministic: bool, -) -> Result { - #[cfg(unix)] - if let Some(sock_path) = unix_sock { - match timeout(connect_timeout, UnixStream::connect(sock_path)).await { - Ok(Ok(stream)) => { - debug!( - sni = %sni, - sock = %sock_path, - "Raw TLS fetch using mask unix socket" - ); - let proxy_header = build_tls_fetch_proxy_header(proxy_protocol, None, None); - return fetch_via_raw_tls_stream( - stream, - sni, - connect_timeout, - proxy_header, - profile, - grease_enabled, - deterministic, - ) - .await; - } - Ok(Err(e)) => { - warn!( - sni = %sni, - sock = %sock_path, - error = %e, - "Raw TLS unix socket connect failed, falling back to TCP" - ); - } - Err(_) => { - warn!( - sni = %sni, - sock = %sock_path, - "Raw TLS unix socket connect timed out, falling back to TCP" - ); - } - } - } - - #[cfg(not(unix))] - let _ = unix_sock; - - let stream = - connect_tcp_with_upstream(host, port, connect_timeout, upstream, scope, strict_route) - .await?; - let (src_addr, dst_addr) = socket_addrs_from_upstream_stream(&stream); - let proxy_header = build_tls_fetch_proxy_header(proxy_protocol, src_addr, dst_addr); - fetch_via_raw_tls_stream( - stream, - sni, - connect_timeout, - proxy_header, - profile, - grease_enabled, - deterministic, - ) - .await -} - -async fn fetch_via_rustls_stream( - mut stream: S, - host: &str, - sni: &str, - proxy_header: Option>, - alpn_protocols: &[&[u8]], -) -> Result -where - S: AsyncRead + AsyncWrite + Unpin, -{ - // rustls handshake path for certificate and basic negotiated metadata. - if let Some(header) = proxy_header.as_ref() { - stream.write_all(&header).await?; - stream.flush().await?; - } - - let config = build_client_config(alpn_protocols); - let connector = TlsConnector::from(config); - - let server_name = ServerName::try_from(sni.to_owned()) - .or_else(|_| ServerName::try_from(host.to_owned())) - .map_err(|_| RustlsError::General("invalid SNI".into()))?; - - let tls_stream: TlsStream = connector.connect(server_name, stream).await?; - - // Extract negotiated parameters and certificates - let (_io, session) = tls_stream.get_ref(); - let cipher_suite = session - .negotiated_cipher_suite() - .map(|s| u16::from(s.suite()).to_be_bytes()) - .unwrap_or([0x13, 0x01]); - - let certs: Vec> = session - .peer_certificates() - .map(|slice| slice.to_vec()) - .unwrap_or_default(); - let cert_chain_der: Vec> = certs.iter().map(|c| c.as_ref().to_vec()).collect(); - let cert_payload = - encode_tls13_certificate_message(&cert_chain_der).map(|certificate_message| { - TlsCertPayload { - cert_chain_der: cert_chain_der.clone(), - certificate_message, - } - }); - - let total_cert_len = cert_payload - .as_ref() - .map(|payload| payload.certificate_message.len()) - .unwrap_or_else(|| cert_chain_der.iter().map(Vec::len).sum::()) - .max(1024); - let cert_info = parse_cert_info(&certs); - - // Heuristic: split across two records if large to mimic real servers a bit. - let app_data_records_sizes = if total_cert_len > 3000 { - vec![total_cert_len / 2, total_cert_len - total_cert_len / 2] - } else { - vec![total_cert_len] - }; - - let parsed = ParsedServerHello { - version: [0x03, 0x03], - random: [0u8; 32], - session_id: Vec::new(), - cipher_suite, - compression: 0, - extensions: Vec::new(), - }; - - debug!( - sni = %sni, - len = total_cert_len, - cipher = format!("0x{:04x}", u16::from_be_bytes(cipher_suite)), - has_cert_payload = cert_payload.is_some(), - "Fetched TLS metadata via rustls" - ); - - Ok(TlsFetchResult { - server_hello_parsed: parsed, - app_data_records_sizes: app_data_records_sizes.clone(), - total_app_data_len: app_data_records_sizes.iter().sum(), - behavior_profile: TlsBehaviorProfile { - change_cipher_spec_count: 1, - app_data_record_sizes: app_data_records_sizes, - ticket_record_sizes: Vec::new(), - source: TlsProfileSource::Rustls, - ..TlsBehaviorProfile::default() - }, - cert_info, - cert_payload, - }) -} - -async fn fetch_via_rustls( - host: &str, - port: u16, - sni: &str, - connect_timeout: Duration, - upstream: Option>, - scope: Option<&str>, - proxy_protocol: u8, - unix_sock: Option<&str>, - strict_route: bool, - alpn_protocols: &[&[u8]], -) -> Result { - #[cfg(unix)] - if let Some(sock_path) = unix_sock { - match timeout(connect_timeout, UnixStream::connect(sock_path)).await { - Ok(Ok(stream)) => { - debug!( - sni = %sni, - sock = %sock_path, - "Rustls fetch using mask unix socket" - ); - let proxy_header = build_tls_fetch_proxy_header(proxy_protocol, None, None); - return fetch_via_rustls_stream(stream, host, sni, proxy_header, alpn_protocols) - .await; - } - Ok(Err(e)) => { - warn!( - sni = %sni, - sock = %sock_path, - error = %e, - "Rustls unix socket connect failed, falling back to TCP" - ); - } - Err(_) => { - warn!( - sni = %sni, - sock = %sock_path, - "Rustls unix socket connect timed out, falling back to TCP" - ); - } - } - } - - #[cfg(not(unix))] - let _ = unix_sock; - - let stream = - connect_tcp_with_upstream(host, port, connect_timeout, upstream, scope, strict_route) - .await?; - let (src_addr, dst_addr) = socket_addrs_from_upstream_stream(&stream); - let proxy_header = build_tls_fetch_proxy_header(proxy_protocol, src_addr, dst_addr); - fetch_via_rustls_stream(stream, host, sni, proxy_header, alpn_protocols).await -} - -/// Fetch real TLS metadata with an adaptive multi-profile strategy. -pub async fn fetch_real_tls_with_strategy( - host: &str, - port: u16, - sni: &str, - strategy: &TlsFetchStrategy, - upstream: Option>, - scope: Option<&str>, - proxy_protocol: u8, - unix_sock: Option<&str>, -) -> Result { - let attempt_timeout = strategy.attempt_timeout.max(Duration::from_millis(1)); - let total_budget = strategy.total_budget.max(Duration::from_millis(1)); - let started_at = Instant::now(); - let cache_key = profile_cache_key( - host, - port, - sni, - upstream.as_ref(), - scope, - proxy_protocol, - unix_sock, - ); - let profiles = order_profiles(strategy, Some(&cache_key), started_at); - - let mut raw_result = None; - let mut raw_last_error: Option = None; - let mut raw_last_error_kind = FetchErrorKind::Other; - let mut selected_profile = None; - - for profile in profiles { - let elapsed = started_at.elapsed(); - if elapsed >= total_budget { - break; - } - let timeout_for_attempt = attempt_timeout.min(total_budget - elapsed); - debug!( - sni = %sni, - profile = profile.as_str(), - alpn = ?profile_alpn_labels(profile), - grease_enabled = strategy.grease_enabled, - deterministic = strategy.deterministic, - "TLS fetch ClientHello params (raw)" - ); - - match fetch_via_raw_tls( - host, - port, - sni, - timeout_for_attempt, - upstream.clone(), - scope, - proxy_protocol, - unix_sock, - strategy.strict_route, - profile, - strategy.grease_enabled, - strategy.deterministic, - ) - .await - { - Ok(res) => { - selected_profile = Some(profile); - raw_result = Some(res); - break; - } - Err(err) => { - let kind = classify_fetch_error(&err); - warn!( - sni = %sni, - profile = profile.as_str(), - error_kind = ?kind, - error = %err, - "Raw TLS fetch attempt failed" - ); - raw_last_error_kind = kind; - raw_last_error = Some(err); - if strategy.strict_route && matches!(kind, FetchErrorKind::Route) { - break; - } - } - } - } - - if let Some(profile) = selected_profile { - remember_profile_success(strategy, Some(cache_key), profile, Instant::now()); - } - - if raw_result.is_none() - && strategy.strict_route - && matches!(raw_last_error_kind, FetchErrorKind::Route) - { - if let Some(err) = raw_last_error { - return Err(err); - } - return Err(anyhow!("TLS fetch strict-route failure")); - } - - let elapsed = started_at.elapsed(); - if elapsed >= total_budget { - return match raw_result { - Some(raw) => Ok(raw), - None => { - Err(raw_last_error.unwrap_or_else(|| anyhow!("TLS fetch total budget exhausted"))) - } - }; - } - - let rustls_timeout = attempt_timeout.min(total_budget - elapsed); - let rustls_profile = selected_profile.unwrap_or(TlsFetchProfile::ModernChromeLike); - let rustls_alpn_protocols = profile_alpn(rustls_profile); - debug!( - sni = %sni, - profile = rustls_profile.as_str(), - alpn = ?profile_alpn_labels(rustls_profile), - grease_enabled = strategy.grease_enabled, - deterministic = strategy.deterministic, - "TLS fetch ClientHello params (rustls)" - ); - let rustls_result = fetch_via_rustls( - host, - port, - sni, - rustls_timeout, - upstream, - scope, - proxy_protocol, - unix_sock, - strategy.strict_route, - rustls_alpn_protocols, - ) - .await; - - match rustls_result { - Ok(rustls) => { - if let Some(mut raw) = raw_result { - raw.cert_info = rustls.cert_info; - raw.cert_payload = rustls.cert_payload; - raw.behavior_profile.source = TlsProfileSource::Merged; - raw.behavior_profile - .refresh_server_hello_summary(&raw.server_hello_parsed); - debug!(sni = %sni, "Fetched TLS metadata via adaptive raw probe + rustls cert chain"); - Ok(raw) - } else { - Ok(rustls) - } - } - Err(err) => { - if let Some(raw) = raw_result { - warn!(sni = %sni, error = %err, "Rustls cert fetch failed, using raw TLS metadata only"); - Ok(raw) - } else if let Some(raw_err) = raw_last_error { - Err(anyhow!("TLS fetch failed (raw: {raw_err}; rustls: {err})")) - } else { - Err(err) - } - } - } -} +// TLS client configuration and wire-compatible ClientHello construction. +mod client_hello; +// TLS record and certificate metadata parsing. +mod records; +// TCP, upstream, and PROXY-protocol connection setup. +mod connection; +// Raw TLS fetch transport. +mod raw_fetch; +// Rustls-backed fetch transport. +mod rustls_fetch; +// Adaptive profile selection and public fetch entry points. +mod strategy; + +use client_hello::*; +use connection::*; +use raw_fetch::*; +use records::*; +use rustls_fetch::*; +pub use strategy::fetch_real_tls_with_strategy; /// Fetch real TLS metadata for the given SNI using a single-attempt compatibility strategy. #[allow(dead_code)] @@ -1517,510 +392,4 @@ pub async fn fetch_real_tls( } #[cfg(test)] -mod tests { - use std::net::SocketAddr; - use std::time::{Duration, Instant}; - - use super::{ - MLKEM768_CLIENT_ENCAPSULATION_KEY_LEN, ProfileCacheValue, TLS_NAMED_GROUP_X25519, - TLS_NAMED_GROUP_X25519MLKEM768, TlsFetchStrategy, X25519_KEY_SHARE_LEN, build_client_hello, - build_tls_fetch_proxy_header, derive_behavior_profile, encode_tls13_certificate_message, - fetch_via_rustls_stream, order_profiles, profile_alpn, profile_cache, profile_cache_key, - }; - use crate::config::TlsFetchProfile; - use crate::crypto::SecureRandom; - use crate::protocol::constants::{ - TLS_RECORD_APPLICATION, TLS_RECORD_CHANGE_CIPHER, TLS_RECORD_HANDSHAKE, - }; - use crate::tls_front::types::TlsProfileSource; - use tokio::io::AsyncReadExt; - - struct ParsedClientHelloForTest { - session_id: Vec, - extensions: Vec<(u16, Vec)>, - } - - fn read_u24(bytes: &[u8]) -> usize { - ((bytes[0] as usize) << 16) | ((bytes[1] as usize) << 8) | (bytes[2] as usize) - } - - fn parse_client_hello_for_test(record: &[u8]) -> ParsedClientHelloForTest { - assert!(record.len() >= 9, "record too short"); - assert_eq!(record[0], TLS_RECORD_HANDSHAKE, "not a handshake record"); - let record_len = u16::from_be_bytes([record[3], record[4]]) as usize; - assert_eq!(record.len(), 5 + record_len, "record length mismatch"); - - let handshake = &record[5..]; - assert_eq!(handshake[0], 0x01, "not a ClientHello handshake"); - let hello_len = read_u24(&handshake[1..4]); - assert_eq!(handshake.len(), 4 + hello_len, "handshake length mismatch"); - let hello = &handshake[4..]; - - let mut pos = 0usize; - pos += 2; - pos += 32; - - let session_len = hello[pos] as usize; - pos += 1; - let session_id = hello[pos..pos + session_len].to_vec(); - pos += session_len; - - let cipher_len = u16::from_be_bytes([hello[pos], hello[pos + 1]]) as usize; - pos += 2 + cipher_len; - - let compression_len = hello[pos] as usize; - pos += 1 + compression_len; - - let ext_len = u16::from_be_bytes([hello[pos], hello[pos + 1]]) as usize; - pos += 2; - let ext_end = pos + ext_len; - assert_eq!(ext_end, hello.len(), "extensions length mismatch"); - - let mut extensions = Vec::new(); - while pos + 4 <= ext_end { - let ext_type = u16::from_be_bytes([hello[pos], hello[pos + 1]]); - let data_len = u16::from_be_bytes([hello[pos + 2], hello[pos + 3]]) as usize; - pos += 4; - let data = hello[pos..pos + data_len].to_vec(); - pos += data_len; - extensions.push((ext_type, data)); - } - assert_eq!(pos, ext_end, "extension parse did not consume all bytes"); - - ParsedClientHelloForTest { - session_id, - extensions, - } - } - - fn parse_alpn_protocols(data: &[u8]) -> Vec> { - assert!(data.len() >= 2, "ALPN extension is too short"); - let protocols_len = u16::from_be_bytes([data[0], data[1]]) as usize; - assert_eq!(protocols_len + 2, data.len(), "ALPN list length mismatch"); - let mut pos = 2usize; - let mut out = Vec::new(); - while pos < data.len() { - let len = data[pos] as usize; - pos += 1; - out.push(data[pos..pos + len].to_vec()); - pos += len; - } - out - } - - async fn capture_rustls_client_hello_record( - alpn_protocols: &'static [&'static [u8]], - ) -> Vec { - let (client, mut server) = tokio::io::duplex(32 * 1024); - let fetch_task = tokio::spawn(async move { - fetch_via_rustls_stream(client, "example.com", "example.com", None, alpn_protocols) - .await - }); - - let mut header = [0u8; 5]; - server - .read_exact(&mut header) - .await - .expect("must read client hello record header"); - let body_len = u16::from_be_bytes([header[3], header[4]]) as usize; - let mut body = vec![0u8; body_len]; - server - .read_exact(&mut body) - .await - .expect("must read client hello record body"); - drop(server); - - let result = fetch_task.await.expect("fetch task must join"); - assert!( - result.is_err(), - "capture task should end with handshake error" - ); - - let mut record = Vec::with_capacity(5 + body_len); - record.extend_from_slice(&header); - record.extend_from_slice(&body); - record - } - - #[test] - fn test_encode_tls13_certificate_message_single_cert() { - let cert = vec![0x30, 0x03, 0x02, 0x01, 0x01]; - let message = - encode_tls13_certificate_message(std::slice::from_ref(&cert)).expect("message"); - - assert_eq!(message[0], 0x0b); - assert_eq!(read_u24(&message[1..4]), message.len() - 4); - assert_eq!(message[4], 0x00); - - let cert_list_len = read_u24(&message[5..8]); - assert_eq!(cert_list_len, cert.len() + 5); - - let cert_len = read_u24(&message[8..11]); - assert_eq!(cert_len, cert.len()); - assert_eq!(&message[11..11 + cert.len()], cert.as_slice()); - assert_eq!(&message[11 + cert.len()..13 + cert.len()], &[0x00, 0x00]); - } - - #[test] - fn test_encode_tls13_certificate_message_empty_chain() { - assert!(encode_tls13_certificate_message(&[]).is_none()); - } - - #[test] - fn test_derive_behavior_profile_splits_ticket_like_tail_records() { - let profile = derive_behavior_profile(&[ - (TLS_RECORD_HANDSHAKE, vec![0u8; 90]), - (TLS_RECORD_CHANGE_CIPHER, vec![0x01]), - (TLS_RECORD_APPLICATION, vec![0u8; 1400]), - (TLS_RECORD_APPLICATION, vec![0u8; 220]), - (TLS_RECORD_APPLICATION, vec![0u8; 180]), - ]); - - assert_eq!(profile.change_cipher_spec_count, 1); - assert_eq!(profile.app_data_record_sizes, vec![1400]); - assert_eq!(profile.ticket_record_sizes, vec![220, 180]); - assert_eq!(profile.source, TlsProfileSource::Raw); - } - - #[test] - fn test_order_profiles_prioritizes_fresh_cached_winner() { - let strategy = TlsFetchStrategy { - profiles: vec![ - TlsFetchProfile::ModernChromeLike, - TlsFetchProfile::CompatTls12, - TlsFetchProfile::LegacyMinimal, - ], - strict_route: true, - attempt_timeout: Duration::from_secs(1), - total_budget: Duration::from_secs(2), - grease_enabled: false, - deterministic: false, - profile_cache_ttl: Duration::from_secs(60), - }; - let cache_key = profile_cache_key( - "mask.example", - 443, - "tls.example", - None, - Some("tls"), - 0, - None, - ); - profile_cache().remove(&cache_key); - profile_cache().insert( - cache_key.clone(), - ProfileCacheValue { - profile: TlsFetchProfile::CompatTls12, - updated_at: Instant::now(), - }, - ); - - let ordered = order_profiles(&strategy, Some(&cache_key), Instant::now()); - assert_eq!(ordered[0], TlsFetchProfile::CompatTls12); - profile_cache().remove(&cache_key); - } - - #[test] - fn test_order_profiles_drops_expired_cached_winner() { - let strategy = TlsFetchStrategy { - profiles: vec![ - TlsFetchProfile::ModernFirefoxLike, - TlsFetchProfile::CompatTls12, - ], - strict_route: true, - attempt_timeout: Duration::from_secs(1), - total_budget: Duration::from_secs(2), - grease_enabled: false, - deterministic: false, - profile_cache_ttl: Duration::from_secs(5), - }; - let cache_key = - profile_cache_key("mask2.example", 443, "tls2.example", None, None, 0, None); - profile_cache().remove(&cache_key); - profile_cache().insert( - cache_key.clone(), - ProfileCacheValue { - profile: TlsFetchProfile::CompatTls12, - updated_at: Instant::now() - Duration::from_secs(6), - }, - ); - - let ordered = order_profiles(&strategy, Some(&cache_key), Instant::now()); - assert_eq!(ordered[0], TlsFetchProfile::ModernFirefoxLike); - assert!(profile_cache().get(&cache_key).is_none()); - } - - #[test] - fn test_deterministic_client_hello_is_stable() { - let rng = SecureRandom::new(); - let first = build_client_hello( - "stable.example", - &rng, - TlsFetchProfile::ModernChromeLike, - true, - true, - ); - let second = build_client_hello( - "stable.example", - &rng, - TlsFetchProfile::ModernChromeLike, - true, - true, - ); - - assert_eq!(first, second); - } - - #[test] - fn test_raw_client_hello_alpn_matches_profile() { - let rng = SecureRandom::new(); - for profile in [ - TlsFetchProfile::ModernChromeLike, - TlsFetchProfile::ModernFirefoxLike, - TlsFetchProfile::CompatTls12, - TlsFetchProfile::LegacyMinimal, - ] { - let hello = build_client_hello("alpn.example", &rng, profile, false, true); - let parsed = parse_client_hello_for_test(&hello); - let alpn_ext = parsed - .extensions - .iter() - .find(|(ext_type, _)| *ext_type == 0x0010) - .expect("ALPN extension must exist"); - let parsed_alpn = parse_alpn_protocols(&alpn_ext.1); - let expected_alpn = profile_alpn(profile) - .iter() - .map(|proto| proto.to_vec()) - .collect::>(); - assert_eq!( - parsed_alpn, - expected_alpn, - "ALPN mismatch for {}", - profile.as_str() - ); - } - } - - #[test] - fn test_modern_chrome_like_browser_extension_layout() { - let rng = SecureRandom::new(); - let hello = build_client_hello( - "chrome.example", - &rng, - TlsFetchProfile::ModernChromeLike, - false, - true, - ); - let parsed = parse_client_hello_for_test(&hello); - assert_eq!( - parsed.session_id.len(), - 32, - "modern chrome must use non-empty session id" - ); - - let extension_ids = parsed - .extensions - .iter() - .map(|(ext_type, _)| *ext_type) - .collect::>(); - let expected_prefix = [ - 0x0000, 0x000b, 0x000a, 0x0023, 0x000d, 0x002b, 0x002d, 0x0033, 0x0010, - ]; - assert!( - extension_ids.as_slice().starts_with(&expected_prefix), - "unexpected extension order: {extension_ids:?}" - ); - assert!( - extension_ids.contains(&0x0015), - "modern chrome profile should include padding extension" - ); - - let key_share = parsed - .extensions - .iter() - .find(|(ext_type, _)| *ext_type == 0x0033) - .expect("key_share extension must exist"); - let key_share_data = &key_share.1; - assert!( - key_share_data.len() >= 2 + 4 + 32, - "key_share payload is too short" - ); - let entry_len = u16::from_be_bytes([key_share_data[0], key_share_data[1]]) as usize; - assert_eq!( - entry_len, - key_share_data.len() - 2, - "key_share list length mismatch" - ); - let mut pos = 2usize; - let hybrid_group = u16::from_be_bytes([key_share_data[pos], key_share_data[pos + 1]]); - let hybrid_len = - u16::from_be_bytes([key_share_data[pos + 2], key_share_data[pos + 3]]) as usize; - pos += 4; - let hybrid_key = &key_share_data[pos..pos + hybrid_len]; - pos += hybrid_len; - assert_eq!( - hybrid_group, TLS_NAMED_GROUP_X25519MLKEM768, - "first key_share group must be X25519MLKEM768" - ); - assert_eq!( - hybrid_len, - MLKEM768_CLIENT_ENCAPSULATION_KEY_LEN + X25519_KEY_SHARE_LEN, - "hybrid key length must match X25519MLKEM768" - ); - assert!( - hybrid_key.iter().any(|b| *b != 0), - "hybrid key must not be all zero" - ); - - let group = u16::from_be_bytes([key_share_data[pos], key_share_data[pos + 1]]); - let key_len = - u16::from_be_bytes([key_share_data[pos + 2], key_share_data[pos + 3]]) as usize; - pos += 4; - let key = &key_share_data[pos..pos + key_len]; - assert_eq!( - group, TLS_NAMED_GROUP_X25519, - "second key_share group must be x25519" - ); - assert_eq!( - key_len, X25519_KEY_SHARE_LEN, - "x25519 key length must be 32" - ); - assert!( - key.iter().any(|b| *b != 0), - "x25519 key must not be all zero" - ); - } - - #[test] - fn test_fallback_profiles_keep_compat_extension_set() { - let rng = SecureRandom::new(); - for profile in [ - TlsFetchProfile::ModernFirefoxLike, - TlsFetchProfile::CompatTls12, - TlsFetchProfile::LegacyMinimal, - ] { - let hello = build_client_hello("fallback.example", &rng, profile, false, true); - let parsed = parse_client_hello_for_test(&hello); - let extension_ids = parsed - .extensions - .iter() - .map(|(ext_type, _)| *ext_type) - .collect::>(); - - assert!(extension_ids.contains(&0x0000), "SNI extension must exist"); - assert!( - extension_ids.contains(&0x000a), - "supported_groups extension must exist" - ); - assert!( - extension_ids.contains(&0x000d), - "signature_algorithms extension must exist" - ); - assert!( - extension_ids.contains(&0x002b), - "supported_versions extension must exist" - ); - assert!( - extension_ids.contains(&0x0033), - "key_share extension must exist" - ); - assert!(extension_ids.contains(&0x0010), "ALPN extension must exist"); - assert!( - !extension_ids.contains(&0x000b), - "ec_point_formats must stay chrome-only" - ); - assert!( - !extension_ids.contains(&0x0023), - "session_ticket must stay chrome-only" - ); - assert!( - !extension_ids.contains(&0x002d), - "psk_key_exchange_modes must stay chrome-only" - ); - - let expected_session_len = if matches!(profile, TlsFetchProfile::ModernFirefoxLike) { - 32 - } else { - 0 - }; - assert_eq!( - parsed.session_id.len(), - expected_session_len, - "unexpected session id length for {}", - profile.as_str() - ); - } - } - - #[tokio::test(flavor = "current_thread")] - async fn test_rustls_client_hello_alpn_matches_selected_profile() { - for profile in [ - TlsFetchProfile::ModernChromeLike, - TlsFetchProfile::CompatTls12, - TlsFetchProfile::LegacyMinimal, - ] { - let record = capture_rustls_client_hello_record(profile_alpn(profile)).await; - let parsed = parse_client_hello_for_test(&record); - let alpn_ext = parsed - .extensions - .iter() - .find(|(ext_type, _)| *ext_type == 0x0010) - .expect("ALPN extension must exist"); - let parsed_alpn = parse_alpn_protocols(&alpn_ext.1); - let expected_alpn = profile_alpn(profile) - .iter() - .map(|proto| proto.to_vec()) - .collect::>(); - assert_eq!( - parsed_alpn, - expected_alpn, - "rustls ALPN mismatch for {}", - profile.as_str() - ); - } - } - - #[test] - fn test_build_tls_fetch_proxy_header_v2_with_tcp_addrs() { - let src: SocketAddr = "198.51.100.10:42000".parse().expect("valid src"); - let dst: SocketAddr = "203.0.113.20:443".parse().expect("valid dst"); - let header = build_tls_fetch_proxy_header(2, Some(src), Some(dst)).expect("header"); - - assert_eq!( - &header[..12], - &[ - 0x0d, 0x0a, 0x0d, 0x0a, 0x00, 0x0d, 0x0a, 0x51, 0x55, 0x49, 0x54, 0x0a - ] - ); - assert_eq!(header[12], 0x21); - assert_eq!(header[13], 0x11); - assert_eq!(u16::from_be_bytes([header[14], header[15]]), 12); - assert_eq!(&header[16..20], &[198, 51, 100, 10]); - assert_eq!(&header[20..24], &[203, 0, 113, 20]); - assert_eq!(u16::from_be_bytes([header[24], header[25]]), 42000); - assert_eq!(u16::from_be_bytes([header[26], header[27]]), 443); - } - - #[test] - fn test_build_tls_fetch_proxy_header_v2_mixed_family_falls_back_to_local_command() { - let src: SocketAddr = "198.51.100.10:42000".parse().expect("valid src"); - let dst: SocketAddr = "[2001:db8::20]:443".parse().expect("valid dst"); - let header = build_tls_fetch_proxy_header(2, Some(src), Some(dst)).expect("header"); - - assert_eq!(header[12], 0x20); - assert_eq!(header[13], 0x00); - assert_eq!(u16::from_be_bytes([header[14], header[15]]), 0); - } - - #[test] - fn test_build_tls_fetch_proxy_header_v1_with_tcp_addrs() { - let src: SocketAddr = "198.51.100.10:42000".parse().expect("valid src"); - let dst: SocketAddr = "203.0.113.20:443".parse().expect("valid dst"); - let header = build_tls_fetch_proxy_header(1, Some(src), Some(dst)).expect("header"); - - assert_eq!( - header, - b"PROXY TCP4 198.51.100.10 203.0.113.20 42000 443\r\n" - ); - } -} +mod tests; diff --git a/src/tls_front/fetcher/client_hello.rs b/src/tls_front/fetcher/client_hello.rs new file mode 100644 index 0000000..2bbfecf --- /dev/null +++ b/src/tls_front/fetcher/client_hello.rs @@ -0,0 +1,387 @@ +use super::*; + +pub(super) fn build_client_config(alpn_protocols: &[&[u8]]) -> Arc { + let root = rustls::RootCertStore::empty(); + + let provider = rustls::crypto::ring::default_provider(); + let mut config = ClientConfig::builder_with_provider(Arc::new(provider)) + .with_protocol_versions(&[&rustls::version::TLS13, &rustls::version::TLS12]) + .expect("protocol versions") + .with_root_certificates(root) + .with_no_client_auth(); + + config + .dangerous() + .set_certificate_verifier(Arc::new(NoVerify)); + config.alpn_protocols = alpn_protocols.iter().map(|proto| proto.to_vec()).collect(); + + Arc::new(config) +} + +pub(super) fn deterministic_bytes(seed: &str, len: usize) -> Vec { + let mut out = Vec::with_capacity(len); + let mut counter: u32 = 0; + while out.len() < len { + let mut chunk_seed = Vec::with_capacity(seed.len() + std::mem::size_of::()); + chunk_seed.extend_from_slice(seed.as_bytes()); + chunk_seed.extend_from_slice(&counter.to_le_bytes()); + out.extend_from_slice(&sha256(&chunk_seed)); + counter = counter.wrapping_add(1); + } + out.truncate(len); + out +} + +pub(super) fn profile_cipher_suites(profile: TlsFetchProfile) -> &'static [u16] { + const MODERN_CHROME: &[u16] = &[ + 0x1301, 0x1302, 0x1303, 0xc02b, 0xc02c, 0xcca9, 0xc02f, 0xc030, 0xcca8, 0x009e, 0x00ff, + ]; + const MODERN_FIREFOX: &[u16] = &[ + 0x1301, 0x1303, 0x1302, 0xc02b, 0xcca9, 0xc02c, 0xc02f, 0xcca8, 0xc030, 0x009e, 0x00ff, + ]; + const COMPAT_TLS12: &[u16] = &[ + 0xc02b, 0xc02c, 0xc02f, 0xc030, 0xcca9, 0xcca8, 0x1301, 0x1302, 0x1303, 0x009e, 0x00ff, + ]; + const LEGACY_MINIMAL: &[u16] = &[0xc02b, 0xc02f, 0x1301, 0x1302, 0x00ff]; + + match profile { + TlsFetchProfile::ModernChromeLike => MODERN_CHROME, + TlsFetchProfile::ModernFirefoxLike => MODERN_FIREFOX, + TlsFetchProfile::CompatTls12 => COMPAT_TLS12, + TlsFetchProfile::LegacyMinimal => LEGACY_MINIMAL, + } +} + +pub(super) fn profile_groups(profile: TlsFetchProfile) -> &'static [u16] { + const MODERN: &[u16] = &[ + TLS_NAMED_GROUP_X25519MLKEM768, + TLS_NAMED_GROUP_X25519, + 0x0017, + 0x0018, + ]; + const COMPAT: &[u16] = &[TLS_NAMED_GROUP_X25519, 0x0017]; + const LEGACY: &[u16] = &[0x0017]; + + match profile { + TlsFetchProfile::ModernChromeLike | TlsFetchProfile::ModernFirefoxLike => MODERN, + TlsFetchProfile::CompatTls12 => COMPAT, + TlsFetchProfile::LegacyMinimal => LEGACY, + } +} + +pub(super) fn profile_sig_algs(profile: TlsFetchProfile) -> &'static [u16] { + const MODERN: &[u16] = &[0x0804, 0x0805, 0x0403, 0x0503, 0x0806]; + const COMPAT: &[u16] = &[0x0403, 0x0503, 0x0804, 0x0805]; + const LEGACY: &[u16] = &[0x0403, 0x0804]; + + match profile { + TlsFetchProfile::ModernChromeLike | TlsFetchProfile::ModernFirefoxLike => MODERN, + TlsFetchProfile::CompatTls12 => COMPAT, + TlsFetchProfile::LegacyMinimal => LEGACY, + } +} + +pub(super) fn profile_alpn(profile: TlsFetchProfile) -> &'static [&'static [u8]] { + const H2_HTTP11: &[&[u8]] = &[b"h2", b"http/1.1"]; + const HTTP11: &[&[u8]] = &[b"http/1.1"]; + match profile { + TlsFetchProfile::ModernChromeLike | TlsFetchProfile::ModernFirefoxLike => H2_HTTP11, + TlsFetchProfile::CompatTls12 | TlsFetchProfile::LegacyMinimal => HTTP11, + } +} + +pub(super) fn profile_alpn_labels(profile: TlsFetchProfile) -> &'static [&'static str] { + const H2_HTTP11: &[&str] = &["h2", "http/1.1"]; + const HTTP11: &[&str] = &["http/1.1"]; + match profile { + TlsFetchProfile::ModernChromeLike | TlsFetchProfile::ModernFirefoxLike => H2_HTTP11, + TlsFetchProfile::CompatTls12 | TlsFetchProfile::LegacyMinimal => HTTP11, + } +} + +pub(super) fn profile_session_id_len(profile: TlsFetchProfile) -> usize { + match profile { + TlsFetchProfile::ModernChromeLike | TlsFetchProfile::ModernFirefoxLike => 32, + TlsFetchProfile::CompatTls12 | TlsFetchProfile::LegacyMinimal => 0, + } +} + +pub(super) fn profile_supported_versions(profile: TlsFetchProfile) -> &'static [u16] { + const MODERN: &[u16] = &[0x0304, 0x0303]; + const COMPAT: &[u16] = &[0x0303, 0x0304]; + const LEGACY: &[u16] = &[0x0303]; + match profile { + TlsFetchProfile::ModernChromeLike | TlsFetchProfile::ModernFirefoxLike => MODERN, + TlsFetchProfile::CompatTls12 => COMPAT, + TlsFetchProfile::LegacyMinimal => LEGACY, + } +} + +pub(super) fn profile_padding_target(profile: TlsFetchProfile) -> usize { + match profile { + // X25519MLKEM768 makes the Chrome-like ClientHello much larger than + // legacy pre-hybrid profiles; keep enough headroom for padding. + TlsFetchProfile::ModernChromeLike => 1450, + TlsFetchProfile::ModernFirefoxLike => 200, + TlsFetchProfile::CompatTls12 => 180, + TlsFetchProfile::LegacyMinimal => 64, + } +} + +pub(super) fn grease_value(rng: &SecureRandom, deterministic: bool, seed: &str) -> u16 { + const GREASE_VALUES: [u16; 16] = [ + 0x0a0a, 0x1a1a, 0x2a2a, 0x3a3a, 0x4a4a, 0x5a5a, 0x6a6a, 0x7a7a, 0x8a8a, 0x9a9a, 0xaaaa, + 0xbaba, 0xcaca, 0xdada, 0xeaea, 0xfafa, + ]; + if deterministic { + let idx = deterministic_bytes(seed, 1)[0] as usize % GREASE_VALUES.len(); + GREASE_VALUES[idx] + } else { + let idx = (rng.bytes(1)[0] as usize) % GREASE_VALUES.len(); + GREASE_VALUES[idx] + } +} + +pub(super) fn gen_mlkem768_client_encapsulation_key( + rng: &SecureRandom, + deterministic: bool, + seed: &str, +) -> Option> { + let seed_bytes = if deterministic { + deterministic_bytes(seed, 64) + } else { + rng.bytes(64) + }; + let seed = MlKemSeed::try_from(seed_bytes.as_slice()).ok()?; + let decapsulation_key = MlKemDecapsulationKey::::from_seed(seed); + let encapsulation_key = decapsulation_key.encapsulation_key().to_bytes(); + let bytes = encapsulation_key.as_slice(); + if bytes.len() == MLKEM768_CLIENT_ENCAPSULATION_KEY_LEN { + Some(bytes.to_vec()) + } else { + None + } +} + +pub(super) fn gen_x25519mlkem768_client_key_share( + rng: &SecureRandom, + deterministic: bool, + seed: &str, +) -> Option> { + let mlkem_key = + gen_mlkem768_client_encapsulation_key(rng, deterministic, &format!("{seed}:mlkem768"))?; + let x25519_key = gen_key_share(rng, deterministic, &format!("{seed}:x25519")); + let mut key_share = + Vec::with_capacity(MLKEM768_CLIENT_ENCAPSULATION_KEY_LEN + x25519_key.len()); + key_share.extend_from_slice(&mlkem_key); + key_share.extend_from_slice(&x25519_key); + Some(key_share) +} + +pub(super) fn push_client_key_share_entry(keyshare: &mut Vec, group: u16, key: &[u8]) { + keyshare.extend_from_slice(&group.to_be_bytes()); + keyshare.extend_from_slice(&(key.len() as u16).to_be_bytes()); + keyshare.extend_from_slice(key); +} + +pub(super) fn build_client_hello( + sni: &str, + rng: &SecureRandom, + profile: TlsFetchProfile, + grease_enabled: bool, + deterministic: bool, +) -> Vec { + // === ClientHello body === + let mut body = Vec::new(); + + // Legacy version (TLS 1.0) as in real ClientHello headers + body.extend_from_slice(&[0x03, 0x03]); + + // Random + if deterministic { + body.extend_from_slice(&deterministic_bytes(&format!("tls-fetch-random:{sni}"), 32)); + } else { + body.extend_from_slice(&rng.bytes(32)); + } + + // Use non-empty Session ID for modern TLS 1.3-like profiles to reduce middlebox friction. + let session_id_len = profile_session_id_len(profile); + let session_id = if session_id_len == 0 { + Vec::new() + } else if deterministic { + deterministic_bytes( + &format!("tls-fetch-session:{sni}:{}", profile.as_str()), + session_id_len, + ) + } else { + rng.bytes(session_id_len) + }; + body.push(session_id.len() as u8); + body.extend_from_slice(&session_id); + + let mut cipher_suites = profile_cipher_suites(profile).to_vec(); + if grease_enabled { + let grease = grease_value(rng, deterministic, &format!("cipher:{sni}")); + cipher_suites.insert(0, grease); + } + body.extend_from_slice(&((cipher_suites.len() * 2) as u16).to_be_bytes()); + for suite in cipher_suites { + body.extend_from_slice(&suite.to_be_bytes()); + } + + // Compression methods: null only + body.push(1); + body.push(0); + + // === Extensions === + let mut exts = Vec::new(); + + let mut push_extension = |ext_type: u16, data: &[u8]| { + exts.extend_from_slice(&ext_type.to_be_bytes()); + exts.extend_from_slice(&(data.len() as u16).to_be_bytes()); + exts.extend_from_slice(data); + }; + + // server_name (SNI) + let sni_bytes = sni.as_bytes(); + let mut sni_ext = Vec::with_capacity(5 + sni_bytes.len()); + sni_ext.extend_from_slice(&(sni_bytes.len() as u16 + 3).to_be_bytes()); + sni_ext.push(0); + sni_ext.extend_from_slice(&(sni_bytes.len() as u16).to_be_bytes()); + sni_ext.extend_from_slice(sni_bytes); + push_extension(0x0000, &sni_ext); + + // Chrome-like profile keeps browser-like ordering and extension set. + if matches!(profile, TlsFetchProfile::ModernChromeLike) { + // ec_point_formats: uncompressed only. + push_extension(0x000b, &[0x01, 0x00]); + } + + // supported_groups + let mut groups = profile_groups(profile).to_vec(); + if grease_enabled { + let grease = grease_value(rng, deterministic, &format!("group:{sni}")); + groups.insert(0, grease); + } + let mut groups_ext = Vec::with_capacity(2 + groups.len() * 2); + groups_ext.extend_from_slice(&(groups.len() as u16 * 2).to_be_bytes()); + for g in groups { + groups_ext.extend_from_slice(&g.to_be_bytes()); + } + push_extension(0x000a, &groups_ext); + + if matches!(profile, TlsFetchProfile::ModernChromeLike) { + // session_ticket + push_extension(0x0023, &[]); + } + + // signature_algorithms + let mut sig_algs = profile_sig_algs(profile).to_vec(); + if grease_enabled { + let grease = grease_value(rng, deterministic, &format!("sigalg:{sni}")); + sig_algs.insert(0, grease); + } + let mut sig_algs_ext = Vec::with_capacity(2 + sig_algs.len() * 2); + sig_algs_ext.extend_from_slice(&(sig_algs.len() as u16 * 2).to_be_bytes()); + for a in sig_algs { + sig_algs_ext.extend_from_slice(&a.to_be_bytes()); + } + push_extension(0x000d, &sig_algs_ext); + + // supported_versions + let mut versions = profile_supported_versions(profile).to_vec(); + if grease_enabled { + let grease = grease_value(rng, deterministic, &format!("version:{sni}")); + versions.insert(0, grease); + } + let mut versions_ext = Vec::with_capacity(1 + versions.len() * 2); + versions_ext.push((versions.len() * 2) as u8); + for v in versions { + versions_ext.extend_from_slice(&v.to_be_bytes()); + } + push_extension(0x002b, &versions_ext); + + if matches!(profile, TlsFetchProfile::ModernChromeLike) { + // psk_key_exchange_modes: psk_dhe_ke + push_extension(0x002d, &[0x01, 0x01]); + } + + // key_share + let key_share_seed = format!("tls-fetch-keyshare:{sni}:{}", profile.as_str()); + let mut keyshare = Vec::new(); + if matches!( + profile, + TlsFetchProfile::ModernChromeLike | TlsFetchProfile::ModernFirefoxLike + ) { + if let Some(key) = gen_x25519mlkem768_client_key_share(rng, deterministic, &key_share_seed) + { + push_client_key_share_entry(&mut keyshare, TLS_NAMED_GROUP_X25519MLKEM768, &key); + } + } + let key = gen_key_share(rng, deterministic, &key_share_seed); + push_client_key_share_entry(&mut keyshare, TLS_NAMED_GROUP_X25519, &key); + let mut keyshare_ext = Vec::with_capacity(2 + keyshare.len()); + keyshare_ext.extend_from_slice(&(keyshare.len() as u16).to_be_bytes()); + keyshare_ext.extend_from_slice(&keyshare); + push_extension(0x0033, &keyshare_ext); + + // ALPN + let mut alpn_list = Vec::new(); + for proto in profile_alpn(profile) { + alpn_list.push(proto.len() as u8); + alpn_list.extend_from_slice(proto); + } + if !alpn_list.is_empty() { + let mut alpn_ext = Vec::with_capacity(2 + alpn_list.len()); + alpn_ext.extend_from_slice(&(alpn_list.len() as u16).to_be_bytes()); + alpn_ext.extend_from_slice(&alpn_list); + push_extension(0x0010, &alpn_ext); + } + + if grease_enabled { + let grease = grease_value(rng, deterministic, &format!("ext:{sni}")); + push_extension(grease, &[]); + } + + // padding to reduce recognizability and keep length ~500 bytes + let target_ext_len = profile_padding_target(profile); + if exts.len() < target_ext_len { + let remaining = target_ext_len - exts.len(); + if remaining > 4 { + let pad_len = remaining - 4; // minus type+len + exts.extend_from_slice(&0x0015u16.to_be_bytes()); // padding extension + exts.extend_from_slice(&(pad_len as u16).to_be_bytes()); + exts.resize(exts.len() + pad_len, 0); + } + } + + // Extensions length prefix + body.extend_from_slice(&(exts.len() as u16).to_be_bytes()); + body.extend_from_slice(&exts); + + // === Handshake wrapper === + let mut handshake = Vec::new(); + handshake.push(0x01); // ClientHello + let len_bytes = (body.len() as u32).to_be_bytes(); + handshake.extend_from_slice(&len_bytes[1..4]); + handshake.extend_from_slice(&body); + + // === Record === + let mut record = Vec::new(); + record.push(TLS_RECORD_HANDSHAKE); + record.extend_from_slice(&[0x03, 0x01]); // legacy record version + record.extend_from_slice(&(handshake.len() as u16).to_be_bytes()); + record.extend_from_slice(&handshake); + + record +} + +pub(super) fn gen_key_share(rng: &SecureRandom, deterministic: bool, seed: &str) -> [u8; 32] { + let mut scalar = [0u8; 32]; + if deterministic { + scalar.copy_from_slice(&deterministic_bytes(seed, 32)); + } else { + scalar.copy_from_slice(&rng.bytes(32)); + } + x25519(scalar, X25519_BASEPOINT_BYTES) +} diff --git a/src/tls_front/fetcher/connection.rs b/src/tls_front/fetcher/connection.rs new file mode 100644 index 0000000..c79fd7d --- /dev/null +++ b/src/tls_front/fetcher/connection.rs @@ -0,0 +1,114 @@ +use super::*; + +pub(super) async fn connect_with_dns_override( + host: &str, + port: u16, + connect_timeout: Duration, +) -> Result { + Ok(timeout(connect_timeout, TcpStream::connect((host, port))).await??) +} + +pub(super) async fn connect_tcp_with_upstream( + host: &str, + port: u16, + connect_timeout: Duration, + upstream: Option>, + scope: Option<&str>, + strict_route: bool, +) -> Result { + if let Some(manager) = upstream { + let resolved = match manager.resolve_hostname(host, port).await { + Ok(addr) => Some(addr), + Err(e) => { + if strict_route { + return Err(anyhow!( + "upstream route DNS resolution failed for {host}:{port}: {e}" + )); + } + warn!( + host = %host, + port = port, + scope = ?scope, + error = %e, + "Upstream DNS resolution failed, using direct connect" + ); + None + } + }; + + if let Some(addr) = resolved { + match manager.connect(addr, None, scope).await { + Ok(stream) => return Ok(stream), + Err(e) => { + if strict_route { + return Err(anyhow!( + "upstream route connect failed for {host}:{port}: {e}" + )); + } + warn!( + host = %host, + port = port, + scope = ?scope, + error = %e, + "Upstream connect failed, using direct connect" + ); + return Ok(UpstreamStream::Tcp( + timeout(connect_timeout, TcpStream::connect(addr)).await??, + )); + } + } + } else if strict_route { + return Err(anyhow!( + "upstream route resolution produced no usable address for {host}:{port}" + )); + } + } + Ok(UpstreamStream::Tcp( + connect_with_dns_override(host, port, connect_timeout).await?, + )) +} + +pub(super) fn socket_addrs_from_upstream_stream( + stream: &UpstreamStream, +) -> (Option, Option) { + match stream { + UpstreamStream::Tcp(tcp) => (tcp.local_addr().ok(), tcp.peer_addr().ok()), + UpstreamStream::Shadowsocks(_) => (None, None), + } +} + +pub(super) fn build_tls_fetch_proxy_header( + proxy_protocol: u8, + src_addr: Option, + dst_addr: Option, +) -> Option> { + match proxy_protocol { + 0 => None, + 2 => { + let header = match (src_addr, dst_addr) { + (Some(src @ SocketAddr::V4(_)), Some(dst @ SocketAddr::V4(_))) + | (Some(src @ SocketAddr::V6(_)), Some(dst @ SocketAddr::V6(_))) => { + ProxyProtocolV2Builder::new().with_addrs(src, dst).build() + } + _ => ProxyProtocolV2Builder::new().build(), + }; + Some(header) + } + _ => { + let header = match (src_addr, dst_addr) { + (Some(SocketAddr::V4(src)), Some(SocketAddr::V4(dst))) => { + ProxyProtocolV1Builder::new() + .tcp4(src.into(), dst.into()) + .build() + } + (Some(SocketAddr::V6(src)), Some(SocketAddr::V6(dst))) => { + ProxyProtocolV1Builder::new() + .tcp6(src.into(), dst.into()) + .build() + } + _ => ProxyProtocolV1Builder::new().build(), + }; + Some(header) + } + } +} diff --git a/src/tls_front/fetcher/raw_fetch.rs b/src/tls_front/fetcher/raw_fetch.rs new file mode 100644 index 0000000..14c4665 --- /dev/null +++ b/src/tls_front/fetcher/raw_fetch.rs @@ -0,0 +1,177 @@ +use super::*; + +pub(super) fn encode_tls13_certificate_message(cert_chain_der: &[Vec]) -> Option> { + if cert_chain_der.is_empty() { + return None; + } + + let mut certificate_list = Vec::new(); + for cert in cert_chain_der { + if cert.is_empty() { + return None; + } + certificate_list.extend_from_slice(&u24_bytes(cert.len())?); + certificate_list.extend_from_slice(cert); + certificate_list.extend_from_slice(&0u16.to_be_bytes()); // cert_entry extensions + } + + // Certificate = context_len(1) + certificate_list_len(3) + entries + let body_len = 1usize.checked_add(3)?.checked_add(certificate_list.len())?; + + let mut message = Vec::with_capacity(4 + body_len); + message.push(0x0b); // HandshakeType::certificate + message.extend_from_slice(&u24_bytes(body_len)?); + message.push(0x00); // certificate_request_context length + message.extend_from_slice(&u24_bytes(certificate_list.len())?); + message.extend_from_slice(&certificate_list); + Some(message) +} + +pub(super) async fn fetch_via_raw_tls_stream( + mut stream: S, + sni: &str, + connect_timeout: Duration, + proxy_header: Option>, + profile: TlsFetchProfile, + grease_enabled: bool, + deterministic: bool, +) -> Result +where + S: AsyncRead + AsyncWrite + Unpin, +{ + let rng = SecureRandom::new(); + let client_hello = build_client_hello(sni, &rng, profile, grease_enabled, deterministic); + timeout(connect_timeout, async { + if let Some(header) = proxy_header.as_ref() { + stream.write_all(&header).await?; + } + stream.write_all(&client_hello).await?; + stream.flush().await?; + Ok::<(), std::io::Error>(()) + }) + .await??; + + let mut records = Vec::new(); + let mut app_records_seen = 0usize; + // Read a bounded encrypted flight: ServerHello, CCS, certificate-like data, + // and a small number of ticket-like tail records. + for _ in 0..8 { + match timeout(connect_timeout, read_tls_record(&mut stream)).await { + Ok(Ok(rec)) => { + if rec.0 == TLS_RECORD_APPLICATION { + app_records_seen += 1; + } + records.push(rec); + } + Ok(Err(e)) => return Err(e), + Err(_) => break, + } + if app_records_seen >= 4 { + break; + } + } + + let mut server_hello = None; + let mut server_hello_record_len = 0usize; + for (t, body) in &records { + if *t == TLS_RECORD_HANDSHAKE && server_hello.is_none() { + server_hello = parse_server_hello(body); + server_hello_record_len = body.len(); + } + } + + let parsed = server_hello.ok_or_else(|| anyhow!("ServerHello not received"))?; + let mut behavior_profile = derive_behavior_profile(&records); + behavior_profile.server_hello_record_len = server_hello_record_len; + behavior_profile.refresh_server_hello_summary(&parsed); + let mut app_sizes = behavior_profile.app_data_record_sizes.clone(); + app_sizes.extend_from_slice(&behavior_profile.ticket_record_sizes); + let total_app_data_len = app_sizes.iter().sum::().max(1024); + let app_data_records_sizes = if app_sizes.is_empty() { + vec![total_app_data_len] + } else { + app_sizes + }; + + Ok(TlsFetchResult { + server_hello_parsed: parsed, + app_data_records_sizes, + total_app_data_len, + behavior_profile, + cert_info: None, + cert_payload: None, + }) +} + +pub(super) async fn fetch_via_raw_tls( + host: &str, + port: u16, + sni: &str, + connect_timeout: Duration, + upstream: Option>, + scope: Option<&str>, + proxy_protocol: u8, + unix_sock: Option<&str>, + strict_route: bool, + profile: TlsFetchProfile, + grease_enabled: bool, + deterministic: bool, +) -> Result { + #[cfg(unix)] + if let Some(sock_path) = unix_sock { + match timeout(connect_timeout, UnixStream::connect(sock_path)).await { + Ok(Ok(stream)) => { + debug!( + sni = %sni, + sock = %sock_path, + "Raw TLS fetch using mask unix socket" + ); + let proxy_header = build_tls_fetch_proxy_header(proxy_protocol, None, None); + return fetch_via_raw_tls_stream( + stream, + sni, + connect_timeout, + proxy_header, + profile, + grease_enabled, + deterministic, + ) + .await; + } + Ok(Err(e)) => { + warn!( + sni = %sni, + sock = %sock_path, + error = %e, + "Raw TLS unix socket connect failed, falling back to TCP" + ); + } + Err(_) => { + warn!( + sni = %sni, + sock = %sock_path, + "Raw TLS unix socket connect timed out, falling back to TCP" + ); + } + } + } + + #[cfg(not(unix))] + let _ = unix_sock; + + let stream = + connect_tcp_with_upstream(host, port, connect_timeout, upstream, scope, strict_route) + .await?; + let (src_addr, dst_addr) = socket_addrs_from_upstream_stream(&stream); + let proxy_header = build_tls_fetch_proxy_header(proxy_protocol, src_addr, dst_addr); + fetch_via_raw_tls_stream( + stream, + sni, + connect_timeout, + proxy_header, + profile, + grease_enabled, + deterministic, + ) + .await +} diff --git a/src/tls_front/fetcher/records.rs b/src/tls_front/fetcher/records.rs new file mode 100644 index 0000000..8b9aabb --- /dev/null +++ b/src/tls_front/fetcher/records.rs @@ -0,0 +1,165 @@ +use super::*; + +pub(super) async fn read_tls_record(stream: &mut S) -> Result<(u8, Vec)> +where + S: AsyncRead + Unpin, +{ + let mut header = [0u8; 5]; + stream.read_exact(&mut header).await?; + let len = u16::from_be_bytes([header[3], header[4]]) as usize; + let mut body = vec![0u8; len]; + stream.read_exact(&mut body).await?; + Ok((header[0], body)) +} + +pub(super) fn parse_server_hello(body: &[u8]) -> Option { + if body.len() < 4 || body[0] != 0x02 { + return None; + } + + let msg_len = u32::from_be_bytes([0, body[1], body[2], body[3]]) as usize; + if msg_len + 4 > body.len() { + return None; + } + + let mut pos = 4; + let version = [*body.get(pos)?, *body.get(pos + 1)?]; + pos += 2; + + let mut random = [0u8; 32]; + random.copy_from_slice(body.get(pos..pos + 32)?); + pos += 32; + + let session_len = *body.get(pos)? as usize; + pos += 1; + let session_id = body.get(pos..pos + session_len)?.to_vec(); + pos += session_len; + + let cipher_suite = [*body.get(pos)?, *body.get(pos + 1)?]; + pos += 2; + + let compression = *body.get(pos)?; + pos += 1; + + let ext_len = u16::from_be_bytes([*body.get(pos)?, *body.get(pos + 1)?]) as usize; + pos += 2; + let ext_end = pos.checked_add(ext_len)?; + if ext_end > body.len() { + return None; + } + + let mut extensions = Vec::new(); + while pos + 4 <= ext_end { + let etype = u16::from_be_bytes([body[pos], body[pos + 1]]); + let elen = u16::from_be_bytes([body[pos + 2], body[pos + 3]]) as usize; + pos += 4; + let data = body.get(pos..pos + elen)?.to_vec(); + pos += elen; + extensions.push(TlsExtension { + ext_type: etype, + data, + }); + } + + Some(ParsedServerHello { + version, + random, + session_id, + cipher_suite, + compression, + extensions, + }) +} + +pub(super) fn derive_behavior_profile(records: &[(u8, Vec)]) -> TlsBehaviorProfile { + let mut change_cipher_spec_count = 0u8; + let mut app_data_record_sizes = Vec::new(); + + for (record_type, body) in records { + match *record_type { + TLS_RECORD_CHANGE_CIPHER => { + change_cipher_spec_count = change_cipher_spec_count.saturating_add(1); + } + TLS_RECORD_APPLICATION => { + app_data_record_sizes.push(body.len()); + } + _ => {} + } + } + + let mut ticket_record_sizes = Vec::new(); + while app_data_record_sizes + .last() + .is_some_and(|size| *size <= 256 && ticket_record_sizes.len() < 2) + { + if let Some(size) = app_data_record_sizes.pop() { + ticket_record_sizes.push(size); + } + } + ticket_record_sizes.reverse(); + + TlsBehaviorProfile { + change_cipher_spec_count: change_cipher_spec_count.max(1), + app_data_record_sizes, + ticket_record_sizes, + source: TlsProfileSource::Raw, + ..TlsBehaviorProfile::default() + } +} + +pub(super) fn parse_cert_info(certs: &[CertificateDer<'static>]) -> Option { + let first = certs.first()?; + let (_rem, cert) = X509Certificate::from_der(first.as_ref()).ok()?; + + let not_before = Some(cert.validity().not_before.to_datetime().unix_timestamp()); + let not_after = Some(cert.validity().not_after.to_datetime().unix_timestamp()); + + let issuer_cn = cert + .issuer() + .iter_common_name() + .next() + .and_then(|cn| cn.as_str().ok()) + .map(|s| s.to_string()); + + let subject_cn = cert + .subject() + .iter_common_name() + .next() + .and_then(|cn| cn.as_str().ok()) + .map(|s| s.to_string()); + + let san_names = cert + .subject_alternative_name() + .ok() + .flatten() + .map(|san| { + san.value + .general_names + .iter() + .filter_map(|gn| match gn { + x509_parser::extensions::GeneralName::DNSName(n) => Some(n.to_string()), + _ => None, + }) + .collect::>() + }) + .unwrap_or_default(); + + Some(ParsedCertificateInfo { + not_after_unix: not_after, + not_before_unix: not_before, + issuer_cn, + subject_cn, + san_names, + }) +} + +pub(super) fn u24_bytes(value: usize) -> Option<[u8; 3]> { + if value > 0x00ff_ffff { + return None; + } + Some([ + ((value >> 16) & 0xff) as u8, + ((value >> 8) & 0xff) as u8, + (value & 0xff) as u8, + ]) +} diff --git a/src/tls_front/fetcher/rustls_fetch.rs b/src/tls_front/fetcher/rustls_fetch.rs new file mode 100644 index 0000000..fd22e8d --- /dev/null +++ b/src/tls_front/fetcher/rustls_fetch.rs @@ -0,0 +1,147 @@ +use super::*; + +pub(super) async fn fetch_via_rustls_stream( + mut stream: S, + host: &str, + sni: &str, + proxy_header: Option>, + alpn_protocols: &[&[u8]], +) -> Result +where + S: AsyncRead + AsyncWrite + Unpin, +{ + // rustls handshake path for certificate and basic negotiated metadata. + if let Some(header) = proxy_header.as_ref() { + stream.write_all(&header).await?; + stream.flush().await?; + } + + let config = build_client_config(alpn_protocols); + let connector = TlsConnector::from(config); + + let server_name = ServerName::try_from(sni.to_owned()) + .or_else(|_| ServerName::try_from(host.to_owned())) + .map_err(|_| RustlsError::General("invalid SNI".into()))?; + + let tls_stream: TlsStream = connector.connect(server_name, stream).await?; + + // Extract negotiated parameters and certificates + let (_io, session) = tls_stream.get_ref(); + let cipher_suite = session + .negotiated_cipher_suite() + .map(|s| u16::from(s.suite()).to_be_bytes()) + .unwrap_or([0x13, 0x01]); + + let certs: Vec> = session + .peer_certificates() + .map(|slice| slice.to_vec()) + .unwrap_or_default(); + let cert_chain_der: Vec> = certs.iter().map(|c| c.as_ref().to_vec()).collect(); + let cert_payload = + encode_tls13_certificate_message(&cert_chain_der).map(|certificate_message| { + TlsCertPayload { + cert_chain_der: cert_chain_der.clone(), + certificate_message, + } + }); + + let total_cert_len = cert_payload + .as_ref() + .map(|payload| payload.certificate_message.len()) + .unwrap_or_else(|| cert_chain_der.iter().map(Vec::len).sum::()) + .max(1024); + let cert_info = parse_cert_info(&certs); + + // Heuristic: split across two records if large to mimic real servers a bit. + let app_data_records_sizes = if total_cert_len > 3000 { + vec![total_cert_len / 2, total_cert_len - total_cert_len / 2] + } else { + vec![total_cert_len] + }; + + let parsed = ParsedServerHello { + version: [0x03, 0x03], + random: [0u8; 32], + session_id: Vec::new(), + cipher_suite, + compression: 0, + extensions: Vec::new(), + }; + + debug!( + sni = %sni, + len = total_cert_len, + cipher = format!("0x{:04x}", u16::from_be_bytes(cipher_suite)), + has_cert_payload = cert_payload.is_some(), + "Fetched TLS metadata via rustls" + ); + + Ok(TlsFetchResult { + server_hello_parsed: parsed, + app_data_records_sizes: app_data_records_sizes.clone(), + total_app_data_len: app_data_records_sizes.iter().sum(), + behavior_profile: TlsBehaviorProfile { + change_cipher_spec_count: 1, + app_data_record_sizes: app_data_records_sizes, + ticket_record_sizes: Vec::new(), + source: TlsProfileSource::Rustls, + ..TlsBehaviorProfile::default() + }, + cert_info, + cert_payload, + }) +} + +pub(super) async fn fetch_via_rustls( + host: &str, + port: u16, + sni: &str, + connect_timeout: Duration, + upstream: Option>, + scope: Option<&str>, + proxy_protocol: u8, + unix_sock: Option<&str>, + strict_route: bool, + alpn_protocols: &[&[u8]], +) -> Result { + #[cfg(unix)] + if let Some(sock_path) = unix_sock { + match timeout(connect_timeout, UnixStream::connect(sock_path)).await { + Ok(Ok(stream)) => { + debug!( + sni = %sni, + sock = %sock_path, + "Rustls fetch using mask unix socket" + ); + let proxy_header = build_tls_fetch_proxy_header(proxy_protocol, None, None); + return fetch_via_rustls_stream(stream, host, sni, proxy_header, alpn_protocols) + .await; + } + Ok(Err(e)) => { + warn!( + sni = %sni, + sock = %sock_path, + error = %e, + "Rustls unix socket connect failed, falling back to TCP" + ); + } + Err(_) => { + warn!( + sni = %sni, + sock = %sock_path, + "Rustls unix socket connect timed out, falling back to TCP" + ); + } + } + } + + #[cfg(not(unix))] + let _ = unix_sock; + + let stream = + connect_tcp_with_upstream(host, port, connect_timeout, upstream, scope, strict_route) + .await?; + let (src_addr, dst_addr) = socket_addrs_from_upstream_stream(&stream); + let proxy_header = build_tls_fetch_proxy_header(proxy_protocol, src_addr, dst_addr); + fetch_via_rustls_stream(stream, host, sni, proxy_header, alpn_protocols).await +} diff --git a/src/tls_front/fetcher/strategy.rs b/src/tls_front/fetcher/strategy.rs new file mode 100644 index 0000000..b370eba --- /dev/null +++ b/src/tls_front/fetcher/strategy.rs @@ -0,0 +1,161 @@ +use super::*; + +/// Fetch real TLS metadata with an adaptive multi-profile strategy. +pub async fn fetch_real_tls_with_strategy( + host: &str, + port: u16, + sni: &str, + strategy: &TlsFetchStrategy, + upstream: Option>, + scope: Option<&str>, + proxy_protocol: u8, + unix_sock: Option<&str>, +) -> Result { + let attempt_timeout = strategy.attempt_timeout.max(Duration::from_millis(1)); + let total_budget = strategy.total_budget.max(Duration::from_millis(1)); + let started_at = Instant::now(); + let cache_key = profile_cache_key( + host, + port, + sni, + upstream.as_ref(), + scope, + proxy_protocol, + unix_sock, + ); + let profiles = order_profiles(strategy, Some(&cache_key), started_at); + + let mut raw_result = None; + let mut raw_last_error: Option = None; + let mut raw_last_error_kind = FetchErrorKind::Other; + let mut selected_profile = None; + + for profile in profiles { + let elapsed = started_at.elapsed(); + if elapsed >= total_budget { + break; + } + let timeout_for_attempt = attempt_timeout.min(total_budget - elapsed); + debug!( + sni = %sni, + profile = profile.as_str(), + alpn = ?profile_alpn_labels(profile), + grease_enabled = strategy.grease_enabled, + deterministic = strategy.deterministic, + "TLS fetch ClientHello params (raw)" + ); + + match fetch_via_raw_tls( + host, + port, + sni, + timeout_for_attempt, + upstream.clone(), + scope, + proxy_protocol, + unix_sock, + strategy.strict_route, + profile, + strategy.grease_enabled, + strategy.deterministic, + ) + .await + { + Ok(res) => { + selected_profile = Some(profile); + raw_result = Some(res); + break; + } + Err(err) => { + let kind = classify_fetch_error(&err); + warn!( + sni = %sni, + profile = profile.as_str(), + error_kind = ?kind, + error = %err, + "Raw TLS fetch attempt failed" + ); + raw_last_error_kind = kind; + raw_last_error = Some(err); + if strategy.strict_route && matches!(kind, FetchErrorKind::Route) { + break; + } + } + } + } + + if let Some(profile) = selected_profile { + remember_profile_success(strategy, Some(cache_key), profile, Instant::now()); + } + + if raw_result.is_none() + && strategy.strict_route + && matches!(raw_last_error_kind, FetchErrorKind::Route) + { + if let Some(err) = raw_last_error { + return Err(err); + } + return Err(anyhow!("TLS fetch strict-route failure")); + } + + let elapsed = started_at.elapsed(); + if elapsed >= total_budget { + return match raw_result { + Some(raw) => Ok(raw), + None => { + Err(raw_last_error.unwrap_or_else(|| anyhow!("TLS fetch total budget exhausted"))) + } + }; + } + + let rustls_timeout = attempt_timeout.min(total_budget - elapsed); + let rustls_profile = selected_profile.unwrap_or(TlsFetchProfile::ModernChromeLike); + let rustls_alpn_protocols = profile_alpn(rustls_profile); + debug!( + sni = %sni, + profile = rustls_profile.as_str(), + alpn = ?profile_alpn_labels(rustls_profile), + grease_enabled = strategy.grease_enabled, + deterministic = strategy.deterministic, + "TLS fetch ClientHello params (rustls)" + ); + let rustls_result = fetch_via_rustls( + host, + port, + sni, + rustls_timeout, + upstream, + scope, + proxy_protocol, + unix_sock, + strategy.strict_route, + rustls_alpn_protocols, + ) + .await; + + match rustls_result { + Ok(rustls) => { + if let Some(mut raw) = raw_result { + raw.cert_info = rustls.cert_info; + raw.cert_payload = rustls.cert_payload; + raw.behavior_profile.source = TlsProfileSource::Merged; + raw.behavior_profile + .refresh_server_hello_summary(&raw.server_hello_parsed); + debug!(sni = %sni, "Fetched TLS metadata via adaptive raw probe + rustls cert chain"); + Ok(raw) + } else { + Ok(rustls) + } + } + Err(err) => { + if let Some(raw) = raw_result { + warn!(sni = %sni, error = %err, "Rustls cert fetch failed, using raw TLS metadata only"); + Ok(raw) + } else if let Some(raw_err) = raw_last_error { + Err(anyhow!("TLS fetch failed (raw: {raw_err}; rustls: {err})")) + } else { + Err(err) + } + } + } +} diff --git a/src/tls_front/fetcher/tests.rs b/src/tls_front/fetcher/tests.rs new file mode 100644 index 0000000..2309754 --- /dev/null +++ b/src/tls_front/fetcher/tests.rs @@ -0,0 +1,499 @@ +use std::net::SocketAddr; +use std::time::{Duration, Instant}; + +use super::{ + MLKEM768_CLIENT_ENCAPSULATION_KEY_LEN, ProfileCacheValue, TLS_NAMED_GROUP_X25519, + TLS_NAMED_GROUP_X25519MLKEM768, TlsFetchStrategy, X25519_KEY_SHARE_LEN, build_client_hello, + build_tls_fetch_proxy_header, derive_behavior_profile, encode_tls13_certificate_message, + fetch_via_rustls_stream, order_profiles, profile_alpn, profile_cache, profile_cache_key, +}; +use crate::config::TlsFetchProfile; +use crate::crypto::SecureRandom; +use crate::protocol::constants::{ + TLS_RECORD_APPLICATION, TLS_RECORD_CHANGE_CIPHER, TLS_RECORD_HANDSHAKE, +}; +use crate::tls_front::types::TlsProfileSource; +use tokio::io::AsyncReadExt; + +struct ParsedClientHelloForTest { + session_id: Vec, + extensions: Vec<(u16, Vec)>, +} + +fn read_u24(bytes: &[u8]) -> usize { + ((bytes[0] as usize) << 16) | ((bytes[1] as usize) << 8) | (bytes[2] as usize) +} + +fn parse_client_hello_for_test(record: &[u8]) -> ParsedClientHelloForTest { + assert!(record.len() >= 9, "record too short"); + assert_eq!(record[0], TLS_RECORD_HANDSHAKE, "not a handshake record"); + let record_len = u16::from_be_bytes([record[3], record[4]]) as usize; + assert_eq!(record.len(), 5 + record_len, "record length mismatch"); + + let handshake = &record[5..]; + assert_eq!(handshake[0], 0x01, "not a ClientHello handshake"); + let hello_len = read_u24(&handshake[1..4]); + assert_eq!(handshake.len(), 4 + hello_len, "handshake length mismatch"); + let hello = &handshake[4..]; + + let mut pos = 0usize; + pos += 2; + pos += 32; + + let session_len = hello[pos] as usize; + pos += 1; + let session_id = hello[pos..pos + session_len].to_vec(); + pos += session_len; + + let cipher_len = u16::from_be_bytes([hello[pos], hello[pos + 1]]) as usize; + pos += 2 + cipher_len; + + let compression_len = hello[pos] as usize; + pos += 1 + compression_len; + + let ext_len = u16::from_be_bytes([hello[pos], hello[pos + 1]]) as usize; + pos += 2; + let ext_end = pos + ext_len; + assert_eq!(ext_end, hello.len(), "extensions length mismatch"); + + let mut extensions = Vec::new(); + while pos + 4 <= ext_end { + let ext_type = u16::from_be_bytes([hello[pos], hello[pos + 1]]); + let data_len = u16::from_be_bytes([hello[pos + 2], hello[pos + 3]]) as usize; + pos += 4; + let data = hello[pos..pos + data_len].to_vec(); + pos += data_len; + extensions.push((ext_type, data)); + } + assert_eq!(pos, ext_end, "extension parse did not consume all bytes"); + + ParsedClientHelloForTest { + session_id, + extensions, + } +} + +fn parse_alpn_protocols(data: &[u8]) -> Vec> { + assert!(data.len() >= 2, "ALPN extension is too short"); + let protocols_len = u16::from_be_bytes([data[0], data[1]]) as usize; + assert_eq!(protocols_len + 2, data.len(), "ALPN list length mismatch"); + let mut pos = 2usize; + let mut out = Vec::new(); + while pos < data.len() { + let len = data[pos] as usize; + pos += 1; + out.push(data[pos..pos + len].to_vec()); + pos += len; + } + out +} + +async fn capture_rustls_client_hello_record(alpn_protocols: &'static [&'static [u8]]) -> Vec { + let (client, mut server) = tokio::io::duplex(32 * 1024); + let fetch_task = tokio::spawn(async move { + fetch_via_rustls_stream(client, "example.com", "example.com", None, alpn_protocols).await + }); + + let mut header = [0u8; 5]; + server + .read_exact(&mut header) + .await + .expect("must read client hello record header"); + let body_len = u16::from_be_bytes([header[3], header[4]]) as usize; + let mut body = vec![0u8; body_len]; + server + .read_exact(&mut body) + .await + .expect("must read client hello record body"); + drop(server); + + let result = fetch_task.await.expect("fetch task must join"); + assert!( + result.is_err(), + "capture task should end with handshake error" + ); + + let mut record = Vec::with_capacity(5 + body_len); + record.extend_from_slice(&header); + record.extend_from_slice(&body); + record +} + +#[test] +fn test_encode_tls13_certificate_message_single_cert() { + let cert = vec![0x30, 0x03, 0x02, 0x01, 0x01]; + let message = encode_tls13_certificate_message(std::slice::from_ref(&cert)).expect("message"); + + assert_eq!(message[0], 0x0b); + assert_eq!(read_u24(&message[1..4]), message.len() - 4); + assert_eq!(message[4], 0x00); + + let cert_list_len = read_u24(&message[5..8]); + assert_eq!(cert_list_len, cert.len() + 5); + + let cert_len = read_u24(&message[8..11]); + assert_eq!(cert_len, cert.len()); + assert_eq!(&message[11..11 + cert.len()], cert.as_slice()); + assert_eq!(&message[11 + cert.len()..13 + cert.len()], &[0x00, 0x00]); +} + +#[test] +fn test_encode_tls13_certificate_message_empty_chain() { + assert!(encode_tls13_certificate_message(&[]).is_none()); +} + +#[test] +fn test_derive_behavior_profile_splits_ticket_like_tail_records() { + let profile = derive_behavior_profile(&[ + (TLS_RECORD_HANDSHAKE, vec![0u8; 90]), + (TLS_RECORD_CHANGE_CIPHER, vec![0x01]), + (TLS_RECORD_APPLICATION, vec![0u8; 1400]), + (TLS_RECORD_APPLICATION, vec![0u8; 220]), + (TLS_RECORD_APPLICATION, vec![0u8; 180]), + ]); + + assert_eq!(profile.change_cipher_spec_count, 1); + assert_eq!(profile.app_data_record_sizes, vec![1400]); + assert_eq!(profile.ticket_record_sizes, vec![220, 180]); + assert_eq!(profile.source, TlsProfileSource::Raw); +} + +#[test] +fn test_order_profiles_prioritizes_fresh_cached_winner() { + let strategy = TlsFetchStrategy { + profiles: vec![ + TlsFetchProfile::ModernChromeLike, + TlsFetchProfile::CompatTls12, + TlsFetchProfile::LegacyMinimal, + ], + strict_route: true, + attempt_timeout: Duration::from_secs(1), + total_budget: Duration::from_secs(2), + grease_enabled: false, + deterministic: false, + profile_cache_ttl: Duration::from_secs(60), + }; + let cache_key = profile_cache_key( + "mask.example", + 443, + "tls.example", + None, + Some("tls"), + 0, + None, + ); + profile_cache().remove(&cache_key); + profile_cache().insert( + cache_key.clone(), + ProfileCacheValue { + profile: TlsFetchProfile::CompatTls12, + updated_at: Instant::now(), + }, + ); + + let ordered = order_profiles(&strategy, Some(&cache_key), Instant::now()); + assert_eq!(ordered[0], TlsFetchProfile::CompatTls12); + profile_cache().remove(&cache_key); +} + +#[test] +fn test_order_profiles_drops_expired_cached_winner() { + let strategy = TlsFetchStrategy { + profiles: vec![ + TlsFetchProfile::ModernFirefoxLike, + TlsFetchProfile::CompatTls12, + ], + strict_route: true, + attempt_timeout: Duration::from_secs(1), + total_budget: Duration::from_secs(2), + grease_enabled: false, + deterministic: false, + profile_cache_ttl: Duration::from_secs(5), + }; + let cache_key = profile_cache_key("mask2.example", 443, "tls2.example", None, None, 0, None); + profile_cache().remove(&cache_key); + profile_cache().insert( + cache_key.clone(), + ProfileCacheValue { + profile: TlsFetchProfile::CompatTls12, + updated_at: Instant::now() - Duration::from_secs(6), + }, + ); + + let ordered = order_profiles(&strategy, Some(&cache_key), Instant::now()); + assert_eq!(ordered[0], TlsFetchProfile::ModernFirefoxLike); + assert!(profile_cache().get(&cache_key).is_none()); +} + +#[test] +fn test_deterministic_client_hello_is_stable() { + let rng = SecureRandom::new(); + let first = build_client_hello( + "stable.example", + &rng, + TlsFetchProfile::ModernChromeLike, + true, + true, + ); + let second = build_client_hello( + "stable.example", + &rng, + TlsFetchProfile::ModernChromeLike, + true, + true, + ); + + assert_eq!(first, second); +} + +#[test] +fn test_raw_client_hello_alpn_matches_profile() { + let rng = SecureRandom::new(); + for profile in [ + TlsFetchProfile::ModernChromeLike, + TlsFetchProfile::ModernFirefoxLike, + TlsFetchProfile::CompatTls12, + TlsFetchProfile::LegacyMinimal, + ] { + let hello = build_client_hello("alpn.example", &rng, profile, false, true); + let parsed = parse_client_hello_for_test(&hello); + let alpn_ext = parsed + .extensions + .iter() + .find(|(ext_type, _)| *ext_type == 0x0010) + .expect("ALPN extension must exist"); + let parsed_alpn = parse_alpn_protocols(&alpn_ext.1); + let expected_alpn = profile_alpn(profile) + .iter() + .map(|proto| proto.to_vec()) + .collect::>(); + assert_eq!( + parsed_alpn, + expected_alpn, + "ALPN mismatch for {}", + profile.as_str() + ); + } +} + +#[test] +fn test_modern_chrome_like_browser_extension_layout() { + let rng = SecureRandom::new(); + let hello = build_client_hello( + "chrome.example", + &rng, + TlsFetchProfile::ModernChromeLike, + false, + true, + ); + let parsed = parse_client_hello_for_test(&hello); + assert_eq!( + parsed.session_id.len(), + 32, + "modern chrome must use non-empty session id" + ); + + let extension_ids = parsed + .extensions + .iter() + .map(|(ext_type, _)| *ext_type) + .collect::>(); + let expected_prefix = [ + 0x0000, 0x000b, 0x000a, 0x0023, 0x000d, 0x002b, 0x002d, 0x0033, 0x0010, + ]; + assert!( + extension_ids.as_slice().starts_with(&expected_prefix), + "unexpected extension order: {extension_ids:?}" + ); + assert!( + extension_ids.contains(&0x0015), + "modern chrome profile should include padding extension" + ); + + let key_share = parsed + .extensions + .iter() + .find(|(ext_type, _)| *ext_type == 0x0033) + .expect("key_share extension must exist"); + let key_share_data = &key_share.1; + assert!( + key_share_data.len() >= 2 + 4 + 32, + "key_share payload is too short" + ); + let entry_len = u16::from_be_bytes([key_share_data[0], key_share_data[1]]) as usize; + assert_eq!( + entry_len, + key_share_data.len() - 2, + "key_share list length mismatch" + ); + let mut pos = 2usize; + let hybrid_group = u16::from_be_bytes([key_share_data[pos], key_share_data[pos + 1]]); + let hybrid_len = + u16::from_be_bytes([key_share_data[pos + 2], key_share_data[pos + 3]]) as usize; + pos += 4; + let hybrid_key = &key_share_data[pos..pos + hybrid_len]; + pos += hybrid_len; + assert_eq!( + hybrid_group, TLS_NAMED_GROUP_X25519MLKEM768, + "first key_share group must be X25519MLKEM768" + ); + assert_eq!( + hybrid_len, + MLKEM768_CLIENT_ENCAPSULATION_KEY_LEN + X25519_KEY_SHARE_LEN, + "hybrid key length must match X25519MLKEM768" + ); + assert!( + hybrid_key.iter().any(|b| *b != 0), + "hybrid key must not be all zero" + ); + + let group = u16::from_be_bytes([key_share_data[pos], key_share_data[pos + 1]]); + let key_len = u16::from_be_bytes([key_share_data[pos + 2], key_share_data[pos + 3]]) as usize; + pos += 4; + let key = &key_share_data[pos..pos + key_len]; + assert_eq!( + group, TLS_NAMED_GROUP_X25519, + "second key_share group must be x25519" + ); + assert_eq!( + key_len, X25519_KEY_SHARE_LEN, + "x25519 key length must be 32" + ); + assert!( + key.iter().any(|b| *b != 0), + "x25519 key must not be all zero" + ); +} + +#[test] +fn test_fallback_profiles_keep_compat_extension_set() { + let rng = SecureRandom::new(); + for profile in [ + TlsFetchProfile::ModernFirefoxLike, + TlsFetchProfile::CompatTls12, + TlsFetchProfile::LegacyMinimal, + ] { + let hello = build_client_hello("fallback.example", &rng, profile, false, true); + let parsed = parse_client_hello_for_test(&hello); + let extension_ids = parsed + .extensions + .iter() + .map(|(ext_type, _)| *ext_type) + .collect::>(); + + assert!(extension_ids.contains(&0x0000), "SNI extension must exist"); + assert!( + extension_ids.contains(&0x000a), + "supported_groups extension must exist" + ); + assert!( + extension_ids.contains(&0x000d), + "signature_algorithms extension must exist" + ); + assert!( + extension_ids.contains(&0x002b), + "supported_versions extension must exist" + ); + assert!( + extension_ids.contains(&0x0033), + "key_share extension must exist" + ); + assert!(extension_ids.contains(&0x0010), "ALPN extension must exist"); + assert!( + !extension_ids.contains(&0x000b), + "ec_point_formats must stay chrome-only" + ); + assert!( + !extension_ids.contains(&0x0023), + "session_ticket must stay chrome-only" + ); + assert!( + !extension_ids.contains(&0x002d), + "psk_key_exchange_modes must stay chrome-only" + ); + + let expected_session_len = if matches!(profile, TlsFetchProfile::ModernFirefoxLike) { + 32 + } else { + 0 + }; + assert_eq!( + parsed.session_id.len(), + expected_session_len, + "unexpected session id length for {}", + profile.as_str() + ); + } +} + +#[tokio::test(flavor = "current_thread")] +async fn test_rustls_client_hello_alpn_matches_selected_profile() { + for profile in [ + TlsFetchProfile::ModernChromeLike, + TlsFetchProfile::CompatTls12, + TlsFetchProfile::LegacyMinimal, + ] { + let record = capture_rustls_client_hello_record(profile_alpn(profile)).await; + let parsed = parse_client_hello_for_test(&record); + let alpn_ext = parsed + .extensions + .iter() + .find(|(ext_type, _)| *ext_type == 0x0010) + .expect("ALPN extension must exist"); + let parsed_alpn = parse_alpn_protocols(&alpn_ext.1); + let expected_alpn = profile_alpn(profile) + .iter() + .map(|proto| proto.to_vec()) + .collect::>(); + assert_eq!( + parsed_alpn, + expected_alpn, + "rustls ALPN mismatch for {}", + profile.as_str() + ); + } +} + +#[test] +fn test_build_tls_fetch_proxy_header_v2_with_tcp_addrs() { + let src: SocketAddr = "198.51.100.10:42000".parse().expect("valid src"); + let dst: SocketAddr = "203.0.113.20:443".parse().expect("valid dst"); + let header = build_tls_fetch_proxy_header(2, Some(src), Some(dst)).expect("header"); + + assert_eq!( + &header[..12], + &[ + 0x0d, 0x0a, 0x0d, 0x0a, 0x00, 0x0d, 0x0a, 0x51, 0x55, 0x49, 0x54, 0x0a + ] + ); + assert_eq!(header[12], 0x21); + assert_eq!(header[13], 0x11); + assert_eq!(u16::from_be_bytes([header[14], header[15]]), 12); + assert_eq!(&header[16..20], &[198, 51, 100, 10]); + assert_eq!(&header[20..24], &[203, 0, 113, 20]); + assert_eq!(u16::from_be_bytes([header[24], header[25]]), 42000); + assert_eq!(u16::from_be_bytes([header[26], header[27]]), 443); +} + +#[test] +fn test_build_tls_fetch_proxy_header_v2_mixed_family_falls_back_to_local_command() { + let src: SocketAddr = "198.51.100.10:42000".parse().expect("valid src"); + let dst: SocketAddr = "[2001:db8::20]:443".parse().expect("valid dst"); + let header = build_tls_fetch_proxy_header(2, Some(src), Some(dst)).expect("header"); + + assert_eq!(header[12], 0x20); + assert_eq!(header[13], 0x00); + assert_eq!(u16::from_be_bytes([header[14], header[15]]), 0); +} + +#[test] +fn test_build_tls_fetch_proxy_header_v1_with_tcp_addrs() { + let src: SocketAddr = "198.51.100.10:42000".parse().expect("valid src"); + let dst: SocketAddr = "203.0.113.20:443".parse().expect("valid dst"); + let header = build_tls_fetch_proxy_header(1, Some(src), Some(dst)).expect("header"); + + assert_eq!( + header, + b"PROXY TCP4 198.51.100.10 203.0.113.20 42000 443\r\n" + ); +} diff --git a/src/transport/middle_proxy/config_updater.rs b/src/transport/middle_proxy/config_updater.rs index 070e020..50278b5 100644 --- a/src/transport/middle_proxy/config_updater.rs +++ b/src/transport/middle_proxy/config_updater.rs @@ -279,202 +279,9 @@ pub async fn fetch_proxy_config_via_upstream( .map(|(parsed, _raw)| parsed) } -fn snapshot_passes_guards( - cfg: &ProxyConfig, - snapshot: &ProxyConfigData, - snapshot_name: &'static str, -) -> bool { - if cfg.general.me_snapshot_require_http_2xx && !(200..=299).contains(&snapshot.http_status) { - warn!( - snapshot = snapshot_name, - http_status = snapshot.http_status, - "ME snapshot rejected by non-2xx HTTP status" - ); - return false; - } - - let min_proxy_for = cfg.general.me_snapshot_min_proxy_for_lines; - if snapshot.proxy_for_lines < min_proxy_for { - warn!( - snapshot = snapshot_name, - parsed_proxy_for_lines = snapshot.proxy_for_lines, - min_proxy_for_lines = min_proxy_for, - "ME snapshot rejected by proxy_for line floor" - ); - return false; - } - - true -} - -async fn run_update_cycle( - pool: &Arc, - cfg: &ProxyConfig, - state: &mut UpdaterState, - reinit_tx: &mpsc::Sender, -) { - let upstream = pool.upstream.clone(); - - let required_cfg_snapshots = cfg.general.me_config_stable_snapshots.max(1); - let required_secret_snapshots = cfg.general.proxy_secret_stable_snapshots.max(1); - let apply_cooldown = Duration::from_secs(cfg.general.me_config_apply_cooldown_secs); - let mut maps_changed = false; - - let mut ready_v4: Option<(ProxyConfigData, u64)> = None; - let cfg_v4 = retry_fetch( - cfg.general - .proxy_config_v4_url - .as_deref() - .unwrap_or("https://core.telegram.org/getProxyConfig"), - upstream.clone(), - ) - .await; - if let Some(cfg_v4) = cfg_v4 - && snapshot_passes_guards(cfg, &cfg_v4, "getProxyConfig") - { - let cfg_v4_hash = hash_proxy_config(&cfg_v4); - let stable_hits = state.config_v4.observe(cfg_v4_hash); - if stable_hits < required_cfg_snapshots { - debug!( - stable_hits, - required_cfg_snapshots, - snapshot = format_args!("0x{cfg_v4_hash:016x}"), - "ME config v4 candidate observed" - ); - } else if state.config_v4.is_applied(cfg_v4_hash) { - debug!( - snapshot = format_args!("0x{cfg_v4_hash:016x}"), - "ME config v4 stable snapshot already applied" - ); - } else { - ready_v4 = Some((cfg_v4, cfg_v4_hash)); - } - } - - let mut ready_v6: Option<(ProxyConfigData, u64)> = None; - let cfg_v6 = retry_fetch( - cfg.general - .proxy_config_v6_url - .as_deref() - .unwrap_or("https://core.telegram.org/getProxyConfigV6"), - upstream.clone(), - ) - .await; - if let Some(cfg_v6) = cfg_v6 - && snapshot_passes_guards(cfg, &cfg_v6, "getProxyConfigV6") - { - let cfg_v6_hash = hash_proxy_config(&cfg_v6); - let stable_hits = state.config_v6.observe(cfg_v6_hash); - if stable_hits < required_cfg_snapshots { - debug!( - stable_hits, - required_cfg_snapshots, - snapshot = format_args!("0x{cfg_v6_hash:016x}"), - "ME config v6 candidate observed" - ); - } else if state.config_v6.is_applied(cfg_v6_hash) { - debug!( - snapshot = format_args!("0x{cfg_v6_hash:016x}"), - "ME config v6 stable snapshot already applied" - ); - } else { - ready_v6 = Some((cfg_v6, cfg_v6_hash)); - } - } - - if ready_v4.is_some() || ready_v6.is_some() { - if map_apply_cooldown_ready(state.last_map_apply_at, apply_cooldown) { - let update_v4 = ready_v4 - .as_ref() - .map(|(snapshot, _)| snapshot.map.clone()) - .unwrap_or_default(); - let update_v6 = ready_v6.as_ref().map(|(snapshot, _)| snapshot.map.clone()); - let update_is_empty = - update_v4.is_empty() && update_v6.as_ref().is_none_or(|v| v.is_empty()); - let apply_outcome = if update_is_empty && !cfg.general.me_snapshot_reject_empty_map { - super::pool_config::SnapshotApplyOutcome::AppliedNoDelta - } else { - pool.update_proxy_maps(update_v4, update_v6).await - }; - - if matches!( - apply_outcome, - super::pool_config::SnapshotApplyOutcome::RejectedEmpty - ) { - warn!("ME config stable snapshot rejected (empty endpoint map)"); - } else { - if let Some((snapshot, hash)) = ready_v4 { - if let Some(dc) = snapshot.default_dc { - pool.default_dc - .store(dc, std::sync::atomic::Ordering::Relaxed); - } - state.config_v4.mark_applied(hash); - } - - if let Some((_snapshot, hash)) = ready_v6 { - state.config_v6.mark_applied(hash); - } - - state.last_map_apply_at = Some(tokio::time::Instant::now()); - - if apply_outcome.changed() { - maps_changed = true; - info!("ME config update applied after stable-gate"); - } else { - debug!("ME config stable-gate applied with no map delta"); - } - } - } else if let Some(last) = state.last_map_apply_at { - let wait_secs = map_apply_cooldown_remaining_secs(last, apply_cooldown); - debug!(wait_secs, "ME config stable snapshot deferred by cooldown"); - } - } - - if maps_changed { - enqueue_reinit_trigger(reinit_tx, MeReinitTrigger::MapChanged); - } - - pool.reset_stun_state(); - - if cfg.general.proxy_secret_rotate_runtime { - match download_proxy_secret_with_max_len_via_upstream( - cfg.general.proxy_secret_len_max, - upstream, - cfg.general.proxy_secret_url.as_deref(), - ) - .await - { - Ok(secret) => { - let secret_hash = hash_secret(&secret); - let stable_hits = state.secret.observe(secret_hash); - if stable_hits < required_secret_snapshots { - debug!( - stable_hits, - required_secret_snapshots, - snapshot = format_args!("0x{secret_hash:016x}"), - "proxy-secret candidate observed" - ); - } else if state.secret.is_applied(secret_hash) { - debug!( - snapshot = format_args!("0x{secret_hash:016x}"), - "proxy-secret stable snapshot already applied" - ); - } else { - let rotated = pool.update_secret(secret).await; - state.secret.mark_applied(secret_hash); - if rotated { - info!("proxy-secret rotated after stable-gate"); - } else { - debug!("proxy-secret stable snapshot confirmed as unchanged"); - } - } - } - Err(e) => warn!(error = %e, "proxy-secret update failed"), - } - } else { - debug!("proxy-secret runtime rotation disabled by config"); - } -} +// Stable-snapshot guards and one bounded update cycle. +mod cycle; +use cycle::run_update_cycle; pub async fn me_config_updater( pool: Arc, diff --git a/src/transport/middle_proxy/config_updater/cycle.rs b/src/transport/middle_proxy/config_updater/cycle.rs new file mode 100644 index 0000000..1f60ddf --- /dev/null +++ b/src/transport/middle_proxy/config_updater/cycle.rs @@ -0,0 +1,198 @@ +use super::*; + +pub(super) fn snapshot_passes_guards( + cfg: &ProxyConfig, + snapshot: &ProxyConfigData, + snapshot_name: &'static str, +) -> bool { + if cfg.general.me_snapshot_require_http_2xx && !(200..=299).contains(&snapshot.http_status) { + warn!( + snapshot = snapshot_name, + http_status = snapshot.http_status, + "ME snapshot rejected by non-2xx HTTP status" + ); + return false; + } + + let min_proxy_for = cfg.general.me_snapshot_min_proxy_for_lines; + if snapshot.proxy_for_lines < min_proxy_for { + warn!( + snapshot = snapshot_name, + parsed_proxy_for_lines = snapshot.proxy_for_lines, + min_proxy_for_lines = min_proxy_for, + "ME snapshot rejected by proxy_for line floor" + ); + return false; + } + + true +} + +pub(super) async fn run_update_cycle( + pool: &Arc, + cfg: &ProxyConfig, + state: &mut UpdaterState, + reinit_tx: &mpsc::Sender, +) { + let upstream = pool.upstream.clone(); + + let required_cfg_snapshots = cfg.general.me_config_stable_snapshots.max(1); + let required_secret_snapshots = cfg.general.proxy_secret_stable_snapshots.max(1); + let apply_cooldown = Duration::from_secs(cfg.general.me_config_apply_cooldown_secs); + let mut maps_changed = false; + + let mut ready_v4: Option<(ProxyConfigData, u64)> = None; + let cfg_v4 = retry_fetch( + cfg.general + .proxy_config_v4_url + .as_deref() + .unwrap_or("https://core.telegram.org/getProxyConfig"), + upstream.clone(), + ) + .await; + if let Some(cfg_v4) = cfg_v4 + && snapshot_passes_guards(cfg, &cfg_v4, "getProxyConfig") + { + let cfg_v4_hash = hash_proxy_config(&cfg_v4); + let stable_hits = state.config_v4.observe(cfg_v4_hash); + if stable_hits < required_cfg_snapshots { + debug!( + stable_hits, + required_cfg_snapshots, + snapshot = format_args!("0x{cfg_v4_hash:016x}"), + "ME config v4 candidate observed" + ); + } else if state.config_v4.is_applied(cfg_v4_hash) { + debug!( + snapshot = format_args!("0x{cfg_v4_hash:016x}"), + "ME config v4 stable snapshot already applied" + ); + } else { + ready_v4 = Some((cfg_v4, cfg_v4_hash)); + } + } + + let mut ready_v6: Option<(ProxyConfigData, u64)> = None; + let cfg_v6 = retry_fetch( + cfg.general + .proxy_config_v6_url + .as_deref() + .unwrap_or("https://core.telegram.org/getProxyConfigV6"), + upstream.clone(), + ) + .await; + if let Some(cfg_v6) = cfg_v6 + && snapshot_passes_guards(cfg, &cfg_v6, "getProxyConfigV6") + { + let cfg_v6_hash = hash_proxy_config(&cfg_v6); + let stable_hits = state.config_v6.observe(cfg_v6_hash); + if stable_hits < required_cfg_snapshots { + debug!( + stable_hits, + required_cfg_snapshots, + snapshot = format_args!("0x{cfg_v6_hash:016x}"), + "ME config v6 candidate observed" + ); + } else if state.config_v6.is_applied(cfg_v6_hash) { + debug!( + snapshot = format_args!("0x{cfg_v6_hash:016x}"), + "ME config v6 stable snapshot already applied" + ); + } else { + ready_v6 = Some((cfg_v6, cfg_v6_hash)); + } + } + + if ready_v4.is_some() || ready_v6.is_some() { + if map_apply_cooldown_ready(state.last_map_apply_at, apply_cooldown) { + let update_v4 = ready_v4 + .as_ref() + .map(|(snapshot, _)| snapshot.map.clone()) + .unwrap_or_default(); + let update_v6 = ready_v6.as_ref().map(|(snapshot, _)| snapshot.map.clone()); + let update_is_empty = + update_v4.is_empty() && update_v6.as_ref().is_none_or(|v| v.is_empty()); + let apply_outcome = if update_is_empty && !cfg.general.me_snapshot_reject_empty_map { + crate::transport::middle_proxy::pool_config::SnapshotApplyOutcome::AppliedNoDelta + } else { + pool.update_proxy_maps(update_v4, update_v6).await + }; + + if matches!( + apply_outcome, + crate::transport::middle_proxy::pool_config::SnapshotApplyOutcome::RejectedEmpty + ) { + warn!("ME config stable snapshot rejected (empty endpoint map)"); + } else { + if let Some((snapshot, hash)) = ready_v4 { + if let Some(dc) = snapshot.default_dc { + pool.default_dc + .store(dc, std::sync::atomic::Ordering::Relaxed); + } + state.config_v4.mark_applied(hash); + } + + if let Some((_snapshot, hash)) = ready_v6 { + state.config_v6.mark_applied(hash); + } + + state.last_map_apply_at = Some(tokio::time::Instant::now()); + + if apply_outcome.changed() { + maps_changed = true; + info!("ME config update applied after stable-gate"); + } else { + debug!("ME config stable-gate applied with no map delta"); + } + } + } else if let Some(last) = state.last_map_apply_at { + let wait_secs = map_apply_cooldown_remaining_secs(last, apply_cooldown); + debug!(wait_secs, "ME config stable snapshot deferred by cooldown"); + } + } + + if maps_changed { + enqueue_reinit_trigger(reinit_tx, MeReinitTrigger::MapChanged); + } + + pool.reset_stun_state(); + + if cfg.general.proxy_secret_rotate_runtime { + match download_proxy_secret_with_max_len_via_upstream( + cfg.general.proxy_secret_len_max, + upstream, + cfg.general.proxy_secret_url.as_deref(), + ) + .await + { + Ok(secret) => { + let secret_hash = hash_secret(&secret); + let stable_hits = state.secret.observe(secret_hash); + if stable_hits < required_secret_snapshots { + debug!( + stable_hits, + required_secret_snapshots, + snapshot = format_args!("0x{secret_hash:016x}"), + "proxy-secret candidate observed" + ); + } else if state.secret.is_applied(secret_hash) { + debug!( + snapshot = format_args!("0x{secret_hash:016x}"), + "proxy-secret stable snapshot already applied" + ); + } else { + let rotated = pool.update_secret(secret).await; + state.secret.mark_applied(secret_hash); + if rotated { + info!("proxy-secret rotated after stable-gate"); + } else { + debug!("proxy-secret stable snapshot confirmed as unchanged"); + } + } + } + Err(e) => warn!(error = %e, "proxy-secret update failed"), + } + } else { + debug!("proxy-secret runtime rotation disabled by config"); + } +} diff --git a/src/transport/middle_proxy/health.rs b/src/transport/middle_proxy/health.rs index d40430c..f51b7f5 100644 --- a/src/transport/middle_proxy/health.rs +++ b/src/transport/middle_proxy/health.rs @@ -90,8 +90,7 @@ impl ScheduledReconnects<'_> { impl Drop for ScheduledReconnects<'_> { fn drop(&mut self) { for key in self.keys.drain(..) { - let std::collections::hash_map::Entry::Occupied(mut entry) = - self.inflight.entry(key) + let std::collections::hash_map::Entry::Occupied(mut entry) = self.inflight.entry(key) else { continue; }; @@ -105,1949 +104,36 @@ impl Drop for ScheduledReconnects<'_> { } } -pub async fn me_health_monitor(pool: Arc, rng: Arc, _min_connections: usize) { - let mut backoff: HashMap<(i32, IpFamily), u64> = HashMap::new(); - let mut next_attempt: HashMap<(i32, IpFamily), Instant> = HashMap::new(); - let mut inflight: HashMap<(i32, IpFamily), usize> = HashMap::new(); - let mut outage_backoff: HashMap<(i32, IpFamily), u64> = HashMap::new(); - let mut outage_next_attempt: HashMap<(i32, IpFamily), Instant> = HashMap::new(); - let mut single_endpoint_outage: HashSet<(i32, IpFamily)> = HashSet::new(); - let mut shadow_rotate_deadline: HashMap<(i32, IpFamily), Instant> = HashMap::new(); - let mut idle_refresh_next_attempt: HashMap<(i32, IpFamily), Instant> = HashMap::new(); - let mut floor_warn_next_allowed: HashMap<(i32, IpFamily), Instant> = HashMap::new(); - let mut drain_warn_next_allowed: HashMap = HashMap::new(); - let mut degraded_interval = true; - loop { - let interval = if degraded_interval { - pool.health_interval_unhealthy() - } else { - pool.health_interval_healthy() - }; - tokio::time::sleep(interval).await; - pool.prune_closed_writers().await; - pool.sweep_endpoint_quarantine().await; - reap_draining_writers(&pool, &mut drain_warn_next_allowed).await; - let v4_degraded = check_family( - IpFamily::V4, - &pool, - &rng, - &mut backoff, - &mut next_attempt, - &mut inflight, - &mut outage_backoff, - &mut outage_next_attempt, - &mut single_endpoint_outage, - &mut shadow_rotate_deadline, - &mut idle_refresh_next_attempt, - &mut floor_warn_next_allowed, - ) - .await; - let v6_degraded = check_family( - IpFamily::V6, - &pool, - &rng, - &mut backoff, - &mut next_attempt, - &mut inflight, - &mut outage_backoff, - &mut outage_next_attempt, - &mut single_endpoint_outage, - &mut shadow_rotate_deadline, - &mut idle_refresh_next_attempt, - &mut floor_warn_next_allowed, - ) - .await; - update_family_runtime_state(&pool, IpFamily::V4, v4_degraded); - update_family_runtime_state(&pool, IpFamily::V6, v6_degraded); - degraded_interval = v4_degraded || v6_degraded; - } -} +// Periodic family health monitor orchestration. +mod monitor; +// Draining-writer deadline enforcement. +mod drain; +// Per-family reconnect scheduling and health state. +mod family; +// Adaptive-floor planning and warning rate limits. +mod floor_plan; +// Idle-writer refresh and cap-aware swaps. +mod idle_refresh; +// Single-endpoint recovery and shadow rotation. +mod recovery; +// Independent stale draining-writer watchdog. +mod zombie_watchdog; + +pub use drain::me_drain_timeout_enforcer; +pub(in crate::transport::middle_proxy) use drain::{ + health_drain_close_budget, reap_draining_writers, +}; +use family::{check_family, update_family_runtime_state}; +use floor_plan::{ + build_family_floor_plan, live_active_writers_for_dc_family, should_emit_rate_limited_warn, +}; +use idle_refresh::{maybe_refresh_idle_writer_for_dc, maybe_swap_idle_writer_for_cap}; +pub use monitor::me_health_monitor; +use recovery::{ + has_bound_clients_on_endpoint, maybe_rotate_single_endpoint_shadow, + recover_single_endpoint_outage, +}; +pub use zombie_watchdog::me_zombie_writer_watchdog; -pub async fn me_drain_timeout_enforcer(pool: Arc) { - let mut drain_warn_next_allowed: HashMap = HashMap::new(); - loop { - tokio::time::sleep(Duration::from_secs( - HEALTH_DRAIN_TIMEOUT_ENFORCER_INTERVAL_SECS, - )) - .await; - reap_draining_writers(&pool, &mut drain_warn_next_allowed).await; - } -} - -pub(super) async fn reap_draining_writers( - pool: &Arc, - warn_next_allowed: &mut HashMap, -) { - let now_epoch_secs = MePool::now_epoch_secs(); - let now = Instant::now(); - let drain_ttl_secs = pool - .drain_runtime - .me_pool_drain_ttl_secs - .load(std::sync::atomic::Ordering::Relaxed); - let drain_threshold = pool - .drain_runtime - .me_pool_drain_threshold - .load(std::sync::atomic::Ordering::Relaxed); - let activity = pool.registry.writer_activity_snapshot().await; - let mut draining_writers = Vec::::new(); - let mut empty_writer_ids = Vec::::new(); - let mut force_close_writer_ids = Vec::::new(); - let writers = pool.writers.read().await; - for writer in writers.iter() { - if !writer.draining.load(std::sync::atomic::Ordering::Relaxed) { - continue; - } - if activity - .bound_clients_by_writer - .get(&writer.id) - .copied() - .unwrap_or(0) - == 0 - { - empty_writer_ids.push(writer.id); - continue; - } - draining_writers.push(DrainingWriterSnapshot { - id: writer.id, - writer_dc: writer.writer_dc, - addr: writer.addr, - generation: writer.generation, - created_at: writer.created_at, - draining_started_at_epoch_secs: writer - .draining_started_at_epoch_secs - .load(std::sync::atomic::Ordering::Relaxed), - drain_deadline_epoch_secs: writer - .drain_deadline_epoch_secs - .load(std::sync::atomic::Ordering::Relaxed), - allow_drain_fallback: writer - .allow_drain_fallback - .load(std::sync::atomic::Ordering::Relaxed), - }); - } - drop(writers); - - let overflow = if drain_threshold > 0 && draining_writers.len() > drain_threshold as usize { - draining_writers - .len() - .saturating_sub(drain_threshold as usize) - } else { - 0 - }; - - if overflow > 0 { - draining_writers.sort_by(|left, right| { - left.draining_started_at_epoch_secs - .cmp(&right.draining_started_at_epoch_secs) - .then_with(|| left.created_at.cmp(&right.created_at)) - .then_with(|| left.id.cmp(&right.id)) - }); - warn!( - draining_writers = draining_writers.len(), - me_pool_drain_threshold = drain_threshold, - removing_writers = overflow, - "ME draining writer threshold exceeded, force-closing oldest draining writers" - ); - for writer in draining_writers.drain(..overflow) { - force_close_writer_ids.push(writer.id); - } - } - - for writer in draining_writers { - if drain_ttl_secs > 0 - && writer.draining_started_at_epoch_secs != 0 - && now_epoch_secs.saturating_sub(writer.draining_started_at_epoch_secs) > drain_ttl_secs - && should_emit_writer_warn( - warn_next_allowed, - writer.id, - now, - pool.warn_rate_limit_duration(), - ) - { - warn!( - writer_id = writer.id, - writer_dc = writer.writer_dc, - endpoint = %writer.addr, - generation = writer.generation, - drain_ttl_secs, - force_close_secs = pool - .drain_runtime - .me_pool_force_close_secs - .load(std::sync::atomic::Ordering::Relaxed), - allow_drain_fallback = writer.allow_drain_fallback, - "ME draining writer remains non-empty past drain TTL" - ); - } - if writer.drain_deadline_epoch_secs != 0 - && now_epoch_secs >= writer.drain_deadline_epoch_secs - { - warn!(writer_id = writer.id, "Drain timeout, force-closing"); - force_close_writer_ids.push(writer.id); - } - } - - let close_budget = health_drain_close_budget(); - let requested_force_close = force_close_writer_ids.len(); - let requested_empty_close = empty_writer_ids.len(); - let requested_close_total = requested_force_close.saturating_add(requested_empty_close); - let mut closed_writer_ids = HashSet::::new(); - let mut closed_total = 0usize; - for writer_id in force_close_writer_ids { - if closed_total >= close_budget { - break; - } - if !closed_writer_ids.insert(writer_id) { - continue; - } - pool.stats.increment_pool_force_close_total(); - pool.remove_writer_and_close_clients(writer_id).await; - closed_total = closed_total.saturating_add(1); - } - for writer_id in empty_writer_ids { - if closed_total >= close_budget { - break; - } - if !closed_writer_ids.insert(writer_id) { - continue; - } - pool.remove_writer_and_close_clients(writer_id).await; - closed_total = closed_total.saturating_add(1); - } - - let pending_close_total = requested_close_total.saturating_sub(closed_total); - if pending_close_total > 0 { - warn!( - close_budget, - closed_total, - pending_close_total, - "ME draining close backlog deferred to next health cycle" - ); - } - - // Keep warn cooldown state for draining writers still present in the pool; - // drop state only once a writer is actually removed. - let active_draining_writer_ids = { - let writers = pool.writers.read().await; - writers - .iter() - .filter(|writer| writer.draining.load(std::sync::atomic::Ordering::Relaxed)) - .map(|writer| writer.id) - .collect::>() - }; - warn_next_allowed.retain(|writer_id, _| active_draining_writer_ids.contains(writer_id)); -} - -pub(super) fn health_drain_close_budget() -> usize { - let cpu_cores = std::thread::available_parallelism() - .map(std::num::NonZeroUsize::get) - .unwrap_or(1); - cpu_cores - .saturating_mul(HEALTH_DRAIN_CLOSE_BUDGET_PER_CORE) - .clamp(HEALTH_DRAIN_CLOSE_BUDGET_MIN, HEALTH_DRAIN_CLOSE_BUDGET_MAX) -} - -#[derive(Debug, Clone)] -struct DrainingWriterSnapshot { - id: u64, - writer_dc: i32, - addr: SocketAddr, - generation: u64, - created_at: Instant, - draining_started_at_epoch_secs: u64, - drain_deadline_epoch_secs: u64, - allow_drain_fallback: bool, -} - -fn should_emit_writer_warn( - next_allowed: &mut HashMap, - writer_id: u64, - now: Instant, - cooldown: Duration, -) -> bool { - let Some(ready_at) = next_allowed.get(&writer_id).copied() else { - next_allowed.insert(writer_id, now + cooldown); - return true; - }; - if now >= ready_at { - next_allowed.insert(writer_id, now + cooldown); - return true; - } - false -} - -async fn check_family( - family: IpFamily, - pool: &Arc, - rng: &Arc, - backoff: &mut HashMap<(i32, IpFamily), u64>, - next_attempt: &mut HashMap<(i32, IpFamily), Instant>, - inflight: &mut HashMap<(i32, IpFamily), usize>, - outage_backoff: &mut HashMap<(i32, IpFamily), u64>, - outage_next_attempt: &mut HashMap<(i32, IpFamily), Instant>, - single_endpoint_outage: &mut HashSet<(i32, IpFamily)>, - shadow_rotate_deadline: &mut HashMap<(i32, IpFamily), Instant>, - idle_refresh_next_attempt: &mut HashMap<(i32, IpFamily), Instant>, - floor_warn_next_allowed: &mut HashMap<(i32, IpFamily), Instant>, -) -> bool { - let enabled = match family { - IpFamily::V4 => pool.decision.ipv4_me, - IpFamily::V6 => pool.decision.ipv6_me, - }; - if !enabled { - return false; - } - - let mut family_degraded = false; - - let mut dc_endpoints = HashMap::>::new(); - let map_guard = match family { - IpFamily::V4 => pool.proxy_map_v4.read().await, - IpFamily::V6 => pool.proxy_map_v6.read().await, - }; - for (dc, addrs) in map_guard.iter() { - let entry = dc_endpoints.entry(*dc).or_default(); - for (ip, port) in addrs.iter().copied() { - entry.push(SocketAddr::new(ip, port)); - } - } - drop(map_guard); - for endpoints in dc_endpoints.values_mut() { - endpoints.sort_unstable(); - endpoints.dedup(); - } - let reconnect_budget = health_reconnect_budget(pool, dc_endpoints.len()); - let reconnect_sem = Arc::new(Semaphore::new(reconnect_budget)); - - if pool.floor_mode() == MeFloorMode::Static {} - - let mut live_addr_counts = HashMap::<(i32, SocketAddr), usize>::new(); - let mut live_writer_ids_by_addr = HashMap::<(i32, SocketAddr), Vec>::new(); - for writer in pool - .writers - .read() - .await - .iter() - .filter(|w| !w.draining.load(std::sync::atomic::Ordering::Relaxed)) - { - if !matches!( - super::pool::WriterContour::from_u8( - writer.contour.load(std::sync::atomic::Ordering::Relaxed), - ), - super::pool::WriterContour::Active - ) { - continue; - } - let key = (writer.writer_dc, writer.addr); - *live_addr_counts.entry(key).or_insert(0) += 1; - live_writer_ids_by_addr - .entry(key) - .or_default() - .push(writer.id); - } - let writer_idle_since = pool.registry.writer_idle_since_snapshot().await; - let bound_clients_by_writer = pool - .registry - .writer_activity_snapshot() - .await - .bound_clients_by_writer; - let floor_plan = build_family_floor_plan( - pool, - family, - &dc_endpoints, - &live_addr_counts, - &live_writer_ids_by_addr, - &bound_clients_by_writer, - ) - .await; - pool.set_adaptive_floor_runtime_caps( - floor_plan.active_cap_configured_total, - floor_plan.active_cap_effective_total, - floor_plan.warm_cap_configured_total, - floor_plan.warm_cap_effective_total, - floor_plan.target_writers_total, - floor_plan.active_writers_current, - floor_plan.warm_writers_current, - ); - let live_writer_ids_by_addr = Arc::new(live_writer_ids_by_addr); - let writer_idle_since = Arc::new(writer_idle_since); - let bound_clients_by_writer = Arc::new(bound_clients_by_writer); - let mut reconnect_set = JoinSet::::new(); - let mut scheduled_reconnects = ScheduledReconnects { - inflight, - keys: Vec::new(), - }; - - for (dc, endpoints) in dc_endpoints { - if endpoints.is_empty() { - continue; - } - let key = (dc, family); - let required = floor_plan - .by_dc - .get(&dc) - .map(|entry| entry.target_required) - .unwrap_or_else(|| { - pool.required_writers_for_dc_with_floor_mode(endpoints.len(), false) - }); - let alive = endpoints - .iter() - .map(|addr| *live_addr_counts.get(&(dc, *addr)).unwrap_or(&0)) - .sum::(); - - if endpoints.len() == 1 && pool.single_endpoint_outage_mode_enabled() && alive == 0 { - family_degraded = true; - if single_endpoint_outage.insert(key) { - pool.stats.increment_me_single_endpoint_outage_enter_total(); - warn!( - dc = %dc, - ?family, - required, - endpoint_count = endpoints.len(), - "Single-endpoint DC outage detected" - ); - } - - recover_single_endpoint_outage( - pool, - rng, - key, - endpoints[0], - required, - outage_backoff, - outage_next_attempt, - &reconnect_sem, - ) - .await; - continue; - } - - if single_endpoint_outage.remove(&key) { - pool.stats.increment_me_single_endpoint_outage_exit_total(); - outage_backoff.remove(&key); - outage_next_attempt.remove(&key); - shadow_rotate_deadline.remove(&key); - idle_refresh_next_attempt.remove(&key); - info!( - dc = %dc, - ?family, - alive, - required, - endpoint_count = endpoints.len(), - "Single-endpoint DC outage recovered" - ); - } - - if alive >= required { - maybe_refresh_idle_writer_for_dc( - pool, - rng, - key, - dc, - family, - &endpoints, - alive, - required, - live_writer_ids_by_addr.as_ref(), - writer_idle_since.as_ref(), - bound_clients_by_writer.as_ref(), - idle_refresh_next_attempt, - ) - .await; - maybe_rotate_single_endpoint_shadow( - pool, - rng, - key, - dc, - family, - &endpoints, - alive, - required, - live_writer_ids_by_addr.as_ref(), - bound_clients_by_writer.as_ref(), - shadow_rotate_deadline, - ) - .await; - continue; - } - let missing = required - alive; - family_degraded = true; - - let now = Instant::now(); - if reconnect_sem.available_permits() == 0 { - let base_ms = pool.reconnect_runtime.me_reconnect_backoff_base.as_millis() as u64; - let next_ms = (*backoff.get(&key).unwrap_or(&base_ms)).max(base_ms); - let jitter = next_ms / JITTER_FRAC_NUM; - let wait = Duration::from_millis(next_ms) - + Duration::from_millis(rand::rng().random_range(0..=jitter.max(1))); - next_attempt.insert(key, now + wait); - debug!( - dc = %dc, - ?family, - alive, - required, - endpoint_count = endpoints.len(), - reconnect_budget, - "Skipping reconnect due to per-tick health reconnect budget" - ); - continue; - } - if let Some(ts) = next_attempt.get(&key) - && now < *ts - { - continue; - } - - let max_concurrent = pool - .reconnect_runtime - .me_reconnect_max_concurrent_per_dc - .max(1) as usize; - if scheduled_reconnects.current(&key) >= max_concurrent { - continue; - } - if pool - .has_refill_inflight_for_dc_key(super::pool::RefillDcKey { dc, family }) - .await - { - debug!( - dc = %dc, - ?family, - alive, - required, - endpoint_count = endpoints.len(), - "Skipping health reconnect: immediate refill is already in flight for this DC group" - ); - continue; - } - scheduled_reconnects.reserve(key); - let pool_for_reconnect = pool.clone(); - let rng_for_reconnect = rng.clone(); - let reconnect_sem_for_dc = reconnect_sem.clone(); - let endpoints_for_dc = endpoints.clone(); - let live_writer_ids_by_addr_for_dc = live_writer_ids_by_addr.clone(); - let writer_idle_since_for_dc = writer_idle_since.clone(); - let bound_clients_by_writer_for_dc = bound_clients_by_writer.clone(); - let active_cap_effective_total = floor_plan.active_cap_effective_total; - reconnect_set.spawn(async move { - let mut restored = 0usize; - for _ in 0..missing { - let Ok(reconnect_permit) = reconnect_sem_for_dc.clone().try_acquire_owned() else { - break; - }; - if pool_for_reconnect.active_contour_writer_count_total().await - >= active_cap_effective_total - { - let swapped = maybe_swap_idle_writer_for_cap( - &pool_for_reconnect, - &rng_for_reconnect, - dc, - family, - &endpoints_for_dc, - live_writer_ids_by_addr_for_dc.as_ref(), - writer_idle_since_for_dc.as_ref(), - bound_clients_by_writer_for_dc.as_ref(), - ) - .await; - if swapped { - pool_for_reconnect - .stats - .increment_me_floor_swap_idle_total(); - restored += 1; - continue; - } - - let base_req = pool_for_reconnect - .required_writers_for_dc_with_floor_mode(endpoints_for_dc.len(), false); - if alive + restored >= base_req { - pool_for_reconnect - .stats - .increment_me_floor_cap_block_total(); - pool_for_reconnect - .stats - .increment_me_floor_swap_idle_failed_total(); - debug!( - dc = %dc, - ?family, - alive, - required, - active_cap_effective_total, - "Adaptive floor cap reached, reconnect attempt blocked" - ); - break; - } - } - pool_for_reconnect.stats.increment_me_reconnect_attempt(); - let res = tokio::time::timeout( - pool_for_reconnect.reconnect_runtime.me_one_timeout, - pool_for_reconnect.connect_endpoints_round_robin( - dc, - &endpoints_for_dc, - rng_for_reconnect.as_ref(), - ), - ) - .await; - match res { - Ok(true) => { - restored += 1; - pool_for_reconnect.stats.increment_me_reconnect_success(); - } - Ok(false) => { - debug!(dc = %dc, ?family, "ME round-robin reconnect failed") - } - Err(_) => { - debug!(dc = %dc, ?family, "ME reconnect timed out"); - } - } - drop(reconnect_permit); - } - - FamilyReconnectOutcome { - key, - dc, - family, - required, - endpoint_count: endpoints_for_dc.len(), - } - }); - } - - while let Some(joined) = reconnect_set.join_next().await { - let outcome = match joined { - Ok(outcome) => outcome, - Err(join_error) => { - debug!(error = %join_error, "Health reconnect task failed"); - continue; - } - }; - let now = Instant::now(); - let now_alive = live_active_writers_for_dc_family(pool, outcome.dc, outcome.family).await; - if now_alive >= outcome.required { - info!( - dc = %outcome.dc, - family = ?outcome.family, - alive = now_alive, - required = outcome.required, - endpoint_count = outcome.endpoint_count, - "ME writer floor restored for DC" - ); - backoff.insert( - outcome.key, - pool.reconnect_runtime.me_reconnect_backoff_base.as_millis() as u64, - ); - let jitter = pool.reconnect_runtime.me_reconnect_backoff_base.as_millis() as u64 - / JITTER_FRAC_NUM; - let wait = pool.reconnect_runtime.me_reconnect_backoff_base - + Duration::from_millis(rand::rng().random_range(0..=jitter.max(1))); - next_attempt.insert(outcome.key, now + wait); - } else { - let curr = *backoff - .get(&outcome.key) - .unwrap_or(&(pool.reconnect_runtime.me_reconnect_backoff_base.as_millis() as u64)); - let next_ms = (curr.saturating_mul(2)) - .min(pool.reconnect_runtime.me_reconnect_backoff_cap.as_millis() as u64); - backoff.insert(outcome.key, next_ms); - let jitter = next_ms / JITTER_FRAC_NUM; - let wait = Duration::from_millis(next_ms) - + Duration::from_millis(rand::rng().random_range(0..=jitter.max(1))); - next_attempt.insert(outcome.key, now + wait); - if pool.is_runtime_ready() { - let warn_cooldown = pool.warn_rate_limit_duration(); - if should_emit_rate_limited_warn( - floor_warn_next_allowed, - outcome.key, - now, - warn_cooldown, - ) { - warn!( - dc = %outcome.dc, - family = ?outcome.family, - alive = now_alive, - required = outcome.required, - endpoint_count = outcome.endpoint_count, - backoff_ms = next_ms, - "DC writer floor is below required level, scheduled reconnect" - ); - } - } else { - info!( - dc = %outcome.dc, - family = ?outcome.family, - alive = now_alive, - required = outcome.required, - endpoint_count = outcome.endpoint_count, - backoff_ms = next_ms, - "DC writer floor is below required level during startup, scheduled reconnect" - ); - } - } - } - - family_degraded -} - -fn health_reconnect_budget(pool: &Arc, dc_groups: usize) -> usize { - let cpu_cores = pool.adaptive_floor_effective_cpu_cores().max(1); - let by_cpu = cpu_cores.saturating_mul(HEALTH_RECONNECT_BUDGET_PER_CORE); - let by_dc = dc_groups.saturating_mul(HEALTH_RECONNECT_BUDGET_PER_DC); - by_cpu - .saturating_add(by_dc) - .clamp(HEALTH_RECONNECT_BUDGET_MIN, HEALTH_RECONNECT_BUDGET_MAX) -} - -fn update_family_runtime_state(pool: &Arc, family: IpFamily, degraded: bool) { - let now_epoch_secs = MePool::now_epoch_secs(); - let previous_state = pool.family_runtime_state(family); - let mut state_since_epoch_secs = pool.family_runtime_state_since_epoch_secs(family); - let previous_suppressed_until_epoch_secs = pool.family_suppressed_until_epoch_secs(family); - let previous_fail_streak = pool.family_fail_streak(family); - let previous_recover_success_streak = pool.family_recover_success_streak(family); - - let (next_state, suppressed_until_epoch_secs, fail_streak, recover_success_streak) = - if previous_suppressed_until_epoch_secs > now_epoch_secs { - let fail_streak = if degraded { - previous_fail_streak.saturating_add(1) - } else { - previous_fail_streak - }; - ( - MeFamilyRuntimeState::Suppressed, - previous_suppressed_until_epoch_secs, - fail_streak, - 0, - ) - } else if degraded { - let fail_streak = previous_fail_streak.saturating_add(1); - if fail_streak >= FAMILY_SUPPRESS_FAIL_STREAK_THRESHOLD { - ( - MeFamilyRuntimeState::Suppressed, - now_epoch_secs.saturating_add(FAMILY_SUPPRESS_DURATION_SECS), - fail_streak, - 0, - ) - } else { - (MeFamilyRuntimeState::Degraded, 0, fail_streak, 0) - } - } else if matches!(previous_state, MeFamilyRuntimeState::Healthy) { - (MeFamilyRuntimeState::Healthy, 0, 0, 0) - } else { - let recover_success_streak = previous_recover_success_streak.saturating_add(1); - if recover_success_streak >= FAMILY_RECOVER_SUCCESS_STREAK_TARGET { - (MeFamilyRuntimeState::Healthy, 0, 0, 0) - } else { - ( - MeFamilyRuntimeState::Recovering, - 0, - 0, - recover_success_streak, - ) - } - }; - - if next_state != previous_state || state_since_epoch_secs == 0 { - state_since_epoch_secs = now_epoch_secs; - } - pool.set_family_runtime_state( - family, - next_state, - state_since_epoch_secs, - suppressed_until_epoch_secs, - fail_streak, - recover_success_streak, - ); -} - -fn should_emit_rate_limited_warn( - next_allowed: &mut HashMap<(i32, IpFamily), Instant>, - key: (i32, IpFamily), - now: Instant, - cooldown: Duration, -) -> bool { - let Some(ready_at) = next_allowed.get(&key).copied() else { - next_allowed.insert(key, now + cooldown); - return true; - }; - if now >= ready_at { - next_allowed.insert(key, now + cooldown); - return true; - } - false -} - -async fn live_active_writers_for_dc_family(pool: &Arc, dc: i32, family: IpFamily) -> usize { - let writers = pool.writers.read().await; - writers - .iter() - .filter(|writer| { - if writer.draining.load(std::sync::atomic::Ordering::Relaxed) { - return false; - } - if writer.writer_dc != dc { - return false; - } - if !matches!( - super::pool::WriterContour::from_u8( - writer.contour.load(std::sync::atomic::Ordering::Relaxed), - ), - super::pool::WriterContour::Active - ) { - return false; - } - match family { - IpFamily::V4 => writer.addr.is_ipv4(), - IpFamily::V6 => writer.addr.is_ipv6(), - } - }) - .count() -} - -fn adaptive_floor_class_min( - pool: &Arc, - endpoint_count: usize, - base_required: usize, -) -> usize { - if endpoint_count <= 1 { - let min_single = (pool - .floor_runtime - .me_adaptive_floor_min_writers_single_endpoint - .load(std::sync::atomic::Ordering::Relaxed) as usize) - .max(1); - min_single.min(base_required.max(1)) - } else { - pool.adaptive_floor_min_writers_multi_endpoint() - .min(base_required.max(1)) - } -} - -fn adaptive_floor_class_max( - pool: &Arc, - endpoint_count: usize, - base_required: usize, - cpu_cores: usize, -) -> usize { - let extra_per_core = if endpoint_count <= 1 { - pool.adaptive_floor_max_extra_single_per_core() - } else { - pool.adaptive_floor_max_extra_multi_per_core() - }; - base_required.saturating_add(cpu_cores.saturating_mul(extra_per_core)) -} - -fn list_writer_ids_for_endpoints( - dc: i32, - endpoints: &[SocketAddr], - live_writer_ids_by_addr: &HashMap<(i32, SocketAddr), Vec>, -) -> Vec { - let mut out = Vec::::new(); - for endpoint in endpoints { - if let Some(ids) = live_writer_ids_by_addr.get(&(dc, *endpoint)) { - out.extend(ids.iter().copied()); - } - } - out -} - -async fn build_family_floor_plan( - pool: &Arc, - family: IpFamily, - dc_endpoints: &HashMap>, - live_addr_counts: &HashMap<(i32, SocketAddr), usize>, - live_writer_ids_by_addr: &HashMap<(i32, SocketAddr), Vec>, - bound_clients_by_writer: &HashMap, -) -> FamilyFloorPlan { - let mut entries = Vec::::new(); - let mut by_dc = HashMap::::new(); - let mut family_active_total = 0usize; - - let floor_mode = pool.floor_mode(); - let is_adaptive = floor_mode == MeFloorMode::Adaptive; - let cpu_cores = pool.adaptive_floor_effective_cpu_cores().max(1); - let (active_writers_current, warm_writers_current, _) = - pool.non_draining_writer_counts_by_contour().await; - - for (dc, endpoints) in dc_endpoints { - if endpoints.is_empty() { - continue; - } - let _key = (*dc, family); - let base_required = pool.required_writers_for_dc(endpoints.len()).max(1); - let min_required = if is_adaptive { - adaptive_floor_class_min(pool, endpoints.len(), base_required) - } else { - base_required - }; - let mut max_required = if is_adaptive { - adaptive_floor_class_max(pool, endpoints.len(), base_required, cpu_cores) - } else { - base_required - }; - if max_required < min_required { - max_required = min_required; - } - // We initialize target_required at base_required to prevent 0-writer blackouts - // caused by proactively dropping an idle DC to a single fragile connection. - // The Adaptive Floor constraint loop below will gracefully compress idle DCs - // (prioritized via has_bound_clients = false) to min_required only when global capacity is reached. - let desired_raw = base_required; - let target_required = desired_raw.clamp(min_required, max_required); - let alive = endpoints - .iter() - .map(|endpoint| { - live_addr_counts - .get(&(*dc, *endpoint)) - .copied() - .unwrap_or(0) - }) - .sum::(); - family_active_total = family_active_total.saturating_add(alive); - let writer_ids = list_writer_ids_for_endpoints(*dc, endpoints, live_writer_ids_by_addr); - let has_bound_clients = has_bound_clients_on_endpoint(&writer_ids, bound_clients_by_writer); - - entries.push(DcFloorPlanEntry { - dc: *dc, - endpoints: endpoints.clone(), - alive, - min_required, - target_required, - max_required, - has_bound_clients, - floor_capped: false, - }); - } - - if entries.is_empty() { - let active_cap_configured_total = pool.adaptive_floor_active_cap_configured_total(); - let warm_cap_configured_total = pool.adaptive_floor_warm_cap_configured_total(); - return FamilyFloorPlan { - by_dc, - active_cap_configured_total, - active_cap_effective_total: active_cap_configured_total, - warm_cap_configured_total, - warm_cap_effective_total: warm_cap_configured_total, - active_writers_current, - warm_writers_current, - target_writers_total: 0, - }; - } - - if !is_adaptive { - let target_total = entries - .iter() - .map(|entry| entry.target_required) - .sum::(); - let active_cap_configured_total = pool.adaptive_floor_active_cap_configured_total(); - let warm_cap_configured_total = pool.adaptive_floor_warm_cap_configured_total(); - for entry in entries { - by_dc.insert(entry.dc, entry); - } - return FamilyFloorPlan { - by_dc, - active_cap_configured_total, - active_cap_effective_total: active_cap_configured_total.max(target_total), - warm_cap_configured_total, - warm_cap_effective_total: warm_cap_configured_total, - active_writers_current, - warm_writers_current, - target_writers_total: target_total, - }; - } - - let active_cap_configured_total = pool.adaptive_floor_active_cap_configured_total(); - let warm_cap_configured_total = pool.adaptive_floor_warm_cap_configured_total(); - let other_active = active_writers_current.saturating_sub(family_active_total); - let min_sum = entries - .iter() - .map(|entry| entry.min_required) - .sum::(); - let mut target_sum = entries - .iter() - .map(|entry| entry.target_required) - .sum::(); - let family_cap = active_cap_configured_total - .saturating_sub(other_active) - .max(min_sum); - if target_sum > family_cap { - entries.sort_by_key(|entry| { - ( - entry.has_bound_clients, - std::cmp::Reverse(entry.target_required.saturating_sub(entry.min_required)), - std::cmp::Reverse(entry.alive), - entry.dc.abs(), - entry.dc, - entry.endpoints.len(), - entry.max_required, - ) - }); - let mut changed = true; - while target_sum > family_cap && changed { - changed = false; - for entry in &mut entries { - if target_sum <= family_cap { - break; - } - if entry.target_required > entry.min_required { - entry.target_required -= 1; - entry.floor_capped = true; - target_sum -= 1; - changed = true; - } - } - } - } - - for entry in entries { - by_dc.insert(entry.dc, entry); - } - let active_cap_effective_total = - active_cap_configured_total.max(other_active.saturating_add(min_sum)); - let target_writers_total = other_active.saturating_add(target_sum); - FamilyFloorPlan { - by_dc, - active_cap_configured_total, - active_cap_effective_total, - warm_cap_configured_total, - warm_cap_effective_total: warm_cap_configured_total, - active_writers_current, - warm_writers_current, - target_writers_total, - } -} - -async fn maybe_swap_idle_writer_for_cap( - pool: &Arc, - rng: &Arc, - dc: i32, - family: IpFamily, - endpoints: &[SocketAddr], - live_writer_ids_by_addr: &HashMap<(i32, SocketAddr), Vec>, - writer_idle_since: &HashMap, - bound_clients_by_writer: &HashMap, -) -> bool { - let now_epoch_secs = MePool::now_epoch_secs(); - let mut candidate: Option<(u64, SocketAddr, u64)> = None; - for endpoint in endpoints { - let Some(writer_ids) = live_writer_ids_by_addr.get(&(dc, *endpoint)) else { - continue; - }; - for writer_id in writer_ids { - if bound_clients_by_writer.get(writer_id).copied().unwrap_or(0) > 0 { - continue; - } - let Some(idle_since_epoch_secs) = writer_idle_since.get(writer_id).copied() else { - continue; - }; - let idle_age_secs = now_epoch_secs.saturating_sub(idle_since_epoch_secs); - if candidate - .as_ref() - .map(|(_, _, age)| idle_age_secs > *age) - .unwrap_or(true) - { - candidate = Some((*writer_id, *endpoint, idle_age_secs)); - } - } - } - - let Some((old_writer_id, endpoint, idle_age_secs)) = candidate else { - return false; - }; - - let connected = match tokio::time::timeout( - pool.reconnect_runtime.me_one_timeout, - pool.connect_one_for_dc(endpoint, dc, rng.as_ref()), - ) - .await - { - Ok(Ok(())) => true, - Ok(Err(error)) => { - debug!( - dc = %dc, - ?family, - %endpoint, - old_writer_id, - idle_age_secs, - %error, - "Adaptive floor cap swap connect failed" - ); - false - } - Err(_) => { - debug!( - dc = %dc, - ?family, - %endpoint, - old_writer_id, - idle_age_secs, - "Adaptive floor cap swap connect timed out" - ); - false - } - }; - if !connected { - return false; - } - - pool.mark_writer_draining_with_timeout(old_writer_id, pool.force_close_timeout(), false) - .await; - info!( - dc = %dc, - ?family, - %endpoint, - old_writer_id, - idle_age_secs, - "Adaptive floor cap swap: idle writer rotated" - ); - true -} - -async fn maybe_refresh_idle_writer_for_dc( - pool: &Arc, - rng: &Arc, - key: (i32, IpFamily), - dc: i32, - family: IpFamily, - endpoints: &[SocketAddr], - alive: usize, - required: usize, - live_writer_ids_by_addr: &HashMap<(i32, SocketAddr), Vec>, - writer_idle_since: &HashMap, - bound_clients_by_writer: &HashMap, - idle_refresh_next_attempt: &mut HashMap<(i32, IpFamily), Instant>, -) { - if alive < required { - return; - } - - let now = Instant::now(); - if let Some(next) = idle_refresh_next_attempt.get(&key) - && now < *next - { - return; - } - - let now_epoch_secs = MePool::now_epoch_secs(); - let mut candidate: Option<(u64, SocketAddr, u64, u64)> = None; - for endpoint in endpoints { - let Some(writer_ids) = live_writer_ids_by_addr.get(&(dc, *endpoint)) else { - continue; - }; - for writer_id in writer_ids { - if bound_clients_by_writer.get(writer_id).copied().unwrap_or(0) > 0 { - continue; - } - let Some(idle_since_epoch_secs) = writer_idle_since.get(writer_id).copied() else { - continue; - }; - let idle_age_secs = now_epoch_secs.saturating_sub(idle_since_epoch_secs); - let threshold_secs = IDLE_REFRESH_TRIGGER_BASE_SECS - + (*writer_id % (IDLE_REFRESH_TRIGGER_JITTER_SECS + 1)); - if idle_age_secs < threshold_secs { - continue; - } - if candidate - .as_ref() - .map(|(_, _, age, _)| idle_age_secs > *age) - .unwrap_or(true) - { - candidate = Some((*writer_id, *endpoint, idle_age_secs, threshold_secs)); - } - } - } - - let Some((old_writer_id, endpoint, idle_age_secs, threshold_secs)) = candidate else { - return; - }; - - let rotate_ok = match tokio::time::timeout( - pool.reconnect_runtime.me_one_timeout, - pool.connect_one_for_dc(endpoint, dc, rng.as_ref()), - ) - .await - { - Ok(Ok(())) => true, - Ok(Err(error)) => { - debug!( - dc = %dc, - ?family, - %endpoint, - old_writer_id, - idle_age_secs, - threshold_secs, - %error, - "Idle writer pre-refresh connect failed" - ); - false - } - Err(_) => { - debug!( - dc = %dc, - ?family, - %endpoint, - old_writer_id, - idle_age_secs, - threshold_secs, - "Idle writer pre-refresh connect timed out" - ); - false - } - }; - - if !rotate_ok { - idle_refresh_next_attempt.insert(key, now + Duration::from_secs(IDLE_REFRESH_RETRY_SECS)); - return; - } - - pool.mark_writer_draining_with_timeout(old_writer_id, pool.force_close_timeout(), false) - .await; - idle_refresh_next_attempt.insert( - key, - now + Duration::from_secs(IDLE_REFRESH_SUCCESS_GUARD_SECS), - ); - info!( - dc = %dc, - ?family, - %endpoint, - old_writer_id, - idle_age_secs, - threshold_secs, - alive, - required, - "Idle writer refreshed before upstream idle timeout" - ); -} - -fn has_bound_clients_on_endpoint( - writer_ids: &[u64], - bound_clients_by_writer: &HashMap, -) -> bool { - writer_ids - .iter() - .any(|writer_id| bound_clients_by_writer.get(writer_id).copied().unwrap_or(0) > 0) -} - -async fn recover_single_endpoint_outage( - pool: &Arc, - rng: &Arc, - key: (i32, IpFamily), - endpoint: SocketAddr, - required: usize, - outage_backoff: &mut HashMap<(i32, IpFamily), u64>, - outage_next_attempt: &mut HashMap<(i32, IpFamily), Instant>, - reconnect_sem: &Arc, -) { - let now = Instant::now(); - if let Some(ts) = outage_next_attempt.get(&key) - && now < *ts - { - return; - } - - let (min_backoff_ms, max_backoff_ms) = pool.single_endpoint_outage_backoff_bounds_ms(); - if reconnect_sem.available_permits() == 0 { - outage_next_attempt.insert(key, now + Duration::from_millis(min_backoff_ms.max(250))); - debug!( - dc = %key.0, - family = ?key.1, - %endpoint, - required, - "Single-endpoint outage reconnect deferred by health reconnect budget" - ); - return; - } - let Ok(_reconnect_permit) = reconnect_sem.clone().try_acquire_owned() else { - outage_next_attempt.insert(key, now + Duration::from_millis(min_backoff_ms.max(250))); - debug!( - dc = %key.0, - family = ?key.1, - %endpoint, - required, - "Single-endpoint outage reconnect deferred by semaphore saturation" - ); - return; - }; - pool.stats.increment_me_reconnect_attempt(); - pool.stats - .increment_me_single_endpoint_outage_reconnect_attempt_total(); - - let bypass_quarantine = pool.single_endpoint_outage_disable_quarantine(); - let attempt_ok = if bypass_quarantine { - pool.stats - .increment_me_single_endpoint_quarantine_bypass_total(); - match tokio::time::timeout( - pool.reconnect_runtime.me_one_timeout, - pool.connect_one_for_dc(endpoint, key.0, rng.as_ref()), - ) - .await - { - Ok(Ok(())) => true, - Ok(Err(e)) => { - debug!( - dc = %key.0, - family = ?key.1, - %endpoint, - error = %e, - "Single-endpoint outage reconnect failed (quarantine bypass path)" - ); - false - } - Err(_) => { - debug!( - dc = %key.0, - family = ?key.1, - %endpoint, - "Single-endpoint outage reconnect timed out (quarantine bypass path)" - ); - false - } - } - } else { - let one_endpoint = [endpoint]; - match tokio::time::timeout( - pool.reconnect_runtime.me_one_timeout, - pool.connect_endpoints_round_robin(key.0, &one_endpoint, rng.as_ref()), - ) - .await - { - Ok(ok) => ok, - Err(_) => { - debug!( - dc = %key.0, - family = ?key.1, - %endpoint, - "Single-endpoint outage reconnect timed out" - ); - false - } - } - }; - - if attempt_ok { - pool.stats - .increment_me_single_endpoint_outage_reconnect_success_total(); - pool.stats.increment_me_reconnect_success(); - outage_backoff.insert(key, min_backoff_ms); - let jitter = min_backoff_ms / JITTER_FRAC_NUM; - let wait = Duration::from_millis(min_backoff_ms) - + Duration::from_millis(rand::rng().random_range(0..=jitter.max(1))); - outage_next_attempt.insert(key, now + wait); - info!( - dc = %key.0, - family = ?key.1, - %endpoint, - required, - backoff_ms = min_backoff_ms, - "Single-endpoint outage reconnect succeeded" - ); - return; - } - - let current_ms = *outage_backoff.get(&key).unwrap_or(&min_backoff_ms); - let next_ms = current_ms.saturating_mul(2).min(max_backoff_ms); - outage_backoff.insert(key, next_ms); - let jitter = next_ms / JITTER_FRAC_NUM; - let wait = Duration::from_millis(next_ms) - + Duration::from_millis(rand::rng().random_range(0..=jitter.max(1))); - outage_next_attempt.insert(key, now + wait); - warn!( - dc = %key.0, - family = ?key.1, - %endpoint, - required, - backoff_ms = next_ms, - "Single-endpoint outage reconnect scheduled" - ); -} - -async fn maybe_rotate_single_endpoint_shadow( - pool: &Arc, - rng: &Arc, - key: (i32, IpFamily), - dc: i32, - family: IpFamily, - endpoints: &[SocketAddr], - alive: usize, - required: usize, - live_writer_ids_by_addr: &HashMap<(i32, SocketAddr), Vec>, - bound_clients_by_writer: &HashMap, - shadow_rotate_deadline: &mut HashMap<(i32, IpFamily), Instant>, -) { - if endpoints.len() != 1 || alive < required { - return; - } - - let Some(interval) = pool.single_endpoint_shadow_rotate_interval() else { - return; - }; - - let now = Instant::now(); - if let Some(deadline) = shadow_rotate_deadline.get(&key) - && now < *deadline - { - return; - } - - let endpoint = endpoints[0]; - if pool.is_endpoint_quarantined(endpoint).await { - pool.stats - .increment_me_single_endpoint_shadow_rotate_skipped_quarantine_total(); - shadow_rotate_deadline.insert(key, now + Duration::from_secs(SHADOW_ROTATE_RETRY_SECS)); - debug!( - dc = %dc, - ?family, - %endpoint, - "Single-endpoint shadow rotation skipped: endpoint is quarantined" - ); - return; - } - - let Some(writer_ids) = live_writer_ids_by_addr.get(&(dc, endpoint)) else { - shadow_rotate_deadline.insert(key, now + Duration::from_secs(SHADOW_ROTATE_RETRY_SECS)); - return; - }; - - let mut candidate_writer_id = None; - for writer_id in writer_ids { - if bound_clients_by_writer.get(writer_id).copied().unwrap_or(0) == 0 { - candidate_writer_id = Some(*writer_id); - break; - } - } - - let Some(old_writer_id) = candidate_writer_id else { - shadow_rotate_deadline.insert(key, now + Duration::from_secs(SHADOW_ROTATE_RETRY_SECS)); - debug!( - dc = %dc, - ?family, - %endpoint, - alive, - required, - "Single-endpoint shadow rotation skipped: no empty writer candidate" - ); - return; - }; - - let rotate_ok = match tokio::time::timeout( - pool.reconnect_runtime.me_one_timeout, - pool.connect_one_for_dc(endpoint, dc, rng.as_ref()), - ) - .await - { - Ok(Ok(())) => true, - Ok(Err(e)) => { - debug!( - dc = %dc, - ?family, - %endpoint, - error = %e, - "Single-endpoint shadow rotation connect failed" - ); - false - } - Err(_) => { - debug!( - dc = %dc, - ?family, - %endpoint, - "Single-endpoint shadow rotation connect timed out" - ); - false - } - }; - - if !rotate_ok { - shadow_rotate_deadline.insert( - key, - now + interval.min(Duration::from_secs(SHADOW_ROTATE_RETRY_SECS)), - ); - return; - } - - pool.mark_writer_draining_with_timeout(old_writer_id, pool.force_close_timeout(), false) - .await; - pool.stats - .increment_me_single_endpoint_shadow_rotate_total(); - shadow_rotate_deadline.insert(key, now + interval); - info!( - dc = %dc, - ?family, - %endpoint, - old_writer_id, - rotate_every_secs = interval.as_secs(), - "Single-endpoint shadow writer rotated" - ); -} - -/// Last-resort safety net for draining writers stuck past their deadline. -/// -/// Runs every `TICK_SECS` and force-closes any draining writer whose -/// `drain_deadline_epoch_secs` has been exceeded by more than a threshold. -/// -/// Two thresholds: -/// - `SOFT_THRESHOLD_SECS` (60s): writers with no bound clients -/// - `HARD_THRESHOLD_SECS` (300s): writers WITH bound clients (unconditional) -/// -/// Intentionally kept trivial and independent of pool config to minimise -/// the probability of panicking itself. Uses `SystemTime` directly -/// as a fallback clock source and timeouts on every lock acquisition -/// and writer removal so one stuck writer cannot block the rest. -pub async fn me_zombie_writer_watchdog(pool: Arc) { - use std::time::{SystemTime, UNIX_EPOCH}; - - const TICK_SECS: u64 = 30; - const SOFT_THRESHOLD_SECS: u64 = 60; - const HARD_THRESHOLD_SECS: u64 = 300; - const LOCK_TIMEOUT_SECS: u64 = 5; - const REMOVE_TIMEOUT_SECS: u64 = 10; - const HARD_DETACH_TIMEOUT_STREAK: u8 = 3; - - let mut removal_timeout_streak = HashMap::::new(); - - loop { - tokio::time::sleep(Duration::from_secs(TICK_SECS)).await; - - let now = match SystemTime::now().duration_since(UNIX_EPOCH) { - Ok(d) => d.as_secs(), - Err(_) => continue, - }; - - // Phase 1: collect zombie IDs under a short read-lock with timeout. - let zombie_ids_with_meta: Vec<(u64, bool)> = { - let Ok(ws) = - tokio::time::timeout(Duration::from_secs(LOCK_TIMEOUT_SECS), pool.writers.read()) - .await - else { - warn!("zombie_watchdog: writers read-lock timeout, skipping tick"); - continue; - }; - ws.iter() - .filter(|w| w.draining.load(std::sync::atomic::Ordering::Relaxed)) - .filter_map(|w| { - let deadline = w - .drain_deadline_epoch_secs - .load(std::sync::atomic::Ordering::Relaxed); - if deadline == 0 { - return None; - } - let overdue = now.saturating_sub(deadline); - if overdue == 0 { - return None; - } - let started = w - .draining_started_at_epoch_secs - .load(std::sync::atomic::Ordering::Relaxed); - let drain_age = now.saturating_sub(started); - if drain_age > HARD_THRESHOLD_SECS { - return Some((w.id, true)); - } - if overdue > SOFT_THRESHOLD_SECS { - return Some((w.id, false)); - } - None - }) - .collect() - }; - // read lock released here - - if zombie_ids_with_meta.is_empty() { - removal_timeout_streak.clear(); - continue; - } - - let mut active_zombie_ids = HashSet::::with_capacity(zombie_ids_with_meta.len()); - for (writer_id, _) in &zombie_ids_with_meta { - active_zombie_ids.insert(*writer_id); - } - removal_timeout_streak.retain(|writer_id, _| active_zombie_ids.contains(writer_id)); - - warn!( - zombie_count = zombie_ids_with_meta.len(), - soft_threshold_secs = SOFT_THRESHOLD_SECS, - hard_threshold_secs = HARD_THRESHOLD_SECS, - "Zombie draining writers detected by watchdog, force-closing" - ); - - // Phase 2: remove each writer individually with a timeout. - // One stuck removal cannot block the rest. - for (writer_id, had_clients) in &zombie_ids_with_meta { - let result = tokio::time::timeout( - Duration::from_secs(REMOVE_TIMEOUT_SECS), - pool.remove_writer_and_close_clients(*writer_id), - ) - .await; - match result { - Ok(()) => { - removal_timeout_streak.remove(writer_id); - pool.stats.increment_pool_force_close_total(); - info!(writer_id, had_clients, "Zombie writer removed by watchdog"); - } - Err(_) => { - let streak = removal_timeout_streak - .entry(*writer_id) - .and_modify(|value| *value = value.saturating_add(1)) - .or_insert(1); - warn!( - writer_id, - had_clients, - timeout_streak = *streak, - "Zombie writer removal timed out" - ); - if *streak < HARD_DETACH_TIMEOUT_STREAK { - continue; - } - - let hard_detach = tokio::time::timeout( - Duration::from_secs(REMOVE_TIMEOUT_SECS), - pool.remove_draining_writer_hard_detach(*writer_id), - ) - .await; - match hard_detach { - Ok(true) => { - removal_timeout_streak.remove(writer_id); - pool.stats.increment_pool_force_close_total(); - info!( - writer_id, - had_clients, "Zombie writer hard-detached after repeated timeouts" - ); - } - Ok(false) => { - removal_timeout_streak.remove(writer_id); - debug!( - writer_id, - had_clients, - "Zombie hard-detach skipped (writer already gone or no longer draining)" - ); - } - Err(_) => { - warn!( - writer_id, - had_clients, "Zombie hard-detach timed out, will retry next tick" - ); - } - } - } - } - } - } -} #[cfg(test)] -mod tests { - use std::collections::HashMap; - use std::net::{IpAddr, Ipv4Addr, SocketAddr}; - use std::sync::Arc; - use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU32, AtomicU64, Ordering}; - use std::time::{Duration, Instant}; - - use tokio::sync::mpsc; - use tokio_util::sync::CancellationToken; - - use super::{ScheduledReconnects, reap_draining_writers}; - use crate::config::{GeneralConfig, MeRouteNoWriterMode, MeSocksKdfPolicy, MeWriterPickMode}; - use crate::crypto::SecureRandom; - use crate::network::probe::NetworkDecision; - use crate::network::IpFamily; - use crate::stats::Stats; - use crate::transport::middle_proxy::codec::WriterCommand; - use crate::transport::middle_proxy::pool::{MePool, MeWriter, WriterContour}; - use crate::transport::middle_proxy::registry::ConnMeta; - - #[test] - fn reconnect_batch_releases_every_reserved_key_after_join_failures() { - let retained = (1, IpFamily::V4); - let removed = (2, IpFamily::V6); - let mut inflight = HashMap::from([(retained, 1)]); - - { - let mut scheduled = ScheduledReconnects { - inflight: &mut inflight, - keys: Vec::new(), - }; - scheduled.reserve(retained); - scheduled.reserve(removed); - } - - assert_eq!(inflight.get(&retained), Some(&1)); - assert!(!inflight.contains_key(&removed)); - } - - async fn make_pool(me_pool_drain_threshold: u64) -> Arc { - let general = GeneralConfig { - me_pool_drain_threshold, - ..GeneralConfig::default() - }; - let mut proxy_map_v4 = HashMap::new(); - proxy_map_v4.insert(2, vec![(IpAddr::V4(Ipv4Addr::new(203, 0, 113, 10)), 443)]); - let decision = NetworkDecision { - ipv4_me: true, - ..NetworkDecision::default() - }; - MePool::new( - None, - vec![1u8; 32], - None, - false, - None, - Vec::new(), - false, - Vec::new(), - 1, - None, - 12, - 1200, - proxy_map_v4, - HashMap::new(), - None, - decision, - None, - Arc::new(SecureRandom::new()), - Arc::new(Stats::default()), - general.me_keepalive_enabled, - general.me_keepalive_interval_secs, - general.me_keepalive_jitter_secs, - general.me_keepalive_payload_random, - general.rpc_proxy_req_every, - general.me_warmup_stagger_enabled, - general.me_warmup_step_delay_ms, - general.me_warmup_step_jitter_ms, - general.me_reconnect_max_concurrent_per_dc, - general.me_reconnect_backoff_base_ms, - general.me_reconnect_backoff_cap_ms, - general.me_reconnect_fast_retry_count, - general.me_single_endpoint_shadow_writers, - general.me_single_endpoint_outage_mode_enabled, - general.me_single_endpoint_outage_disable_quarantine, - general.me_single_endpoint_outage_backoff_min_ms, - general.me_single_endpoint_outage_backoff_max_ms, - general.me_single_endpoint_shadow_rotate_every_secs, - general.me_floor_mode, - general.me_adaptive_floor_idle_secs, - general.me_adaptive_floor_min_writers_single_endpoint, - general.me_adaptive_floor_min_writers_multi_endpoint, - general.me_adaptive_floor_recover_grace_secs, - general.me_adaptive_floor_writers_per_core_total, - general.me_adaptive_floor_cpu_cores_override, - general.me_adaptive_floor_max_extra_writers_single_per_core, - general.me_adaptive_floor_max_extra_writers_multi_per_core, - general.me_adaptive_floor_max_active_writers_per_core, - general.me_adaptive_floor_max_warm_writers_per_core, - general.me_adaptive_floor_max_active_writers_global, - general.me_adaptive_floor_max_warm_writers_global, - general.hardswap, - general.me_pool_drain_ttl_secs, - general.me_instadrain, - general.me_pool_drain_threshold, - general.me_pool_drain_soft_evict_enabled, - general.me_pool_drain_soft_evict_grace_secs, - general.me_pool_drain_soft_evict_per_writer, - general.me_pool_drain_soft_evict_budget_per_core, - general.me_pool_drain_soft_evict_cooldown_ms, - general.effective_me_pool_force_close_secs(), - general.me_pool_min_fresh_ratio, - general.me_hardswap_warmup_delay_min_ms, - general.me_hardswap_warmup_delay_max_ms, - general.me_hardswap_warmup_extra_passes, - general.me_hardswap_warmup_pass_backoff_base_ms, - general.me_bind_stale_mode, - general.me_bind_stale_ttl_secs, - general.me_secret_atomic_snapshot, - general.me_deterministic_writer_sort, - MeWriterPickMode::default(), - general.me_writer_pick_sample_size, - MeSocksKdfPolicy::default(), - general.me_writer_cmd_channel_capacity, - general.me_writer_byte_budget_bytes, - general.me_route_channel_capacity, - general.me_route_backpressure_enabled, - general.me_route_fairshare_enabled, - general.me_route_backpressure_base_timeout_ms, - general.me_route_backpressure_high_timeout_ms, - general.me_route_backpressure_high_watermark_pct, - general.me_reader_route_data_wait_ms, - general.me_health_interval_ms_unhealthy, - general.me_health_interval_ms_healthy, - general.me_warn_rate_limit_ms, - MeRouteNoWriterMode::default(), - general.me_route_no_writer_wait_ms, - general.me_route_hybrid_max_wait_ms, - general.me_route_blocking_send_timeout_ms, - general.me_route_inline_recovery_attempts, - general.me_route_inline_recovery_wait_ms, - 16_384, - ) - } - - async fn insert_draining_writer( - pool: &Arc, - writer_id: u64, - drain_started_at_epoch_secs: u64, - ) -> u64 { - let (conn_id, _rx) = pool.registry.register().await; - let (tx, _writer_rx) = mpsc::channel::(8); - let byte_budget = pool.new_writer_byte_budget(); - let writer = MeWriter { - id: writer_id, - addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 4000 + writer_id as u16), - source_ip: IpAddr::V4(Ipv4Addr::LOCALHOST), - writer_dc: 2, - generation: 1, - contour: Arc::new(AtomicU8::new(WriterContour::Draining.as_u8())), - created_at: Instant::now() - Duration::from_secs(writer_id), - tx: tx.clone(), - byte_budget: byte_budget.clone(), - cancel: CancellationToken::new(), - degraded: Arc::new(AtomicBool::new(false)), - rtt_ema_ms_x10: Arc::new(AtomicU32::new(0)), - draining: Arc::new(AtomicBool::new(true)), - draining_started_at_epoch_secs: Arc::new(AtomicU64::new(drain_started_at_epoch_secs)), - drain_deadline_epoch_secs: Arc::new(AtomicU64::new(0)), - allow_drain_fallback: Arc::new(AtomicBool::new(false)), - }; - pool.writers.write().await.push(writer); - pool.registry - .register_writer(writer_id, tx, byte_budget) - .await; - pool.conn_count.fetch_add(1, Ordering::Relaxed); - assert!( - pool.registry - .bind_writer( - conn_id, - writer_id, - ConnMeta { - target_dc: 2, - client_addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 6000), - our_addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443), - proto_flags: 0, - }, - ) - .await - ); - conn_id - } - - async fn insert_live_writer(pool: &Arc, writer_id: u64, writer_dc: i32) { - let (tx, _writer_rx) = mpsc::channel::(8); - let byte_budget = pool.new_writer_byte_budget(); - let writer = MeWriter { - id: writer_id, - addr: SocketAddr::new( - IpAddr::V4(Ipv4Addr::new( - 203, - 0, - 113, - (writer_id as u8).saturating_add(1), - )), - 4000 + writer_id as u16, - ), - source_ip: IpAddr::V4(Ipv4Addr::LOCALHOST), - writer_dc, - generation: 2, - contour: Arc::new(AtomicU8::new(WriterContour::Active.as_u8())), - created_at: Instant::now(), - tx: tx.clone(), - byte_budget: byte_budget.clone(), - cancel: CancellationToken::new(), - degraded: Arc::new(AtomicBool::new(false)), - rtt_ema_ms_x10: Arc::new(AtomicU32::new(0)), - draining: Arc::new(AtomicBool::new(false)), - draining_started_at_epoch_secs: Arc::new(AtomicU64::new(0)), - drain_deadline_epoch_secs: Arc::new(AtomicU64::new(0)), - allow_drain_fallback: Arc::new(AtomicBool::new(false)), - }; - pool.writers.write().await.push(writer); - pool.registry - .register_writer(writer_id, tx, byte_budget) - .await; - pool.conn_count.fetch_add(1, Ordering::Relaxed); - } - - #[tokio::test] - async fn reap_draining_writers_force_closes_oldest_over_threshold() { - let pool = make_pool(2).await; - insert_live_writer(&pool, 1, 2).await; - let now_epoch_secs = MePool::now_epoch_secs(); - let conn_a = insert_draining_writer(&pool, 10, now_epoch_secs.saturating_sub(30)).await; - let conn_b = insert_draining_writer(&pool, 20, now_epoch_secs.saturating_sub(20)).await; - let conn_c = insert_draining_writer(&pool, 30, now_epoch_secs.saturating_sub(10)).await; - let mut warn_next_allowed = HashMap::new(); - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - - let mut writer_ids: Vec = pool - .writers - .read() - .await - .iter() - .map(|writer| writer.id) - .collect(); - writer_ids.sort_unstable(); - assert_eq!(writer_ids, vec![1, 20, 30]); - assert!(pool.registry.get_writer(conn_a).await.is_none()); - assert_eq!( - pool.registry.get_writer(conn_b).await.unwrap().writer_id, - 20 - ); - assert_eq!( - pool.registry.get_writer(conn_c).await.unwrap().writer_id, - 30 - ); - } - - #[tokio::test] - async fn reap_draining_writers_force_closes_overflow_without_replacement() { - let pool = make_pool(2).await; - let now_epoch_secs = MePool::now_epoch_secs(); - let conn_a = insert_draining_writer(&pool, 10, now_epoch_secs.saturating_sub(30)).await; - let conn_b = insert_draining_writer(&pool, 20, now_epoch_secs.saturating_sub(20)).await; - let conn_c = insert_draining_writer(&pool, 30, now_epoch_secs.saturating_sub(10)).await; - let mut warn_next_allowed = HashMap::new(); - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - - let mut writer_ids: Vec = pool - .writers - .read() - .await - .iter() - .map(|writer| writer.id) - .collect(); - writer_ids.sort_unstable(); - assert_eq!(writer_ids, vec![20, 30]); - assert!(pool.registry.get_writer(conn_a).await.is_none()); - assert_eq!( - pool.registry.get_writer(conn_b).await.unwrap().writer_id, - 20 - ); - assert_eq!( - pool.registry.get_writer(conn_c).await.unwrap().writer_id, - 30 - ); - } - - #[tokio::test] - async fn reap_draining_writers_keeps_timeout_only_behavior_when_threshold_disabled() { - let pool = make_pool(0).await; - let now_epoch_secs = MePool::now_epoch_secs(); - let conn_a = insert_draining_writer(&pool, 10, now_epoch_secs.saturating_sub(30)).await; - let conn_b = insert_draining_writer(&pool, 20, now_epoch_secs.saturating_sub(20)).await; - let conn_c = insert_draining_writer(&pool, 30, now_epoch_secs.saturating_sub(10)).await; - let mut warn_next_allowed = HashMap::new(); - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - - let writer_ids: Vec = pool - .writers - .read() - .await - .iter() - .map(|writer| writer.id) - .collect(); - assert_eq!(writer_ids, vec![10, 20, 30]); - assert_eq!( - pool.registry.get_writer(conn_a).await.unwrap().writer_id, - 10 - ); - assert_eq!( - pool.registry.get_writer(conn_b).await.unwrap().writer_id, - 20 - ); - assert_eq!( - pool.registry.get_writer(conn_c).await.unwrap().writer_id, - 30 - ); - } -} +mod tests; diff --git a/src/transport/middle_proxy/health/drain.rs b/src/transport/middle_proxy/health/drain.rs new file mode 100644 index 0000000..336517b --- /dev/null +++ b/src/transport/middle_proxy/health/drain.rs @@ -0,0 +1,212 @@ +use super::*; + +pub async fn me_drain_timeout_enforcer(pool: Arc) { + let mut drain_warn_next_allowed: HashMap = HashMap::new(); + loop { + tokio::time::sleep(Duration::from_secs( + HEALTH_DRAIN_TIMEOUT_ENFORCER_INTERVAL_SECS, + )) + .await; + reap_draining_writers(&pool, &mut drain_warn_next_allowed).await; + } +} + +pub(in crate::transport::middle_proxy) async fn reap_draining_writers( + pool: &Arc, + warn_next_allowed: &mut HashMap, +) { + let now_epoch_secs = MePool::now_epoch_secs(); + let now = Instant::now(); + let drain_ttl_secs = pool + .drain_runtime + .me_pool_drain_ttl_secs + .load(std::sync::atomic::Ordering::Relaxed); + let drain_threshold = pool + .drain_runtime + .me_pool_drain_threshold + .load(std::sync::atomic::Ordering::Relaxed); + let activity = pool.registry.writer_activity_snapshot().await; + let mut draining_writers = Vec::::new(); + let mut empty_writer_ids = Vec::::new(); + let mut force_close_writer_ids = Vec::::new(); + let writers = pool.writers.read().await; + for writer in writers.iter() { + if !writer.draining.load(std::sync::atomic::Ordering::Relaxed) { + continue; + } + if activity + .bound_clients_by_writer + .get(&writer.id) + .copied() + .unwrap_or(0) + == 0 + { + empty_writer_ids.push(writer.id); + continue; + } + draining_writers.push(DrainingWriterSnapshot { + id: writer.id, + writer_dc: writer.writer_dc, + addr: writer.addr, + generation: writer.generation, + created_at: writer.created_at, + draining_started_at_epoch_secs: writer + .draining_started_at_epoch_secs + .load(std::sync::atomic::Ordering::Relaxed), + drain_deadline_epoch_secs: writer + .drain_deadline_epoch_secs + .load(std::sync::atomic::Ordering::Relaxed), + allow_drain_fallback: writer + .allow_drain_fallback + .load(std::sync::atomic::Ordering::Relaxed), + }); + } + drop(writers); + + let overflow = if drain_threshold > 0 && draining_writers.len() > drain_threshold as usize { + draining_writers + .len() + .saturating_sub(drain_threshold as usize) + } else { + 0 + }; + + if overflow > 0 { + draining_writers.sort_by(|left, right| { + left.draining_started_at_epoch_secs + .cmp(&right.draining_started_at_epoch_secs) + .then_with(|| left.created_at.cmp(&right.created_at)) + .then_with(|| left.id.cmp(&right.id)) + }); + warn!( + draining_writers = draining_writers.len(), + me_pool_drain_threshold = drain_threshold, + removing_writers = overflow, + "ME draining writer threshold exceeded, force-closing oldest draining writers" + ); + for writer in draining_writers.drain(..overflow) { + force_close_writer_ids.push(writer.id); + } + } + + for writer in draining_writers { + if drain_ttl_secs > 0 + && writer.draining_started_at_epoch_secs != 0 + && now_epoch_secs.saturating_sub(writer.draining_started_at_epoch_secs) > drain_ttl_secs + && should_emit_writer_warn( + warn_next_allowed, + writer.id, + now, + pool.warn_rate_limit_duration(), + ) + { + warn!( + writer_id = writer.id, + writer_dc = writer.writer_dc, + endpoint = %writer.addr, + generation = writer.generation, + drain_ttl_secs, + force_close_secs = pool + .drain_runtime + .me_pool_force_close_secs + .load(std::sync::atomic::Ordering::Relaxed), + allow_drain_fallback = writer.allow_drain_fallback, + "ME draining writer remains non-empty past drain TTL" + ); + } + if writer.drain_deadline_epoch_secs != 0 + && now_epoch_secs >= writer.drain_deadline_epoch_secs + { + warn!(writer_id = writer.id, "Drain timeout, force-closing"); + force_close_writer_ids.push(writer.id); + } + } + + let close_budget = health_drain_close_budget(); + let requested_force_close = force_close_writer_ids.len(); + let requested_empty_close = empty_writer_ids.len(); + let requested_close_total = requested_force_close.saturating_add(requested_empty_close); + let mut closed_writer_ids = HashSet::::new(); + let mut closed_total = 0usize; + for writer_id in force_close_writer_ids { + if closed_total >= close_budget { + break; + } + if !closed_writer_ids.insert(writer_id) { + continue; + } + pool.stats.increment_pool_force_close_total(); + pool.remove_writer_and_close_clients(writer_id).await; + closed_total = closed_total.saturating_add(1); + } + for writer_id in empty_writer_ids { + if closed_total >= close_budget { + break; + } + if !closed_writer_ids.insert(writer_id) { + continue; + } + pool.remove_writer_and_close_clients(writer_id).await; + closed_total = closed_total.saturating_add(1); + } + + let pending_close_total = requested_close_total.saturating_sub(closed_total); + if pending_close_total > 0 { + warn!( + close_budget, + closed_total, + pending_close_total, + "ME draining close backlog deferred to next health cycle" + ); + } + + // Keep warn cooldown state for draining writers still present in the pool; + // drop state only once a writer is actually removed. + let active_draining_writer_ids = { + let writers = pool.writers.read().await; + writers + .iter() + .filter(|writer| writer.draining.load(std::sync::atomic::Ordering::Relaxed)) + .map(|writer| writer.id) + .collect::>() + }; + warn_next_allowed.retain(|writer_id, _| active_draining_writer_ids.contains(writer_id)); +} + +pub(in crate::transport::middle_proxy) fn health_drain_close_budget() -> usize { + let cpu_cores = std::thread::available_parallelism() + .map(std::num::NonZeroUsize::get) + .unwrap_or(1); + cpu_cores + .saturating_mul(HEALTH_DRAIN_CLOSE_BUDGET_PER_CORE) + .clamp(HEALTH_DRAIN_CLOSE_BUDGET_MIN, HEALTH_DRAIN_CLOSE_BUDGET_MAX) +} + +#[derive(Debug, Clone)] +struct DrainingWriterSnapshot { + id: u64, + writer_dc: i32, + addr: SocketAddr, + generation: u64, + created_at: Instant, + draining_started_at_epoch_secs: u64, + drain_deadline_epoch_secs: u64, + allow_drain_fallback: bool, +} + +pub(super) fn should_emit_writer_warn( + next_allowed: &mut HashMap, + writer_id: u64, + now: Instant, + cooldown: Duration, +) -> bool { + let Some(ready_at) = next_allowed.get(&writer_id).copied() else { + next_allowed.insert(writer_id, now + cooldown); + return true; + }; + if now >= ready_at { + next_allowed.insert(writer_id, now + cooldown); + return true; + } + false +} diff --git a/src/transport/middle_proxy/health/family.rs b/src/transport/middle_proxy/health/family.rs new file mode 100644 index 0000000..c454675 --- /dev/null +++ b/src/transport/middle_proxy/health/family.rs @@ -0,0 +1,484 @@ +use super::*; + +pub(super) async fn check_family( + family: IpFamily, + pool: &Arc, + rng: &Arc, + backoff: &mut HashMap<(i32, IpFamily), u64>, + next_attempt: &mut HashMap<(i32, IpFamily), Instant>, + inflight: &mut HashMap<(i32, IpFamily), usize>, + outage_backoff: &mut HashMap<(i32, IpFamily), u64>, + outage_next_attempt: &mut HashMap<(i32, IpFamily), Instant>, + single_endpoint_outage: &mut HashSet<(i32, IpFamily)>, + shadow_rotate_deadline: &mut HashMap<(i32, IpFamily), Instant>, + idle_refresh_next_attempt: &mut HashMap<(i32, IpFamily), Instant>, + floor_warn_next_allowed: &mut HashMap<(i32, IpFamily), Instant>, +) -> bool { + let enabled = match family { + IpFamily::V4 => pool.decision.ipv4_me, + IpFamily::V6 => pool.decision.ipv6_me, + }; + if !enabled { + return false; + } + + let mut family_degraded = false; + + let mut dc_endpoints = HashMap::>::new(); + let map_guard = match family { + IpFamily::V4 => pool.proxy_map_v4.read().await, + IpFamily::V6 => pool.proxy_map_v6.read().await, + }; + for (dc, addrs) in map_guard.iter() { + let entry = dc_endpoints.entry(*dc).or_default(); + for (ip, port) in addrs.iter().copied() { + entry.push(SocketAddr::new(ip, port)); + } + } + drop(map_guard); + for endpoints in dc_endpoints.values_mut() { + endpoints.sort_unstable(); + endpoints.dedup(); + } + let reconnect_budget = health_reconnect_budget(pool, dc_endpoints.len()); + let reconnect_sem = Arc::new(Semaphore::new(reconnect_budget)); + + if pool.floor_mode() == MeFloorMode::Static {} + + let mut live_addr_counts = HashMap::<(i32, SocketAddr), usize>::new(); + let mut live_writer_ids_by_addr = HashMap::<(i32, SocketAddr), Vec>::new(); + for writer in pool + .writers + .read() + .await + .iter() + .filter(|w| !w.draining.load(std::sync::atomic::Ordering::Relaxed)) + { + if !matches!( + crate::transport::middle_proxy::pool::WriterContour::from_u8( + writer.contour.load(std::sync::atomic::Ordering::Relaxed), + ), + crate::transport::middle_proxy::pool::WriterContour::Active + ) { + continue; + } + let key = (writer.writer_dc, writer.addr); + *live_addr_counts.entry(key).or_insert(0) += 1; + live_writer_ids_by_addr + .entry(key) + .or_default() + .push(writer.id); + } + let writer_idle_since = pool.registry.writer_idle_since_snapshot().await; + let bound_clients_by_writer = pool + .registry + .writer_activity_snapshot() + .await + .bound_clients_by_writer; + let floor_plan = build_family_floor_plan( + pool, + family, + &dc_endpoints, + &live_addr_counts, + &live_writer_ids_by_addr, + &bound_clients_by_writer, + ) + .await; + pool.set_adaptive_floor_runtime_caps( + floor_plan.active_cap_configured_total, + floor_plan.active_cap_effective_total, + floor_plan.warm_cap_configured_total, + floor_plan.warm_cap_effective_total, + floor_plan.target_writers_total, + floor_plan.active_writers_current, + floor_plan.warm_writers_current, + ); + let live_writer_ids_by_addr = Arc::new(live_writer_ids_by_addr); + let writer_idle_since = Arc::new(writer_idle_since); + let bound_clients_by_writer = Arc::new(bound_clients_by_writer); + let mut reconnect_set = JoinSet::::new(); + let mut scheduled_reconnects = ScheduledReconnects { + inflight, + keys: Vec::new(), + }; + + for (dc, endpoints) in dc_endpoints { + if endpoints.is_empty() { + continue; + } + let key = (dc, family); + let required = floor_plan + .by_dc + .get(&dc) + .map(|entry| entry.target_required) + .unwrap_or_else(|| { + pool.required_writers_for_dc_with_floor_mode(endpoints.len(), false) + }); + let alive = endpoints + .iter() + .map(|addr| *live_addr_counts.get(&(dc, *addr)).unwrap_or(&0)) + .sum::(); + + if endpoints.len() == 1 && pool.single_endpoint_outage_mode_enabled() && alive == 0 { + family_degraded = true; + if single_endpoint_outage.insert(key) { + pool.stats.increment_me_single_endpoint_outage_enter_total(); + warn!( + dc = %dc, + ?family, + required, + endpoint_count = endpoints.len(), + "Single-endpoint DC outage detected" + ); + } + + recover_single_endpoint_outage( + pool, + rng, + key, + endpoints[0], + required, + outage_backoff, + outage_next_attempt, + &reconnect_sem, + ) + .await; + continue; + } + + if single_endpoint_outage.remove(&key) { + pool.stats.increment_me_single_endpoint_outage_exit_total(); + outage_backoff.remove(&key); + outage_next_attempt.remove(&key); + shadow_rotate_deadline.remove(&key); + idle_refresh_next_attempt.remove(&key); + info!( + dc = %dc, + ?family, + alive, + required, + endpoint_count = endpoints.len(), + "Single-endpoint DC outage recovered" + ); + } + + if alive >= required { + maybe_refresh_idle_writer_for_dc( + pool, + rng, + key, + dc, + family, + &endpoints, + alive, + required, + live_writer_ids_by_addr.as_ref(), + writer_idle_since.as_ref(), + bound_clients_by_writer.as_ref(), + idle_refresh_next_attempt, + ) + .await; + maybe_rotate_single_endpoint_shadow( + pool, + rng, + key, + dc, + family, + &endpoints, + alive, + required, + live_writer_ids_by_addr.as_ref(), + bound_clients_by_writer.as_ref(), + shadow_rotate_deadline, + ) + .await; + continue; + } + let missing = required - alive; + family_degraded = true; + + let now = Instant::now(); + if reconnect_sem.available_permits() == 0 { + let base_ms = pool.reconnect_runtime.me_reconnect_backoff_base.as_millis() as u64; + let next_ms = (*backoff.get(&key).unwrap_or(&base_ms)).max(base_ms); + let jitter = next_ms / JITTER_FRAC_NUM; + let wait = Duration::from_millis(next_ms) + + Duration::from_millis(rand::rng().random_range(0..=jitter.max(1))); + next_attempt.insert(key, now + wait); + debug!( + dc = %dc, + ?family, + alive, + required, + endpoint_count = endpoints.len(), + reconnect_budget, + "Skipping reconnect due to per-tick health reconnect budget" + ); + continue; + } + if let Some(ts) = next_attempt.get(&key) + && now < *ts + { + continue; + } + + let max_concurrent = pool + .reconnect_runtime + .me_reconnect_max_concurrent_per_dc + .max(1) as usize; + if scheduled_reconnects.current(&key) >= max_concurrent { + continue; + } + if pool + .has_refill_inflight_for_dc_key(crate::transport::middle_proxy::pool::RefillDcKey { + dc, + family, + }) + .await + { + debug!( + dc = %dc, + ?family, + alive, + required, + endpoint_count = endpoints.len(), + "Skipping health reconnect: immediate refill is already in flight for this DC group" + ); + continue; + } + scheduled_reconnects.reserve(key); + let pool_for_reconnect = pool.clone(); + let rng_for_reconnect = rng.clone(); + let reconnect_sem_for_dc = reconnect_sem.clone(); + let endpoints_for_dc = endpoints.clone(); + let live_writer_ids_by_addr_for_dc = live_writer_ids_by_addr.clone(); + let writer_idle_since_for_dc = writer_idle_since.clone(); + let bound_clients_by_writer_for_dc = bound_clients_by_writer.clone(); + let active_cap_effective_total = floor_plan.active_cap_effective_total; + reconnect_set.spawn(async move { + let mut restored = 0usize; + for _ in 0..missing { + let Ok(reconnect_permit) = reconnect_sem_for_dc.clone().try_acquire_owned() else { + break; + }; + if pool_for_reconnect.active_contour_writer_count_total().await + >= active_cap_effective_total + { + let swapped = maybe_swap_idle_writer_for_cap( + &pool_for_reconnect, + &rng_for_reconnect, + dc, + family, + &endpoints_for_dc, + live_writer_ids_by_addr_for_dc.as_ref(), + writer_idle_since_for_dc.as_ref(), + bound_clients_by_writer_for_dc.as_ref(), + ) + .await; + if swapped { + pool_for_reconnect + .stats + .increment_me_floor_swap_idle_total(); + restored += 1; + continue; + } + + let base_req = pool_for_reconnect + .required_writers_for_dc_with_floor_mode(endpoints_for_dc.len(), false); + if alive + restored >= base_req { + pool_for_reconnect + .stats + .increment_me_floor_cap_block_total(); + pool_for_reconnect + .stats + .increment_me_floor_swap_idle_failed_total(); + debug!( + dc = %dc, + ?family, + alive, + required, + active_cap_effective_total, + "Adaptive floor cap reached, reconnect attempt blocked" + ); + break; + } + } + pool_for_reconnect.stats.increment_me_reconnect_attempt(); + let res = tokio::time::timeout( + pool_for_reconnect.reconnect_runtime.me_one_timeout, + pool_for_reconnect.connect_endpoints_round_robin( + dc, + &endpoints_for_dc, + rng_for_reconnect.as_ref(), + ), + ) + .await; + match res { + Ok(true) => { + restored += 1; + pool_for_reconnect.stats.increment_me_reconnect_success(); + } + Ok(false) => { + debug!(dc = %dc, ?family, "ME round-robin reconnect failed") + } + Err(_) => { + debug!(dc = %dc, ?family, "ME reconnect timed out"); + } + } + drop(reconnect_permit); + } + + FamilyReconnectOutcome { + key, + dc, + family, + required, + endpoint_count: endpoints_for_dc.len(), + } + }); + } + + while let Some(joined) = reconnect_set.join_next().await { + let outcome = match joined { + Ok(outcome) => outcome, + Err(join_error) => { + debug!(error = %join_error, "Health reconnect task failed"); + continue; + } + }; + let now = Instant::now(); + let now_alive = live_active_writers_for_dc_family(pool, outcome.dc, outcome.family).await; + if now_alive >= outcome.required { + info!( + dc = %outcome.dc, + family = ?outcome.family, + alive = now_alive, + required = outcome.required, + endpoint_count = outcome.endpoint_count, + "ME writer floor restored for DC" + ); + backoff.insert( + outcome.key, + pool.reconnect_runtime.me_reconnect_backoff_base.as_millis() as u64, + ); + let jitter = pool.reconnect_runtime.me_reconnect_backoff_base.as_millis() as u64 + / JITTER_FRAC_NUM; + let wait = pool.reconnect_runtime.me_reconnect_backoff_base + + Duration::from_millis(rand::rng().random_range(0..=jitter.max(1))); + next_attempt.insert(outcome.key, now + wait); + } else { + let curr = *backoff + .get(&outcome.key) + .unwrap_or(&(pool.reconnect_runtime.me_reconnect_backoff_base.as_millis() as u64)); + let next_ms = (curr.saturating_mul(2)) + .min(pool.reconnect_runtime.me_reconnect_backoff_cap.as_millis() as u64); + backoff.insert(outcome.key, next_ms); + let jitter = next_ms / JITTER_FRAC_NUM; + let wait = Duration::from_millis(next_ms) + + Duration::from_millis(rand::rng().random_range(0..=jitter.max(1))); + next_attempt.insert(outcome.key, now + wait); + if pool.is_runtime_ready() { + let warn_cooldown = pool.warn_rate_limit_duration(); + if should_emit_rate_limited_warn( + floor_warn_next_allowed, + outcome.key, + now, + warn_cooldown, + ) { + warn!( + dc = %outcome.dc, + family = ?outcome.family, + alive = now_alive, + required = outcome.required, + endpoint_count = outcome.endpoint_count, + backoff_ms = next_ms, + "DC writer floor is below required level, scheduled reconnect" + ); + } + } else { + info!( + dc = %outcome.dc, + family = ?outcome.family, + alive = now_alive, + required = outcome.required, + endpoint_count = outcome.endpoint_count, + backoff_ms = next_ms, + "DC writer floor is below required level during startup, scheduled reconnect" + ); + } + } + } + + family_degraded +} + +pub(super) fn health_reconnect_budget(pool: &Arc, dc_groups: usize) -> usize { + let cpu_cores = pool.adaptive_floor_effective_cpu_cores().max(1); + let by_cpu = cpu_cores.saturating_mul(HEALTH_RECONNECT_BUDGET_PER_CORE); + let by_dc = dc_groups.saturating_mul(HEALTH_RECONNECT_BUDGET_PER_DC); + by_cpu + .saturating_add(by_dc) + .clamp(HEALTH_RECONNECT_BUDGET_MIN, HEALTH_RECONNECT_BUDGET_MAX) +} + +pub(super) fn update_family_runtime_state(pool: &Arc, family: IpFamily, degraded: bool) { + let now_epoch_secs = MePool::now_epoch_secs(); + let previous_state = pool.family_runtime_state(family); + let mut state_since_epoch_secs = pool.family_runtime_state_since_epoch_secs(family); + let previous_suppressed_until_epoch_secs = pool.family_suppressed_until_epoch_secs(family); + let previous_fail_streak = pool.family_fail_streak(family); + let previous_recover_success_streak = pool.family_recover_success_streak(family); + + let (next_state, suppressed_until_epoch_secs, fail_streak, recover_success_streak) = + if previous_suppressed_until_epoch_secs > now_epoch_secs { + let fail_streak = if degraded { + previous_fail_streak.saturating_add(1) + } else { + previous_fail_streak + }; + ( + MeFamilyRuntimeState::Suppressed, + previous_suppressed_until_epoch_secs, + fail_streak, + 0, + ) + } else if degraded { + let fail_streak = previous_fail_streak.saturating_add(1); + if fail_streak >= FAMILY_SUPPRESS_FAIL_STREAK_THRESHOLD { + ( + MeFamilyRuntimeState::Suppressed, + now_epoch_secs.saturating_add(FAMILY_SUPPRESS_DURATION_SECS), + fail_streak, + 0, + ) + } else { + (MeFamilyRuntimeState::Degraded, 0, fail_streak, 0) + } + } else if matches!(previous_state, MeFamilyRuntimeState::Healthy) { + (MeFamilyRuntimeState::Healthy, 0, 0, 0) + } else { + let recover_success_streak = previous_recover_success_streak.saturating_add(1); + if recover_success_streak >= FAMILY_RECOVER_SUCCESS_STREAK_TARGET { + (MeFamilyRuntimeState::Healthy, 0, 0, 0) + } else { + ( + MeFamilyRuntimeState::Recovering, + 0, + 0, + recover_success_streak, + ) + } + }; + + if next_state != previous_state || state_since_epoch_secs == 0 { + state_since_epoch_secs = now_epoch_secs; + } + pool.set_family_runtime_state( + family, + next_state, + state_since_epoch_secs, + suppressed_until_epoch_secs, + fail_streak, + recover_success_streak, + ); +} diff --git a/src/transport/middle_proxy/health/floor_plan.rs b/src/transport/middle_proxy/health/floor_plan.rs new file mode 100644 index 0000000..0ad16eb --- /dev/null +++ b/src/transport/middle_proxy/health/floor_plan.rs @@ -0,0 +1,261 @@ +use super::*; + +pub(super) fn should_emit_rate_limited_warn( + next_allowed: &mut HashMap<(i32, IpFamily), Instant>, + key: (i32, IpFamily), + now: Instant, + cooldown: Duration, +) -> bool { + let Some(ready_at) = next_allowed.get(&key).copied() else { + next_allowed.insert(key, now + cooldown); + return true; + }; + if now >= ready_at { + next_allowed.insert(key, now + cooldown); + return true; + } + false +} + +pub(super) async fn live_active_writers_for_dc_family( + pool: &Arc, + dc: i32, + family: IpFamily, +) -> usize { + let writers = pool.writers.read().await; + writers + .iter() + .filter(|writer| { + if writer.draining.load(std::sync::atomic::Ordering::Relaxed) { + return false; + } + if writer.writer_dc != dc { + return false; + } + if !matches!( + crate::transport::middle_proxy::pool::WriterContour::from_u8( + writer.contour.load(std::sync::atomic::Ordering::Relaxed), + ), + crate::transport::middle_proxy::pool::WriterContour::Active + ) { + return false; + } + match family { + IpFamily::V4 => writer.addr.is_ipv4(), + IpFamily::V6 => writer.addr.is_ipv6(), + } + }) + .count() +} + +pub(super) fn adaptive_floor_class_min( + pool: &Arc, + endpoint_count: usize, + base_required: usize, +) -> usize { + if endpoint_count <= 1 { + let min_single = (pool + .floor_runtime + .me_adaptive_floor_min_writers_single_endpoint + .load(std::sync::atomic::Ordering::Relaxed) as usize) + .max(1); + min_single.min(base_required.max(1)) + } else { + pool.adaptive_floor_min_writers_multi_endpoint() + .min(base_required.max(1)) + } +} + +pub(super) fn adaptive_floor_class_max( + pool: &Arc, + endpoint_count: usize, + base_required: usize, + cpu_cores: usize, +) -> usize { + let extra_per_core = if endpoint_count <= 1 { + pool.adaptive_floor_max_extra_single_per_core() + } else { + pool.adaptive_floor_max_extra_multi_per_core() + }; + base_required.saturating_add(cpu_cores.saturating_mul(extra_per_core)) +} + +pub(super) fn list_writer_ids_for_endpoints( + dc: i32, + endpoints: &[SocketAddr], + live_writer_ids_by_addr: &HashMap<(i32, SocketAddr), Vec>, +) -> Vec { + let mut out = Vec::::new(); + for endpoint in endpoints { + if let Some(ids) = live_writer_ids_by_addr.get(&(dc, *endpoint)) { + out.extend(ids.iter().copied()); + } + } + out +} + +pub(super) async fn build_family_floor_plan( + pool: &Arc, + family: IpFamily, + dc_endpoints: &HashMap>, + live_addr_counts: &HashMap<(i32, SocketAddr), usize>, + live_writer_ids_by_addr: &HashMap<(i32, SocketAddr), Vec>, + bound_clients_by_writer: &HashMap, +) -> FamilyFloorPlan { + let mut entries = Vec::::new(); + let mut by_dc = HashMap::::new(); + let mut family_active_total = 0usize; + + let floor_mode = pool.floor_mode(); + let is_adaptive = floor_mode == MeFloorMode::Adaptive; + let cpu_cores = pool.adaptive_floor_effective_cpu_cores().max(1); + let (active_writers_current, warm_writers_current, _) = + pool.non_draining_writer_counts_by_contour().await; + + for (dc, endpoints) in dc_endpoints { + if endpoints.is_empty() { + continue; + } + let _key = (*dc, family); + let base_required = pool.required_writers_for_dc(endpoints.len()).max(1); + let min_required = if is_adaptive { + adaptive_floor_class_min(pool, endpoints.len(), base_required) + } else { + base_required + }; + let mut max_required = if is_adaptive { + adaptive_floor_class_max(pool, endpoints.len(), base_required, cpu_cores) + } else { + base_required + }; + if max_required < min_required { + max_required = min_required; + } + // We initialize target_required at base_required to prevent 0-writer blackouts + // caused by proactively dropping an idle DC to a single fragile connection. + // The Adaptive Floor constraint loop below will gracefully compress idle DCs + // (prioritized via has_bound_clients = false) to min_required only when global capacity is reached. + let desired_raw = base_required; + let target_required = desired_raw.clamp(min_required, max_required); + let alive = endpoints + .iter() + .map(|endpoint| { + live_addr_counts + .get(&(*dc, *endpoint)) + .copied() + .unwrap_or(0) + }) + .sum::(); + family_active_total = family_active_total.saturating_add(alive); + let writer_ids = list_writer_ids_for_endpoints(*dc, endpoints, live_writer_ids_by_addr); + let has_bound_clients = has_bound_clients_on_endpoint(&writer_ids, bound_clients_by_writer); + + entries.push(DcFloorPlanEntry { + dc: *dc, + endpoints: endpoints.clone(), + alive, + min_required, + target_required, + max_required, + has_bound_clients, + floor_capped: false, + }); + } + + if entries.is_empty() { + let active_cap_configured_total = pool.adaptive_floor_active_cap_configured_total(); + let warm_cap_configured_total = pool.adaptive_floor_warm_cap_configured_total(); + return FamilyFloorPlan { + by_dc, + active_cap_configured_total, + active_cap_effective_total: active_cap_configured_total, + warm_cap_configured_total, + warm_cap_effective_total: warm_cap_configured_total, + active_writers_current, + warm_writers_current, + target_writers_total: 0, + }; + } + + if !is_adaptive { + let target_total = entries + .iter() + .map(|entry| entry.target_required) + .sum::(); + let active_cap_configured_total = pool.adaptive_floor_active_cap_configured_total(); + let warm_cap_configured_total = pool.adaptive_floor_warm_cap_configured_total(); + for entry in entries { + by_dc.insert(entry.dc, entry); + } + return FamilyFloorPlan { + by_dc, + active_cap_configured_total, + active_cap_effective_total: active_cap_configured_total.max(target_total), + warm_cap_configured_total, + warm_cap_effective_total: warm_cap_configured_total, + active_writers_current, + warm_writers_current, + target_writers_total: target_total, + }; + } + + let active_cap_configured_total = pool.adaptive_floor_active_cap_configured_total(); + let warm_cap_configured_total = pool.adaptive_floor_warm_cap_configured_total(); + let other_active = active_writers_current.saturating_sub(family_active_total); + let min_sum = entries + .iter() + .map(|entry| entry.min_required) + .sum::(); + let mut target_sum = entries + .iter() + .map(|entry| entry.target_required) + .sum::(); + let family_cap = active_cap_configured_total + .saturating_sub(other_active) + .max(min_sum); + if target_sum > family_cap { + entries.sort_by_key(|entry| { + ( + entry.has_bound_clients, + std::cmp::Reverse(entry.target_required.saturating_sub(entry.min_required)), + std::cmp::Reverse(entry.alive), + entry.dc.abs(), + entry.dc, + entry.endpoints.len(), + entry.max_required, + ) + }); + let mut changed = true; + while target_sum > family_cap && changed { + changed = false; + for entry in &mut entries { + if target_sum <= family_cap { + break; + } + if entry.target_required > entry.min_required { + entry.target_required -= 1; + entry.floor_capped = true; + target_sum -= 1; + changed = true; + } + } + } + } + + for entry in entries { + by_dc.insert(entry.dc, entry); + } + let active_cap_effective_total = + active_cap_configured_total.max(other_active.saturating_add(min_sum)); + let target_writers_total = other_active.saturating_add(target_sum); + FamilyFloorPlan { + by_dc, + active_cap_configured_total, + active_cap_effective_total, + warm_cap_configured_total, + warm_cap_effective_total: warm_cap_configured_total, + active_writers_current, + warm_writers_current, + target_writers_total, + } +} diff --git a/src/transport/middle_proxy/health/idle_refresh.rs b/src/transport/middle_proxy/health/idle_refresh.rs new file mode 100644 index 0000000..3cc26dc --- /dev/null +++ b/src/transport/middle_proxy/health/idle_refresh.rs @@ -0,0 +1,203 @@ +use super::*; + +pub(super) async fn maybe_swap_idle_writer_for_cap( + pool: &Arc, + rng: &Arc, + dc: i32, + family: IpFamily, + endpoints: &[SocketAddr], + live_writer_ids_by_addr: &HashMap<(i32, SocketAddr), Vec>, + writer_idle_since: &HashMap, + bound_clients_by_writer: &HashMap, +) -> bool { + let now_epoch_secs = MePool::now_epoch_secs(); + let mut candidate: Option<(u64, SocketAddr, u64)> = None; + for endpoint in endpoints { + let Some(writer_ids) = live_writer_ids_by_addr.get(&(dc, *endpoint)) else { + continue; + }; + for writer_id in writer_ids { + if bound_clients_by_writer.get(writer_id).copied().unwrap_or(0) > 0 { + continue; + } + let Some(idle_since_epoch_secs) = writer_idle_since.get(writer_id).copied() else { + continue; + }; + let idle_age_secs = now_epoch_secs.saturating_sub(idle_since_epoch_secs); + if candidate + .as_ref() + .map(|(_, _, age)| idle_age_secs > *age) + .unwrap_or(true) + { + candidate = Some((*writer_id, *endpoint, idle_age_secs)); + } + } + } + + let Some((old_writer_id, endpoint, idle_age_secs)) = candidate else { + return false; + }; + + let connected = match tokio::time::timeout( + pool.reconnect_runtime.me_one_timeout, + pool.connect_one_for_dc(endpoint, dc, rng.as_ref()), + ) + .await + { + Ok(Ok(())) => true, + Ok(Err(error)) => { + debug!( + dc = %dc, + ?family, + %endpoint, + old_writer_id, + idle_age_secs, + %error, + "Adaptive floor cap swap connect failed" + ); + false + } + Err(_) => { + debug!( + dc = %dc, + ?family, + %endpoint, + old_writer_id, + idle_age_secs, + "Adaptive floor cap swap connect timed out" + ); + false + } + }; + if !connected { + return false; + } + + pool.mark_writer_draining_with_timeout(old_writer_id, pool.force_close_timeout(), false) + .await; + info!( + dc = %dc, + ?family, + %endpoint, + old_writer_id, + idle_age_secs, + "Adaptive floor cap swap: idle writer rotated" + ); + true +} + +pub(super) async fn maybe_refresh_idle_writer_for_dc( + pool: &Arc, + rng: &Arc, + key: (i32, IpFamily), + dc: i32, + family: IpFamily, + endpoints: &[SocketAddr], + alive: usize, + required: usize, + live_writer_ids_by_addr: &HashMap<(i32, SocketAddr), Vec>, + writer_idle_since: &HashMap, + bound_clients_by_writer: &HashMap, + idle_refresh_next_attempt: &mut HashMap<(i32, IpFamily), Instant>, +) { + if alive < required { + return; + } + + let now = Instant::now(); + if let Some(next) = idle_refresh_next_attempt.get(&key) + && now < *next + { + return; + } + + let now_epoch_secs = MePool::now_epoch_secs(); + let mut candidate: Option<(u64, SocketAddr, u64, u64)> = None; + for endpoint in endpoints { + let Some(writer_ids) = live_writer_ids_by_addr.get(&(dc, *endpoint)) else { + continue; + }; + for writer_id in writer_ids { + if bound_clients_by_writer.get(writer_id).copied().unwrap_or(0) > 0 { + continue; + } + let Some(idle_since_epoch_secs) = writer_idle_since.get(writer_id).copied() else { + continue; + }; + let idle_age_secs = now_epoch_secs.saturating_sub(idle_since_epoch_secs); + let threshold_secs = IDLE_REFRESH_TRIGGER_BASE_SECS + + (*writer_id % (IDLE_REFRESH_TRIGGER_JITTER_SECS + 1)); + if idle_age_secs < threshold_secs { + continue; + } + if candidate + .as_ref() + .map(|(_, _, age, _)| idle_age_secs > *age) + .unwrap_or(true) + { + candidate = Some((*writer_id, *endpoint, idle_age_secs, threshold_secs)); + } + } + } + + let Some((old_writer_id, endpoint, idle_age_secs, threshold_secs)) = candidate else { + return; + }; + + let rotate_ok = match tokio::time::timeout( + pool.reconnect_runtime.me_one_timeout, + pool.connect_one_for_dc(endpoint, dc, rng.as_ref()), + ) + .await + { + Ok(Ok(())) => true, + Ok(Err(error)) => { + debug!( + dc = %dc, + ?family, + %endpoint, + old_writer_id, + idle_age_secs, + threshold_secs, + %error, + "Idle writer pre-refresh connect failed" + ); + false + } + Err(_) => { + debug!( + dc = %dc, + ?family, + %endpoint, + old_writer_id, + idle_age_secs, + threshold_secs, + "Idle writer pre-refresh connect timed out" + ); + false + } + }; + + if !rotate_ok { + idle_refresh_next_attempt.insert(key, now + Duration::from_secs(IDLE_REFRESH_RETRY_SECS)); + return; + } + + pool.mark_writer_draining_with_timeout(old_writer_id, pool.force_close_timeout(), false) + .await; + idle_refresh_next_attempt.insert( + key, + now + Duration::from_secs(IDLE_REFRESH_SUCCESS_GUARD_SECS), + ); + info!( + dc = %dc, + ?family, + %endpoint, + old_writer_id, + idle_age_secs, + threshold_secs, + alive, + required, + "Idle writer refreshed before upstream idle timeout" + ); +} diff --git a/src/transport/middle_proxy/health/monitor.rs b/src/transport/middle_proxy/health/monitor.rs new file mode 100644 index 0000000..1388613 --- /dev/null +++ b/src/transport/middle_proxy/health/monitor.rs @@ -0,0 +1,59 @@ +use super::*; + +pub async fn me_health_monitor(pool: Arc, rng: Arc, _min_connections: usize) { + let mut backoff: HashMap<(i32, IpFamily), u64> = HashMap::new(); + let mut next_attempt: HashMap<(i32, IpFamily), Instant> = HashMap::new(); + let mut inflight: HashMap<(i32, IpFamily), usize> = HashMap::new(); + let mut outage_backoff: HashMap<(i32, IpFamily), u64> = HashMap::new(); + let mut outage_next_attempt: HashMap<(i32, IpFamily), Instant> = HashMap::new(); + let mut single_endpoint_outage: HashSet<(i32, IpFamily)> = HashSet::new(); + let mut shadow_rotate_deadline: HashMap<(i32, IpFamily), Instant> = HashMap::new(); + let mut idle_refresh_next_attempt: HashMap<(i32, IpFamily), Instant> = HashMap::new(); + let mut floor_warn_next_allowed: HashMap<(i32, IpFamily), Instant> = HashMap::new(); + let mut drain_warn_next_allowed: HashMap = HashMap::new(); + let mut degraded_interval = true; + loop { + let interval = if degraded_interval { + pool.health_interval_unhealthy() + } else { + pool.health_interval_healthy() + }; + tokio::time::sleep(interval).await; + pool.prune_closed_writers().await; + pool.sweep_endpoint_quarantine().await; + reap_draining_writers(&pool, &mut drain_warn_next_allowed).await; + let v4_degraded = check_family( + IpFamily::V4, + &pool, + &rng, + &mut backoff, + &mut next_attempt, + &mut inflight, + &mut outage_backoff, + &mut outage_next_attempt, + &mut single_endpoint_outage, + &mut shadow_rotate_deadline, + &mut idle_refresh_next_attempt, + &mut floor_warn_next_allowed, + ) + .await; + let v6_degraded = check_family( + IpFamily::V6, + &pool, + &rng, + &mut backoff, + &mut next_attempt, + &mut inflight, + &mut outage_backoff, + &mut outage_next_attempt, + &mut single_endpoint_outage, + &mut shadow_rotate_deadline, + &mut idle_refresh_next_attempt, + &mut floor_warn_next_allowed, + ) + .await; + update_family_runtime_state(&pool, IpFamily::V4, v4_degraded); + update_family_runtime_state(&pool, IpFamily::V6, v6_degraded); + degraded_interval = v4_degraded || v6_degraded; + } +} diff --git a/src/transport/middle_proxy/health/recovery.rs b/src/transport/middle_proxy/health/recovery.rs new file mode 100644 index 0000000..cfeee67 --- /dev/null +++ b/src/transport/middle_proxy/health/recovery.rs @@ -0,0 +1,262 @@ +use super::*; + +pub(super) fn has_bound_clients_on_endpoint( + writer_ids: &[u64], + bound_clients_by_writer: &HashMap, +) -> bool { + writer_ids + .iter() + .any(|writer_id| bound_clients_by_writer.get(writer_id).copied().unwrap_or(0) > 0) +} + +pub(super) async fn recover_single_endpoint_outage( + pool: &Arc, + rng: &Arc, + key: (i32, IpFamily), + endpoint: SocketAddr, + required: usize, + outage_backoff: &mut HashMap<(i32, IpFamily), u64>, + outage_next_attempt: &mut HashMap<(i32, IpFamily), Instant>, + reconnect_sem: &Arc, +) { + let now = Instant::now(); + if let Some(ts) = outage_next_attempt.get(&key) + && now < *ts + { + return; + } + + let (min_backoff_ms, max_backoff_ms) = pool.single_endpoint_outage_backoff_bounds_ms(); + if reconnect_sem.available_permits() == 0 { + outage_next_attempt.insert(key, now + Duration::from_millis(min_backoff_ms.max(250))); + debug!( + dc = %key.0, + family = ?key.1, + %endpoint, + required, + "Single-endpoint outage reconnect deferred by health reconnect budget" + ); + return; + } + let Ok(_reconnect_permit) = reconnect_sem.clone().try_acquire_owned() else { + outage_next_attempt.insert(key, now + Duration::from_millis(min_backoff_ms.max(250))); + debug!( + dc = %key.0, + family = ?key.1, + %endpoint, + required, + "Single-endpoint outage reconnect deferred by semaphore saturation" + ); + return; + }; + pool.stats.increment_me_reconnect_attempt(); + pool.stats + .increment_me_single_endpoint_outage_reconnect_attempt_total(); + + let bypass_quarantine = pool.single_endpoint_outage_disable_quarantine(); + let attempt_ok = if bypass_quarantine { + pool.stats + .increment_me_single_endpoint_quarantine_bypass_total(); + match tokio::time::timeout( + pool.reconnect_runtime.me_one_timeout, + pool.connect_one_for_dc(endpoint, key.0, rng.as_ref()), + ) + .await + { + Ok(Ok(())) => true, + Ok(Err(e)) => { + debug!( + dc = %key.0, + family = ?key.1, + %endpoint, + error = %e, + "Single-endpoint outage reconnect failed (quarantine bypass path)" + ); + false + } + Err(_) => { + debug!( + dc = %key.0, + family = ?key.1, + %endpoint, + "Single-endpoint outage reconnect timed out (quarantine bypass path)" + ); + false + } + } + } else { + let one_endpoint = [endpoint]; + match tokio::time::timeout( + pool.reconnect_runtime.me_one_timeout, + pool.connect_endpoints_round_robin(key.0, &one_endpoint, rng.as_ref()), + ) + .await + { + Ok(ok) => ok, + Err(_) => { + debug!( + dc = %key.0, + family = ?key.1, + %endpoint, + "Single-endpoint outage reconnect timed out" + ); + false + } + } + }; + + if attempt_ok { + pool.stats + .increment_me_single_endpoint_outage_reconnect_success_total(); + pool.stats.increment_me_reconnect_success(); + outage_backoff.insert(key, min_backoff_ms); + let jitter = min_backoff_ms / JITTER_FRAC_NUM; + let wait = Duration::from_millis(min_backoff_ms) + + Duration::from_millis(rand::rng().random_range(0..=jitter.max(1))); + outage_next_attempt.insert(key, now + wait); + info!( + dc = %key.0, + family = ?key.1, + %endpoint, + required, + backoff_ms = min_backoff_ms, + "Single-endpoint outage reconnect succeeded" + ); + return; + } + + let current_ms = *outage_backoff.get(&key).unwrap_or(&min_backoff_ms); + let next_ms = current_ms.saturating_mul(2).min(max_backoff_ms); + outage_backoff.insert(key, next_ms); + let jitter = next_ms / JITTER_FRAC_NUM; + let wait = Duration::from_millis(next_ms) + + Duration::from_millis(rand::rng().random_range(0..=jitter.max(1))); + outage_next_attempt.insert(key, now + wait); + warn!( + dc = %key.0, + family = ?key.1, + %endpoint, + required, + backoff_ms = next_ms, + "Single-endpoint outage reconnect scheduled" + ); +} + +pub(super) async fn maybe_rotate_single_endpoint_shadow( + pool: &Arc, + rng: &Arc, + key: (i32, IpFamily), + dc: i32, + family: IpFamily, + endpoints: &[SocketAddr], + alive: usize, + required: usize, + live_writer_ids_by_addr: &HashMap<(i32, SocketAddr), Vec>, + bound_clients_by_writer: &HashMap, + shadow_rotate_deadline: &mut HashMap<(i32, IpFamily), Instant>, +) { + if endpoints.len() != 1 || alive < required { + return; + } + + let Some(interval) = pool.single_endpoint_shadow_rotate_interval() else { + return; + }; + + let now = Instant::now(); + if let Some(deadline) = shadow_rotate_deadline.get(&key) + && now < *deadline + { + return; + } + + let endpoint = endpoints[0]; + if pool.is_endpoint_quarantined(endpoint).await { + pool.stats + .increment_me_single_endpoint_shadow_rotate_skipped_quarantine_total(); + shadow_rotate_deadline.insert(key, now + Duration::from_secs(SHADOW_ROTATE_RETRY_SECS)); + debug!( + dc = %dc, + ?family, + %endpoint, + "Single-endpoint shadow rotation skipped: endpoint is quarantined" + ); + return; + } + + let Some(writer_ids) = live_writer_ids_by_addr.get(&(dc, endpoint)) else { + shadow_rotate_deadline.insert(key, now + Duration::from_secs(SHADOW_ROTATE_RETRY_SECS)); + return; + }; + + let mut candidate_writer_id = None; + for writer_id in writer_ids { + if bound_clients_by_writer.get(writer_id).copied().unwrap_or(0) == 0 { + candidate_writer_id = Some(*writer_id); + break; + } + } + + let Some(old_writer_id) = candidate_writer_id else { + shadow_rotate_deadline.insert(key, now + Duration::from_secs(SHADOW_ROTATE_RETRY_SECS)); + debug!( + dc = %dc, + ?family, + %endpoint, + alive, + required, + "Single-endpoint shadow rotation skipped: no empty writer candidate" + ); + return; + }; + + let rotate_ok = match tokio::time::timeout( + pool.reconnect_runtime.me_one_timeout, + pool.connect_one_for_dc(endpoint, dc, rng.as_ref()), + ) + .await + { + Ok(Ok(())) => true, + Ok(Err(e)) => { + debug!( + dc = %dc, + ?family, + %endpoint, + error = %e, + "Single-endpoint shadow rotation connect failed" + ); + false + } + Err(_) => { + debug!( + dc = %dc, + ?family, + %endpoint, + "Single-endpoint shadow rotation connect timed out" + ); + false + } + }; + + if !rotate_ok { + shadow_rotate_deadline.insert( + key, + now + interval.min(Duration::from_secs(SHADOW_ROTATE_RETRY_SECS)), + ); + return; + } + + pool.mark_writer_draining_with_timeout(old_writer_id, pool.force_close_timeout(), false) + .await; + pool.stats + .increment_me_single_endpoint_shadow_rotate_total(); + shadow_rotate_deadline.insert(key, now + interval); + info!( + dc = %dc, + ?family, + %endpoint, + old_writer_id, + rotate_every_secs = interval.as_secs(), + "Single-endpoint shadow writer rotated" + ); +} diff --git a/src/transport/middle_proxy/health/tests.rs b/src/transport/middle_proxy/health/tests.rs new file mode 100644 index 0000000..0e612ef --- /dev/null +++ b/src/transport/middle_proxy/health/tests.rs @@ -0,0 +1,323 @@ +use std::collections::HashMap; +use std::net::{IpAddr, Ipv4Addr, SocketAddr}; +use std::sync::Arc; +use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU32, AtomicU64, Ordering}; +use std::time::{Duration, Instant}; + +use tokio::sync::mpsc; +use tokio_util::sync::CancellationToken; + +use super::{ScheduledReconnects, reap_draining_writers}; +use crate::config::{GeneralConfig, MeRouteNoWriterMode, MeSocksKdfPolicy, MeWriterPickMode}; +use crate::crypto::SecureRandom; +use crate::network::IpFamily; +use crate::network::probe::NetworkDecision; +use crate::stats::Stats; +use crate::transport::middle_proxy::codec::WriterCommand; +use crate::transport::middle_proxy::pool::{MePool, MeWriter, WriterContour}; +use crate::transport::middle_proxy::registry::ConnMeta; + +#[test] +fn reconnect_batch_releases_every_reserved_key_after_join_failures() { + let retained = (1, IpFamily::V4); + let removed = (2, IpFamily::V6); + let mut inflight = HashMap::from([(retained, 1)]); + + { + let mut scheduled = ScheduledReconnects { + inflight: &mut inflight, + keys: Vec::new(), + }; + scheduled.reserve(retained); + scheduled.reserve(removed); + } + + assert_eq!(inflight.get(&retained), Some(&1)); + assert!(!inflight.contains_key(&removed)); +} + +async fn make_pool(me_pool_drain_threshold: u64) -> Arc { + let general = GeneralConfig { + me_pool_drain_threshold, + ..GeneralConfig::default() + }; + let mut proxy_map_v4 = HashMap::new(); + proxy_map_v4.insert(2, vec![(IpAddr::V4(Ipv4Addr::new(203, 0, 113, 10)), 443)]); + let decision = NetworkDecision { + ipv4_me: true, + ..NetworkDecision::default() + }; + MePool::new( + None, + vec![1u8; 32], + None, + false, + None, + Vec::new(), + false, + Vec::new(), + 1, + None, + 12, + 1200, + proxy_map_v4, + HashMap::new(), + None, + decision, + None, + Arc::new(SecureRandom::new()), + Arc::new(Stats::default()), + general.me_keepalive_enabled, + general.me_keepalive_interval_secs, + general.me_keepalive_jitter_secs, + general.me_keepalive_payload_random, + general.rpc_proxy_req_every, + general.me_warmup_stagger_enabled, + general.me_warmup_step_delay_ms, + general.me_warmup_step_jitter_ms, + general.me_reconnect_max_concurrent_per_dc, + general.me_reconnect_backoff_base_ms, + general.me_reconnect_backoff_cap_ms, + general.me_reconnect_fast_retry_count, + general.me_single_endpoint_shadow_writers, + general.me_single_endpoint_outage_mode_enabled, + general.me_single_endpoint_outage_disable_quarantine, + general.me_single_endpoint_outage_backoff_min_ms, + general.me_single_endpoint_outage_backoff_max_ms, + general.me_single_endpoint_shadow_rotate_every_secs, + general.me_floor_mode, + general.me_adaptive_floor_idle_secs, + general.me_adaptive_floor_min_writers_single_endpoint, + general.me_adaptive_floor_min_writers_multi_endpoint, + general.me_adaptive_floor_recover_grace_secs, + general.me_adaptive_floor_writers_per_core_total, + general.me_adaptive_floor_cpu_cores_override, + general.me_adaptive_floor_max_extra_writers_single_per_core, + general.me_adaptive_floor_max_extra_writers_multi_per_core, + general.me_adaptive_floor_max_active_writers_per_core, + general.me_adaptive_floor_max_warm_writers_per_core, + general.me_adaptive_floor_max_active_writers_global, + general.me_adaptive_floor_max_warm_writers_global, + general.hardswap, + general.me_pool_drain_ttl_secs, + general.me_instadrain, + general.me_pool_drain_threshold, + general.me_pool_drain_soft_evict_enabled, + general.me_pool_drain_soft_evict_grace_secs, + general.me_pool_drain_soft_evict_per_writer, + general.me_pool_drain_soft_evict_budget_per_core, + general.me_pool_drain_soft_evict_cooldown_ms, + general.effective_me_pool_force_close_secs(), + general.me_pool_min_fresh_ratio, + general.me_hardswap_warmup_delay_min_ms, + general.me_hardswap_warmup_delay_max_ms, + general.me_hardswap_warmup_extra_passes, + general.me_hardswap_warmup_pass_backoff_base_ms, + general.me_bind_stale_mode, + general.me_bind_stale_ttl_secs, + general.me_secret_atomic_snapshot, + general.me_deterministic_writer_sort, + MeWriterPickMode::default(), + general.me_writer_pick_sample_size, + MeSocksKdfPolicy::default(), + general.me_writer_cmd_channel_capacity, + general.me_writer_byte_budget_bytes, + general.me_route_channel_capacity, + general.me_route_backpressure_enabled, + general.me_route_fairshare_enabled, + general.me_route_backpressure_base_timeout_ms, + general.me_route_backpressure_high_timeout_ms, + general.me_route_backpressure_high_watermark_pct, + general.me_reader_route_data_wait_ms, + general.me_health_interval_ms_unhealthy, + general.me_health_interval_ms_healthy, + general.me_warn_rate_limit_ms, + MeRouteNoWriterMode::default(), + general.me_route_no_writer_wait_ms, + general.me_route_hybrid_max_wait_ms, + general.me_route_blocking_send_timeout_ms, + general.me_route_inline_recovery_attempts, + general.me_route_inline_recovery_wait_ms, + 16_384, + ) +} + +async fn insert_draining_writer( + pool: &Arc, + writer_id: u64, + drain_started_at_epoch_secs: u64, +) -> u64 { + let (conn_id, _rx) = pool.registry.register().await; + let (tx, _writer_rx) = mpsc::channel::(8); + let byte_budget = pool.new_writer_byte_budget(); + let writer = MeWriter { + id: writer_id, + addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 4000 + writer_id as u16), + source_ip: IpAddr::V4(Ipv4Addr::LOCALHOST), + writer_dc: 2, + generation: 1, + contour: Arc::new(AtomicU8::new(WriterContour::Draining.as_u8())), + created_at: Instant::now() - Duration::from_secs(writer_id), + tx: tx.clone(), + byte_budget: byte_budget.clone(), + cancel: CancellationToken::new(), + degraded: Arc::new(AtomicBool::new(false)), + rtt_ema_ms_x10: Arc::new(AtomicU32::new(0)), + draining: Arc::new(AtomicBool::new(true)), + draining_started_at_epoch_secs: Arc::new(AtomicU64::new(drain_started_at_epoch_secs)), + drain_deadline_epoch_secs: Arc::new(AtomicU64::new(0)), + allow_drain_fallback: Arc::new(AtomicBool::new(false)), + }; + pool.writers.write().await.push(writer); + pool.registry + .register_writer(writer_id, tx, byte_budget) + .await; + pool.conn_count.fetch_add(1, Ordering::Relaxed); + assert!( + pool.registry + .bind_writer( + conn_id, + writer_id, + ConnMeta { + target_dc: 2, + client_addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 6000), + our_addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443), + proto_flags: 0, + }, + ) + .await + ); + conn_id +} + +async fn insert_live_writer(pool: &Arc, writer_id: u64, writer_dc: i32) { + let (tx, _writer_rx) = mpsc::channel::(8); + let byte_budget = pool.new_writer_byte_budget(); + let writer = MeWriter { + id: writer_id, + addr: SocketAddr::new( + IpAddr::V4(Ipv4Addr::new( + 203, + 0, + 113, + (writer_id as u8).saturating_add(1), + )), + 4000 + writer_id as u16, + ), + source_ip: IpAddr::V4(Ipv4Addr::LOCALHOST), + writer_dc, + generation: 2, + contour: Arc::new(AtomicU8::new(WriterContour::Active.as_u8())), + created_at: Instant::now(), + tx: tx.clone(), + byte_budget: byte_budget.clone(), + cancel: CancellationToken::new(), + degraded: Arc::new(AtomicBool::new(false)), + rtt_ema_ms_x10: Arc::new(AtomicU32::new(0)), + draining: Arc::new(AtomicBool::new(false)), + draining_started_at_epoch_secs: Arc::new(AtomicU64::new(0)), + drain_deadline_epoch_secs: Arc::new(AtomicU64::new(0)), + allow_drain_fallback: Arc::new(AtomicBool::new(false)), + }; + pool.writers.write().await.push(writer); + pool.registry + .register_writer(writer_id, tx, byte_budget) + .await; + pool.conn_count.fetch_add(1, Ordering::Relaxed); +} + +#[tokio::test] +async fn reap_draining_writers_force_closes_oldest_over_threshold() { + let pool = make_pool(2).await; + insert_live_writer(&pool, 1, 2).await; + let now_epoch_secs = MePool::now_epoch_secs(); + let conn_a = insert_draining_writer(&pool, 10, now_epoch_secs.saturating_sub(30)).await; + let conn_b = insert_draining_writer(&pool, 20, now_epoch_secs.saturating_sub(20)).await; + let conn_c = insert_draining_writer(&pool, 30, now_epoch_secs.saturating_sub(10)).await; + let mut warn_next_allowed = HashMap::new(); + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + + let mut writer_ids: Vec = pool + .writers + .read() + .await + .iter() + .map(|writer| writer.id) + .collect(); + writer_ids.sort_unstable(); + assert_eq!(writer_ids, vec![1, 20, 30]); + assert!(pool.registry.get_writer(conn_a).await.is_none()); + assert_eq!( + pool.registry.get_writer(conn_b).await.unwrap().writer_id, + 20 + ); + assert_eq!( + pool.registry.get_writer(conn_c).await.unwrap().writer_id, + 30 + ); +} + +#[tokio::test] +async fn reap_draining_writers_force_closes_overflow_without_replacement() { + let pool = make_pool(2).await; + let now_epoch_secs = MePool::now_epoch_secs(); + let conn_a = insert_draining_writer(&pool, 10, now_epoch_secs.saturating_sub(30)).await; + let conn_b = insert_draining_writer(&pool, 20, now_epoch_secs.saturating_sub(20)).await; + let conn_c = insert_draining_writer(&pool, 30, now_epoch_secs.saturating_sub(10)).await; + let mut warn_next_allowed = HashMap::new(); + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + + let mut writer_ids: Vec = pool + .writers + .read() + .await + .iter() + .map(|writer| writer.id) + .collect(); + writer_ids.sort_unstable(); + assert_eq!(writer_ids, vec![20, 30]); + assert!(pool.registry.get_writer(conn_a).await.is_none()); + assert_eq!( + pool.registry.get_writer(conn_b).await.unwrap().writer_id, + 20 + ); + assert_eq!( + pool.registry.get_writer(conn_c).await.unwrap().writer_id, + 30 + ); +} + +#[tokio::test] +async fn reap_draining_writers_keeps_timeout_only_behavior_when_threshold_disabled() { + let pool = make_pool(0).await; + let now_epoch_secs = MePool::now_epoch_secs(); + let conn_a = insert_draining_writer(&pool, 10, now_epoch_secs.saturating_sub(30)).await; + let conn_b = insert_draining_writer(&pool, 20, now_epoch_secs.saturating_sub(20)).await; + let conn_c = insert_draining_writer(&pool, 30, now_epoch_secs.saturating_sub(10)).await; + let mut warn_next_allowed = HashMap::new(); + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + + let writer_ids: Vec = pool + .writers + .read() + .await + .iter() + .map(|writer| writer.id) + .collect(); + assert_eq!(writer_ids, vec![10, 20, 30]); + assert_eq!( + pool.registry.get_writer(conn_a).await.unwrap().writer_id, + 10 + ); + assert_eq!( + pool.registry.get_writer(conn_b).await.unwrap().writer_id, + 20 + ); + assert_eq!( + pool.registry.get_writer(conn_c).await.unwrap().writer_id, + 30 + ); +} diff --git a/src/transport/middle_proxy/health/zombie_watchdog.rs b/src/transport/middle_proxy/health/zombie_watchdog.rs new file mode 100644 index 0000000..074cad4 --- /dev/null +++ b/src/transport/middle_proxy/health/zombie_watchdog.rs @@ -0,0 +1,146 @@ +use super::*; + +/// Last-resort safety net for draining writers stuck past their deadline. +/// +/// Runs periodically and force-closes draining writers that remain past their deadline. +/// The watchdog keeps independent lock and removal timeouts so a stalled writer cannot +/// prevent subsequent writers from being inspected. +pub async fn me_zombie_writer_watchdog(pool: Arc) { + use std::time::{SystemTime, UNIX_EPOCH}; + + const TICK_SECS: u64 = 30; + const SOFT_THRESHOLD_SECS: u64 = 60; + const HARD_THRESHOLD_SECS: u64 = 300; + const LOCK_TIMEOUT_SECS: u64 = 5; + const REMOVE_TIMEOUT_SECS: u64 = 10; + const HARD_DETACH_TIMEOUT_STREAK: u8 = 3; + + let mut removal_timeout_streak = HashMap::::new(); + + loop { + tokio::time::sleep(Duration::from_secs(TICK_SECS)).await; + + let now = match SystemTime::now().duration_since(UNIX_EPOCH) { + Ok(d) => d.as_secs(), + Err(_) => continue, + }; + + // Phase 1: collect zombie IDs under a short read-lock with timeout. + let zombie_ids_with_meta: Vec<(u64, bool)> = { + let Ok(ws) = + tokio::time::timeout(Duration::from_secs(LOCK_TIMEOUT_SECS), pool.writers.read()) + .await + else { + warn!("zombie_watchdog: writers read-lock timeout, skipping tick"); + continue; + }; + ws.iter() + .filter(|w| w.draining.load(std::sync::atomic::Ordering::Relaxed)) + .filter_map(|w| { + let deadline = w + .drain_deadline_epoch_secs + .load(std::sync::atomic::Ordering::Relaxed); + if deadline == 0 { + return None; + } + let overdue = now.saturating_sub(deadline); + if overdue == 0 { + return None; + } + let started = w + .draining_started_at_epoch_secs + .load(std::sync::atomic::Ordering::Relaxed); + let drain_age = now.saturating_sub(started); + if drain_age > HARD_THRESHOLD_SECS { + return Some((w.id, true)); + } + if overdue > SOFT_THRESHOLD_SECS { + return Some((w.id, false)); + } + None + }) + .collect() + }; + // read lock released here + + if zombie_ids_with_meta.is_empty() { + removal_timeout_streak.clear(); + continue; + } + + let mut active_zombie_ids = HashSet::::with_capacity(zombie_ids_with_meta.len()); + for (writer_id, _) in &zombie_ids_with_meta { + active_zombie_ids.insert(*writer_id); + } + removal_timeout_streak.retain(|writer_id, _| active_zombie_ids.contains(writer_id)); + + warn!( + zombie_count = zombie_ids_with_meta.len(), + soft_threshold_secs = SOFT_THRESHOLD_SECS, + hard_threshold_secs = HARD_THRESHOLD_SECS, + "Zombie draining writers detected by watchdog, force-closing" + ); + + // Phase 2: remove each writer individually with a timeout. + // One stuck removal cannot block the rest. + for (writer_id, had_clients) in &zombie_ids_with_meta { + let result = tokio::time::timeout( + Duration::from_secs(REMOVE_TIMEOUT_SECS), + pool.remove_writer_and_close_clients(*writer_id), + ) + .await; + match result { + Ok(()) => { + removal_timeout_streak.remove(writer_id); + pool.stats.increment_pool_force_close_total(); + info!(writer_id, had_clients, "Zombie writer removed by watchdog"); + } + Err(_) => { + let streak = removal_timeout_streak + .entry(*writer_id) + .and_modify(|value| *value = value.saturating_add(1)) + .or_insert(1); + warn!( + writer_id, + had_clients, + timeout_streak = *streak, + "Zombie writer removal timed out" + ); + if *streak < HARD_DETACH_TIMEOUT_STREAK { + continue; + } + + let hard_detach = tokio::time::timeout( + Duration::from_secs(REMOVE_TIMEOUT_SECS), + pool.remove_draining_writer_hard_detach(*writer_id), + ) + .await; + match hard_detach { + Ok(true) => { + removal_timeout_streak.remove(writer_id); + pool.stats.increment_pool_force_close_total(); + info!( + writer_id, + had_clients, "Zombie writer hard-detached after repeated timeouts" + ); + } + Ok(false) => { + removal_timeout_streak.remove(writer_id); + debug!( + writer_id, + had_clients, + "Zombie hard-detach skipped (writer already gone or no longer draining)" + ); + } + Err(_) => { + warn!( + writer_id, + had_clients, "Zombie hard-detach timed out, will retry next tick" + ); + } + } + } + } + } + } +} diff --git a/src/transport/middle_proxy/http_fetch.rs b/src/transport/middle_proxy/http_fetch.rs index 938dfac..ab747b5 100644 --- a/src/transport/middle_proxy/http_fetch.rs +++ b/src/transport/middle_proxy/http_fetch.rs @@ -159,15 +159,15 @@ pub(crate) async fn https_get( HTTP_REQUEST_TIMEOUT, Limited::new(response.into_body(), max_body_bytes).collect(), ) - .await - .map_err(|_| ProxyError::Proxy(format!("HTTP body read timeout for {url}")))? - .map_err(|e| { - ProxyError::Proxy(format!( - "HTTP body read failed or exceeded {max_body_bytes} bytes for {url}: {e}" - )) - })? - .to_bytes() - .to_vec(); + .await + .map_err(|_| ProxyError::Proxy(format!("HTTP body read timeout for {url}")))? + .map_err(|e| { + ProxyError::Proxy(format!( + "HTTP body read failed or exceeded {max_body_bytes} bytes for {url}: {e}" + )) + })? + .to_bytes() + .to_vec(); Ok(HttpsGetResponse { status, diff --git a/src/transport/middle_proxy/mod.rs b/src/transport/middle_proxy/mod.rs index feb658d..2c38fb5 100644 --- a/src/transport/middle_proxy/mod.rs +++ b/src/transport/middle_proxy/mod.rs @@ -61,9 +61,9 @@ pub use ping::{ MePingFamily, MePingReport, MePingSample, format_me_route, format_sample_line, run_me_ping, }; pub use pool::MePool; -pub(crate) use registry::ConnLease; #[allow(unused_imports)] pub use pool_nat::{detect_public_ip, stun_probe}; +pub(crate) use registry::ConnLease; pub use registry::ConnRegistry; pub use rotation::{MeReinitTrigger, me_reinit_scheduler, me_rotation_task}; #[allow(unused_imports)] diff --git a/src/transport/middle_proxy/pool.rs b/src/transport/middle_proxy/pool.rs index 30e6d94..232b556 100644 --- a/src/transport/middle_proxy/pool.rs +++ b/src/transport/middle_proxy/pool.rs @@ -506,1660 +506,15 @@ pub struct NatReflectionCache { pub v6: Option<(std::time::Instant, std::net::SocketAddr)>, } -impl MePool { - fn ratio_to_permille(ratio: f32) -> u32 { - let clamped = ratio.clamp(0.0, 1.0); - (clamped * 1000.0).round() as u32 - } - - pub(super) fn permille_to_ratio(permille: u32) -> f32 { - (permille.min(1000) as f32) / 1000.0 - } - - pub(super) fn now_epoch_secs() -> u64 { - SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap_or_default() - .as_secs() - } - - fn normalize_force_close_secs(force_close_secs: u64) -> u64 { - if force_close_secs == 0 { - ME_FORCE_CLOSE_SAFETY_FALLBACK_SECS - } else { - force_close_secs - } - } - - pub fn new( - proxy_tag: Option>, - proxy_secret: Vec, - nat_ip: Option, - nat_probe: bool, - nat_stun: Option, - nat_stun_servers: Vec, - stun_tcp_fallback: bool, - http_ip_detect_urls: Vec, - nat_probe_concurrency: usize, - detected_ipv6: Option, - me_one_retry: u8, - me_one_timeout_ms: u64, - proxy_map_v4: HashMap>, - proxy_map_v6: HashMap>, - default_dc: Option, - decision: NetworkDecision, - upstream: Option>, - rng: Arc, - stats: Arc, - me_keepalive_enabled: bool, - me_keepalive_interval_secs: u64, - me_keepalive_jitter_secs: u64, - me_keepalive_payload_random: bool, - rpc_proxy_req_every_secs: u64, - me_warmup_stagger_enabled: bool, - me_warmup_step_delay_ms: u64, - me_warmup_step_jitter_ms: u64, - me_reconnect_max_concurrent_per_dc: u32, - me_reconnect_backoff_base_ms: u64, - me_reconnect_backoff_cap_ms: u64, - me_reconnect_fast_retry_count: u32, - me_single_endpoint_shadow_writers: u8, - me_single_endpoint_outage_mode_enabled: bool, - me_single_endpoint_outage_disable_quarantine: bool, - me_single_endpoint_outage_backoff_min_ms: u64, - me_single_endpoint_outage_backoff_max_ms: u64, - me_single_endpoint_shadow_rotate_every_secs: u64, - me_floor_mode: MeFloorMode, - me_adaptive_floor_idle_secs: u64, - me_adaptive_floor_min_writers_single_endpoint: u8, - me_adaptive_floor_min_writers_multi_endpoint: u8, - me_adaptive_floor_recover_grace_secs: u64, - me_adaptive_floor_writers_per_core_total: u16, - me_adaptive_floor_cpu_cores_override: u16, - me_adaptive_floor_max_extra_writers_single_per_core: u16, - me_adaptive_floor_max_extra_writers_multi_per_core: u16, - me_adaptive_floor_max_active_writers_per_core: u16, - me_adaptive_floor_max_warm_writers_per_core: u16, - me_adaptive_floor_max_active_writers_global: u32, - me_adaptive_floor_max_warm_writers_global: u32, - hardswap: bool, - me_pool_drain_ttl_secs: u64, - me_instadrain: bool, - me_pool_drain_threshold: u64, - me_pool_drain_soft_evict_enabled: bool, - me_pool_drain_soft_evict_grace_secs: u64, - me_pool_drain_soft_evict_per_writer: u8, - me_pool_drain_soft_evict_budget_per_core: u16, - me_pool_drain_soft_evict_cooldown_ms: u64, - me_pool_force_close_secs: u64, - me_pool_min_fresh_ratio: f32, - me_hardswap_warmup_delay_min_ms: u64, - me_hardswap_warmup_delay_max_ms: u64, - me_hardswap_warmup_extra_passes: u8, - me_hardswap_warmup_pass_backoff_base_ms: u64, - me_bind_stale_mode: MeBindStaleMode, - me_bind_stale_ttl_secs: u64, - me_secret_atomic_snapshot: bool, - me_deterministic_writer_sort: bool, - me_writer_pick_mode: MeWriterPickMode, - me_writer_pick_sample_size: u8, - me_socks_kdf_policy: MeSocksKdfPolicy, - me_writer_cmd_channel_capacity: usize, - me_writer_byte_budget_bytes: usize, - me_route_channel_capacity: usize, - me_route_backpressure_enabled: bool, - me_route_fairshare_enabled: bool, - me_route_backpressure_base_timeout_ms: u64, - me_route_backpressure_high_timeout_ms: u64, - me_route_backpressure_high_watermark_pct: u8, - me_reader_route_data_wait_ms: u64, - me_health_interval_ms_unhealthy: u64, - me_health_interval_ms_healthy: u64, - me_warn_rate_limit_ms: u64, - me_route_no_writer_mode: MeRouteNoWriterMode, - me_route_no_writer_wait_ms: u64, - me_route_hybrid_max_wait_ms: u64, - me_route_blocking_send_timeout_ms: u64, - me_route_inline_recovery_attempts: u32, - me_route_inline_recovery_wait_ms: u64, - me_connection_cleanup_capacity: usize, - ) -> Arc { - let endpoint_dc_map = Self::build_endpoint_dc_map_from_maps(&proxy_map_v4, &proxy_map_v6); - let preferred_endpoints_by_dc = - Self::build_preferred_endpoints_by_dc(&decision, &proxy_map_v4, &proxy_map_v6); - let registry = Arc::new(ConnRegistry::with_route_and_cleanup_capacity( - me_route_channel_capacity, - me_connection_cleanup_capacity, - )); - registry.update_route_backpressure_policy( - me_route_backpressure_base_timeout_ms, - me_route_backpressure_high_timeout_ms, - me_route_backpressure_high_watermark_pct, - ); - let (writer_epoch, _) = watch::channel(0u64); - let now_epoch_secs = Self::now_epoch_secs(); - let reinit_status = ReinitStatusSnapshot { - active_generation: 1, - warm_generations: Vec::new(), - pending_hardswap_generation: 0, - pending_hardswap_started_at_epoch_secs: 0, - pending_hardswap_map_hash: 0, - inflight: 0, - }; - stats.set_me_writer_byte_budget_limit_bytes(me_writer_byte_budget_bytes); - Arc::new(Self { - routing: Arc::new(RoutingCore { - registry, - writers: Arc::new(WritersState::new()), - rr: AtomicU64::new(0), - writer_epoch, - preferred_endpoints_by_dc: ArcSwap::from_pointee(preferred_endpoints_by_dc), - }), - reinit: Arc::new(ReinitCore { - generation: AtomicU64::new(1), - active_generation: AtomicU64::new(1), - warm_generation: AtomicU64::new(0), - pending_hardswap_generation: AtomicU64::new(0), - pending_hardswap_started_at_epoch_secs: AtomicU64::new(0), - pending_hardswap_map_hash: AtomicU64::new(0), - scheduler_inflight: AtomicUsize::new(0), - max_concurrency_effective: AtomicUsize::new(1), - coordinator: ParkingMutex::new(ReinitCoordinatorState { - next_attempt_id: 1, - active_generation: 1, - desired_map_hash: 0, - pending: None, - attempts: HashMap::new(), - }), - status: ArcSwap::from_pointee(reinit_status), - hardswap: AtomicBool::new(hardswap), - me_hardswap_warmup_delay_min_ms: AtomicU64::new(me_hardswap_warmup_delay_min_ms), - me_hardswap_warmup_delay_max_ms: AtomicU64::new(me_hardswap_warmup_delay_max_ms), - me_hardswap_warmup_extra_passes: AtomicU32::new( - me_hardswap_warmup_extra_passes as u32, - ), - me_hardswap_warmup_pass_backoff_base_ms: AtomicU64::new( - me_hardswap_warmup_pass_backoff_base_ms, - ), - }), - writer_lifecycle: Arc::new(WriterLifecycleCore { - me_keepalive_enabled, - me_keepalive_interval: Duration::from_secs(me_keepalive_interval_secs), - me_keepalive_jitter: Duration::from_secs(me_keepalive_jitter_secs), - me_keepalive_payload_random, - rpc_proxy_req_every_secs: AtomicU64::new(rpc_proxy_req_every_secs), - writer_cmd_channel_capacity: me_writer_cmd_channel_capacity.max(1), - writer_byte_budget_permits: me_writer_byte_budget_bytes - .div_ceil(crate::config::defaults::ME_WRITER_BYTE_PERMIT_UNIT_BYTES) - .max(1), - }), - route_runtime: Arc::new(RouteRuntimeCore { - me_route_no_writer_mode: AtomicU8::new(me_route_no_writer_mode.as_u8()), - me_route_no_writer_wait: Duration::from_millis(me_route_no_writer_wait_ms), - me_route_hybrid_max_wait: Duration::from_millis( - me_route_hybrid_max_wait_ms.max(50), - ), - me_route_blocking_send_timeout: Some(Duration::from_millis( - me_route_blocking_send_timeout_ms.clamp(1, 5_000), - )), - me_route_last_success_epoch_ms: AtomicU64::new(0), - me_route_hybrid_timeout_warn_epoch_ms: AtomicU64::new(0), - me_async_recovery_last_trigger_epoch_ms: AtomicU64::new(0), - me_route_inline_recovery_attempts, - me_route_inline_recovery_wait: Duration::from_millis( - me_route_inline_recovery_wait_ms, - ), - }), - health_runtime: Arc::new(HealthRuntimeCore { - me_health_interval_ms_unhealthy: AtomicU64::new( - me_health_interval_ms_unhealthy.max(1), - ), - me_health_interval_ms_healthy: AtomicU64::new(me_health_interval_ms_healthy.max(1)), - me_warn_rate_limit_ms: AtomicU64::new(me_warn_rate_limit_ms.max(1)), - family_health_v4: ArcSwap::from_pointee(FamilyHealthSnapshot::new( - MeFamilyRuntimeState::Healthy, - now_epoch_secs, - 0, - 0, - 0, - )), - family_health_v6: ArcSwap::from_pointee(FamilyHealthSnapshot::new( - MeFamilyRuntimeState::Healthy, - now_epoch_secs, - 0, - 0, - 0, - )), - }), - drain_runtime: Arc::new(DrainRuntimeCore { - me_pool_drain_ttl_secs: AtomicU64::new(me_pool_drain_ttl_secs), - me_instadrain: AtomicBool::new(me_instadrain), - me_pool_drain_threshold: AtomicU64::new(me_pool_drain_threshold), - me_pool_drain_soft_evict_enabled: AtomicBool::new(me_pool_drain_soft_evict_enabled), - me_pool_drain_soft_evict_grace_secs: AtomicU64::new( - me_pool_drain_soft_evict_grace_secs, - ), - me_pool_drain_soft_evict_per_writer: AtomicU8::new( - me_pool_drain_soft_evict_per_writer.max(1), - ), - me_pool_drain_soft_evict_budget_per_core: AtomicU32::new( - me_pool_drain_soft_evict_budget_per_core.max(1) as u32, - ), - me_pool_drain_soft_evict_cooldown_ms: AtomicU64::new( - me_pool_drain_soft_evict_cooldown_ms.max(1), - ), - me_pool_force_close_secs: AtomicU64::new(Self::normalize_force_close_secs( - me_pool_force_close_secs, - )), - me_pool_min_fresh_ratio_permille: AtomicU32::new(Self::ratio_to_permille( - me_pool_min_fresh_ratio, - )), - me_last_drain_gate_route_quorum_ok: AtomicBool::new(false), - me_last_drain_gate_redundancy_ok: AtomicBool::new(false), - me_last_drain_gate_block_reason: AtomicU8::new(MeDrainGateReason::Open as u8), - me_last_drain_gate_updated_at_epoch_secs: AtomicU64::new(now_epoch_secs), - }), - single_endpoint_runtime: Arc::new(SingleEndpointRuntimeCore { - me_single_endpoint_shadow_writers: AtomicU8::new(me_single_endpoint_shadow_writers), - me_single_endpoint_outage_mode_enabled: AtomicBool::new( - me_single_endpoint_outage_mode_enabled, - ), - me_single_endpoint_outage_disable_quarantine: AtomicBool::new( - me_single_endpoint_outage_disable_quarantine, - ), - me_single_endpoint_outage_backoff_min_ms: AtomicU64::new( - me_single_endpoint_outage_backoff_min_ms, - ), - me_single_endpoint_outage_backoff_max_ms: AtomicU64::new( - me_single_endpoint_outage_backoff_max_ms, - ), - me_single_endpoint_shadow_rotate_every_secs: AtomicU64::new( - me_single_endpoint_shadow_rotate_every_secs, - ), - }), - binding_policy: Arc::new(BindingPolicyCore { - me_bind_stale_mode: AtomicU8::new(me_bind_stale_mode.as_u8()), - me_bind_stale_ttl_secs: AtomicU64::new(me_bind_stale_ttl_secs), - }), - nat_runtime: Arc::new(NatRuntimeCore { - nat_ip_cfg: nat_ip, - nat_ip_detected: Arc::new(RwLock::new(None)), - nat_probe, - nat_stun, - nat_stun_servers, - stun_tcp_fallback, - http_ip_detect_urls, - nat_stun_live_servers: Arc::new(RwLock::new(Vec::new())), - nat_probe_concurrency: nat_probe_concurrency.max(1), - detected_ipv6, - nat_probe_attempts: std::sync::atomic::AtomicU8::new(0), - nat_probe_disabled: std::sync::atomic::AtomicBool::new(false), - stun_backoff_until: Arc::new(RwLock::new(None)), - nat_reflection_cache: Arc::new(Mutex::new(NatReflectionCache::default())), - nat_reflection_singleflight_v4: Arc::new(Mutex::new(())), - nat_reflection_singleflight_v6: Arc::new(Mutex::new(())), - }), - reconnect_runtime: Arc::new(ReconnectRuntimeCore { - me_one_retry, - me_one_timeout: Duration::from_millis(me_one_timeout_ms), - me_warmup_stagger_enabled, - me_warmup_step_delay: Duration::from_millis(me_warmup_step_delay_ms), - me_warmup_step_jitter: Duration::from_millis(me_warmup_step_jitter_ms), - me_reconnect_max_concurrent_per_dc, - me_reconnect_backoff_base: Duration::from_millis(me_reconnect_backoff_base_ms), - me_reconnect_backoff_cap: Duration::from_millis(me_reconnect_backoff_cap_ms), - me_reconnect_fast_retry_count, - }), - floor_runtime: Arc::new(FloorRuntimeCore { - me_floor_mode: AtomicU8::new(me_floor_mode.as_u8()), - me_adaptive_floor_idle_secs: AtomicU64::new(me_adaptive_floor_idle_secs), - me_adaptive_floor_min_writers_single_endpoint: AtomicU8::new( - me_adaptive_floor_min_writers_single_endpoint, - ), - me_adaptive_floor_min_writers_multi_endpoint: AtomicU8::new( - me_adaptive_floor_min_writers_multi_endpoint, - ), - me_adaptive_floor_recover_grace_secs: AtomicU64::new( - me_adaptive_floor_recover_grace_secs, - ), - me_adaptive_floor_writers_per_core_total: AtomicU32::new( - me_adaptive_floor_writers_per_core_total as u32, - ), - me_adaptive_floor_cpu_cores_override: AtomicU32::new( - me_adaptive_floor_cpu_cores_override as u32, - ), - me_adaptive_floor_max_extra_writers_single_per_core: AtomicU32::new( - me_adaptive_floor_max_extra_writers_single_per_core as u32, - ), - me_adaptive_floor_max_extra_writers_multi_per_core: AtomicU32::new( - me_adaptive_floor_max_extra_writers_multi_per_core as u32, - ), - me_adaptive_floor_max_active_writers_per_core: AtomicU32::new( - me_adaptive_floor_max_active_writers_per_core as u32, - ), - me_adaptive_floor_max_warm_writers_per_core: AtomicU32::new( - me_adaptive_floor_max_warm_writers_per_core as u32, - ), - me_adaptive_floor_max_active_writers_global: AtomicU32::new( - me_adaptive_floor_max_active_writers_global, - ), - me_adaptive_floor_max_warm_writers_global: AtomicU32::new( - me_adaptive_floor_max_warm_writers_global, - ), - me_adaptive_floor_cpu_cores_detected: AtomicU32::new(1), - me_adaptive_floor_cpu_cores_effective: AtomicU32::new(1), - me_adaptive_floor_global_cap_raw: AtomicU64::new(0), - me_adaptive_floor_global_cap_effective: AtomicU64::new(0), - me_adaptive_floor_target_writers_total: AtomicU64::new(0), - me_adaptive_floor_active_cap_configured: AtomicU64::new(0), - me_adaptive_floor_active_cap_effective: AtomicU64::new(0), - me_adaptive_floor_warm_cap_configured: AtomicU64::new(0), - me_adaptive_floor_warm_cap_effective: AtomicU64::new(0), - me_adaptive_floor_active_writers_current: AtomicU64::new(0), - me_adaptive_floor_warm_writers_current: AtomicU64::new(0), - }), - writer_selection_policy: Arc::new(WriterSelectionPolicyCore { - secret_atomic_snapshot: AtomicBool::new(me_secret_atomic_snapshot), - me_deterministic_writer_sort: AtomicBool::new(me_deterministic_writer_sort), - me_writer_pick_mode: AtomicU8::new(me_writer_pick_mode.as_u8()), - me_writer_pick_sample_size: AtomicU8::new(me_writer_pick_sample_size.clamp(2, 4)), - }), - transport_policy: Arc::new(TransportPolicyCore { - me_socks_kdf_policy: AtomicU8::new(me_socks_kdf_policy.as_u8()), - me_route_backpressure_enabled: Arc::new(AtomicBool::new( - me_route_backpressure_enabled, - )), - me_route_fairshare_enabled: Arc::new(AtomicBool::new(me_route_fairshare_enabled)), - me_reader_route_data_wait_ms: Arc::new(AtomicU64::new( - me_reader_route_data_wait_ms, - )), - }), - lifecycle: MePoolLifecycle::new(), - decision, - upstream, - rng, - proxy_tag, - proxy_secret: Arc::new(RwLock::new(SecretSnapshot { - epoch: 1, - key_selector: if proxy_secret.len() >= 4 { - u32::from_le_bytes([ - proxy_secret[0], - proxy_secret[1], - proxy_secret[2], - proxy_secret[3], - ]) - } else { - 0 - }, - secret: proxy_secret, - })), - stats, - pool_size: 2, - proxy_map_v4: Arc::new(RwLock::new(proxy_map_v4)), - proxy_map_v6: Arc::new(RwLock::new(proxy_map_v6)), - endpoint_dc_map: Arc::new(RwLock::new(endpoint_dc_map)), - default_dc: AtomicI32::new(default_dc.unwrap_or(2)), - next_writer_id: AtomicU64::new(1), - writer_connect_active_reserved: AtomicUsize::new(0), - writer_connect_warm_reserved: AtomicUsize::new(0), - rtt_stats: Arc::new(Mutex::new(HashMap::new())), - refill_states: Arc::new(ParkingMutex::new(HashMap::new())), - refill_running: AtomicUsize::new(0), - refill_pending: AtomicUsize::new(0), - conn_count: AtomicUsize::new(0), - draining_active_runtime: AtomicU64::new(0), - endpoint_quarantine: Arc::new(Mutex::new(HashMap::new())), - kdf_material_fingerprint: Arc::new(RwLock::new(HashMap::new())), - runtime_ready: AtomicBool::new(false), - }) - } - - /// Creates the immutable byte semaphore assigned to one ME writer generation. - pub(crate) fn new_writer_byte_budget(&self) -> Arc { - Arc::new(Semaphore::new( - self.writer_lifecycle.writer_byte_budget_permits, - )) - } - - pub fn current_generation(&self) -> u64 { - self.reinit.active_generation.load(Ordering::Relaxed) - } - - pub fn set_runtime_ready(&self, ready: bool) { - self.runtime_ready.store(ready, Ordering::Relaxed); - } - - pub fn is_runtime_ready(&self) -> bool { - self.runtime_ready.load(Ordering::Relaxed) - } - - pub(super) fn now_epoch_millis() -> u64 { - SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap_or_default() - .as_millis() as u64 - } - - pub(super) fn notify_writer_epoch(&self) { - self.writer_epoch.send_modify(|epoch| { - *epoch = epoch.wrapping_add(1); - }); - } - - pub(super) fn set_family_runtime_state( - &self, - family: IpFamily, - state: MeFamilyRuntimeState, - state_since_epoch_secs: u64, - suppressed_until_epoch_secs: u64, - fail_streak: u32, - recover_success_streak: u32, - ) { - let snapshot = Arc::new(FamilyHealthSnapshot::new( - state, - state_since_epoch_secs, - suppressed_until_epoch_secs, - fail_streak, - recover_success_streak, - )); - match family { - IpFamily::V4 => self.health_runtime.family_health_v4.store(snapshot), - IpFamily::V6 => self.health_runtime.family_health_v6.store(snapshot), - } - } - - pub(crate) fn family_runtime_state(&self, family: IpFamily) -> MeFamilyRuntimeState { - match family { - IpFamily::V4 => self.health_runtime.family_health_v4.load().state, - IpFamily::V6 => self.health_runtime.family_health_v6.load().state, - } - } - - pub(crate) fn family_runtime_state_since_epoch_secs(&self, family: IpFamily) -> u64 { - match family { - IpFamily::V4 => { - self.health_runtime - .family_health_v4 - .load() - .state_since_epoch_secs - } - IpFamily::V6 => { - self.health_runtime - .family_health_v6 - .load() - .state_since_epoch_secs - } - } - } - - pub(crate) fn family_suppressed_until_epoch_secs(&self, family: IpFamily) -> u64 { - match family { - IpFamily::V4 => { - self.health_runtime - .family_health_v4 - .load() - .suppressed_until_epoch_secs - } - IpFamily::V6 => { - self.health_runtime - .family_health_v6 - .load() - .suppressed_until_epoch_secs - } - } - } - - pub(crate) fn family_fail_streak(&self, family: IpFamily) -> u32 { - match family { - IpFamily::V4 => self.health_runtime.family_health_v4.load().fail_streak, - IpFamily::V6 => self.health_runtime.family_health_v6.load().fail_streak, - } - } - - pub(crate) fn family_recover_success_streak(&self, family: IpFamily) -> u32 { - match family { - IpFamily::V4 => { - self.health_runtime - .family_health_v4 - .load() - .recover_success_streak - } - IpFamily::V6 => { - self.health_runtime - .family_health_v6 - .load() - .recover_success_streak - } - } - } - - pub(crate) fn is_family_temporarily_suppressed( - &self, - family: IpFamily, - now_epoch_secs: u64, - ) -> bool { - self.family_suppressed_until_epoch_secs(family) > now_epoch_secs - } - - pub(super) fn family_enabled_for_drain_coverage( - &self, - family: IpFamily, - now_epoch_secs: u64, - ) -> bool { - let configured = match family { - IpFamily::V4 => self.decision.ipv4_me, - IpFamily::V6 => self.decision.ipv6_me, - }; - configured && !self.is_family_temporarily_suppressed(family, now_epoch_secs) - } - - pub(super) fn set_last_drain_gate( - &self, - route_quorum_ok: bool, - redundancy_ok: bool, - block_reason: MeDrainGateReason, - updated_at_epoch_secs: u64, - ) { - self.drain_runtime - .me_last_drain_gate_route_quorum_ok - .store(route_quorum_ok, Ordering::Relaxed); - self.drain_runtime - .me_last_drain_gate_redundancy_ok - .store(redundancy_ok, Ordering::Relaxed); - self.drain_runtime - .me_last_drain_gate_block_reason - .store(block_reason as u8, Ordering::Relaxed); - self.drain_runtime - .me_last_drain_gate_updated_at_epoch_secs - .store(updated_at_epoch_secs, Ordering::Relaxed); - } - - pub(crate) fn last_drain_gate_route_quorum_ok(&self) -> bool { - self.drain_runtime - .me_last_drain_gate_route_quorum_ok - .load(Ordering::Relaxed) - } - - pub(crate) fn last_drain_gate_redundancy_ok(&self) -> bool { - self.drain_runtime - .me_last_drain_gate_redundancy_ok - .load(Ordering::Relaxed) - } - - pub(crate) fn last_drain_gate_block_reason(&self) -> MeDrainGateReason { - MeDrainGateReason::from_u8( - self.drain_runtime - .me_last_drain_gate_block_reason - .load(Ordering::Relaxed), - ) - } - - pub(crate) fn last_drain_gate_updated_at_epoch_secs(&self) -> u64 { - self.drain_runtime - .me_last_drain_gate_updated_at_epoch_secs - .load(Ordering::Relaxed) - } - - pub fn update_runtime_reinit_policy( - &self, - hardswap: bool, - drain_ttl_secs: u64, - instadrain: bool, - pool_drain_threshold: u64, - pool_drain_soft_evict_enabled: bool, - pool_drain_soft_evict_grace_secs: u64, - pool_drain_soft_evict_per_writer: u8, - pool_drain_soft_evict_budget_per_core: u16, - pool_drain_soft_evict_cooldown_ms: u64, - force_close_secs: u64, - min_fresh_ratio: f32, - hardswap_warmup_delay_min_ms: u64, - hardswap_warmup_delay_max_ms: u64, - hardswap_warmup_extra_passes: u8, - hardswap_warmup_pass_backoff_base_ms: u64, - bind_stale_mode: MeBindStaleMode, - bind_stale_ttl_secs: u64, - secret_atomic_snapshot: bool, - deterministic_writer_sort: bool, - writer_pick_mode: MeWriterPickMode, - writer_pick_sample_size: u8, - single_endpoint_shadow_writers: u8, - single_endpoint_outage_mode_enabled: bool, - single_endpoint_outage_disable_quarantine: bool, - single_endpoint_outage_backoff_min_ms: u64, - single_endpoint_outage_backoff_max_ms: u64, - single_endpoint_shadow_rotate_every_secs: u64, - floor_mode: MeFloorMode, - adaptive_floor_idle_secs: u64, - adaptive_floor_min_writers_single_endpoint: u8, - adaptive_floor_min_writers_multi_endpoint: u8, - adaptive_floor_recover_grace_secs: u64, - adaptive_floor_writers_per_core_total: u16, - adaptive_floor_cpu_cores_override: u16, - adaptive_floor_max_extra_writers_single_per_core: u16, - adaptive_floor_max_extra_writers_multi_per_core: u16, - adaptive_floor_max_active_writers_per_core: u16, - adaptive_floor_max_warm_writers_per_core: u16, - adaptive_floor_max_active_writers_global: u32, - adaptive_floor_max_warm_writers_global: u32, - me_health_interval_ms_unhealthy: u64, - me_health_interval_ms_healthy: u64, - me_warn_rate_limit_ms: u64, - ) { - self.reinit.hardswap.store(hardswap, Ordering::Relaxed); - self.drain_runtime - .me_pool_drain_ttl_secs - .store(drain_ttl_secs, Ordering::Relaxed); - self.drain_runtime - .me_instadrain - .store(instadrain, Ordering::Relaxed); - self.drain_runtime - .me_pool_drain_threshold - .store(pool_drain_threshold, Ordering::Relaxed); - // Runtime soft-evict knobs are updated lock-free to keep control-plane - // writes non-blocking; readers observe a short eventual-consistency - // window by design. - self.drain_runtime - .me_pool_drain_soft_evict_enabled - .store(pool_drain_soft_evict_enabled, Ordering::Relaxed); - self.drain_runtime - .me_pool_drain_soft_evict_grace_secs - .store(pool_drain_soft_evict_grace_secs, Ordering::Relaxed); - self.drain_runtime - .me_pool_drain_soft_evict_per_writer - .store(pool_drain_soft_evict_per_writer.max(1), Ordering::Relaxed); - self.drain_runtime - .me_pool_drain_soft_evict_budget_per_core - .store( - pool_drain_soft_evict_budget_per_core.max(1) as u32, - Ordering::Relaxed, - ); - self.drain_runtime - .me_pool_drain_soft_evict_cooldown_ms - .store(pool_drain_soft_evict_cooldown_ms.max(1), Ordering::Relaxed); - self.drain_runtime.me_pool_force_close_secs.store( - Self::normalize_force_close_secs(force_close_secs), - Ordering::Relaxed, - ); - self.drain_runtime - .me_pool_min_fresh_ratio_permille - .store(Self::ratio_to_permille(min_fresh_ratio), Ordering::Relaxed); - self.reinit - .me_hardswap_warmup_delay_min_ms - .store(hardswap_warmup_delay_min_ms, Ordering::Relaxed); - self.reinit - .me_hardswap_warmup_delay_max_ms - .store(hardswap_warmup_delay_max_ms, Ordering::Relaxed); - self.reinit - .me_hardswap_warmup_extra_passes - .store(hardswap_warmup_extra_passes as u32, Ordering::Relaxed); - self.reinit - .me_hardswap_warmup_pass_backoff_base_ms - .store(hardswap_warmup_pass_backoff_base_ms, Ordering::Relaxed); - self.binding_policy - .me_bind_stale_mode - .store(bind_stale_mode.as_u8(), Ordering::Relaxed); - self.binding_policy - .me_bind_stale_ttl_secs - .store(bind_stale_ttl_secs, Ordering::Relaxed); - self.writer_selection_policy - .secret_atomic_snapshot - .store(secret_atomic_snapshot, Ordering::Relaxed); - self.writer_selection_policy - .me_deterministic_writer_sort - .store(deterministic_writer_sort, Ordering::Relaxed); - let previous_writer_pick_mode = self.writer_pick_mode(); - self.writer_selection_policy - .me_writer_pick_mode - .store(writer_pick_mode.as_u8(), Ordering::Relaxed); - self.writer_selection_policy - .me_writer_pick_sample_size - .store(writer_pick_sample_size.clamp(2, 4), Ordering::Relaxed); - if previous_writer_pick_mode != writer_pick_mode { - self.stats.increment_me_writer_pick_mode_switch_total(); - } - self.single_endpoint_runtime - .me_single_endpoint_shadow_writers - .store(single_endpoint_shadow_writers, Ordering::Relaxed); - self.single_endpoint_runtime - .me_single_endpoint_outage_mode_enabled - .store(single_endpoint_outage_mode_enabled, Ordering::Relaxed); - self.single_endpoint_runtime - .me_single_endpoint_outage_disable_quarantine - .store(single_endpoint_outage_disable_quarantine, Ordering::Relaxed); - self.single_endpoint_runtime - .me_single_endpoint_outage_backoff_min_ms - .store(single_endpoint_outage_backoff_min_ms, Ordering::Relaxed); - self.single_endpoint_runtime - .me_single_endpoint_outage_backoff_max_ms - .store(single_endpoint_outage_backoff_max_ms, Ordering::Relaxed); - self.single_endpoint_runtime - .me_single_endpoint_shadow_rotate_every_secs - .store(single_endpoint_shadow_rotate_every_secs, Ordering::Relaxed); - let previous_floor_mode = self.floor_mode(); - self.floor_runtime - .me_floor_mode - .store(floor_mode.as_u8(), Ordering::Relaxed); - self.floor_runtime - .me_adaptive_floor_idle_secs - .store(adaptive_floor_idle_secs, Ordering::Relaxed); - self.floor_runtime - .me_adaptive_floor_min_writers_single_endpoint - .store( - adaptive_floor_min_writers_single_endpoint, - Ordering::Relaxed, - ); - self.floor_runtime - .me_adaptive_floor_min_writers_multi_endpoint - .store(adaptive_floor_min_writers_multi_endpoint, Ordering::Relaxed); - self.floor_runtime - .me_adaptive_floor_recover_grace_secs - .store(adaptive_floor_recover_grace_secs, Ordering::Relaxed); - self.floor_runtime - .me_adaptive_floor_writers_per_core_total - .store( - adaptive_floor_writers_per_core_total as u32, - Ordering::Relaxed, - ); - self.floor_runtime - .me_adaptive_floor_cpu_cores_override - .store(adaptive_floor_cpu_cores_override as u32, Ordering::Relaxed); - self.floor_runtime - .me_adaptive_floor_max_extra_writers_single_per_core - .store( - adaptive_floor_max_extra_writers_single_per_core as u32, - Ordering::Relaxed, - ); - self.floor_runtime - .me_adaptive_floor_max_extra_writers_multi_per_core - .store( - adaptive_floor_max_extra_writers_multi_per_core as u32, - Ordering::Relaxed, - ); - self.floor_runtime - .me_adaptive_floor_max_active_writers_per_core - .store( - adaptive_floor_max_active_writers_per_core as u32, - Ordering::Relaxed, - ); - self.floor_runtime - .me_adaptive_floor_max_warm_writers_per_core - .store( - adaptive_floor_max_warm_writers_per_core as u32, - Ordering::Relaxed, - ); - self.floor_runtime - .me_adaptive_floor_max_active_writers_global - .store(adaptive_floor_max_active_writers_global, Ordering::Relaxed); - self.floor_runtime - .me_adaptive_floor_max_warm_writers_global - .store(adaptive_floor_max_warm_writers_global, Ordering::Relaxed); - self.health_runtime - .me_health_interval_ms_unhealthy - .store(me_health_interval_ms_unhealthy.max(1), Ordering::Relaxed); - self.health_runtime - .me_health_interval_ms_healthy - .store(me_health_interval_ms_healthy.max(1), Ordering::Relaxed); - self.health_runtime - .me_warn_rate_limit_ms - .store(me_warn_rate_limit_ms.max(1), Ordering::Relaxed); - if previous_floor_mode != floor_mode { - self.stats.increment_me_floor_mode_switch_total(); - match (previous_floor_mode, floor_mode) { - (MeFloorMode::Static, MeFloorMode::Adaptive) => { - self.stats - .increment_me_floor_mode_switch_static_to_adaptive_total(); - } - (MeFloorMode::Adaptive, MeFloorMode::Static) => { - self.stats - .increment_me_floor_mode_switch_adaptive_to_static_total(); - } - _ => {} - } - } - } - - pub fn reset_stun_state(&self) { - self.nat_runtime - .nat_probe_attempts - .store(0, Ordering::Relaxed); - self.nat_runtime - .nat_probe_disabled - .store(false, Ordering::Relaxed); - if let Ok(mut live) = self.nat_runtime.nat_stun_live_servers.try_write() { - live.clear(); - } - } - - /// Translate the local ME address into the address material sent to the proxy. - pub fn translate_our_addr(&self, addr: SocketAddr) -> SocketAddr { - self.translate_our_addr_with_reflection(addr, None) - } - - #[allow(dead_code)] - pub fn registry(&self) -> &Arc { - &self.registry - } - - pub fn update_runtime_transport_policy( - &self, - socks_kdf_policy: MeSocksKdfPolicy, - route_backpressure_enabled: bool, - route_fairshare_enabled: bool, - route_backpressure_base_timeout_ms: u64, - route_backpressure_high_timeout_ms: u64, - route_backpressure_high_watermark_pct: u8, - reader_route_data_wait_ms: u64, - ) { - self.transport_policy - .me_socks_kdf_policy - .store(socks_kdf_policy.as_u8(), Ordering::Relaxed); - self.transport_policy - .me_route_backpressure_enabled - .store(route_backpressure_enabled, Ordering::Relaxed); - self.transport_policy - .me_route_fairshare_enabled - .store(route_fairshare_enabled, Ordering::Relaxed); - self.transport_policy - .me_reader_route_data_wait_ms - .store(reader_route_data_wait_ms, Ordering::Relaxed); - self.registry.update_route_backpressure_policy( - route_backpressure_base_timeout_ms, - route_backpressure_high_timeout_ms, - route_backpressure_high_watermark_pct, - ); - } - - pub(super) fn socks_kdf_policy(&self) -> MeSocksKdfPolicy { - MeSocksKdfPolicy::from_u8( - self.transport_policy - .me_socks_kdf_policy - .load(Ordering::Relaxed), - ) - } - - pub(super) fn writers_arc(&self) -> Arc { - self.writers.clone() - } - - pub(super) fn force_close_timeout(&self) -> Option { - let secs = Self::normalize_force_close_secs( - self.drain_runtime - .me_pool_force_close_secs - .load(Ordering::Relaxed), - ); - Some(Duration::from_secs(secs)) - } - - #[allow(dead_code)] - pub(super) fn drain_soft_evict_enabled(&self) -> bool { - self.drain_runtime - .me_pool_drain_soft_evict_enabled - .load(Ordering::Relaxed) - } - - #[allow(dead_code)] - pub(super) fn drain_soft_evict_grace_secs(&self) -> u64 { - self.drain_runtime - .me_pool_drain_soft_evict_grace_secs - .load(Ordering::Relaxed) - } - - #[allow(dead_code)] - pub(super) fn drain_soft_evict_per_writer(&self) -> usize { - self.drain_runtime - .me_pool_drain_soft_evict_per_writer - .load(Ordering::Relaxed) - .max(1) as usize - } - - #[allow(dead_code)] - pub(super) fn drain_soft_evict_budget_per_core(&self) -> usize { - self.drain_runtime - .me_pool_drain_soft_evict_budget_per_core - .load(Ordering::Relaxed) - .max(1) as usize - } - - #[allow(dead_code)] - pub(super) fn drain_soft_evict_cooldown(&self) -> Duration { - Duration::from_millis( - self.drain_runtime - .me_pool_drain_soft_evict_cooldown_ms - .load(Ordering::Relaxed) - .max(1), - ) - } - - #[allow(dead_code)] - pub(super) fn draining_active_runtime(&self) -> u64 { - self.draining_active_runtime.load(Ordering::Relaxed) - } - - pub(super) fn increment_draining_active_runtime(&self) { - self.draining_active_runtime.fetch_add(1, Ordering::Relaxed); - } - - pub(super) fn decrement_draining_active_runtime(&self) { - let mut current = self.draining_active_runtime.load(Ordering::Relaxed); - loop { - if current == 0 { - break; - } - match self.draining_active_runtime.compare_exchange_weak( - current, - current - 1, - Ordering::Relaxed, - Ordering::Relaxed, - ) { - Ok(_) => break, - Err(actual) => current = actual, - } - } - } - - pub(super) async fn key_selector(&self) -> u32 { - self.proxy_secret.read().await.key_selector - } - - pub(super) async fn non_draining_writer_counts_by_contour(&self) -> (usize, usize, usize) { - let ws = self.writers.read().await; - let mut active = 0usize; - let mut warm = 0usize; - for writer in ws.iter() { - if writer.draining.load(Ordering::Relaxed) { - continue; - } - match WriterContour::from_u8(writer.contour.load(Ordering::Relaxed)) { - WriterContour::Active => active = active.saturating_add(1), - WriterContour::Warm => warm = warm.saturating_add(1), - WriterContour::Draining => {} - } - } - (active, warm, active.saturating_add(warm)) - } - - pub(super) async fn active_contour_writer_count_total(&self) -> usize { - let (active, _, _) = self.non_draining_writer_counts_by_contour().await; - active - } - - pub(super) async fn secret_snapshot(&self) -> SecretSnapshot { - self.proxy_secret.read().await.clone() - } - - pub(super) fn bind_stale_mode(&self) -> MeBindStaleMode { - MeBindStaleMode::from_u8( - self.binding_policy - .me_bind_stale_mode - .load(Ordering::Relaxed), - ) - } - - pub(super) fn writer_pick_mode(&self) -> MeWriterPickMode { - MeWriterPickMode::from_u8( - self.writer_selection_policy - .me_writer_pick_mode - .load(Ordering::Relaxed), - ) - } - - pub(super) fn writer_pick_sample_size(&self) -> usize { - self.writer_selection_policy - .me_writer_pick_sample_size - .load(Ordering::Relaxed) - .clamp(2, 4) as usize - } - - pub(super) fn required_writers_for_dc(&self, endpoint_count: usize) -> usize { - if endpoint_count == 0 { - return 0; - } - if endpoint_count == 1 { - let shadow = self - .single_endpoint_runtime - .me_single_endpoint_shadow_writers - .load(Ordering::Relaxed) as usize; - return (1 + shadow).max(3); - } - endpoint_count.max(3) - } - - pub(super) fn floor_mode(&self) -> MeFloorMode { - MeFloorMode::from_u8(self.floor_runtime.me_floor_mode.load(Ordering::Relaxed)) - } - - pub(super) fn adaptive_floor_min_writers_multi_endpoint(&self) -> usize { - (self - .floor_runtime - .me_adaptive_floor_min_writers_multi_endpoint - .load(Ordering::Relaxed) as usize) - .max(1) - } - - pub(super) fn adaptive_floor_max_extra_single_per_core(&self) -> usize { - self.floor_runtime - .me_adaptive_floor_max_extra_writers_single_per_core - .load(Ordering::Relaxed) as usize - } - - pub(super) fn adaptive_floor_max_extra_multi_per_core(&self) -> usize { - self.floor_runtime - .me_adaptive_floor_max_extra_writers_multi_per_core - .load(Ordering::Relaxed) as usize - } - - pub(super) fn adaptive_floor_max_active_writers_per_core(&self) -> usize { - (self - .floor_runtime - .me_adaptive_floor_max_active_writers_per_core - .load(Ordering::Relaxed) as usize) - .max(1) - } - - pub(super) fn adaptive_floor_max_warm_writers_per_core(&self) -> usize { - (self - .floor_runtime - .me_adaptive_floor_max_warm_writers_per_core - .load(Ordering::Relaxed) as usize) - .max(1) - } - - pub(super) fn adaptive_floor_max_active_writers_global(&self) -> usize { - (self - .floor_runtime - .me_adaptive_floor_max_active_writers_global - .load(Ordering::Relaxed) as usize) - .max(1) - } - - pub(super) fn adaptive_floor_max_warm_writers_global(&self) -> usize { - (self - .floor_runtime - .me_adaptive_floor_max_warm_writers_global - .load(Ordering::Relaxed) as usize) - .max(1) - } - - pub(super) fn adaptive_floor_detected_cpu_cores(&self) -> usize { - std::thread::available_parallelism() - .map(|value| value.get()) - .unwrap_or(1) - .max(1) - } - - pub(super) fn adaptive_floor_effective_cpu_cores(&self) -> usize { - let detected = self.adaptive_floor_detected_cpu_cores(); - let override_cores = self - .floor_runtime - .me_adaptive_floor_cpu_cores_override - .load(Ordering::Relaxed) as usize; - let effective = if override_cores == 0 { - detected - } else { - override_cores.max(1) - }; - self.floor_runtime - .me_adaptive_floor_cpu_cores_detected - .store(detected as u32, Ordering::Relaxed); - self.floor_runtime - .me_adaptive_floor_cpu_cores_effective - .store(effective as u32, Ordering::Relaxed); - self.stats - .set_me_floor_cpu_cores_detected_gauge(detected as u64); - self.stats - .set_me_floor_cpu_cores_effective_gauge(effective as u64); - effective - } - - // Keeps per-contour (active/warm) writer budget bounded by CPU count. - // Baseline is 86 writers on the first core and +48 for each extra core. - fn adaptive_floor_cpu_budget_per_contour_cap(&self, cores: usize) -> usize { - const FIRST_CORE_WRITER_BUDGET: usize = 86; - const EXTRA_CORE_WRITER_BUDGET: usize = 48; - if cores == 0 { - return FIRST_CORE_WRITER_BUDGET; - } - FIRST_CORE_WRITER_BUDGET.saturating_add( - cores - .saturating_sub(1) - .saturating_mul(EXTRA_CORE_WRITER_BUDGET), - ) - } - - pub(super) fn adaptive_floor_active_cap_configured_total(&self) -> usize { - let cores = self.adaptive_floor_effective_cpu_cores(); - let per_contour_budget = self.adaptive_floor_cpu_budget_per_contour_cap(cores); - let configured = cores - .saturating_mul(self.adaptive_floor_max_active_writers_per_core()) - .min(self.adaptive_floor_max_active_writers_global()) - .min(per_contour_budget) - .max(1); - self.floor_runtime - .me_adaptive_floor_active_cap_configured - .store(configured as u64, Ordering::Relaxed); - self.stats - .set_me_floor_active_cap_configured_gauge(configured as u64); - configured - } - - pub(super) fn adaptive_floor_warm_cap_configured_total(&self) -> usize { - let cores = self.adaptive_floor_effective_cpu_cores(); - let per_contour_budget = self.adaptive_floor_cpu_budget_per_contour_cap(cores); - let configured = cores - .saturating_mul(self.adaptive_floor_max_warm_writers_per_core()) - .min(self.adaptive_floor_max_warm_writers_global()) - .min(per_contour_budget) - .max(1); - self.floor_runtime - .me_adaptive_floor_warm_cap_configured - .store(configured as u64, Ordering::Relaxed); - self.stats - .set_me_floor_warm_cap_configured_gauge(configured as u64); - configured - } - - pub(super) fn set_adaptive_floor_runtime_caps( - &self, - active_cap_configured: usize, - active_cap_effective: usize, - warm_cap_configured: usize, - warm_cap_effective: usize, - target_writers_total: usize, - active_writers_current: usize, - warm_writers_current: usize, - ) { - self.floor_runtime - .me_adaptive_floor_global_cap_raw - .store(active_cap_configured as u64, Ordering::Relaxed); - self.floor_runtime - .me_adaptive_floor_global_cap_effective - .store(active_cap_effective as u64, Ordering::Relaxed); - self.floor_runtime - .me_adaptive_floor_target_writers_total - .store(target_writers_total as u64, Ordering::Relaxed); - self.floor_runtime - .me_adaptive_floor_active_cap_configured - .store(active_cap_configured as u64, Ordering::Relaxed); - self.floor_runtime - .me_adaptive_floor_active_cap_effective - .store(active_cap_effective as u64, Ordering::Relaxed); - self.floor_runtime - .me_adaptive_floor_warm_cap_configured - .store(warm_cap_configured as u64, Ordering::Relaxed); - self.floor_runtime - .me_adaptive_floor_warm_cap_effective - .store(warm_cap_effective as u64, Ordering::Relaxed); - self.floor_runtime - .me_adaptive_floor_active_writers_current - .store(active_writers_current as u64, Ordering::Relaxed); - self.floor_runtime - .me_adaptive_floor_warm_writers_current - .store(warm_writers_current as u64, Ordering::Relaxed); - self.stats - .set_me_floor_global_cap_raw_gauge(active_cap_configured as u64); - self.stats - .set_me_floor_global_cap_effective_gauge(active_cap_effective as u64); - self.stats - .set_me_floor_target_writers_total_gauge(target_writers_total as u64); - self.stats - .set_me_floor_active_cap_configured_gauge(active_cap_configured as u64); - self.stats - .set_me_floor_active_cap_effective_gauge(active_cap_effective as u64); - self.stats - .set_me_floor_warm_cap_configured_gauge(warm_cap_configured as u64); - self.stats - .set_me_floor_warm_cap_effective_gauge(warm_cap_effective as u64); - self.stats - .set_me_writers_active_current_gauge(active_writers_current as u64); - self.stats - .set_me_writers_warm_current_gauge(warm_writers_current as u64); - } - - pub(super) async fn active_coverage_required_total(&self) -> usize { - let now_epoch_secs = Self::now_epoch_secs(); - let mut endpoints_by_dc = HashMap::>::new(); - - if self.family_enabled_for_drain_coverage(IpFamily::V4, now_epoch_secs) { - let map = self.proxy_map_v4.read().await; - for (dc, addrs) in map.iter() { - let entry = endpoints_by_dc.entry(*dc).or_default(); - for (ip, port) in addrs.iter().copied() { - entry.insert(SocketAddr::new(ip, port)); - } - } - } - - if self.family_enabled_for_drain_coverage(IpFamily::V6, now_epoch_secs) { - let map = self.proxy_map_v6.read().await; - for (dc, addrs) in map.iter() { - let entry = endpoints_by_dc.entry(*dc).or_default(); - for (ip, port) in addrs.iter().copied() { - entry.insert(SocketAddr::new(ip, port)); - } - } - } - - endpoints_by_dc - .values() - .map(|endpoints| self.required_writers_for_dc_with_floor_mode(endpoints.len(), false)) - .sum() - } - - pub(super) async fn can_open_writer_for_contour( - &self, - contour: WriterContour, - allow_coverage_override: bool, - writer_dc: i32, - ) -> bool { - let (active_writers, warm_writers, _) = self.non_draining_writer_counts_by_contour().await; - match contour { - WriterContour::Active => { - let active_cap = self.adaptive_floor_active_cap_configured_total(); - if active_writers < active_cap { - return true; - } - if !allow_coverage_override { - return false; - } - - let mut endpoints_len = 0; - let now_epoch = Self::now_epoch_secs(); - if self.family_enabled_for_drain_coverage(IpFamily::V4, now_epoch) { - if let Some(addrs) = self.proxy_map_v4.read().await.get(&writer_dc) { - endpoints_len += addrs.len(); - } - } - if self.family_enabled_for_drain_coverage(IpFamily::V6, now_epoch) { - if let Some(addrs) = self.proxy_map_v6.read().await.get(&writer_dc) { - endpoints_len += addrs.len(); - } - } - - if endpoints_len > 0 { - let base_req = - self.required_writers_for_dc_with_floor_mode(endpoints_len, false); - let active_for_dc = { - let ws = self.writers.read().await; - ws.iter() - .filter(|w| { - !w.draining.load(std::sync::atomic::Ordering::Relaxed) - && w.writer_dc == writer_dc - && matches!( - WriterContour::from_u8( - w.contour.load(std::sync::atomic::Ordering::Relaxed), - ), - WriterContour::Active - ) - }) - .count() - }; - if active_for_dc < base_req { - return true; - } - } - - let coverage_required = self.active_coverage_required_total().await; - active_writers < coverage_required - } - WriterContour::Warm => warm_writers < self.adaptive_floor_warm_cap_configured_total(), - WriterContour::Draining => true, - } - } - - pub(super) async fn reserve_writer_open( - &self, - contour: WriterContour, - allow_coverage_override: bool, - writer_dc: i32, - ) -> Option> { - let counter = match contour { - WriterContour::Active => &self.writer_connect_active_reserved, - WriterContour::Warm => &self.writer_connect_warm_reserved, - WriterContour::Draining => { - return Some(WriterOpenReservation { counter: None }); - } - }; - - loop { - if !self - .can_open_writer_for_contour(contour, allow_coverage_override, writer_dc) - .await - { - return None; - } - let (active_writers, warm_writers, _) = - self.non_draining_writer_counts_by_contour().await; - let live = match contour { - WriterContour::Active => active_writers, - WriterContour::Warm => warm_writers, - WriterContour::Draining => 0, - }; - let mut limit = match contour { - WriterContour::Active => self.adaptive_floor_active_cap_configured_total(), - WriterContour::Warm => self.adaptive_floor_warm_cap_configured_total(), - WriterContour::Draining => usize::MAX, - }; - if contour == WriterContour::Active && allow_coverage_override { - limit = limit - .max(self.active_coverage_required_total().await) - .saturating_add( - self.reconnect_runtime - .me_reconnect_max_concurrent_per_dc - .max(1) as usize, - ); - } - - let reserved = counter.load(Ordering::Acquire); - if live.saturating_add(reserved) >= limit { - return None; - } - if counter - .compare_exchange_weak( - reserved, - reserved + 1, - Ordering::AcqRel, - Ordering::Acquire, - ) - .is_ok() - { - return Some(WriterOpenReservation { - counter: Some(counter), - }); - } - } - } - - pub(super) fn required_writers_for_dc_with_floor_mode( - &self, - endpoint_count: usize, - reduce_for_idle: bool, - ) -> usize { - let base_required = self.required_writers_for_dc(endpoint_count); - if !reduce_for_idle { - return base_required; - } - if self.floor_mode() != MeFloorMode::Adaptive { - return base_required; - } - let min_writers = if endpoint_count == 1 { - (self - .floor_runtime - .me_adaptive_floor_min_writers_single_endpoint - .load(Ordering::Relaxed) as usize) - .max(1) - } else { - (self - .floor_runtime - .me_adaptive_floor_min_writers_multi_endpoint - .load(Ordering::Relaxed) as usize) - .max(1) - }; - base_required.min(min_writers) - } - - pub(super) fn single_endpoint_outage_mode_enabled(&self) -> bool { - self.single_endpoint_runtime - .me_single_endpoint_outage_mode_enabled - .load(Ordering::Relaxed) - } - - pub(super) fn single_endpoint_outage_disable_quarantine(&self) -> bool { - self.single_endpoint_runtime - .me_single_endpoint_outage_disable_quarantine - .load(Ordering::Relaxed) - } - - pub(super) fn single_endpoint_outage_backoff_bounds_ms(&self) -> (u64, u64) { - let min_ms = self - .single_endpoint_runtime - .me_single_endpoint_outage_backoff_min_ms - .load(Ordering::Relaxed); - let max_ms = self - .single_endpoint_runtime - .me_single_endpoint_outage_backoff_max_ms - .load(Ordering::Relaxed); - if min_ms <= max_ms { - (min_ms, max_ms) - } else { - (max_ms, min_ms) - } - } - - pub(super) fn single_endpoint_shadow_rotate_interval(&self) -> Option { - let secs = self - .single_endpoint_runtime - .me_single_endpoint_shadow_rotate_every_secs - .load(Ordering::Relaxed); - if secs == 0 { - None - } else { - Some(Duration::from_secs(secs)) - } - } - - pub(super) fn family_order(&self) -> Vec { - let mut order = Vec::new(); - if self.decision.prefer_ipv6() { - if self.decision.ipv6_me { - order.push(IpFamily::V6); - } - if self.decision.ipv4_me { - order.push(IpFamily::V4); - } - } else { - if self.decision.ipv4_me { - order.push(IpFamily::V4); - } - if self.decision.ipv6_me { - order.push(IpFamily::V6); - } - } - order - } - - pub(super) fn default_dc_for_routing(&self) -> i32 { - let dc = self.default_dc.load(Ordering::Relaxed); - if dc == 0 { 2 } else { dc } - } - - pub(super) async fn has_configured_endpoints_for_dc(&self, dc: i32) -> bool { - if self.decision.ipv4_me { - let map = self.proxy_map_v4.read().await; - if map.get(&dc).is_some_and(|endpoints| !endpoints.is_empty()) { - return true; - } - } - - if self.decision.ipv6_me { - let map = self.proxy_map_v6.read().await; - if map.get(&dc).is_some_and(|endpoints| !endpoints.is_empty()) { - return true; - } - } - - false - } - - pub(super) async fn resolve_target_dc_for_routing(&self, target_dc: i32) -> (i32, bool) { - if target_dc == 0 { - return (self.default_dc_for_routing(), true); - } - - if self.has_configured_endpoints_for_dc(target_dc).await { - return (target_dc, false); - } - - (self.default_dc_for_routing(), true) - } - - pub(super) async fn resolve_dc_for_endpoint(&self, addr: SocketAddr) -> i32 { - if let Some(cached) = self.endpoint_dc_map.read().await.get(&addr).copied() - && let Some(dc) = cached - { - return dc; - } - - self.default_dc_for_routing() - } - - pub(super) async fn proxy_map_for_family( - &self, - family: IpFamily, - ) -> HashMap> { - match family { - IpFamily::V4 => self.proxy_map_v4.read().await.clone(), - IpFamily::V6 => self.proxy_map_v6.read().await.clone(), - } - } - - fn merge_endpoint_dc( - endpoint_dc_map: &mut HashMap>, - dc: i32, - ip: IpAddr, - port: u16, - ) { - let endpoint = SocketAddr::new(ip, port); - match endpoint_dc_map.get_mut(&endpoint) { - None => { - endpoint_dc_map.insert(endpoint, Some(dc)); - } - Some(existing) => { - if existing.is_some_and(|existing_dc| existing_dc != dc) { - *existing = None; - } - } - } - } - - fn build_preferred_endpoints_by_dc( - decision: &NetworkDecision, - map_v4: &HashMap>, - map_v6: &HashMap>, - ) -> HashMap> { - let mut out = HashMap::>::new(); - let mut dcs = HashSet::::new(); - dcs.extend(map_v4.keys().copied()); - dcs.extend(map_v6.keys().copied()); - - for dc in dcs { - let v4 = map_v4 - .get(&dc) - .map(|items| { - items - .iter() - .map(|(ip, port)| SocketAddr::new(*ip, *port)) - .collect::>() - }) - .unwrap_or_default(); - let v6 = map_v6 - .get(&dc) - .map(|items| { - items - .iter() - .map(|(ip, port)| SocketAddr::new(*ip, *port)) - .collect::>() - }) - .unwrap_or_default(); - - let mut selected = if decision.effective_multipath { - let mut both = Vec::::with_capacity(v4.len().saturating_add(v6.len())); - if decision.prefer_ipv6() { - both.extend(v6.iter().copied()); - both.extend(v4.iter().copied()); - } else { - both.extend(v4.iter().copied()); - both.extend(v6.iter().copied()); - } - both - } else if decision.prefer_ipv6() { - if !v6.is_empty() { v6 } else { v4 } - } else if !v4.is_empty() { - v4 - } else { - v6 - }; - - selected.sort_unstable(); - selected.dedup(); - out.insert(dc, selected); - } - - out - } - - fn build_endpoint_dc_map_from_maps( - map_v4: &HashMap>, - map_v6: &HashMap>, - ) -> HashMap> { - let mut endpoint_dc_map = HashMap::>::new(); - for (dc, endpoints) in map_v4 { - for (ip, port) in endpoints { - Self::merge_endpoint_dc(&mut endpoint_dc_map, *dc, *ip, *port); - } - } - for (dc, endpoints) in map_v6 { - for (ip, port) in endpoints { - Self::merge_endpoint_dc(&mut endpoint_dc_map, *dc, *ip, *port); - } - } - endpoint_dc_map - } - - pub(super) async fn rebuild_endpoint_dc_map(&self) { - let map_v4 = self.proxy_map_v4.read().await.clone(); - let map_v6 = self.proxy_map_v6.read().await.clone(); - let rebuilt = Self::build_endpoint_dc_map_from_maps(&map_v4, &map_v6); - let preferred = Self::build_preferred_endpoints_by_dc(&self.decision, &map_v4, &map_v6); - *self.endpoint_dc_map.write().await = rebuilt; - self.preferred_endpoints_by_dc.store(Arc::new(preferred)); - let configured_endpoints = self - .endpoint_dc_map - .read() - .await - .keys() - .copied() - .collect::>(); - { - let mut quarantine = self.endpoint_quarantine.lock().await; - let now = Instant::now(); - quarantine.retain(|addr, expiry| *expiry > now && configured_endpoints.contains(addr)); - } - { - let mut kdf_fp = self.kdf_material_fingerprint.write().await; - kdf_fp.retain(|addr, _| configured_endpoints.contains(addr)); - } - } - - pub(super) async fn preferred_endpoints_for_dc(&self, dc: i32) -> Vec { - let guard = self.preferred_endpoints_by_dc.load(); - guard.get(&dc).cloned().unwrap_or_default() - } - - pub(super) fn health_interval_unhealthy(&self) -> Duration { - Duration::from_millis( - self.health_runtime - .me_health_interval_ms_unhealthy - .load(Ordering::Relaxed) - .max(1), - ) - } - - pub(super) fn health_interval_healthy(&self) -> Duration { - Duration::from_millis( - self.health_runtime - .me_health_interval_ms_healthy - .load(Ordering::Relaxed) - .max(1), - ) - } - - pub(super) fn warn_rate_limit_duration(&self) -> Duration { - Duration::from_millis( - self.health_runtime - .me_warn_rate_limit_ms - .load(Ordering::Relaxed) - .max(1), - ) - } -} +// Pool construction and immutable process-generation resources. +mod construction; +// Mutable generation and family runtime policy. +mod runtime_policy; +// NAT, transport, and draining runtime policy. +mod transport_policy; +// Writer contour and adaptive-floor selection policy. +mod selection_policy; +// Bounded writer-open admission and coverage accounting. +mod writer_admission; +// Endpoint-to-DC routing and health timing policy. +mod routing; diff --git a/src/transport/middle_proxy/pool/construction.rs b/src/transport/middle_proxy/pool/construction.rs new file mode 100644 index 0000000..dd76e56 --- /dev/null +++ b/src/transport/middle_proxy/pool/construction.rs @@ -0,0 +1,412 @@ +use super::*; + +impl MePool { + pub(in crate::transport::middle_proxy) fn ratio_to_permille(ratio: f32) -> u32 { + let clamped = ratio.clamp(0.0, 1.0); + (clamped * 1000.0).round() as u32 + } + + pub(in crate::transport::middle_proxy) fn permille_to_ratio(permille: u32) -> f32 { + (permille.min(1000) as f32) / 1000.0 + } + + pub(in crate::transport::middle_proxy) fn now_epoch_secs() -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs() + } + + pub(in crate::transport::middle_proxy) fn normalize_force_close_secs( + force_close_secs: u64, + ) -> u64 { + if force_close_secs == 0 { + ME_FORCE_CLOSE_SAFETY_FALLBACK_SECS + } else { + force_close_secs + } + } + + pub fn new( + proxy_tag: Option>, + proxy_secret: Vec, + nat_ip: Option, + nat_probe: bool, + nat_stun: Option, + nat_stun_servers: Vec, + stun_tcp_fallback: bool, + http_ip_detect_urls: Vec, + nat_probe_concurrency: usize, + detected_ipv6: Option, + me_one_retry: u8, + me_one_timeout_ms: u64, + proxy_map_v4: HashMap>, + proxy_map_v6: HashMap>, + default_dc: Option, + decision: NetworkDecision, + upstream: Option>, + rng: Arc, + stats: Arc, + me_keepalive_enabled: bool, + me_keepalive_interval_secs: u64, + me_keepalive_jitter_secs: u64, + me_keepalive_payload_random: bool, + rpc_proxy_req_every_secs: u64, + me_warmup_stagger_enabled: bool, + me_warmup_step_delay_ms: u64, + me_warmup_step_jitter_ms: u64, + me_reconnect_max_concurrent_per_dc: u32, + me_reconnect_backoff_base_ms: u64, + me_reconnect_backoff_cap_ms: u64, + me_reconnect_fast_retry_count: u32, + me_single_endpoint_shadow_writers: u8, + me_single_endpoint_outage_mode_enabled: bool, + me_single_endpoint_outage_disable_quarantine: bool, + me_single_endpoint_outage_backoff_min_ms: u64, + me_single_endpoint_outage_backoff_max_ms: u64, + me_single_endpoint_shadow_rotate_every_secs: u64, + me_floor_mode: MeFloorMode, + me_adaptive_floor_idle_secs: u64, + me_adaptive_floor_min_writers_single_endpoint: u8, + me_adaptive_floor_min_writers_multi_endpoint: u8, + me_adaptive_floor_recover_grace_secs: u64, + me_adaptive_floor_writers_per_core_total: u16, + me_adaptive_floor_cpu_cores_override: u16, + me_adaptive_floor_max_extra_writers_single_per_core: u16, + me_adaptive_floor_max_extra_writers_multi_per_core: u16, + me_adaptive_floor_max_active_writers_per_core: u16, + me_adaptive_floor_max_warm_writers_per_core: u16, + me_adaptive_floor_max_active_writers_global: u32, + me_adaptive_floor_max_warm_writers_global: u32, + hardswap: bool, + me_pool_drain_ttl_secs: u64, + me_instadrain: bool, + me_pool_drain_threshold: u64, + me_pool_drain_soft_evict_enabled: bool, + me_pool_drain_soft_evict_grace_secs: u64, + me_pool_drain_soft_evict_per_writer: u8, + me_pool_drain_soft_evict_budget_per_core: u16, + me_pool_drain_soft_evict_cooldown_ms: u64, + me_pool_force_close_secs: u64, + me_pool_min_fresh_ratio: f32, + me_hardswap_warmup_delay_min_ms: u64, + me_hardswap_warmup_delay_max_ms: u64, + me_hardswap_warmup_extra_passes: u8, + me_hardswap_warmup_pass_backoff_base_ms: u64, + me_bind_stale_mode: MeBindStaleMode, + me_bind_stale_ttl_secs: u64, + me_secret_atomic_snapshot: bool, + me_deterministic_writer_sort: bool, + me_writer_pick_mode: MeWriterPickMode, + me_writer_pick_sample_size: u8, + me_socks_kdf_policy: MeSocksKdfPolicy, + me_writer_cmd_channel_capacity: usize, + me_writer_byte_budget_bytes: usize, + me_route_channel_capacity: usize, + me_route_backpressure_enabled: bool, + me_route_fairshare_enabled: bool, + me_route_backpressure_base_timeout_ms: u64, + me_route_backpressure_high_timeout_ms: u64, + me_route_backpressure_high_watermark_pct: u8, + me_reader_route_data_wait_ms: u64, + me_health_interval_ms_unhealthy: u64, + me_health_interval_ms_healthy: u64, + me_warn_rate_limit_ms: u64, + me_route_no_writer_mode: MeRouteNoWriterMode, + me_route_no_writer_wait_ms: u64, + me_route_hybrid_max_wait_ms: u64, + me_route_blocking_send_timeout_ms: u64, + me_route_inline_recovery_attempts: u32, + me_route_inline_recovery_wait_ms: u64, + me_connection_cleanup_capacity: usize, + ) -> Arc { + let endpoint_dc_map = Self::build_endpoint_dc_map_from_maps(&proxy_map_v4, &proxy_map_v6); + let preferred_endpoints_by_dc = + Self::build_preferred_endpoints_by_dc(&decision, &proxy_map_v4, &proxy_map_v6); + let registry = Arc::new(ConnRegistry::with_route_and_cleanup_capacity( + me_route_channel_capacity, + me_connection_cleanup_capacity, + )); + registry.update_route_backpressure_policy( + me_route_backpressure_base_timeout_ms, + me_route_backpressure_high_timeout_ms, + me_route_backpressure_high_watermark_pct, + ); + let (writer_epoch, _) = watch::channel(0u64); + let now_epoch_secs = Self::now_epoch_secs(); + let reinit_status = ReinitStatusSnapshot { + active_generation: 1, + warm_generations: Vec::new(), + pending_hardswap_generation: 0, + pending_hardswap_started_at_epoch_secs: 0, + pending_hardswap_map_hash: 0, + inflight: 0, + }; + stats.set_me_writer_byte_budget_limit_bytes(me_writer_byte_budget_bytes); + Arc::new(Self { + routing: Arc::new(RoutingCore { + registry, + writers: Arc::new(WritersState::new()), + rr: AtomicU64::new(0), + writer_epoch, + preferred_endpoints_by_dc: ArcSwap::from_pointee(preferred_endpoints_by_dc), + }), + reinit: Arc::new(ReinitCore { + generation: AtomicU64::new(1), + active_generation: AtomicU64::new(1), + warm_generation: AtomicU64::new(0), + pending_hardswap_generation: AtomicU64::new(0), + pending_hardswap_started_at_epoch_secs: AtomicU64::new(0), + pending_hardswap_map_hash: AtomicU64::new(0), + scheduler_inflight: AtomicUsize::new(0), + max_concurrency_effective: AtomicUsize::new(1), + coordinator: ParkingMutex::new(ReinitCoordinatorState { + next_attempt_id: 1, + active_generation: 1, + desired_map_hash: 0, + pending: None, + attempts: HashMap::new(), + }), + status: ArcSwap::from_pointee(reinit_status), + hardswap: AtomicBool::new(hardswap), + me_hardswap_warmup_delay_min_ms: AtomicU64::new(me_hardswap_warmup_delay_min_ms), + me_hardswap_warmup_delay_max_ms: AtomicU64::new(me_hardswap_warmup_delay_max_ms), + me_hardswap_warmup_extra_passes: AtomicU32::new( + me_hardswap_warmup_extra_passes as u32, + ), + me_hardswap_warmup_pass_backoff_base_ms: AtomicU64::new( + me_hardswap_warmup_pass_backoff_base_ms, + ), + }), + writer_lifecycle: Arc::new(WriterLifecycleCore { + me_keepalive_enabled, + me_keepalive_interval: Duration::from_secs(me_keepalive_interval_secs), + me_keepalive_jitter: Duration::from_secs(me_keepalive_jitter_secs), + me_keepalive_payload_random, + rpc_proxy_req_every_secs: AtomicU64::new(rpc_proxy_req_every_secs), + writer_cmd_channel_capacity: me_writer_cmd_channel_capacity.max(1), + writer_byte_budget_permits: me_writer_byte_budget_bytes + .div_ceil(crate::config::defaults::ME_WRITER_BYTE_PERMIT_UNIT_BYTES) + .max(1), + }), + route_runtime: Arc::new(RouteRuntimeCore { + me_route_no_writer_mode: AtomicU8::new(me_route_no_writer_mode.as_u8()), + me_route_no_writer_wait: Duration::from_millis(me_route_no_writer_wait_ms), + me_route_hybrid_max_wait: Duration::from_millis( + me_route_hybrid_max_wait_ms.max(50), + ), + me_route_blocking_send_timeout: Some(Duration::from_millis( + me_route_blocking_send_timeout_ms.clamp(1, 5_000), + )), + me_route_last_success_epoch_ms: AtomicU64::new(0), + me_route_hybrid_timeout_warn_epoch_ms: AtomicU64::new(0), + me_async_recovery_last_trigger_epoch_ms: AtomicU64::new(0), + me_route_inline_recovery_attempts, + me_route_inline_recovery_wait: Duration::from_millis( + me_route_inline_recovery_wait_ms, + ), + }), + health_runtime: Arc::new(HealthRuntimeCore { + me_health_interval_ms_unhealthy: AtomicU64::new( + me_health_interval_ms_unhealthy.max(1), + ), + me_health_interval_ms_healthy: AtomicU64::new(me_health_interval_ms_healthy.max(1)), + me_warn_rate_limit_ms: AtomicU64::new(me_warn_rate_limit_ms.max(1)), + family_health_v4: ArcSwap::from_pointee(FamilyHealthSnapshot::new( + MeFamilyRuntimeState::Healthy, + now_epoch_secs, + 0, + 0, + 0, + )), + family_health_v6: ArcSwap::from_pointee(FamilyHealthSnapshot::new( + MeFamilyRuntimeState::Healthy, + now_epoch_secs, + 0, + 0, + 0, + )), + }), + drain_runtime: Arc::new(DrainRuntimeCore { + me_pool_drain_ttl_secs: AtomicU64::new(me_pool_drain_ttl_secs), + me_instadrain: AtomicBool::new(me_instadrain), + me_pool_drain_threshold: AtomicU64::new(me_pool_drain_threshold), + me_pool_drain_soft_evict_enabled: AtomicBool::new(me_pool_drain_soft_evict_enabled), + me_pool_drain_soft_evict_grace_secs: AtomicU64::new( + me_pool_drain_soft_evict_grace_secs, + ), + me_pool_drain_soft_evict_per_writer: AtomicU8::new( + me_pool_drain_soft_evict_per_writer.max(1), + ), + me_pool_drain_soft_evict_budget_per_core: AtomicU32::new( + me_pool_drain_soft_evict_budget_per_core.max(1) as u32, + ), + me_pool_drain_soft_evict_cooldown_ms: AtomicU64::new( + me_pool_drain_soft_evict_cooldown_ms.max(1), + ), + me_pool_force_close_secs: AtomicU64::new(Self::normalize_force_close_secs( + me_pool_force_close_secs, + )), + me_pool_min_fresh_ratio_permille: AtomicU32::new(Self::ratio_to_permille( + me_pool_min_fresh_ratio, + )), + me_last_drain_gate_route_quorum_ok: AtomicBool::new(false), + me_last_drain_gate_redundancy_ok: AtomicBool::new(false), + me_last_drain_gate_block_reason: AtomicU8::new(MeDrainGateReason::Open as u8), + me_last_drain_gate_updated_at_epoch_secs: AtomicU64::new(now_epoch_secs), + }), + single_endpoint_runtime: Arc::new(SingleEndpointRuntimeCore { + me_single_endpoint_shadow_writers: AtomicU8::new(me_single_endpoint_shadow_writers), + me_single_endpoint_outage_mode_enabled: AtomicBool::new( + me_single_endpoint_outage_mode_enabled, + ), + me_single_endpoint_outage_disable_quarantine: AtomicBool::new( + me_single_endpoint_outage_disable_quarantine, + ), + me_single_endpoint_outage_backoff_min_ms: AtomicU64::new( + me_single_endpoint_outage_backoff_min_ms, + ), + me_single_endpoint_outage_backoff_max_ms: AtomicU64::new( + me_single_endpoint_outage_backoff_max_ms, + ), + me_single_endpoint_shadow_rotate_every_secs: AtomicU64::new( + me_single_endpoint_shadow_rotate_every_secs, + ), + }), + binding_policy: Arc::new(BindingPolicyCore { + me_bind_stale_mode: AtomicU8::new(me_bind_stale_mode.as_u8()), + me_bind_stale_ttl_secs: AtomicU64::new(me_bind_stale_ttl_secs), + }), + nat_runtime: Arc::new(NatRuntimeCore { + nat_ip_cfg: nat_ip, + nat_ip_detected: Arc::new(RwLock::new(None)), + nat_probe, + nat_stun, + nat_stun_servers, + stun_tcp_fallback, + http_ip_detect_urls, + nat_stun_live_servers: Arc::new(RwLock::new(Vec::new())), + nat_probe_concurrency: nat_probe_concurrency.max(1), + detected_ipv6, + nat_probe_attempts: std::sync::atomic::AtomicU8::new(0), + nat_probe_disabled: std::sync::atomic::AtomicBool::new(false), + stun_backoff_until: Arc::new(RwLock::new(None)), + nat_reflection_cache: Arc::new(Mutex::new(NatReflectionCache::default())), + nat_reflection_singleflight_v4: Arc::new(Mutex::new(())), + nat_reflection_singleflight_v6: Arc::new(Mutex::new(())), + }), + reconnect_runtime: Arc::new(ReconnectRuntimeCore { + me_one_retry, + me_one_timeout: Duration::from_millis(me_one_timeout_ms), + me_warmup_stagger_enabled, + me_warmup_step_delay: Duration::from_millis(me_warmup_step_delay_ms), + me_warmup_step_jitter: Duration::from_millis(me_warmup_step_jitter_ms), + me_reconnect_max_concurrent_per_dc, + me_reconnect_backoff_base: Duration::from_millis(me_reconnect_backoff_base_ms), + me_reconnect_backoff_cap: Duration::from_millis(me_reconnect_backoff_cap_ms), + me_reconnect_fast_retry_count, + }), + floor_runtime: Arc::new(FloorRuntimeCore { + me_floor_mode: AtomicU8::new(me_floor_mode.as_u8()), + me_adaptive_floor_idle_secs: AtomicU64::new(me_adaptive_floor_idle_secs), + me_adaptive_floor_min_writers_single_endpoint: AtomicU8::new( + me_adaptive_floor_min_writers_single_endpoint, + ), + me_adaptive_floor_min_writers_multi_endpoint: AtomicU8::new( + me_adaptive_floor_min_writers_multi_endpoint, + ), + me_adaptive_floor_recover_grace_secs: AtomicU64::new( + me_adaptive_floor_recover_grace_secs, + ), + me_adaptive_floor_writers_per_core_total: AtomicU32::new( + me_adaptive_floor_writers_per_core_total as u32, + ), + me_adaptive_floor_cpu_cores_override: AtomicU32::new( + me_adaptive_floor_cpu_cores_override as u32, + ), + me_adaptive_floor_max_extra_writers_single_per_core: AtomicU32::new( + me_adaptive_floor_max_extra_writers_single_per_core as u32, + ), + me_adaptive_floor_max_extra_writers_multi_per_core: AtomicU32::new( + me_adaptive_floor_max_extra_writers_multi_per_core as u32, + ), + me_adaptive_floor_max_active_writers_per_core: AtomicU32::new( + me_adaptive_floor_max_active_writers_per_core as u32, + ), + me_adaptive_floor_max_warm_writers_per_core: AtomicU32::new( + me_adaptive_floor_max_warm_writers_per_core as u32, + ), + me_adaptive_floor_max_active_writers_global: AtomicU32::new( + me_adaptive_floor_max_active_writers_global, + ), + me_adaptive_floor_max_warm_writers_global: AtomicU32::new( + me_adaptive_floor_max_warm_writers_global, + ), + me_adaptive_floor_cpu_cores_detected: AtomicU32::new(1), + me_adaptive_floor_cpu_cores_effective: AtomicU32::new(1), + me_adaptive_floor_global_cap_raw: AtomicU64::new(0), + me_adaptive_floor_global_cap_effective: AtomicU64::new(0), + me_adaptive_floor_target_writers_total: AtomicU64::new(0), + me_adaptive_floor_active_cap_configured: AtomicU64::new(0), + me_adaptive_floor_active_cap_effective: AtomicU64::new(0), + me_adaptive_floor_warm_cap_configured: AtomicU64::new(0), + me_adaptive_floor_warm_cap_effective: AtomicU64::new(0), + me_adaptive_floor_active_writers_current: AtomicU64::new(0), + me_adaptive_floor_warm_writers_current: AtomicU64::new(0), + }), + writer_selection_policy: Arc::new(WriterSelectionPolicyCore { + secret_atomic_snapshot: AtomicBool::new(me_secret_atomic_snapshot), + me_deterministic_writer_sort: AtomicBool::new(me_deterministic_writer_sort), + me_writer_pick_mode: AtomicU8::new(me_writer_pick_mode.as_u8()), + me_writer_pick_sample_size: AtomicU8::new(me_writer_pick_sample_size.clamp(2, 4)), + }), + transport_policy: Arc::new(TransportPolicyCore { + me_socks_kdf_policy: AtomicU8::new(me_socks_kdf_policy.as_u8()), + me_route_backpressure_enabled: Arc::new(AtomicBool::new( + me_route_backpressure_enabled, + )), + me_route_fairshare_enabled: Arc::new(AtomicBool::new(me_route_fairshare_enabled)), + me_reader_route_data_wait_ms: Arc::new(AtomicU64::new( + me_reader_route_data_wait_ms, + )), + }), + lifecycle: MePoolLifecycle::new(), + decision, + upstream, + rng, + proxy_tag, + proxy_secret: Arc::new(RwLock::new(SecretSnapshot { + epoch: 1, + key_selector: if proxy_secret.len() >= 4 { + u32::from_le_bytes([ + proxy_secret[0], + proxy_secret[1], + proxy_secret[2], + proxy_secret[3], + ]) + } else { + 0 + }, + secret: proxy_secret, + })), + stats, + pool_size: 2, + proxy_map_v4: Arc::new(RwLock::new(proxy_map_v4)), + proxy_map_v6: Arc::new(RwLock::new(proxy_map_v6)), + endpoint_dc_map: Arc::new(RwLock::new(endpoint_dc_map)), + default_dc: AtomicI32::new(default_dc.unwrap_or(2)), + next_writer_id: AtomicU64::new(1), + writer_connect_active_reserved: AtomicUsize::new(0), + writer_connect_warm_reserved: AtomicUsize::new(0), + rtt_stats: Arc::new(Mutex::new(HashMap::new())), + refill_states: Arc::new(ParkingMutex::new(HashMap::new())), + refill_running: AtomicUsize::new(0), + refill_pending: AtomicUsize::new(0), + conn_count: AtomicUsize::new(0), + draining_active_runtime: AtomicU64::new(0), + endpoint_quarantine: Arc::new(Mutex::new(HashMap::new())), + kdf_material_fingerprint: Arc::new(RwLock::new(HashMap::new())), + runtime_ready: AtomicBool::new(false), + }) + } +} diff --git a/src/transport/middle_proxy/pool/routing.rs b/src/transport/middle_proxy/pool/routing.rs new file mode 100644 index 0000000..c999730 --- /dev/null +++ b/src/transport/middle_proxy/pool/routing.rs @@ -0,0 +1,286 @@ +use super::*; + +impl MePool { + pub(in crate::transport::middle_proxy) fn single_endpoint_outage_mode_enabled(&self) -> bool { + self.single_endpoint_runtime + .me_single_endpoint_outage_mode_enabled + .load(Ordering::Relaxed) + } + + pub(in crate::transport::middle_proxy) fn single_endpoint_outage_disable_quarantine( + &self, + ) -> bool { + self.single_endpoint_runtime + .me_single_endpoint_outage_disable_quarantine + .load(Ordering::Relaxed) + } + + pub(in crate::transport::middle_proxy) fn single_endpoint_outage_backoff_bounds_ms( + &self, + ) -> (u64, u64) { + let min_ms = self + .single_endpoint_runtime + .me_single_endpoint_outage_backoff_min_ms + .load(Ordering::Relaxed); + let max_ms = self + .single_endpoint_runtime + .me_single_endpoint_outage_backoff_max_ms + .load(Ordering::Relaxed); + if min_ms <= max_ms { + (min_ms, max_ms) + } else { + (max_ms, min_ms) + } + } + + pub(in crate::transport::middle_proxy) fn single_endpoint_shadow_rotate_interval( + &self, + ) -> Option { + let secs = self + .single_endpoint_runtime + .me_single_endpoint_shadow_rotate_every_secs + .load(Ordering::Relaxed); + if secs == 0 { + None + } else { + Some(Duration::from_secs(secs)) + } + } + + pub(in crate::transport::middle_proxy) fn family_order(&self) -> Vec { + let mut order = Vec::new(); + if self.decision.prefer_ipv6() { + if self.decision.ipv6_me { + order.push(IpFamily::V6); + } + if self.decision.ipv4_me { + order.push(IpFamily::V4); + } + } else { + if self.decision.ipv4_me { + order.push(IpFamily::V4); + } + if self.decision.ipv6_me { + order.push(IpFamily::V6); + } + } + order + } + + pub(in crate::transport::middle_proxy) fn default_dc_for_routing(&self) -> i32 { + let dc = self.default_dc.load(Ordering::Relaxed); + if dc == 0 { 2 } else { dc } + } + + pub(in crate::transport::middle_proxy) async fn has_configured_endpoints_for_dc( + &self, + dc: i32, + ) -> bool { + if self.decision.ipv4_me { + let map = self.proxy_map_v4.read().await; + if map.get(&dc).is_some_and(|endpoints| !endpoints.is_empty()) { + return true; + } + } + + if self.decision.ipv6_me { + let map = self.proxy_map_v6.read().await; + if map.get(&dc).is_some_and(|endpoints| !endpoints.is_empty()) { + return true; + } + } + + false + } + + pub(in crate::transport::middle_proxy) async fn resolve_target_dc_for_routing( + &self, + target_dc: i32, + ) -> (i32, bool) { + if target_dc == 0 { + return (self.default_dc_for_routing(), true); + } + + if self.has_configured_endpoints_for_dc(target_dc).await { + return (target_dc, false); + } + + (self.default_dc_for_routing(), true) + } + + pub(in crate::transport::middle_proxy) async fn resolve_dc_for_endpoint( + &self, + addr: SocketAddr, + ) -> i32 { + if let Some(cached) = self.endpoint_dc_map.read().await.get(&addr).copied() + && let Some(dc) = cached + { + return dc; + } + + self.default_dc_for_routing() + } + + pub(in crate::transport::middle_proxy) async fn proxy_map_for_family( + &self, + family: IpFamily, + ) -> HashMap> { + match family { + IpFamily::V4 => self.proxy_map_v4.read().await.clone(), + IpFamily::V6 => self.proxy_map_v6.read().await.clone(), + } + } + + pub(in crate::transport::middle_proxy) fn merge_endpoint_dc( + endpoint_dc_map: &mut HashMap>, + dc: i32, + ip: IpAddr, + port: u16, + ) { + let endpoint = SocketAddr::new(ip, port); + match endpoint_dc_map.get_mut(&endpoint) { + None => { + endpoint_dc_map.insert(endpoint, Some(dc)); + } + Some(existing) => { + if existing.is_some_and(|existing_dc| existing_dc != dc) { + *existing = None; + } + } + } + } + + pub(in crate::transport::middle_proxy) fn build_preferred_endpoints_by_dc( + decision: &NetworkDecision, + map_v4: &HashMap>, + map_v6: &HashMap>, + ) -> HashMap> { + let mut out = HashMap::>::new(); + let mut dcs = HashSet::::new(); + dcs.extend(map_v4.keys().copied()); + dcs.extend(map_v6.keys().copied()); + + for dc in dcs { + let v4 = map_v4 + .get(&dc) + .map(|items| { + items + .iter() + .map(|(ip, port)| SocketAddr::new(*ip, *port)) + .collect::>() + }) + .unwrap_or_default(); + let v6 = map_v6 + .get(&dc) + .map(|items| { + items + .iter() + .map(|(ip, port)| SocketAddr::new(*ip, *port)) + .collect::>() + }) + .unwrap_or_default(); + + let mut selected = if decision.effective_multipath { + let mut both = Vec::::with_capacity(v4.len().saturating_add(v6.len())); + if decision.prefer_ipv6() { + both.extend(v6.iter().copied()); + both.extend(v4.iter().copied()); + } else { + both.extend(v4.iter().copied()); + both.extend(v6.iter().copied()); + } + both + } else if decision.prefer_ipv6() { + if !v6.is_empty() { v6 } else { v4 } + } else if !v4.is_empty() { + v4 + } else { + v6 + }; + + selected.sort_unstable(); + selected.dedup(); + out.insert(dc, selected); + } + + out + } + + pub(in crate::transport::middle_proxy) fn build_endpoint_dc_map_from_maps( + map_v4: &HashMap>, + map_v6: &HashMap>, + ) -> HashMap> { + let mut endpoint_dc_map = HashMap::>::new(); + for (dc, endpoints) in map_v4 { + for (ip, port) in endpoints { + Self::merge_endpoint_dc(&mut endpoint_dc_map, *dc, *ip, *port); + } + } + for (dc, endpoints) in map_v6 { + for (ip, port) in endpoints { + Self::merge_endpoint_dc(&mut endpoint_dc_map, *dc, *ip, *port); + } + } + endpoint_dc_map + } + + pub(in crate::transport::middle_proxy) async fn rebuild_endpoint_dc_map(&self) { + let map_v4 = self.proxy_map_v4.read().await.clone(); + let map_v6 = self.proxy_map_v6.read().await.clone(); + let rebuilt = Self::build_endpoint_dc_map_from_maps(&map_v4, &map_v6); + let preferred = Self::build_preferred_endpoints_by_dc(&self.decision, &map_v4, &map_v6); + *self.endpoint_dc_map.write().await = rebuilt; + self.preferred_endpoints_by_dc.store(Arc::new(preferred)); + let configured_endpoints = self + .endpoint_dc_map + .read() + .await + .keys() + .copied() + .collect::>(); + { + let mut quarantine = self.endpoint_quarantine.lock().await; + let now = Instant::now(); + quarantine.retain(|addr, expiry| *expiry > now && configured_endpoints.contains(addr)); + } + { + let mut kdf_fp = self.kdf_material_fingerprint.write().await; + kdf_fp.retain(|addr, _| configured_endpoints.contains(addr)); + } + } + + pub(in crate::transport::middle_proxy) async fn preferred_endpoints_for_dc( + &self, + dc: i32, + ) -> Vec { + let guard = self.preferred_endpoints_by_dc.load(); + guard.get(&dc).cloned().unwrap_or_default() + } + + pub(in crate::transport::middle_proxy) fn health_interval_unhealthy(&self) -> Duration { + Duration::from_millis( + self.health_runtime + .me_health_interval_ms_unhealthy + .load(Ordering::Relaxed) + .max(1), + ) + } + + pub(in crate::transport::middle_proxy) fn health_interval_healthy(&self) -> Duration { + Duration::from_millis( + self.health_runtime + .me_health_interval_ms_healthy + .load(Ordering::Relaxed) + .max(1), + ) + } + + pub(in crate::transport::middle_proxy) fn warn_rate_limit_duration(&self) -> Duration { + Duration::from_millis( + self.health_runtime + .me_warn_rate_limit_ms + .load(Ordering::Relaxed) + .max(1), + ) + } +} diff --git a/src/transport/middle_proxy/pool/runtime_policy.rs b/src/transport/middle_proxy/pool/runtime_policy.rs new file mode 100644 index 0000000..3c11bf8 --- /dev/null +++ b/src/transport/middle_proxy/pool/runtime_policy.rs @@ -0,0 +1,408 @@ +use super::*; + +impl MePool { + /// Creates the immutable byte semaphore assigned to one ME writer generation. + pub(crate) fn new_writer_byte_budget(&self) -> Arc { + Arc::new(Semaphore::new( + self.writer_lifecycle.writer_byte_budget_permits, + )) + } + + pub fn current_generation(&self) -> u64 { + self.reinit.active_generation.load(Ordering::Relaxed) + } + + pub fn set_runtime_ready(&self, ready: bool) { + self.runtime_ready.store(ready, Ordering::Relaxed); + } + + pub fn is_runtime_ready(&self) -> bool { + self.runtime_ready.load(Ordering::Relaxed) + } + + pub(in crate::transport::middle_proxy) fn now_epoch_millis() -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_millis() as u64 + } + + pub(in crate::transport::middle_proxy) fn notify_writer_epoch(&self) { + self.writer_epoch.send_modify(|epoch| { + *epoch = epoch.wrapping_add(1); + }); + } + + pub(in crate::transport::middle_proxy) fn set_family_runtime_state( + &self, + family: IpFamily, + state: MeFamilyRuntimeState, + state_since_epoch_secs: u64, + suppressed_until_epoch_secs: u64, + fail_streak: u32, + recover_success_streak: u32, + ) { + let snapshot = Arc::new(FamilyHealthSnapshot::new( + state, + state_since_epoch_secs, + suppressed_until_epoch_secs, + fail_streak, + recover_success_streak, + )); + match family { + IpFamily::V4 => self.health_runtime.family_health_v4.store(snapshot), + IpFamily::V6 => self.health_runtime.family_health_v6.store(snapshot), + } + } + + pub(crate) fn family_runtime_state(&self, family: IpFamily) -> MeFamilyRuntimeState { + match family { + IpFamily::V4 => self.health_runtime.family_health_v4.load().state, + IpFamily::V6 => self.health_runtime.family_health_v6.load().state, + } + } + + pub(crate) fn family_runtime_state_since_epoch_secs(&self, family: IpFamily) -> u64 { + match family { + IpFamily::V4 => { + self.health_runtime + .family_health_v4 + .load() + .state_since_epoch_secs + } + IpFamily::V6 => { + self.health_runtime + .family_health_v6 + .load() + .state_since_epoch_secs + } + } + } + + pub(crate) fn family_suppressed_until_epoch_secs(&self, family: IpFamily) -> u64 { + match family { + IpFamily::V4 => { + self.health_runtime + .family_health_v4 + .load() + .suppressed_until_epoch_secs + } + IpFamily::V6 => { + self.health_runtime + .family_health_v6 + .load() + .suppressed_until_epoch_secs + } + } + } + + pub(crate) fn family_fail_streak(&self, family: IpFamily) -> u32 { + match family { + IpFamily::V4 => self.health_runtime.family_health_v4.load().fail_streak, + IpFamily::V6 => self.health_runtime.family_health_v6.load().fail_streak, + } + } + + pub(crate) fn family_recover_success_streak(&self, family: IpFamily) -> u32 { + match family { + IpFamily::V4 => { + self.health_runtime + .family_health_v4 + .load() + .recover_success_streak + } + IpFamily::V6 => { + self.health_runtime + .family_health_v6 + .load() + .recover_success_streak + } + } + } + + pub(crate) fn is_family_temporarily_suppressed( + &self, + family: IpFamily, + now_epoch_secs: u64, + ) -> bool { + self.family_suppressed_until_epoch_secs(family) > now_epoch_secs + } + + pub(in crate::transport::middle_proxy) fn family_enabled_for_drain_coverage( + &self, + family: IpFamily, + now_epoch_secs: u64, + ) -> bool { + let configured = match family { + IpFamily::V4 => self.decision.ipv4_me, + IpFamily::V6 => self.decision.ipv6_me, + }; + configured && !self.is_family_temporarily_suppressed(family, now_epoch_secs) + } + + pub(in crate::transport::middle_proxy) fn set_last_drain_gate( + &self, + route_quorum_ok: bool, + redundancy_ok: bool, + block_reason: MeDrainGateReason, + updated_at_epoch_secs: u64, + ) { + self.drain_runtime + .me_last_drain_gate_route_quorum_ok + .store(route_quorum_ok, Ordering::Relaxed); + self.drain_runtime + .me_last_drain_gate_redundancy_ok + .store(redundancy_ok, Ordering::Relaxed); + self.drain_runtime + .me_last_drain_gate_block_reason + .store(block_reason as u8, Ordering::Relaxed); + self.drain_runtime + .me_last_drain_gate_updated_at_epoch_secs + .store(updated_at_epoch_secs, Ordering::Relaxed); + } + + pub(crate) fn last_drain_gate_route_quorum_ok(&self) -> bool { + self.drain_runtime + .me_last_drain_gate_route_quorum_ok + .load(Ordering::Relaxed) + } + + pub(crate) fn last_drain_gate_redundancy_ok(&self) -> bool { + self.drain_runtime + .me_last_drain_gate_redundancy_ok + .load(Ordering::Relaxed) + } + + pub(crate) fn last_drain_gate_block_reason(&self) -> MeDrainGateReason { + MeDrainGateReason::from_u8( + self.drain_runtime + .me_last_drain_gate_block_reason + .load(Ordering::Relaxed), + ) + } + + pub(crate) fn last_drain_gate_updated_at_epoch_secs(&self) -> u64 { + self.drain_runtime + .me_last_drain_gate_updated_at_epoch_secs + .load(Ordering::Relaxed) + } + + pub fn update_runtime_reinit_policy( + &self, + hardswap: bool, + drain_ttl_secs: u64, + instadrain: bool, + pool_drain_threshold: u64, + pool_drain_soft_evict_enabled: bool, + pool_drain_soft_evict_grace_secs: u64, + pool_drain_soft_evict_per_writer: u8, + pool_drain_soft_evict_budget_per_core: u16, + pool_drain_soft_evict_cooldown_ms: u64, + force_close_secs: u64, + min_fresh_ratio: f32, + hardswap_warmup_delay_min_ms: u64, + hardswap_warmup_delay_max_ms: u64, + hardswap_warmup_extra_passes: u8, + hardswap_warmup_pass_backoff_base_ms: u64, + bind_stale_mode: MeBindStaleMode, + bind_stale_ttl_secs: u64, + secret_atomic_snapshot: bool, + deterministic_writer_sort: bool, + writer_pick_mode: MeWriterPickMode, + writer_pick_sample_size: u8, + single_endpoint_shadow_writers: u8, + single_endpoint_outage_mode_enabled: bool, + single_endpoint_outage_disable_quarantine: bool, + single_endpoint_outage_backoff_min_ms: u64, + single_endpoint_outage_backoff_max_ms: u64, + single_endpoint_shadow_rotate_every_secs: u64, + floor_mode: MeFloorMode, + adaptive_floor_idle_secs: u64, + adaptive_floor_min_writers_single_endpoint: u8, + adaptive_floor_min_writers_multi_endpoint: u8, + adaptive_floor_recover_grace_secs: u64, + adaptive_floor_writers_per_core_total: u16, + adaptive_floor_cpu_cores_override: u16, + adaptive_floor_max_extra_writers_single_per_core: u16, + adaptive_floor_max_extra_writers_multi_per_core: u16, + adaptive_floor_max_active_writers_per_core: u16, + adaptive_floor_max_warm_writers_per_core: u16, + adaptive_floor_max_active_writers_global: u32, + adaptive_floor_max_warm_writers_global: u32, + me_health_interval_ms_unhealthy: u64, + me_health_interval_ms_healthy: u64, + me_warn_rate_limit_ms: u64, + ) { + self.reinit.hardswap.store(hardswap, Ordering::Relaxed); + self.drain_runtime + .me_pool_drain_ttl_secs + .store(drain_ttl_secs, Ordering::Relaxed); + self.drain_runtime + .me_instadrain + .store(instadrain, Ordering::Relaxed); + self.drain_runtime + .me_pool_drain_threshold + .store(pool_drain_threshold, Ordering::Relaxed); + // Runtime soft-evict knobs are updated lock-free to keep control-plane + // writes non-blocking; readers observe a short eventual-consistency + // window by design. + self.drain_runtime + .me_pool_drain_soft_evict_enabled + .store(pool_drain_soft_evict_enabled, Ordering::Relaxed); + self.drain_runtime + .me_pool_drain_soft_evict_grace_secs + .store(pool_drain_soft_evict_grace_secs, Ordering::Relaxed); + self.drain_runtime + .me_pool_drain_soft_evict_per_writer + .store(pool_drain_soft_evict_per_writer.max(1), Ordering::Relaxed); + self.drain_runtime + .me_pool_drain_soft_evict_budget_per_core + .store( + pool_drain_soft_evict_budget_per_core.max(1) as u32, + Ordering::Relaxed, + ); + self.drain_runtime + .me_pool_drain_soft_evict_cooldown_ms + .store(pool_drain_soft_evict_cooldown_ms.max(1), Ordering::Relaxed); + self.drain_runtime.me_pool_force_close_secs.store( + Self::normalize_force_close_secs(force_close_secs), + Ordering::Relaxed, + ); + self.drain_runtime + .me_pool_min_fresh_ratio_permille + .store(Self::ratio_to_permille(min_fresh_ratio), Ordering::Relaxed); + self.reinit + .me_hardswap_warmup_delay_min_ms + .store(hardswap_warmup_delay_min_ms, Ordering::Relaxed); + self.reinit + .me_hardswap_warmup_delay_max_ms + .store(hardswap_warmup_delay_max_ms, Ordering::Relaxed); + self.reinit + .me_hardswap_warmup_extra_passes + .store(hardswap_warmup_extra_passes as u32, Ordering::Relaxed); + self.reinit + .me_hardswap_warmup_pass_backoff_base_ms + .store(hardswap_warmup_pass_backoff_base_ms, Ordering::Relaxed); + self.binding_policy + .me_bind_stale_mode + .store(bind_stale_mode.as_u8(), Ordering::Relaxed); + self.binding_policy + .me_bind_stale_ttl_secs + .store(bind_stale_ttl_secs, Ordering::Relaxed); + self.writer_selection_policy + .secret_atomic_snapshot + .store(secret_atomic_snapshot, Ordering::Relaxed); + self.writer_selection_policy + .me_deterministic_writer_sort + .store(deterministic_writer_sort, Ordering::Relaxed); + let previous_writer_pick_mode = self.writer_pick_mode(); + self.writer_selection_policy + .me_writer_pick_mode + .store(writer_pick_mode.as_u8(), Ordering::Relaxed); + self.writer_selection_policy + .me_writer_pick_sample_size + .store(writer_pick_sample_size.clamp(2, 4), Ordering::Relaxed); + if previous_writer_pick_mode != writer_pick_mode { + self.stats.increment_me_writer_pick_mode_switch_total(); + } + self.single_endpoint_runtime + .me_single_endpoint_shadow_writers + .store(single_endpoint_shadow_writers, Ordering::Relaxed); + self.single_endpoint_runtime + .me_single_endpoint_outage_mode_enabled + .store(single_endpoint_outage_mode_enabled, Ordering::Relaxed); + self.single_endpoint_runtime + .me_single_endpoint_outage_disable_quarantine + .store(single_endpoint_outage_disable_quarantine, Ordering::Relaxed); + self.single_endpoint_runtime + .me_single_endpoint_outage_backoff_min_ms + .store(single_endpoint_outage_backoff_min_ms, Ordering::Relaxed); + self.single_endpoint_runtime + .me_single_endpoint_outage_backoff_max_ms + .store(single_endpoint_outage_backoff_max_ms, Ordering::Relaxed); + self.single_endpoint_runtime + .me_single_endpoint_shadow_rotate_every_secs + .store(single_endpoint_shadow_rotate_every_secs, Ordering::Relaxed); + let previous_floor_mode = self.floor_mode(); + self.floor_runtime + .me_floor_mode + .store(floor_mode.as_u8(), Ordering::Relaxed); + self.floor_runtime + .me_adaptive_floor_idle_secs + .store(adaptive_floor_idle_secs, Ordering::Relaxed); + self.floor_runtime + .me_adaptive_floor_min_writers_single_endpoint + .store( + adaptive_floor_min_writers_single_endpoint, + Ordering::Relaxed, + ); + self.floor_runtime + .me_adaptive_floor_min_writers_multi_endpoint + .store(adaptive_floor_min_writers_multi_endpoint, Ordering::Relaxed); + self.floor_runtime + .me_adaptive_floor_recover_grace_secs + .store(adaptive_floor_recover_grace_secs, Ordering::Relaxed); + self.floor_runtime + .me_adaptive_floor_writers_per_core_total + .store( + adaptive_floor_writers_per_core_total as u32, + Ordering::Relaxed, + ); + self.floor_runtime + .me_adaptive_floor_cpu_cores_override + .store(adaptive_floor_cpu_cores_override as u32, Ordering::Relaxed); + self.floor_runtime + .me_adaptive_floor_max_extra_writers_single_per_core + .store( + adaptive_floor_max_extra_writers_single_per_core as u32, + Ordering::Relaxed, + ); + self.floor_runtime + .me_adaptive_floor_max_extra_writers_multi_per_core + .store( + adaptive_floor_max_extra_writers_multi_per_core as u32, + Ordering::Relaxed, + ); + self.floor_runtime + .me_adaptive_floor_max_active_writers_per_core + .store( + adaptive_floor_max_active_writers_per_core as u32, + Ordering::Relaxed, + ); + self.floor_runtime + .me_adaptive_floor_max_warm_writers_per_core + .store( + adaptive_floor_max_warm_writers_per_core as u32, + Ordering::Relaxed, + ); + self.floor_runtime + .me_adaptive_floor_max_active_writers_global + .store(adaptive_floor_max_active_writers_global, Ordering::Relaxed); + self.floor_runtime + .me_adaptive_floor_max_warm_writers_global + .store(adaptive_floor_max_warm_writers_global, Ordering::Relaxed); + self.health_runtime + .me_health_interval_ms_unhealthy + .store(me_health_interval_ms_unhealthy.max(1), Ordering::Relaxed); + self.health_runtime + .me_health_interval_ms_healthy + .store(me_health_interval_ms_healthy.max(1), Ordering::Relaxed); + self.health_runtime + .me_warn_rate_limit_ms + .store(me_warn_rate_limit_ms.max(1), Ordering::Relaxed); + if previous_floor_mode != floor_mode { + self.stats.increment_me_floor_mode_switch_total(); + match (previous_floor_mode, floor_mode) { + (MeFloorMode::Static, MeFloorMode::Adaptive) => { + self.stats + .increment_me_floor_mode_switch_static_to_adaptive_total(); + } + (MeFloorMode::Adaptive, MeFloorMode::Static) => { + self.stats + .increment_me_floor_mode_switch_adaptive_to_static_total(); + } + _ => {} + } + } + } +} diff --git a/src/transport/middle_proxy/pool/selection_policy.rs b/src/transport/middle_proxy/pool/selection_policy.rs new file mode 100644 index 0000000..b0d0b93 --- /dev/null +++ b/src/transport/middle_proxy/pool/selection_policy.rs @@ -0,0 +1,289 @@ +use super::*; + +impl MePool { + pub(in crate::transport::middle_proxy) async fn key_selector(&self) -> u32 { + self.proxy_secret.read().await.key_selector + } + + pub(in crate::transport::middle_proxy) async fn non_draining_writer_counts_by_contour( + &self, + ) -> (usize, usize, usize) { + let ws = self.writers.read().await; + let mut active = 0usize; + let mut warm = 0usize; + for writer in ws.iter() { + if writer.draining.load(Ordering::Relaxed) { + continue; + } + match WriterContour::from_u8(writer.contour.load(Ordering::Relaxed)) { + WriterContour::Active => active = active.saturating_add(1), + WriterContour::Warm => warm = warm.saturating_add(1), + WriterContour::Draining => {} + } + } + (active, warm, active.saturating_add(warm)) + } + + pub(in crate::transport::middle_proxy) async fn active_contour_writer_count_total( + &self, + ) -> usize { + let (active, _, _) = self.non_draining_writer_counts_by_contour().await; + active + } + + pub(in crate::transport::middle_proxy) async fn secret_snapshot(&self) -> SecretSnapshot { + self.proxy_secret.read().await.clone() + } + + pub(in crate::transport::middle_proxy) fn bind_stale_mode(&self) -> MeBindStaleMode { + MeBindStaleMode::from_u8( + self.binding_policy + .me_bind_stale_mode + .load(Ordering::Relaxed), + ) + } + + pub(in crate::transport::middle_proxy) fn writer_pick_mode(&self) -> MeWriterPickMode { + MeWriterPickMode::from_u8( + self.writer_selection_policy + .me_writer_pick_mode + .load(Ordering::Relaxed), + ) + } + + pub(in crate::transport::middle_proxy) fn writer_pick_sample_size(&self) -> usize { + self.writer_selection_policy + .me_writer_pick_sample_size + .load(Ordering::Relaxed) + .clamp(2, 4) as usize + } + + pub(in crate::transport::middle_proxy) fn required_writers_for_dc( + &self, + endpoint_count: usize, + ) -> usize { + if endpoint_count == 0 { + return 0; + } + if endpoint_count == 1 { + let shadow = self + .single_endpoint_runtime + .me_single_endpoint_shadow_writers + .load(Ordering::Relaxed) as usize; + return (1 + shadow).max(3); + } + endpoint_count.max(3) + } + + pub(in crate::transport::middle_proxy) fn floor_mode(&self) -> MeFloorMode { + MeFloorMode::from_u8(self.floor_runtime.me_floor_mode.load(Ordering::Relaxed)) + } + + pub(in crate::transport::middle_proxy) fn adaptive_floor_min_writers_multi_endpoint( + &self, + ) -> usize { + (self + .floor_runtime + .me_adaptive_floor_min_writers_multi_endpoint + .load(Ordering::Relaxed) as usize) + .max(1) + } + + pub(in crate::transport::middle_proxy) fn adaptive_floor_max_extra_single_per_core( + &self, + ) -> usize { + self.floor_runtime + .me_adaptive_floor_max_extra_writers_single_per_core + .load(Ordering::Relaxed) as usize + } + + pub(in crate::transport::middle_proxy) fn adaptive_floor_max_extra_multi_per_core( + &self, + ) -> usize { + self.floor_runtime + .me_adaptive_floor_max_extra_writers_multi_per_core + .load(Ordering::Relaxed) as usize + } + + pub(in crate::transport::middle_proxy) fn adaptive_floor_max_active_writers_per_core( + &self, + ) -> usize { + (self + .floor_runtime + .me_adaptive_floor_max_active_writers_per_core + .load(Ordering::Relaxed) as usize) + .max(1) + } + + pub(in crate::transport::middle_proxy) fn adaptive_floor_max_warm_writers_per_core( + &self, + ) -> usize { + (self + .floor_runtime + .me_adaptive_floor_max_warm_writers_per_core + .load(Ordering::Relaxed) as usize) + .max(1) + } + + pub(in crate::transport::middle_proxy) fn adaptive_floor_max_active_writers_global( + &self, + ) -> usize { + (self + .floor_runtime + .me_adaptive_floor_max_active_writers_global + .load(Ordering::Relaxed) as usize) + .max(1) + } + + pub(in crate::transport::middle_proxy) fn adaptive_floor_max_warm_writers_global( + &self, + ) -> usize { + (self + .floor_runtime + .me_adaptive_floor_max_warm_writers_global + .load(Ordering::Relaxed) as usize) + .max(1) + } + + pub(in crate::transport::middle_proxy) fn adaptive_floor_detected_cpu_cores(&self) -> usize { + std::thread::available_parallelism() + .map(|value| value.get()) + .unwrap_or(1) + .max(1) + } + + pub(in crate::transport::middle_proxy) fn adaptive_floor_effective_cpu_cores(&self) -> usize { + let detected = self.adaptive_floor_detected_cpu_cores(); + let override_cores = self + .floor_runtime + .me_adaptive_floor_cpu_cores_override + .load(Ordering::Relaxed) as usize; + let effective = if override_cores == 0 { + detected + } else { + override_cores.max(1) + }; + self.floor_runtime + .me_adaptive_floor_cpu_cores_detected + .store(detected as u32, Ordering::Relaxed); + self.floor_runtime + .me_adaptive_floor_cpu_cores_effective + .store(effective as u32, Ordering::Relaxed); + self.stats + .set_me_floor_cpu_cores_detected_gauge(detected as u64); + self.stats + .set_me_floor_cpu_cores_effective_gauge(effective as u64); + effective + } + + // Keeps per-contour (active/warm) writer budget bounded by CPU count. + // Baseline is 86 writers on the first core and +48 for each extra core. + pub(in crate::transport::middle_proxy) fn adaptive_floor_cpu_budget_per_contour_cap( + &self, + cores: usize, + ) -> usize { + const FIRST_CORE_WRITER_BUDGET: usize = 86; + const EXTRA_CORE_WRITER_BUDGET: usize = 48; + if cores == 0 { + return FIRST_CORE_WRITER_BUDGET; + } + FIRST_CORE_WRITER_BUDGET.saturating_add( + cores + .saturating_sub(1) + .saturating_mul(EXTRA_CORE_WRITER_BUDGET), + ) + } + + pub(in crate::transport::middle_proxy) fn adaptive_floor_active_cap_configured_total( + &self, + ) -> usize { + let cores = self.adaptive_floor_effective_cpu_cores(); + let per_contour_budget = self.adaptive_floor_cpu_budget_per_contour_cap(cores); + let configured = cores + .saturating_mul(self.adaptive_floor_max_active_writers_per_core()) + .min(self.adaptive_floor_max_active_writers_global()) + .min(per_contour_budget) + .max(1); + self.floor_runtime + .me_adaptive_floor_active_cap_configured + .store(configured as u64, Ordering::Relaxed); + self.stats + .set_me_floor_active_cap_configured_gauge(configured as u64); + configured + } + + pub(in crate::transport::middle_proxy) fn adaptive_floor_warm_cap_configured_total( + &self, + ) -> usize { + let cores = self.adaptive_floor_effective_cpu_cores(); + let per_contour_budget = self.adaptive_floor_cpu_budget_per_contour_cap(cores); + let configured = cores + .saturating_mul(self.adaptive_floor_max_warm_writers_per_core()) + .min(self.adaptive_floor_max_warm_writers_global()) + .min(per_contour_budget) + .max(1); + self.floor_runtime + .me_adaptive_floor_warm_cap_configured + .store(configured as u64, Ordering::Relaxed); + self.stats + .set_me_floor_warm_cap_configured_gauge(configured as u64); + configured + } + + pub(in crate::transport::middle_proxy) fn set_adaptive_floor_runtime_caps( + &self, + active_cap_configured: usize, + active_cap_effective: usize, + warm_cap_configured: usize, + warm_cap_effective: usize, + target_writers_total: usize, + active_writers_current: usize, + warm_writers_current: usize, + ) { + self.floor_runtime + .me_adaptive_floor_global_cap_raw + .store(active_cap_configured as u64, Ordering::Relaxed); + self.floor_runtime + .me_adaptive_floor_global_cap_effective + .store(active_cap_effective as u64, Ordering::Relaxed); + self.floor_runtime + .me_adaptive_floor_target_writers_total + .store(target_writers_total as u64, Ordering::Relaxed); + self.floor_runtime + .me_adaptive_floor_active_cap_configured + .store(active_cap_configured as u64, Ordering::Relaxed); + self.floor_runtime + .me_adaptive_floor_active_cap_effective + .store(active_cap_effective as u64, Ordering::Relaxed); + self.floor_runtime + .me_adaptive_floor_warm_cap_configured + .store(warm_cap_configured as u64, Ordering::Relaxed); + self.floor_runtime + .me_adaptive_floor_warm_cap_effective + .store(warm_cap_effective as u64, Ordering::Relaxed); + self.floor_runtime + .me_adaptive_floor_active_writers_current + .store(active_writers_current as u64, Ordering::Relaxed); + self.floor_runtime + .me_adaptive_floor_warm_writers_current + .store(warm_writers_current as u64, Ordering::Relaxed); + self.stats + .set_me_floor_global_cap_raw_gauge(active_cap_configured as u64); + self.stats + .set_me_floor_global_cap_effective_gauge(active_cap_effective as u64); + self.stats + .set_me_floor_target_writers_total_gauge(target_writers_total as u64); + self.stats + .set_me_floor_active_cap_configured_gauge(active_cap_configured as u64); + self.stats + .set_me_floor_active_cap_effective_gauge(active_cap_effective as u64); + self.stats + .set_me_floor_warm_cap_configured_gauge(warm_cap_configured as u64); + self.stats + .set_me_floor_warm_cap_effective_gauge(warm_cap_effective as u64); + self.stats + .set_me_writers_active_current_gauge(active_writers_current as u64); + self.stats + .set_me_writers_warm_current_gauge(warm_writers_current as u64); + } +} diff --git a/src/transport/middle_proxy/pool/transport_policy.rs b/src/transport/middle_proxy/pool/transport_policy.rs new file mode 100644 index 0000000..4fd021e --- /dev/null +++ b/src/transport/middle_proxy/pool/transport_policy.rs @@ -0,0 +1,142 @@ +use super::*; + +impl MePool { + pub fn reset_stun_state(&self) { + self.nat_runtime + .nat_probe_attempts + .store(0, Ordering::Relaxed); + self.nat_runtime + .nat_probe_disabled + .store(false, Ordering::Relaxed); + if let Ok(mut live) = self.nat_runtime.nat_stun_live_servers.try_write() { + live.clear(); + } + } + + /// Translate the local ME address into the address material sent to the proxy. + pub fn translate_our_addr(&self, addr: SocketAddr) -> SocketAddr { + self.translate_our_addr_with_reflection(addr, None) + } + + #[allow(dead_code)] + pub fn registry(&self) -> &Arc { + &self.registry + } + + pub fn update_runtime_transport_policy( + &self, + socks_kdf_policy: MeSocksKdfPolicy, + route_backpressure_enabled: bool, + route_fairshare_enabled: bool, + route_backpressure_base_timeout_ms: u64, + route_backpressure_high_timeout_ms: u64, + route_backpressure_high_watermark_pct: u8, + reader_route_data_wait_ms: u64, + ) { + self.transport_policy + .me_socks_kdf_policy + .store(socks_kdf_policy.as_u8(), Ordering::Relaxed); + self.transport_policy + .me_route_backpressure_enabled + .store(route_backpressure_enabled, Ordering::Relaxed); + self.transport_policy + .me_route_fairshare_enabled + .store(route_fairshare_enabled, Ordering::Relaxed); + self.transport_policy + .me_reader_route_data_wait_ms + .store(reader_route_data_wait_ms, Ordering::Relaxed); + self.registry.update_route_backpressure_policy( + route_backpressure_base_timeout_ms, + route_backpressure_high_timeout_ms, + route_backpressure_high_watermark_pct, + ); + } + + pub(in crate::transport::middle_proxy) fn socks_kdf_policy(&self) -> MeSocksKdfPolicy { + MeSocksKdfPolicy::from_u8( + self.transport_policy + .me_socks_kdf_policy + .load(Ordering::Relaxed), + ) + } + + pub(in crate::transport::middle_proxy) fn writers_arc(&self) -> Arc { + self.writers.clone() + } + + pub(in crate::transport::middle_proxy) fn force_close_timeout(&self) -> Option { + let secs = Self::normalize_force_close_secs( + self.drain_runtime + .me_pool_force_close_secs + .load(Ordering::Relaxed), + ); + Some(Duration::from_secs(secs)) + } + + #[allow(dead_code)] + pub(in crate::transport::middle_proxy) fn drain_soft_evict_enabled(&self) -> bool { + self.drain_runtime + .me_pool_drain_soft_evict_enabled + .load(Ordering::Relaxed) + } + + #[allow(dead_code)] + pub(in crate::transport::middle_proxy) fn drain_soft_evict_grace_secs(&self) -> u64 { + self.drain_runtime + .me_pool_drain_soft_evict_grace_secs + .load(Ordering::Relaxed) + } + + #[allow(dead_code)] + pub(in crate::transport::middle_proxy) fn drain_soft_evict_per_writer(&self) -> usize { + self.drain_runtime + .me_pool_drain_soft_evict_per_writer + .load(Ordering::Relaxed) + .max(1) as usize + } + + #[allow(dead_code)] + pub(in crate::transport::middle_proxy) fn drain_soft_evict_budget_per_core(&self) -> usize { + self.drain_runtime + .me_pool_drain_soft_evict_budget_per_core + .load(Ordering::Relaxed) + .max(1) as usize + } + + #[allow(dead_code)] + pub(in crate::transport::middle_proxy) fn drain_soft_evict_cooldown(&self) -> Duration { + Duration::from_millis( + self.drain_runtime + .me_pool_drain_soft_evict_cooldown_ms + .load(Ordering::Relaxed) + .max(1), + ) + } + + #[allow(dead_code)] + pub(in crate::transport::middle_proxy) fn draining_active_runtime(&self) -> u64 { + self.draining_active_runtime.load(Ordering::Relaxed) + } + + pub(in crate::transport::middle_proxy) fn increment_draining_active_runtime(&self) { + self.draining_active_runtime.fetch_add(1, Ordering::Relaxed); + } + + pub(in crate::transport::middle_proxy) fn decrement_draining_active_runtime(&self) { + let mut current = self.draining_active_runtime.load(Ordering::Relaxed); + loop { + if current == 0 { + break; + } + match self.draining_active_runtime.compare_exchange_weak( + current, + current - 1, + Ordering::Relaxed, + Ordering::Relaxed, + ) { + Ok(_) => break, + Err(actual) => current = actual, + } + } + } +} diff --git a/src/transport/middle_proxy/pool/writer_admission.rs b/src/transport/middle_proxy/pool/writer_admission.rs new file mode 100644 index 0000000..49b8c01 --- /dev/null +++ b/src/transport/middle_proxy/pool/writer_admission.rs @@ -0,0 +1,180 @@ +use super::*; + +impl MePool { + pub(in crate::transport::middle_proxy) async fn active_coverage_required_total(&self) -> usize { + let now_epoch_secs = Self::now_epoch_secs(); + let mut endpoints_by_dc = HashMap::>::new(); + + if self.family_enabled_for_drain_coverage(IpFamily::V4, now_epoch_secs) { + let map = self.proxy_map_v4.read().await; + for (dc, addrs) in map.iter() { + let entry = endpoints_by_dc.entry(*dc).or_default(); + for (ip, port) in addrs.iter().copied() { + entry.insert(SocketAddr::new(ip, port)); + } + } + } + + if self.family_enabled_for_drain_coverage(IpFamily::V6, now_epoch_secs) { + let map = self.proxy_map_v6.read().await; + for (dc, addrs) in map.iter() { + let entry = endpoints_by_dc.entry(*dc).or_default(); + for (ip, port) in addrs.iter().copied() { + entry.insert(SocketAddr::new(ip, port)); + } + } + } + + endpoints_by_dc + .values() + .map(|endpoints| self.required_writers_for_dc_with_floor_mode(endpoints.len(), false)) + .sum() + } + + pub(in crate::transport::middle_proxy) async fn can_open_writer_for_contour( + &self, + contour: WriterContour, + allow_coverage_override: bool, + writer_dc: i32, + ) -> bool { + let (active_writers, warm_writers, _) = self.non_draining_writer_counts_by_contour().await; + match contour { + WriterContour::Active => { + let active_cap = self.adaptive_floor_active_cap_configured_total(); + if active_writers < active_cap { + return true; + } + if !allow_coverage_override { + return false; + } + + let mut endpoints_len = 0; + let now_epoch = Self::now_epoch_secs(); + if self.family_enabled_for_drain_coverage(IpFamily::V4, now_epoch) { + if let Some(addrs) = self.proxy_map_v4.read().await.get(&writer_dc) { + endpoints_len += addrs.len(); + } + } + if self.family_enabled_for_drain_coverage(IpFamily::V6, now_epoch) { + if let Some(addrs) = self.proxy_map_v6.read().await.get(&writer_dc) { + endpoints_len += addrs.len(); + } + } + + if endpoints_len > 0 { + let base_req = + self.required_writers_for_dc_with_floor_mode(endpoints_len, false); + let active_for_dc = { + let ws = self.writers.read().await; + ws.iter() + .filter(|w| { + !w.draining.load(std::sync::atomic::Ordering::Relaxed) + && w.writer_dc == writer_dc + && matches!( + WriterContour::from_u8( + w.contour.load(std::sync::atomic::Ordering::Relaxed), + ), + WriterContour::Active + ) + }) + .count() + }; + if active_for_dc < base_req { + return true; + } + } + + let coverage_required = self.active_coverage_required_total().await; + active_writers < coverage_required + } + WriterContour::Warm => warm_writers < self.adaptive_floor_warm_cap_configured_total(), + WriterContour::Draining => true, + } + } + + pub(in crate::transport::middle_proxy) async fn reserve_writer_open( + &self, + contour: WriterContour, + allow_coverage_override: bool, + writer_dc: i32, + ) -> Option> { + let counter = match contour { + WriterContour::Active => &self.writer_connect_active_reserved, + WriterContour::Warm => &self.writer_connect_warm_reserved, + WriterContour::Draining => { + return Some(WriterOpenReservation { counter: None }); + } + }; + + loop { + if !self + .can_open_writer_for_contour(contour, allow_coverage_override, writer_dc) + .await + { + return None; + } + let (active_writers, warm_writers, _) = + self.non_draining_writer_counts_by_contour().await; + let live = match contour { + WriterContour::Active => active_writers, + WriterContour::Warm => warm_writers, + WriterContour::Draining => 0, + }; + let mut limit = match contour { + WriterContour::Active => self.adaptive_floor_active_cap_configured_total(), + WriterContour::Warm => self.adaptive_floor_warm_cap_configured_total(), + WriterContour::Draining => usize::MAX, + }; + if contour == WriterContour::Active && allow_coverage_override { + limit = limit + .max(self.active_coverage_required_total().await) + .saturating_add( + self.reconnect_runtime + .me_reconnect_max_concurrent_per_dc + .max(1) as usize, + ); + } + + let reserved = counter.load(Ordering::Acquire); + if live.saturating_add(reserved) >= limit { + return None; + } + if counter + .compare_exchange_weak(reserved, reserved + 1, Ordering::AcqRel, Ordering::Acquire) + .is_ok() + { + return Some(WriterOpenReservation { + counter: Some(counter), + }); + } + } + } + + pub(in crate::transport::middle_proxy) fn required_writers_for_dc_with_floor_mode( + &self, + endpoint_count: usize, + reduce_for_idle: bool, + ) -> usize { + let base_required = self.required_writers_for_dc(endpoint_count); + if !reduce_for_idle { + return base_required; + } + if self.floor_mode() != MeFloorMode::Adaptive { + return base_required; + } + let min_writers = if endpoint_count == 1 { + (self + .floor_runtime + .me_adaptive_floor_min_writers_single_endpoint + .load(Ordering::Relaxed) as usize) + .max(1) + } else { + (self + .floor_runtime + .me_adaptive_floor_min_writers_multi_endpoint + .load(Ordering::Relaxed) as usize) + .max(1) + }; + base_required.min(min_writers) + } +} diff --git a/src/transport/middle_proxy/pool_lifecycle.rs b/src/transport/middle_proxy/pool_lifecycle.rs index c4eb73a..dcfe93c 100644 --- a/src/transport/middle_proxy/pool_lifecycle.rs +++ b/src/transport/middle_proxy/pool_lifecycle.rs @@ -143,11 +143,8 @@ impl MePoolLifecycle { } /// Spawns a writer after its caller completed task registration. - pub(super) fn spawn_registered_writer( - &self, - registration: MeTaskRegistration<'_>, - future: F, - ) where + pub(super) fn spawn_registered_writer(&self, registration: MeTaskRegistration<'_>, future: F) + where F: Future + Send + 'static, { self.writer_tasks.spawn(future); @@ -215,11 +212,7 @@ impl MePoolLifecycle { } /// Joins producers, writers, and cleanup ownership under one deadline. - pub(super) async fn shutdown_pool( - &self, - pool: &Arc, - timeout: Duration, - ) -> bool { + pub(super) async fn shutdown_pool(&self, pool: &Arc, timeout: Duration) -> bool { let deadline = tokio::time::Instant::now() + timeout; self.begin_shutdown(); pool.set_runtime_ready(false); diff --git a/src/transport/middle_proxy/pool_refill.rs b/src/transport/middle_proxy/pool_refill.rs index 4a8c914..30a7809 100644 --- a/src/transport/middle_proxy/pool_refill.rs +++ b/src/transport/middle_proxy/pool_refill.rs @@ -392,22 +392,21 @@ impl MePool { key: dc_key, active: true, }; - self.lifecycle.spawn_registered_producer(registration, async move { - let mut current_addr = addr; - loop { - pool.stats.increment_me_refill_triggered_total(); - let restored = pool - .refill_writer_after_loss(current_addr, writer_dc) - .await; - if !restored { - warn!(%current_addr, dc = writer_dc, "ME immediate refill failed"); - } + self.lifecycle + .spawn_registered_producer(registration, async move { + let mut current_addr = addr; + loop { + pool.stats.increment_me_refill_triggered_total(); + let restored = pool.refill_writer_after_loss(current_addr, writer_dc).await; + if !restored { + warn!(%current_addr, dc = writer_dc, "ME immediate refill failed"); + } - let Some(next_addr) = run_guard.next_or_finish() else { - return; - }; - current_addr = next_addr; - } - }); + let Some(next_addr) = run_guard.next_or_finish() else { + return; + }; + current_addr = next_addr; + } + }); } } diff --git a/src/transport/middle_proxy/pool_reinit.rs b/src/transport/middle_proxy/pool_reinit.rs index de2cc17..46174cd 100644 --- a/src/transport/middle_proxy/pool_reinit.rs +++ b/src/transport/middle_proxy/pool_reinit.rs @@ -18,6 +18,13 @@ use super::pool::{ ReinitPendingState, ReinitStatusSnapshot, WriterContour, }; +// Reinitialization admission, generation state, and coverage checks. +mod coordination; +// Generation warmup and stale-writer reconciliation. +mod reconcile; + +#[cfg(test)] +mod tests; const ME_HARDSWAP_PENDING_TTL_SECS: u64 = 1800; struct ReinitAttemptGuard { @@ -70,10 +77,9 @@ fn publish_reinit_state(reinit: &ReinitCore, state: &ReinitCoordinatorState) { snapshot.warm_generations.last().copied().unwrap_or(0), Ordering::Release, ); - reinit.pending_hardswap_generation.store( - snapshot.pending_hardswap_generation, - Ordering::Release, - ); + reinit + .pending_hardswap_generation + .store(snapshot.pending_hardswap_generation, Ordering::Release); reinit.pending_hardswap_started_at_epoch_secs.store( snapshot.pending_hardswap_started_at_epoch_secs, Ordering::Release, @@ -112,722 +118,3 @@ fn commit_reinit_state( } true } - -impl MePool { - fn desired_map_hash(desired_by_dc: &HashMap>) -> u64 { - let mut hasher = DefaultHasher::new(); - let mut dcs: Vec = desired_by_dc.keys().copied().collect(); - dcs.sort_unstable(); - for dc in dcs { - dc.hash(&mut hasher); - let mut endpoints: Vec = desired_by_dc - .get(&dc) - .map(|set| set.iter().copied().collect()) - .unwrap_or_default(); - endpoints.sort_unstable(); - for endpoint in endpoints { - endpoint.hash(&mut hasher); - } - } - hasher.finish() - } - - fn reserve_reinit_attempt( - self: &Arc, - hardswap: bool, - map_hash: u64, - now_epoch_secs: u64, - ) -> ReinitReservation { - let mut state = self.reinit.coordinator.lock(); - state.desired_map_hash = map_hash; - let previous_generation = state.active_generation; - let mut pending_reused = false; - let mut pending_expired = false; - let mut pending_age_secs = 0; - - let generation = if hardswap { - let reusable = state.pending.filter(|pending| { - pending_age_secs = now_epoch_secs.saturating_sub(pending.started_at_epoch_secs); - pending_expired = pending.started_at_epoch_secs > 0 - && pending_age_secs > ME_HARDSWAP_PENDING_TTL_SECS; - pending.generation >= previous_generation - && pending.map_hash == map_hash - && !pending_expired - }); - if let Some(pending) = reusable { - pending_reused = true; - pending.generation - } else { - let generation = self.reinit.generation.fetch_add(1, Ordering::AcqRel) + 1; - state.pending = Some(ReinitPendingState { - generation, - started_at_epoch_secs: now_epoch_secs, - map_hash, - }); - generation - } - } else { - state.pending = None; - self.reinit.generation.fetch_add(1, Ordering::AcqRel) + 1 - }; - - let attempt_id = state.next_attempt_id; - state.next_attempt_id = state.next_attempt_id.saturating_add(1); - state.attempts.insert( - attempt_id, - ReinitAttemptState { - generation, - map_hash, - hardswap, - committed: false, - }, - ); - publish_reinit_state(self.reinit.as_ref(), &state); - ReinitReservation { - attempt: ReinitAttemptGuard { - reinit: Arc::clone(&self.reinit), - attempt_id, - generation, - previous_generation, - map_hash, - hardswap, - }, - pending_reused, - pending_expired, - pending_age_secs, - } - } - - fn commit_reinit_attempt(&self, attempt: &ReinitAttemptGuard) -> bool { - let mut state = self.reinit.coordinator.lock(); - if !commit_reinit_state( - &mut state, - attempt.attempt_id, - attempt.generation, - attempt.map_hash, - attempt.hardswap, - ) { - return false; - } - if attempt.hardswap { - let writers = self.writers.snapshot(); - for writer in writers.iter() { - if !writer.draining.load(Ordering::Relaxed) - && writer.generation == attempt.generation - { - writer - .contour - .store(WriterContour::Active.as_u8(), Ordering::Release); - } - } - } - publish_reinit_state(self.reinit.as_ref(), &state); - true - } - - fn coverage_ratio( - desired_by_dc: &HashMap>, - active_writer_addrs: &HashSet<(i32, SocketAddr)>, - ) -> (f32, Vec) { - if desired_by_dc.is_empty() { - return (1.0, Vec::new()); - } - - let mut missing_dc = Vec::::new(); - let mut covered = 0usize; - let mut total = 0usize; - for (dc, endpoints) in desired_by_dc { - if endpoints.is_empty() { - continue; - } - total += 1; - if endpoints - .iter() - .any(|addr| active_writer_addrs.contains(&(*dc, *addr))) - { - covered += 1; - } else { - missing_dc.push(*dc); - } - } - - missing_dc.sort_unstable(); - if total == 0 { - return (1.0, missing_dc); - } - let ratio = (covered as f32) / (total as f32); - (ratio, missing_dc) - } - - pub async fn reconcile_connections(self: &Arc, rng: &SecureRandom) { - for family in self.family_order() { - let map = self.proxy_map_for_family(family).await; - for (dc, addrs) in &map { - let dc_addrs: Vec = addrs - .iter() - .map(|(ip, port)| SocketAddr::new(*ip, *port)) - .collect(); - let dc_endpoints: HashSet = dc_addrs.iter().copied().collect(); - if self - .active_writer_count_for_dc_endpoints(*dc, &dc_endpoints) - .await - == 0 - { - let mut shuffled = dc_addrs.clone(); - shuffled.shuffle(&mut rand::rng()); - for addr in shuffled { - if self.connect_one_for_dc(addr, *dc, rng).await.is_ok() { - break; - } - } - } - } - if !self.decision.effective_multipath && self.connection_count() > 0 { - break; - } - } - } - - async fn desired_dc_endpoints(&self) -> HashMap> { - let now_epoch_secs = Self::now_epoch_secs(); - let mut out: HashMap> = HashMap::new(); - - if self.family_enabled_for_drain_coverage(IpFamily::V4, now_epoch_secs) { - let map_v4 = self.proxy_map_v4.read().await.clone(); - for (dc, addrs) in map_v4 { - let entry = out.entry(dc).or_default(); - for (ip, port) in addrs { - entry.insert(SocketAddr::new(ip, port)); - } - } - } - - if self.family_enabled_for_drain_coverage(IpFamily::V6, now_epoch_secs) { - let map_v6 = self.proxy_map_v6.read().await.clone(); - for (dc, addrs) in map_v6 { - let entry = out.entry(dc).or_default(); - for (ip, port) in addrs { - entry.insert(SocketAddr::new(ip, port)); - } - } - } - - out - } - - pub(super) async fn has_non_draining_writer_per_desired_dc_group(&self) -> bool { - let desired_by_dc = self.desired_dc_endpoints().await; - let required_dcs: HashSet = desired_by_dc - .iter() - .filter_map(|(dc, endpoints)| { - if endpoints.is_empty() { - None - } else { - Some(*dc) - } - }) - .collect(); - if required_dcs.is_empty() { - return true; - } - - let ws = self.writers.read().await; - let mut covered_dcs = HashSet::::with_capacity(required_dcs.len()); - for writer in ws.iter() { - if writer.draining.load(Ordering::Relaxed) { - continue; - } - if required_dcs.contains(&writer.writer_dc) { - covered_dcs.insert(writer.writer_dc); - if covered_dcs.len() == required_dcs.len() { - return true; - } - } - } - false - } - - fn hardswap_warmup_connect_delay_ms(&self) -> u64 { - let min_ms = self - .reinit - .me_hardswap_warmup_delay_min_ms - .load(Ordering::Relaxed); - let max_ms = self - .reinit - .me_hardswap_warmup_delay_max_ms - .load(Ordering::Relaxed); - let (min_ms, max_ms) = if min_ms <= max_ms { - (min_ms, max_ms) - } else { - (max_ms, min_ms) - }; - if min_ms == max_ms { - return min_ms; - } - rand::rng().random_range(min_ms..=max_ms) - } - - fn hardswap_warmup_backoff_ms(&self, pass_idx: usize) -> u64 { - let base_ms = self - .reinit - .me_hardswap_warmup_pass_backoff_base_ms - .load(Ordering::Relaxed); - let cap_ms = - (self.reconnect_runtime.me_reconnect_backoff_cap.as_millis() as u64).max(base_ms); - let shift = (pass_idx as u32).min(20); - let scaled = base_ms.saturating_mul(1u64 << shift); - let core = scaled.min(cap_ms); - let jitter = (core / 2).max(1); - core.saturating_add(rand::rng().random_range(0..=jitter)) - } - - async fn fresh_writer_count_for_dc_endpoints( - &self, - generation: u64, - dc: i32, - endpoints: &HashSet, - ) -> usize { - let ws = self.writers.read().await; - ws.iter() - .filter(|w| !w.draining.load(Ordering::Relaxed)) - .filter(|w| w.generation == generation) - .filter(|w| w.writer_dc == dc) - .filter(|w| endpoints.contains(&w.addr)) - .count() - } - - pub(super) async fn active_writer_count_for_dc_endpoints( - &self, - dc: i32, - endpoints: &HashSet, - ) -> usize { - let ws = self.writers.read().await; - ws.iter() - .filter(|w| !w.draining.load(Ordering::Relaxed)) - .filter(|w| w.writer_dc == dc) - .filter(|w| endpoints.contains(&w.addr)) - .count() - } - - async fn warmup_generation_for_all_dcs( - self: &Arc, - rng: &SecureRandom, - generation: u64, - desired_by_dc: &HashMap>, - ) { - let extra_passes = self - .reinit - .me_hardswap_warmup_extra_passes - .load(Ordering::Relaxed) - .min(10) as usize; - let total_passes = 1 + extra_passes; - - for (dc, endpoints) in desired_by_dc { - if endpoints.is_empty() { - continue; - } - - let mut endpoint_list: Vec = endpoints.iter().copied().collect(); - endpoint_list.sort_unstable(); - let required = self.required_writers_for_dc(endpoint_list.len()); - let mut completed = false; - let mut last_fresh_count = self - .fresh_writer_count_for_dc_endpoints(generation, *dc, endpoints) - .await; - - for pass_idx in 0..total_passes { - if last_fresh_count >= required { - completed = true; - break; - } - - let missing = required.saturating_sub(last_fresh_count); - debug!( - dc = *dc, - pass = pass_idx + 1, - total_passes, - fresh_count = last_fresh_count, - required, - missing, - endpoint_count = endpoint_list.len(), - "ME hardswap warmup pass started" - ); - - for attempt_idx in 0..missing { - let delay_ms = self.hardswap_warmup_connect_delay_ms(); - tokio::time::sleep(Duration::from_millis(delay_ms)).await; - - let connected = self - .connect_endpoints_round_robin_with_generation_contour( - *dc, - &endpoint_list, - rng, - generation, - WriterContour::Warm, - false, - ) - .await; - debug!( - dc = *dc, - pass = pass_idx + 1, - total_passes, - attempt = attempt_idx + 1, - delay_ms, - connected, - "ME hardswap warmup connect attempt finished" - ); - } - - last_fresh_count = self - .fresh_writer_count_for_dc_endpoints(generation, *dc, endpoints) - .await; - if last_fresh_count >= required { - completed = true; - info!( - dc = *dc, - pass = pass_idx + 1, - total_passes, - fresh_count = last_fresh_count, - required, - "ME hardswap warmup floor reached for DC" - ); - break; - } - - if pass_idx + 1 < total_passes { - let backoff_ms = self.hardswap_warmup_backoff_ms(pass_idx); - debug!( - dc = *dc, - pass = pass_idx + 1, - total_passes, - fresh_count = last_fresh_count, - required, - backoff_ms, - "ME hardswap warmup pass incomplete, delaying next pass" - ); - tokio::time::sleep(Duration::from_millis(backoff_ms)).await; - } - } - - if !completed { - warn!( - dc = *dc, - fresh_count = last_fresh_count, - required, - endpoint_count = endpoint_list.len(), - total_passes, - "ME warmup stopped: unable to reach required writer floor for DC" - ); - } - } - } - - pub async fn zero_downtime_reinit_after_map_change( - self: &Arc, - rng: &SecureRandom, - ) -> bool { - let desired_by_dc = self.desired_dc_endpoints().await; - let now_epoch_secs = Self::now_epoch_secs(); - let v4_suppressed = self.is_family_temporarily_suppressed(IpFamily::V4, now_epoch_secs); - let v6_suppressed = self.is_family_temporarily_suppressed(IpFamily::V6, now_epoch_secs); - if desired_by_dc.is_empty() { - warn!("ME endpoint map is empty; skipping stale writer drain"); - let reason = if (self.decision.ipv4_me && v4_suppressed) - || (self.decision.ipv6_me && v6_suppressed) - { - MeDrainGateReason::SuppressionActive - } else { - MeDrainGateReason::CoverageQuorum - }; - self.set_last_drain_gate(false, false, reason, now_epoch_secs); - return false; - } - - let desired_map_hash = Self::desired_map_hash(&desired_by_dc); - let hardswap = self.reinit.hardswap.load(Ordering::Relaxed); - let reservation = - self.reserve_reinit_attempt(hardswap, desired_map_hash, now_epoch_secs); - let attempt = reservation.attempt; - let previous_generation = attempt.previous_generation; - let generation = attempt.generation; - if reservation.pending_reused { - self.stats.increment_me_hardswap_pending_reuse_total(); - debug!( - previous_generation, - generation, - pending_age_secs = reservation.pending_age_secs, - "ME hardswap continues with pending generation" - ); - } else if reservation.pending_expired { - self.stats.increment_me_hardswap_pending_ttl_expired_total(); - warn!( - previous_generation, - generation, - pending_age_secs = reservation.pending_age_secs, - pending_ttl_secs = ME_HARDSWAP_PENDING_TTL_SECS, - "ME hardswap pending generation expired by TTL; starting fresh generation" - ); - } - - if hardswap { - self.warmup_generation_for_all_dcs(rng, generation, &desired_by_dc) - .await; - } else { - self.reconcile_connections(rng).await; - } - - let writers = self.writers.read().await; - let active_writer_addrs: HashSet<(i32, SocketAddr)> = writers - .iter() - .filter(|w| !w.draining.load(Ordering::Relaxed)) - .map(|w| (w.writer_dc, w.addr)) - .collect(); - let min_ratio = Self::permille_to_ratio( - self.drain_runtime - .me_pool_min_fresh_ratio_permille - .load(Ordering::Relaxed), - ); - let (coverage_ratio, missing_dc) = - Self::coverage_ratio(&desired_by_dc, &active_writer_addrs); - let mut route_quorum_ok = coverage_ratio >= min_ratio; - let mut redundancy_ok = missing_dc.is_empty(); - let mut redundancy_missing_dc = missing_dc.clone(); - let mut gate_coverage_ratio = coverage_ratio; - if !hardswap && coverage_ratio < min_ratio { - self.set_last_drain_gate( - false, - redundancy_ok, - MeDrainGateReason::CoverageQuorum, - now_epoch_secs, - ); - warn!( - previous_generation, - generation, - coverage_ratio = format_args!("{coverage_ratio:.3}"), - min_ratio = format_args!("{min_ratio:.3}"), - missing_dc = ?missing_dc, - "ME reinit coverage below threshold; keeping stale writers" - ); - return false; - } - - if hardswap { - let fresh_writer_addrs: HashSet<(i32, SocketAddr)> = writers - .iter() - .filter(|w| !w.draining.load(Ordering::Relaxed)) - .filter(|w| w.generation == generation) - .map(|w| (w.writer_dc, w.addr)) - .collect(); - let (fresh_coverage_ratio, fresh_missing_dc) = - Self::coverage_ratio(&desired_by_dc, &fresh_writer_addrs); - route_quorum_ok = fresh_coverage_ratio >= min_ratio; - redundancy_ok = fresh_missing_dc.is_empty(); - redundancy_missing_dc = fresh_missing_dc.clone(); - gate_coverage_ratio = fresh_coverage_ratio; - if fresh_coverage_ratio < min_ratio { - self.set_last_drain_gate( - false, - redundancy_ok, - MeDrainGateReason::CoverageQuorum, - now_epoch_secs, - ); - warn!( - previous_generation, - generation, - fresh_coverage_ratio = format_args!("{fresh_coverage_ratio:.3}"), - missing_dc = ?fresh_missing_dc, - "ME hardswap pending: fresh generation DC coverage incomplete" - ); - return false; - } - } - - self.set_last_drain_gate( - route_quorum_ok, - redundancy_ok, - MeDrainGateReason::Open, - now_epoch_secs, - ); - if !redundancy_ok { - warn!( - missing_dc = ?redundancy_missing_dc, - coverage_ratio = format_args!("{gate_coverage_ratio:.3}"), - min_ratio = format_args!("{min_ratio:.3}"), - "ME reinit proceeds with weighted quorum while some DC groups remain uncovered" - ); - } - - if !self.commit_reinit_attempt(&attempt) { - debug!( - previous_generation, - generation, - "ME reinit result discarded after a newer desired-map attempt" - ); - return false; - } - - let desired_addrs: HashSet<(i32, SocketAddr)> = desired_by_dc - .iter() - .flat_map(|(dc, set)| set.iter().copied().map(|addr| (*dc, addr))) - .collect(); - - let stale_writer_ids: Vec = writers - .iter() - .filter(|w| !w.draining.load(Ordering::Relaxed)) - .filter(|w| { - if hardswap { - w.generation < generation - } else { - !desired_addrs.contains(&(w.writer_dc, w.addr)) - } - }) - .map(|w| w.id) - .collect(); - drop(writers); - - if stale_writer_ids.is_empty() { - debug!("ME reinit cycle completed with no stale writers"); - return true; - } - - let drain_timeout = self.force_close_timeout(); - let drain_timeout_secs = drain_timeout.map(|d| d.as_secs()).unwrap_or(0); - info!( - stale_writers = stale_writer_ids.len(), - previous_generation, - generation, - hardswap, - coverage_ratio = format_args!("{coverage_ratio:.3}"), - min_ratio = format_args!("{min_ratio:.3}"), - drain_timeout_secs, - "ME reinit cycle covered; processing stale writers" - ); - self.stats.increment_pool_swap_total(); - let can_drop_with_replacement = self.has_non_draining_writer_per_desired_dc_group().await; - if can_drop_with_replacement { - info!( - stale_writers = stale_writer_ids.len(), - "ME reinit stale writers: replacement coverage ready, force-closing clients for fast rebind" - ); - } else { - warn!( - stale_writers = stale_writer_ids.len(), - "ME reinit stale writers: replacement coverage incomplete, keeping draining fallback" - ); - } - for writer_id in stale_writer_ids { - self.mark_writer_draining_with_timeout(writer_id, drain_timeout, !hardswap) - .await; - if can_drop_with_replacement { - self.stats.increment_pool_force_close_total(); - self.remove_writer_and_close_clients(writer_id).await; - } - } - true - } - - pub async fn zero_downtime_reinit_periodic(self: &Arc, rng: &SecureRandom) -> bool { - self.zero_downtime_reinit_after_map_change(rng).await - } -} - -#[cfg(test)] -mod tests { - use std::collections::{HashMap, HashSet}; - use std::net::{IpAddr, Ipv4Addr, SocketAddr}; - - use super::{MePool, commit_reinit_state}; - use crate::transport::middle_proxy::pool::{ - ReinitAttemptState, ReinitCoordinatorState, ReinitPendingState, - }; - - fn addr(octet: u8, port: u16) -> SocketAddr { - SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, octet)), port) - } - - #[test] - fn coverage_ratio_counts_dc_coverage_not_floor() { - let dc1 = addr(1, 2001); - let dc2 = addr(2, 2002); - - let mut desired_by_dc = HashMap::>::new(); - desired_by_dc.insert(1, HashSet::from([dc1])); - desired_by_dc.insert(2, HashSet::from([dc2])); - - let active_writer_addrs = HashSet::from([(1, dc1)]); - let (ratio, missing_dc) = MePool::coverage_ratio(&desired_by_dc, &active_writer_addrs); - - assert_eq!(ratio, 0.5); - assert_eq!(missing_dc, vec![2]); - } - - #[test] - fn coverage_ratio_ignores_empty_dc_groups() { - let dc1 = addr(1, 2001); - - let mut desired_by_dc = HashMap::>::new(); - desired_by_dc.insert(1, HashSet::from([dc1])); - desired_by_dc.insert(2, HashSet::new()); - - let active_writer_addrs = HashSet::from([(1, dc1)]); - let (ratio, missing_dc) = MePool::coverage_ratio(&desired_by_dc, &active_writer_addrs); - - assert_eq!(ratio, 1.0); - assert!(missing_dc.is_empty()); - } - - #[test] - fn coverage_ratio_reports_missing_dcs_sorted() { - let dc1 = addr(1, 2001); - let dc2 = addr(2, 2002); - - let mut desired_by_dc = HashMap::>::new(); - desired_by_dc.insert(2, HashSet::from([dc2])); - desired_by_dc.insert(1, HashSet::from([dc1])); - - let (ratio, missing_dc) = MePool::coverage_ratio(&desired_by_dc, &HashSet::new()); - - assert_eq!(ratio, 0.0); - assert_eq!(missing_dc, vec![1, 2]); - } - - #[test] - fn stale_concurrent_attempt_cannot_regress_active_generation() { - let mut state = ReinitCoordinatorState { - next_attempt_id: 3, - active_generation: 1, - desired_map_hash: 22, - pending: Some(ReinitPendingState { - generation: 3, - started_at_epoch_secs: 1, - map_hash: 22, - }), - attempts: HashMap::from([ - ( - 1, - ReinitAttemptState { - generation: 2, - map_hash: 11, - hardswap: true, - committed: false, - }, - ), - ( - 2, - ReinitAttemptState { - generation: 3, - map_hash: 22, - hardswap: true, - committed: false, - }, - ), - ]), - }; - - assert!(commit_reinit_state(&mut state, 2, 3, 22, true)); - assert_eq!(state.active_generation, 3); - assert!(!commit_reinit_state(&mut state, 1, 2, 11, true)); - assert_eq!(state.active_generation, 3); - assert!(state.pending.is_none()); - } -} diff --git a/src/transport/middle_proxy/pool_reinit/coordination.rs b/src/transport/middle_proxy/pool_reinit/coordination.rs new file mode 100644 index 0000000..64e79d3 --- /dev/null +++ b/src/transport/middle_proxy/pool_reinit/coordination.rs @@ -0,0 +1,300 @@ +use super::*; + +impl MePool { + pub(super) fn desired_map_hash(desired_by_dc: &HashMap>) -> u64 { + let mut hasher = DefaultHasher::new(); + let mut dcs: Vec = desired_by_dc.keys().copied().collect(); + dcs.sort_unstable(); + for dc in dcs { + dc.hash(&mut hasher); + let mut endpoints: Vec = desired_by_dc + .get(&dc) + .map(|set| set.iter().copied().collect()) + .unwrap_or_default(); + endpoints.sort_unstable(); + for endpoint in endpoints { + endpoint.hash(&mut hasher); + } + } + hasher.finish() + } + + pub(super) fn reserve_reinit_attempt( + self: &Arc, + hardswap: bool, + map_hash: u64, + now_epoch_secs: u64, + ) -> ReinitReservation { + let mut state = self.reinit.coordinator.lock(); + state.desired_map_hash = map_hash; + let previous_generation = state.active_generation; + let mut pending_reused = false; + let mut pending_expired = false; + let mut pending_age_secs = 0; + + let generation = if hardswap { + let reusable = state.pending.filter(|pending| { + pending_age_secs = now_epoch_secs.saturating_sub(pending.started_at_epoch_secs); + pending_expired = pending.started_at_epoch_secs > 0 + && pending_age_secs > ME_HARDSWAP_PENDING_TTL_SECS; + pending.generation >= previous_generation + && pending.map_hash == map_hash + && !pending_expired + }); + if let Some(pending) = reusable { + pending_reused = true; + pending.generation + } else { + let generation = self.reinit.generation.fetch_add(1, Ordering::AcqRel) + 1; + state.pending = Some(ReinitPendingState { + generation, + started_at_epoch_secs: now_epoch_secs, + map_hash, + }); + generation + } + } else { + state.pending = None; + self.reinit.generation.fetch_add(1, Ordering::AcqRel) + 1 + }; + + let attempt_id = state.next_attempt_id; + state.next_attempt_id = state.next_attempt_id.saturating_add(1); + state.attempts.insert( + attempt_id, + ReinitAttemptState { + generation, + map_hash, + hardswap, + committed: false, + }, + ); + publish_reinit_state(self.reinit.as_ref(), &state); + ReinitReservation { + attempt: ReinitAttemptGuard { + reinit: Arc::clone(&self.reinit), + attempt_id, + generation, + previous_generation, + map_hash, + hardswap, + }, + pending_reused, + pending_expired, + pending_age_secs, + } + } + + pub(super) fn commit_reinit_attempt(&self, attempt: &ReinitAttemptGuard) -> bool { + let mut state = self.reinit.coordinator.lock(); + if !commit_reinit_state( + &mut state, + attempt.attempt_id, + attempt.generation, + attempt.map_hash, + attempt.hardswap, + ) { + return false; + } + if attempt.hardswap { + let writers = self.writers.snapshot(); + for writer in writers.iter() { + if !writer.draining.load(Ordering::Relaxed) + && writer.generation == attempt.generation + { + writer + .contour + .store(WriterContour::Active.as_u8(), Ordering::Release); + } + } + } + publish_reinit_state(self.reinit.as_ref(), &state); + true + } + + pub(super) fn coverage_ratio( + desired_by_dc: &HashMap>, + active_writer_addrs: &HashSet<(i32, SocketAddr)>, + ) -> (f32, Vec) { + if desired_by_dc.is_empty() { + return (1.0, Vec::new()); + } + + let mut missing_dc = Vec::::new(); + let mut covered = 0usize; + let mut total = 0usize; + for (dc, endpoints) in desired_by_dc { + if endpoints.is_empty() { + continue; + } + total += 1; + if endpoints + .iter() + .any(|addr| active_writer_addrs.contains(&(*dc, *addr))) + { + covered += 1; + } else { + missing_dc.push(*dc); + } + } + + missing_dc.sort_unstable(); + if total == 0 { + return (1.0, missing_dc); + } + let ratio = (covered as f32) / (total as f32); + (ratio, missing_dc) + } + + pub async fn reconcile_connections(self: &Arc, rng: &SecureRandom) { + for family in self.family_order() { + let map = self.proxy_map_for_family(family).await; + for (dc, addrs) in &map { + let dc_addrs: Vec = addrs + .iter() + .map(|(ip, port)| SocketAddr::new(*ip, *port)) + .collect(); + let dc_endpoints: HashSet = dc_addrs.iter().copied().collect(); + if self + .active_writer_count_for_dc_endpoints(*dc, &dc_endpoints) + .await + == 0 + { + let mut shuffled = dc_addrs.clone(); + shuffled.shuffle(&mut rand::rng()); + for addr in shuffled { + if self.connect_one_for_dc(addr, *dc, rng).await.is_ok() { + break; + } + } + } + } + if !self.decision.effective_multipath && self.connection_count() > 0 { + break; + } + } + } + + pub(super) async fn desired_dc_endpoints(&self) -> HashMap> { + let now_epoch_secs = Self::now_epoch_secs(); + let mut out: HashMap> = HashMap::new(); + + if self.family_enabled_for_drain_coverage(IpFamily::V4, now_epoch_secs) { + let map_v4 = self.proxy_map_v4.read().await.clone(); + for (dc, addrs) in map_v4 { + let entry = out.entry(dc).or_default(); + for (ip, port) in addrs { + entry.insert(SocketAddr::new(ip, port)); + } + } + } + + if self.family_enabled_for_drain_coverage(IpFamily::V6, now_epoch_secs) { + let map_v6 = self.proxy_map_v6.read().await.clone(); + for (dc, addrs) in map_v6 { + let entry = out.entry(dc).or_default(); + for (ip, port) in addrs { + entry.insert(SocketAddr::new(ip, port)); + } + } + } + + out + } + + pub(in crate::transport::middle_proxy) async fn has_non_draining_writer_per_desired_dc_group( + &self, + ) -> bool { + let desired_by_dc = self.desired_dc_endpoints().await; + let required_dcs: HashSet = desired_by_dc + .iter() + .filter_map(|(dc, endpoints)| { + if endpoints.is_empty() { + None + } else { + Some(*dc) + } + }) + .collect(); + if required_dcs.is_empty() { + return true; + } + + let ws = self.writers.read().await; + let mut covered_dcs = HashSet::::with_capacity(required_dcs.len()); + for writer in ws.iter() { + if writer.draining.load(Ordering::Relaxed) { + continue; + } + if required_dcs.contains(&writer.writer_dc) { + covered_dcs.insert(writer.writer_dc); + if covered_dcs.len() == required_dcs.len() { + return true; + } + } + } + false + } + + pub(super) fn hardswap_warmup_connect_delay_ms(&self) -> u64 { + let min_ms = self + .reinit + .me_hardswap_warmup_delay_min_ms + .load(Ordering::Relaxed); + let max_ms = self + .reinit + .me_hardswap_warmup_delay_max_ms + .load(Ordering::Relaxed); + let (min_ms, max_ms) = if min_ms <= max_ms { + (min_ms, max_ms) + } else { + (max_ms, min_ms) + }; + if min_ms == max_ms { + return min_ms; + } + rand::rng().random_range(min_ms..=max_ms) + } + + pub(super) fn hardswap_warmup_backoff_ms(&self, pass_idx: usize) -> u64 { + let base_ms = self + .reinit + .me_hardswap_warmup_pass_backoff_base_ms + .load(Ordering::Relaxed); + let cap_ms = + (self.reconnect_runtime.me_reconnect_backoff_cap.as_millis() as u64).max(base_ms); + let shift = (pass_idx as u32).min(20); + let scaled = base_ms.saturating_mul(1u64 << shift); + let core = scaled.min(cap_ms); + let jitter = (core / 2).max(1); + core.saturating_add(rand::rng().random_range(0..=jitter)) + } + + pub(super) async fn fresh_writer_count_for_dc_endpoints( + &self, + generation: u64, + dc: i32, + endpoints: &HashSet, + ) -> usize { + let ws = self.writers.read().await; + ws.iter() + .filter(|w| !w.draining.load(Ordering::Relaxed)) + .filter(|w| w.generation == generation) + .filter(|w| w.writer_dc == dc) + .filter(|w| endpoints.contains(&w.addr)) + .count() + } + + pub(in crate::transport::middle_proxy) async fn active_writer_count_for_dc_endpoints( + &self, + dc: i32, + endpoints: &HashSet, + ) -> usize { + let ws = self.writers.read().await; + ws.iter() + .filter(|w| !w.draining.load(Ordering::Relaxed)) + .filter(|w| w.writer_dc == dc) + .filter(|w| endpoints.contains(&w.addr)) + .count() + } +} diff --git a/src/transport/middle_proxy/pool_reinit/reconcile.rs b/src/transport/middle_proxy/pool_reinit/reconcile.rs new file mode 100644 index 0000000..5c900fe --- /dev/null +++ b/src/transport/middle_proxy/pool_reinit/reconcile.rs @@ -0,0 +1,322 @@ +use super::*; + +impl MePool { + async fn warmup_generation_for_all_dcs( + self: &Arc, + rng: &SecureRandom, + generation: u64, + desired_by_dc: &HashMap>, + ) { + let extra_passes = self + .reinit + .me_hardswap_warmup_extra_passes + .load(Ordering::Relaxed) + .min(10) as usize; + let total_passes = 1 + extra_passes; + + for (dc, endpoints) in desired_by_dc { + if endpoints.is_empty() { + continue; + } + + let mut endpoint_list: Vec = endpoints.iter().copied().collect(); + endpoint_list.sort_unstable(); + let required = self.required_writers_for_dc(endpoint_list.len()); + let mut completed = false; + let mut last_fresh_count = self + .fresh_writer_count_for_dc_endpoints(generation, *dc, endpoints) + .await; + + for pass_idx in 0..total_passes { + if last_fresh_count >= required { + completed = true; + break; + } + + let missing = required.saturating_sub(last_fresh_count); + debug!( + dc = *dc, + pass = pass_idx + 1, + total_passes, + fresh_count = last_fresh_count, + required, + missing, + endpoint_count = endpoint_list.len(), + "ME hardswap warmup pass started" + ); + + for attempt_idx in 0..missing { + let delay_ms = self.hardswap_warmup_connect_delay_ms(); + tokio::time::sleep(Duration::from_millis(delay_ms)).await; + + let connected = self + .connect_endpoints_round_robin_with_generation_contour( + *dc, + &endpoint_list, + rng, + generation, + WriterContour::Warm, + false, + ) + .await; + debug!( + dc = *dc, + pass = pass_idx + 1, + total_passes, + attempt = attempt_idx + 1, + delay_ms, + connected, + "ME hardswap warmup connect attempt finished" + ); + } + + last_fresh_count = self + .fresh_writer_count_for_dc_endpoints(generation, *dc, endpoints) + .await; + if last_fresh_count >= required { + completed = true; + info!( + dc = *dc, + pass = pass_idx + 1, + total_passes, + fresh_count = last_fresh_count, + required, + "ME hardswap warmup floor reached for DC" + ); + break; + } + + if pass_idx + 1 < total_passes { + let backoff_ms = self.hardswap_warmup_backoff_ms(pass_idx); + debug!( + dc = *dc, + pass = pass_idx + 1, + total_passes, + fresh_count = last_fresh_count, + required, + backoff_ms, + "ME hardswap warmup pass incomplete, delaying next pass" + ); + tokio::time::sleep(Duration::from_millis(backoff_ms)).await; + } + } + + if !completed { + warn!( + dc = *dc, + fresh_count = last_fresh_count, + required, + endpoint_count = endpoint_list.len(), + total_passes, + "ME warmup stopped: unable to reach required writer floor for DC" + ); + } + } + } + + pub async fn zero_downtime_reinit_after_map_change( + self: &Arc, + rng: &SecureRandom, + ) -> bool { + let desired_by_dc = self.desired_dc_endpoints().await; + let now_epoch_secs = Self::now_epoch_secs(); + let v4_suppressed = self.is_family_temporarily_suppressed(IpFamily::V4, now_epoch_secs); + let v6_suppressed = self.is_family_temporarily_suppressed(IpFamily::V6, now_epoch_secs); + if desired_by_dc.is_empty() { + warn!("ME endpoint map is empty; skipping stale writer drain"); + let reason = if (self.decision.ipv4_me && v4_suppressed) + || (self.decision.ipv6_me && v6_suppressed) + { + MeDrainGateReason::SuppressionActive + } else { + MeDrainGateReason::CoverageQuorum + }; + self.set_last_drain_gate(false, false, reason, now_epoch_secs); + return false; + } + + let desired_map_hash = Self::desired_map_hash(&desired_by_dc); + let hardswap = self.reinit.hardswap.load(Ordering::Relaxed); + let reservation = self.reserve_reinit_attempt(hardswap, desired_map_hash, now_epoch_secs); + let attempt = reservation.attempt; + let previous_generation = attempt.previous_generation; + let generation = attempt.generation; + if reservation.pending_reused { + self.stats.increment_me_hardswap_pending_reuse_total(); + debug!( + previous_generation, + generation, + pending_age_secs = reservation.pending_age_secs, + "ME hardswap continues with pending generation" + ); + } else if reservation.pending_expired { + self.stats.increment_me_hardswap_pending_ttl_expired_total(); + warn!( + previous_generation, + generation, + pending_age_secs = reservation.pending_age_secs, + pending_ttl_secs = ME_HARDSWAP_PENDING_TTL_SECS, + "ME hardswap pending generation expired by TTL; starting fresh generation" + ); + } + + if hardswap { + self.warmup_generation_for_all_dcs(rng, generation, &desired_by_dc) + .await; + } else { + self.reconcile_connections(rng).await; + } + + let writers = self.writers.read().await; + let active_writer_addrs: HashSet<(i32, SocketAddr)> = writers + .iter() + .filter(|w| !w.draining.load(Ordering::Relaxed)) + .map(|w| (w.writer_dc, w.addr)) + .collect(); + let min_ratio = Self::permille_to_ratio( + self.drain_runtime + .me_pool_min_fresh_ratio_permille + .load(Ordering::Relaxed), + ); + let (coverage_ratio, missing_dc) = + Self::coverage_ratio(&desired_by_dc, &active_writer_addrs); + let mut route_quorum_ok = coverage_ratio >= min_ratio; + let mut redundancy_ok = missing_dc.is_empty(); + let mut redundancy_missing_dc = missing_dc.clone(); + let mut gate_coverage_ratio = coverage_ratio; + if !hardswap && coverage_ratio < min_ratio { + self.set_last_drain_gate( + false, + redundancy_ok, + MeDrainGateReason::CoverageQuorum, + now_epoch_secs, + ); + warn!( + previous_generation, + generation, + coverage_ratio = format_args!("{coverage_ratio:.3}"), + min_ratio = format_args!("{min_ratio:.3}"), + missing_dc = ?missing_dc, + "ME reinit coverage below threshold; keeping stale writers" + ); + return false; + } + + if hardswap { + let fresh_writer_addrs: HashSet<(i32, SocketAddr)> = writers + .iter() + .filter(|w| !w.draining.load(Ordering::Relaxed)) + .filter(|w| w.generation == generation) + .map(|w| (w.writer_dc, w.addr)) + .collect(); + let (fresh_coverage_ratio, fresh_missing_dc) = + Self::coverage_ratio(&desired_by_dc, &fresh_writer_addrs); + route_quorum_ok = fresh_coverage_ratio >= min_ratio; + redundancy_ok = fresh_missing_dc.is_empty(); + redundancy_missing_dc = fresh_missing_dc.clone(); + gate_coverage_ratio = fresh_coverage_ratio; + if fresh_coverage_ratio < min_ratio { + self.set_last_drain_gate( + false, + redundancy_ok, + MeDrainGateReason::CoverageQuorum, + now_epoch_secs, + ); + warn!( + previous_generation, + generation, + fresh_coverage_ratio = format_args!("{fresh_coverage_ratio:.3}"), + missing_dc = ?fresh_missing_dc, + "ME hardswap pending: fresh generation DC coverage incomplete" + ); + return false; + } + } + + self.set_last_drain_gate( + route_quorum_ok, + redundancy_ok, + MeDrainGateReason::Open, + now_epoch_secs, + ); + if !redundancy_ok { + warn!( + missing_dc = ?redundancy_missing_dc, + coverage_ratio = format_args!("{gate_coverage_ratio:.3}"), + min_ratio = format_args!("{min_ratio:.3}"), + "ME reinit proceeds with weighted quorum while some DC groups remain uncovered" + ); + } + + if !self.commit_reinit_attempt(&attempt) { + debug!( + previous_generation, + generation, "ME reinit result discarded after a newer desired-map attempt" + ); + return false; + } + + let desired_addrs: HashSet<(i32, SocketAddr)> = desired_by_dc + .iter() + .flat_map(|(dc, set)| set.iter().copied().map(|addr| (*dc, addr))) + .collect(); + + let stale_writer_ids: Vec = writers + .iter() + .filter(|w| !w.draining.load(Ordering::Relaxed)) + .filter(|w| { + if hardswap { + w.generation < generation + } else { + !desired_addrs.contains(&(w.writer_dc, w.addr)) + } + }) + .map(|w| w.id) + .collect(); + drop(writers); + + if stale_writer_ids.is_empty() { + debug!("ME reinit cycle completed with no stale writers"); + return true; + } + + let drain_timeout = self.force_close_timeout(); + let drain_timeout_secs = drain_timeout.map(|d| d.as_secs()).unwrap_or(0); + info!( + stale_writers = stale_writer_ids.len(), + previous_generation, + generation, + hardswap, + coverage_ratio = format_args!("{coverage_ratio:.3}"), + min_ratio = format_args!("{min_ratio:.3}"), + drain_timeout_secs, + "ME reinit cycle covered; processing stale writers" + ); + self.stats.increment_pool_swap_total(); + let can_drop_with_replacement = self.has_non_draining_writer_per_desired_dc_group().await; + if can_drop_with_replacement { + info!( + stale_writers = stale_writer_ids.len(), + "ME reinit stale writers: replacement coverage ready, force-closing clients for fast rebind" + ); + } else { + warn!( + stale_writers = stale_writer_ids.len(), + "ME reinit stale writers: replacement coverage incomplete, keeping draining fallback" + ); + } + for writer_id in stale_writer_ids { + self.mark_writer_draining_with_timeout(writer_id, drain_timeout, !hardswap) + .await; + if can_drop_with_replacement { + self.stats.increment_pool_force_close_total(); + self.remove_writer_and_close_clients(writer_id).await; + } + } + true + } + + pub async fn zero_downtime_reinit_periodic(self: &Arc, rng: &SecureRandom) -> bool { + self.zero_downtime_reinit_after_map_change(rng).await + } +} diff --git a/src/transport/middle_proxy/pool_reinit/tests.rs b/src/transport/middle_proxy/pool_reinit/tests.rs new file mode 100644 index 0000000..cb85871 --- /dev/null +++ b/src/transport/middle_proxy/pool_reinit/tests.rs @@ -0,0 +1,97 @@ +use std::collections::{HashMap, HashSet}; +use std::net::{IpAddr, Ipv4Addr, SocketAddr}; + +use super::{MePool, commit_reinit_state}; +use crate::transport::middle_proxy::pool::{ + ReinitAttemptState, ReinitCoordinatorState, ReinitPendingState, +}; + +fn addr(octet: u8, port: u16) -> SocketAddr { + SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, octet)), port) +} + +#[test] +fn coverage_ratio_counts_dc_coverage_not_floor() { + let dc1 = addr(1, 2001); + let dc2 = addr(2, 2002); + + let mut desired_by_dc = HashMap::>::new(); + desired_by_dc.insert(1, HashSet::from([dc1])); + desired_by_dc.insert(2, HashSet::from([dc2])); + + let active_writer_addrs = HashSet::from([(1, dc1)]); + let (ratio, missing_dc) = MePool::coverage_ratio(&desired_by_dc, &active_writer_addrs); + + assert_eq!(ratio, 0.5); + assert_eq!(missing_dc, vec![2]); +} + +#[test] +fn coverage_ratio_ignores_empty_dc_groups() { + let dc1 = addr(1, 2001); + + let mut desired_by_dc = HashMap::>::new(); + desired_by_dc.insert(1, HashSet::from([dc1])); + desired_by_dc.insert(2, HashSet::new()); + + let active_writer_addrs = HashSet::from([(1, dc1)]); + let (ratio, missing_dc) = MePool::coverage_ratio(&desired_by_dc, &active_writer_addrs); + + assert_eq!(ratio, 1.0); + assert!(missing_dc.is_empty()); +} + +#[test] +fn coverage_ratio_reports_missing_dcs_sorted() { + let dc1 = addr(1, 2001); + let dc2 = addr(2, 2002); + + let mut desired_by_dc = HashMap::>::new(); + desired_by_dc.insert(2, HashSet::from([dc2])); + desired_by_dc.insert(1, HashSet::from([dc1])); + + let (ratio, missing_dc) = MePool::coverage_ratio(&desired_by_dc, &HashSet::new()); + + assert_eq!(ratio, 0.0); + assert_eq!(missing_dc, vec![1, 2]); +} + +#[test] +fn stale_concurrent_attempt_cannot_regress_active_generation() { + let mut state = ReinitCoordinatorState { + next_attempt_id: 3, + active_generation: 1, + desired_map_hash: 22, + pending: Some(ReinitPendingState { + generation: 3, + started_at_epoch_secs: 1, + map_hash: 22, + }), + attempts: HashMap::from([ + ( + 1, + ReinitAttemptState { + generation: 2, + map_hash: 11, + hardswap: true, + committed: false, + }, + ), + ( + 2, + ReinitAttemptState { + generation: 3, + map_hash: 22, + hardswap: true, + committed: false, + }, + ), + ]), + }; + + assert!(commit_reinit_state(&mut state, 2, 3, 22, true)); + assert_eq!(state.active_generation, 3); + assert!(!commit_reinit_state(&mut state, 1, 2, 11, true)); + assert_eq!(state.active_generation, 3); + assert!(state.pending.is_none()); +} diff --git a/src/transport/middle_proxy/pool_runtime_api.rs b/src/transport/middle_proxy/pool_runtime_api.rs index 7b26583..78c50e7 100644 --- a/src/transport/middle_proxy/pool_runtime_api.rs +++ b/src/transport/middle_proxy/pool_runtime_api.rs @@ -68,10 +68,7 @@ impl MePool { .values() .filter(|pending| pending.is_some()) .count(); - let inflight_dc_keys = refill_states - .keys() - .copied() - .collect::>(); + let inflight_dc_keys = refill_states.keys().copied().collect::>(); drop(refill_states); let mut by_dc_map = HashMap::<(i16, &'static str), usize>::new(); diff --git a/src/transport/middle_proxy/pool_status.rs b/src/transport/middle_proxy/pool_status.rs index 8eea3a6..dcf7fc3 100644 --- a/src/transport/middle_proxy/pool_status.rs +++ b/src/transport/middle_proxy/pool_status.rs @@ -8,6 +8,10 @@ use super::pool::{MePool, ReinitStatusSnapshot, WriterContour}; use crate::config::{MeBindStaleMode, MeFloorMode, MeSocksKdfPolicy}; use crate::transport::upstream::IpPreference; +// ME writer and DC coverage snapshots. +mod status_snapshot; +// ME runtime policy and coherent snapshot assembly. +mod runtime_snapshot; #[derive(Clone, Debug)] pub(crate) struct MeApiWriterStatusSnapshot { pub writer_id: u64, @@ -146,562 +150,6 @@ pub(crate) struct MeApiRuntimeSnapshot { pub network_path: Vec, } -impl MePool { - pub(crate) async fn admission_ready_conditional_cast(&self) -> bool { - let mut endpoints_by_dc = BTreeMap::>::new(); - if self.decision.ipv4_me { - let map = self.proxy_map_v4.read().await.clone(); - extend_signed_endpoints(&mut endpoints_by_dc, map); - } - if self.decision.ipv6_me { - let map = self.proxy_map_v6.read().await.clone(); - extend_signed_endpoints(&mut endpoints_by_dc, map); - } - - if endpoints_by_dc.is_empty() { - return false; - } - - let writers = self.writers.read().await.clone(); - let mut live_writers_by_dc = HashMap::::new(); - for writer in writers.iter() { - if writer.draining.load(Ordering::Relaxed) { - continue; - } - if let Ok(dc) = i16::try_from(writer.writer_dc) { - *live_writers_by_dc.entry(dc).or_insert(0) += 1; - } - } - - for dc in endpoints_by_dc.keys() { - let alive = live_writers_by_dc.get(dc).copied().unwrap_or(0); - if alive == 0 { - return false; - } - } - - true - } - - #[allow(dead_code)] - pub(crate) async fn admission_ready_full_floor(&self) -> bool { - let mut endpoints_by_dc = BTreeMap::>::new(); - if self.decision.ipv4_me { - let map = self.proxy_map_v4.read().await.clone(); - extend_signed_endpoints(&mut endpoints_by_dc, map); - } - if self.decision.ipv6_me { - let map = self.proxy_map_v6.read().await.clone(); - extend_signed_endpoints(&mut endpoints_by_dc, map); - } - - if endpoints_by_dc.is_empty() { - return false; - } - - let writers = self.writers.read().await.clone(); - let mut live_writers_by_dc = HashMap::::new(); - for writer in writers.iter() { - if writer.draining.load(Ordering::Relaxed) { - continue; - } - if let Ok(dc) = i16::try_from(writer.writer_dc) { - *live_writers_by_dc.entry(dc).or_insert(0) += 1; - } - } - - for (dc, endpoints) in endpoints_by_dc { - let endpoint_count = endpoints.len(); - if endpoint_count == 0 { - return false; - } - let required = self.required_writers_for_dc_with_floor_mode(endpoint_count, false); - let alive = live_writers_by_dc.get(&dc).copied().unwrap_or(0); - if alive < required { - return false; - } - } - - true - } - - pub(crate) async fn api_status_snapshot(&self) -> MeApiStatusSnapshot { - let reinit = self.reinit.status.load_full(); - self.api_status_snapshot_for_reinit(reinit.as_ref()).await - } - - async fn api_status_snapshot_for_reinit( - &self, - reinit: &ReinitStatusSnapshot, - ) -> MeApiStatusSnapshot { - let now_epoch_secs = Self::now_epoch_secs(); - let active_generation = reinit.active_generation; - let drain_ttl_secs = self - .drain_runtime - .me_pool_drain_ttl_secs - .load(Ordering::Relaxed); - - let mut endpoints_by_dc = BTreeMap::>::new(); - if self.decision.ipv4_me { - let map = self.proxy_map_v4.read().await.clone(); - extend_signed_endpoints(&mut endpoints_by_dc, map); - } - if self.decision.ipv6_me { - let map = self.proxy_map_v6.read().await.clone(); - extend_signed_endpoints(&mut endpoints_by_dc, map); - } - - let configured_dc_groups = endpoints_by_dc.len(); - let configured_endpoints = endpoints_by_dc.values().map(BTreeSet::len).sum(); - - let required_writers = endpoints_by_dc - .values() - .map(|endpoints| self.required_writers_for_dc_with_floor_mode(endpoints.len(), false)) - .sum(); - - let idle_since = self.registry.writer_idle_since_snapshot().await; - let activity = self.registry.writer_activity_snapshot().await; - let rtt = self.rtt_stats.lock().await.clone(); - let writers = self.writers.read().await.clone(); - - let mut live_writers_by_dc_endpoint = HashMap::<(i16, SocketAddr), usize>::new(); - let mut live_writers_by_dc = HashMap::::new(); - let mut fresh_writers_by_dc = HashMap::::new(); - let mut dc_rtt_agg = HashMap::::new(); - let mut writer_rows = Vec::::with_capacity(writers.len()); - - for writer in writers.iter() { - let endpoint = writer.addr; - let dc = i16::try_from(writer.writer_dc).ok(); - let draining = writer.draining.load(Ordering::Relaxed); - let degraded = writer.degraded.load(Ordering::Relaxed); - let matches_active_generation = writer.generation == active_generation; - let in_desired_map = dc - .and_then(|dc_idx| endpoints_by_dc.get(&dc_idx)) - .is_some_and(|endpoints| endpoints.contains(&endpoint)); - let bound_clients = activity - .bound_clients_by_writer - .get(&writer.id) - .copied() - .unwrap_or(0); - let idle_for_secs = idle_since - .get(&writer.id) - .map(|idle_ts| now_epoch_secs.saturating_sub(*idle_ts)); - let rtt_ema_ms = rtt.get(&writer.id).map(|(_, ema)| *ema); - let allow_drain_fallback = writer.allow_drain_fallback.load(Ordering::Relaxed); - let drain_started_at_epoch_secs = writer - .draining_started_at_epoch_secs - .load(Ordering::Relaxed); - let drain_deadline_epoch_secs = - writer.drain_deadline_epoch_secs.load(Ordering::Relaxed); - let drain_started_at_epoch_secs = - (drain_started_at_epoch_secs != 0).then_some(drain_started_at_epoch_secs); - let drain_deadline_epoch_secs = - (drain_deadline_epoch_secs != 0).then_some(drain_deadline_epoch_secs); - let drain_over_ttl = draining - && drain_ttl_secs > 0 - && drain_started_at_epoch_secs - .is_some_and(|started| now_epoch_secs.saturating_sub(started) > drain_ttl_secs); - let state = match WriterContour::from_u8(writer.contour.load(Ordering::Relaxed)) { - WriterContour::Warm => "warm", - WriterContour::Active => "active", - WriterContour::Draining => "draining", - }; - - if !draining && let Some(dc_idx) = dc { - *live_writers_by_dc_endpoint - .entry((dc_idx, endpoint)) - .or_insert(0) += 1; - *live_writers_by_dc.entry(dc_idx).or_insert(0) += 1; - if let Some(ema_ms) = rtt_ema_ms { - let entry = dc_rtt_agg.entry(dc_idx).or_insert((0.0, 0)); - entry.0 += ema_ms; - entry.1 += 1; - } - if matches_active_generation && in_desired_map { - *fresh_writers_by_dc.entry(dc_idx).or_insert(0) += 1; - } - } - - writer_rows.push(MeApiWriterStatusSnapshot { - writer_id: writer.id, - dc, - endpoint, - generation: writer.generation, - state, - draining, - degraded, - bound_clients, - idle_for_secs, - rtt_ema_ms, - matches_active_generation, - in_desired_map, - allow_drain_fallback, - drain_started_at_epoch_secs, - drain_deadline_epoch_secs, - drain_over_ttl, - }); - } - - writer_rows.sort_by_key(|row| (row.dc.unwrap_or(i16::MAX), row.endpoint, row.writer_id)); - - let mut dcs = Vec::::with_capacity(endpoints_by_dc.len()); - let mut available_endpoints = 0usize; - let mut alive_writers = 0usize; - let mut fresh_alive_writers = 0usize; - let floor_mode = self.floor_mode(); - let adaptive_cpu_cores = (self - .floor_runtime - .me_adaptive_floor_cpu_cores_effective - .load(Ordering::Relaxed) as usize) - .max(1); - for (dc, endpoints) in endpoints_by_dc { - let endpoint_count = endpoints.len(); - let dc_available_endpoints = endpoints - .iter() - .filter(|endpoint| live_writers_by_dc_endpoint.contains_key(&(dc, **endpoint))) - .count(); - let base_required = self.required_writers_for_dc(endpoint_count); - let dc_required_writers = - self.required_writers_for_dc_with_floor_mode(endpoint_count, false); - let floor_min = if endpoint_count <= 1 { - (self - .floor_runtime - .me_adaptive_floor_min_writers_single_endpoint - .load(Ordering::Relaxed) as usize) - .max(1) - .min(base_required.max(1)) - } else { - (self - .floor_runtime - .me_adaptive_floor_min_writers_multi_endpoint - .load(Ordering::Relaxed) as usize) - .max(1) - .min(base_required.max(1)) - }; - let extra_per_core = if endpoint_count <= 1 { - self.floor_runtime - .me_adaptive_floor_max_extra_writers_single_per_core - .load(Ordering::Relaxed) as usize - } else { - self.floor_runtime - .me_adaptive_floor_max_extra_writers_multi_per_core - .load(Ordering::Relaxed) as usize - }; - let floor_max = - base_required.saturating_add(adaptive_cpu_cores.saturating_mul(extra_per_core)); - let floor_capped = - matches!(floor_mode, MeFloorMode::Adaptive) && dc_required_writers < base_required; - let dc_alive_writers = live_writers_by_dc.get(&dc).copied().unwrap_or(0); - let dc_fresh_alive_writers = fresh_writers_by_dc.get(&dc).copied().unwrap_or(0); - let dc_load = activity - .active_sessions_by_target_dc - .get(&dc) - .copied() - .unwrap_or(0); - let dc_rtt_ms = dc_rtt_agg - .get(&dc) - .and_then(|(sum, count)| (*count > 0).then_some(*sum / (*count as f64))); - - available_endpoints += dc_available_endpoints; - alive_writers += dc_alive_writers; - fresh_alive_writers += dc_fresh_alive_writers; - - dcs.push(MeApiDcStatusSnapshot { - dc, - endpoint_writers: endpoints - .iter() - .map(|endpoint| MeApiDcEndpointWriterSnapshot { - endpoint: *endpoint, - active_writers: live_writers_by_dc_endpoint - .get(&(dc, *endpoint)) - .copied() - .unwrap_or(0), - }) - .collect(), - endpoints: endpoints.into_iter().collect(), - available_endpoints: dc_available_endpoints, - available_pct: ratio_pct(dc_available_endpoints, endpoint_count), - required_writers: dc_required_writers, - floor_min, - floor_target: dc_required_writers, - floor_max, - floor_capped, - alive_writers: dc_alive_writers, - coverage_pct: ratio_pct(dc_alive_writers, dc_required_writers), - fresh_alive_writers: dc_fresh_alive_writers, - fresh_coverage_pct: ratio_pct(dc_fresh_alive_writers, dc_required_writers), - rtt_ms: dc_rtt_ms, - load: dc_load, - }); - } - - MeApiStatusSnapshot { - generated_at_epoch_secs: now_epoch_secs, - configured_dc_groups, - configured_endpoints, - available_endpoints, - available_pct: ratio_pct(available_endpoints, configured_endpoints), - required_writers, - alive_writers, - coverage_pct: ratio_pct(alive_writers, required_writers), - fresh_alive_writers, - fresh_coverage_pct: ratio_pct(fresh_alive_writers, required_writers), - writers: writer_rows, - dcs, - } - } - - #[allow(dead_code)] - pub(crate) async fn api_runtime_snapshot(&self) -> MeApiRuntimeSnapshot { - let reinit = self.reinit.status.load_full(); - self.api_runtime_snapshot_for_reinit(reinit.as_ref()).await - } - - async fn api_runtime_snapshot_for_reinit( - &self, - reinit: &ReinitStatusSnapshot, - ) -> MeApiRuntimeSnapshot { - let now = Instant::now(); - let now_epoch_secs = Self::now_epoch_secs(); - let pending_started_at = reinit.pending_hardswap_started_at_epoch_secs; - let pending_hardswap_age_secs = - (pending_started_at > 0).then_some(now_epoch_secs.saturating_sub(pending_started_at)); - - let mut quarantined_endpoints = Vec::::new(); - { - let guard = self.endpoint_quarantine.lock().await; - for (endpoint, expires_at) in guard.iter() { - if *expires_at <= now { - continue; - } - let remaining_ms = expires_at.duration_since(now).as_millis() as u64; - quarantined_endpoints.push(MeApiQuarantinedEndpointSnapshot { - endpoint: *endpoint, - remaining_ms, - }); - } - } - quarantined_endpoints.sort_by_key(|entry| entry.endpoint); - - let mut network_path = Vec::::new(); - if let Some(upstream) = &self.upstream { - for dc in 1..=5 { - let dc_idx = dc as i16; - let ip_preference = upstream - .get_dc_ip_preference(dc_idx) - .await - .map(ip_preference_label); - let selected_addr_v4 = upstream.get_dc_addr(dc_idx, false).await; - let selected_addr_v6 = upstream.get_dc_addr(dc_idx, true).await; - network_path.push(MeApiDcPathSnapshot { - dc: dc_idx, - ip_preference, - selected_addr_v4, - selected_addr_v6, - }); - } - } - - MeApiRuntimeSnapshot { - active_generation: reinit.active_generation, - warm_generation: reinit.warm_generations.last().copied().unwrap_or(0), - warm_generations: reinit.warm_generations.clone(), - pending_hardswap_generation: reinit.pending_hardswap_generation, - pending_hardswap_age_secs, - reinit_inflight: reinit.inflight, - reinit_max_concurrency_effective: self - .reinit - .max_concurrency_effective - .load(Ordering::Acquire), - hardswap_enabled: self.reinit.hardswap.load(Ordering::Relaxed), - floor_mode: floor_mode_label(self.floor_mode()), - adaptive_floor_idle_secs: self - .floor_runtime - .me_adaptive_floor_idle_secs - .load(Ordering::Relaxed), - adaptive_floor_min_writers_single_endpoint: self - .floor_runtime - .me_adaptive_floor_min_writers_single_endpoint - .load(Ordering::Relaxed), - adaptive_floor_min_writers_multi_endpoint: self - .floor_runtime - .me_adaptive_floor_min_writers_multi_endpoint - .load(Ordering::Relaxed), - adaptive_floor_recover_grace_secs: self - .floor_runtime - .me_adaptive_floor_recover_grace_secs - .load(Ordering::Relaxed), - adaptive_floor_writers_per_core_total: self - .floor_runtime - .me_adaptive_floor_writers_per_core_total - .load(Ordering::Relaxed) as u16, - adaptive_floor_cpu_cores_override: self - .floor_runtime - .me_adaptive_floor_cpu_cores_override - .load(Ordering::Relaxed) as u16, - adaptive_floor_max_extra_writers_single_per_core: self - .floor_runtime - .me_adaptive_floor_max_extra_writers_single_per_core - .load(Ordering::Relaxed) - as u16, - adaptive_floor_max_extra_writers_multi_per_core: self - .floor_runtime - .me_adaptive_floor_max_extra_writers_multi_per_core - .load(Ordering::Relaxed) - as u16, - adaptive_floor_max_active_writers_per_core: self - .floor_runtime - .me_adaptive_floor_max_active_writers_per_core - .load(Ordering::Relaxed) - as u16, - adaptive_floor_max_warm_writers_per_core: self - .floor_runtime - .me_adaptive_floor_max_warm_writers_per_core - .load(Ordering::Relaxed) - as u16, - adaptive_floor_max_active_writers_global: self - .floor_runtime - .me_adaptive_floor_max_active_writers_global - .load(Ordering::Relaxed), - adaptive_floor_max_warm_writers_global: self - .floor_runtime - .me_adaptive_floor_max_warm_writers_global - .load(Ordering::Relaxed), - adaptive_floor_cpu_cores_detected: self - .floor_runtime - .me_adaptive_floor_cpu_cores_detected - .load(Ordering::Relaxed), - adaptive_floor_cpu_cores_effective: self - .floor_runtime - .me_adaptive_floor_cpu_cores_effective - .load(Ordering::Relaxed), - adaptive_floor_global_cap_raw: self - .floor_runtime - .me_adaptive_floor_global_cap_raw - .load(Ordering::Relaxed), - adaptive_floor_global_cap_effective: self - .floor_runtime - .me_adaptive_floor_global_cap_effective - .load(Ordering::Relaxed), - adaptive_floor_target_writers_total: self - .floor_runtime - .me_adaptive_floor_target_writers_total - .load(Ordering::Relaxed), - adaptive_floor_active_cap_configured: self - .floor_runtime - .me_adaptive_floor_active_cap_configured - .load(Ordering::Relaxed), - adaptive_floor_active_cap_effective: self - .floor_runtime - .me_adaptive_floor_active_cap_effective - .load(Ordering::Relaxed), - adaptive_floor_warm_cap_configured: self - .floor_runtime - .me_adaptive_floor_warm_cap_configured - .load(Ordering::Relaxed), - adaptive_floor_warm_cap_effective: self - .floor_runtime - .me_adaptive_floor_warm_cap_effective - .load(Ordering::Relaxed), - adaptive_floor_active_writers_current: self - .floor_runtime - .me_adaptive_floor_active_writers_current - .load(Ordering::Relaxed), - adaptive_floor_warm_writers_current: self - .floor_runtime - .me_adaptive_floor_warm_writers_current - .load(Ordering::Relaxed), - me_keepalive_enabled: self.writer_lifecycle.me_keepalive_enabled, - me_keepalive_interval_secs: self.writer_lifecycle.me_keepalive_interval.as_secs(), - me_keepalive_jitter_secs: self.writer_lifecycle.me_keepalive_jitter.as_secs(), - me_keepalive_payload_random: self.writer_lifecycle.me_keepalive_payload_random, - rpc_proxy_req_every_secs: self - .writer_lifecycle - .rpc_proxy_req_every_secs - .load(Ordering::Relaxed), - me_reconnect_max_concurrent_per_dc: self - .reconnect_runtime - .me_reconnect_max_concurrent_per_dc, - me_reconnect_backoff_base_ms: self - .reconnect_runtime - .me_reconnect_backoff_base - .as_millis() as u64, - me_reconnect_backoff_cap_ms: self.reconnect_runtime.me_reconnect_backoff_cap.as_millis() - as u64, - me_reconnect_fast_retry_count: self.reconnect_runtime.me_reconnect_fast_retry_count, - me_pool_drain_ttl_secs: self - .drain_runtime - .me_pool_drain_ttl_secs - .load(Ordering::Relaxed), - me_pool_force_close_secs: self - .drain_runtime - .me_pool_force_close_secs - .load(Ordering::Relaxed), - me_pool_min_fresh_ratio: Self::permille_to_ratio( - self.drain_runtime - .me_pool_min_fresh_ratio_permille - .load(Ordering::Relaxed), - ), - me_bind_stale_mode: bind_stale_mode_label(self.bind_stale_mode()), - me_bind_stale_ttl_secs: self - .binding_policy - .me_bind_stale_ttl_secs - .load(Ordering::Relaxed), - me_single_endpoint_shadow_writers: self - .single_endpoint_runtime - .me_single_endpoint_shadow_writers - .load(Ordering::Relaxed), - me_single_endpoint_outage_mode_enabled: self - .single_endpoint_runtime - .me_single_endpoint_outage_mode_enabled - .load(Ordering::Relaxed), - me_single_endpoint_outage_disable_quarantine: self - .single_endpoint_runtime - .me_single_endpoint_outage_disable_quarantine - .load(Ordering::Relaxed), - me_single_endpoint_outage_backoff_min_ms: self - .single_endpoint_runtime - .me_single_endpoint_outage_backoff_min_ms - .load(Ordering::Relaxed), - me_single_endpoint_outage_backoff_max_ms: self - .single_endpoint_runtime - .me_single_endpoint_outage_backoff_max_ms - .load(Ordering::Relaxed), - me_single_endpoint_shadow_rotate_every_secs: self - .single_endpoint_runtime - .me_single_endpoint_shadow_rotate_every_secs - .load(Ordering::Relaxed), - me_deterministic_writer_sort: self - .writer_selection_policy - .me_deterministic_writer_sort - .load(Ordering::Relaxed), - me_writer_pick_mode: writer_pick_mode_label(self.writer_pick_mode()), - me_writer_pick_sample_size: self.writer_pick_sample_size() as u8, - me_socks_kdf_policy: socks_kdf_policy_label(self.socks_kdf_policy()), - quarantined_endpoints, - network_path, - } - } - - pub(crate) async fn api_coherent_snapshots( - &self, - ) -> (MeApiStatusSnapshot, MeApiRuntimeSnapshot) { - let mut attempts = 0usize; - loop { - let reinit = self.reinit.status.load_full(); - let status = self.api_status_snapshot_for_reinit(reinit.as_ref()).await; - let runtime = self - .api_runtime_snapshot_for_reinit(reinit.as_ref()) - .await; - attempts += 1; - if Arc::ptr_eq(&reinit, &self.reinit.status.load_full()) || attempts >= 3 { - return (status, runtime); - } - } - } -} - fn ratio_pct(part: usize, total: usize) -> f64 { if total == 0 { return 0.0; diff --git a/src/transport/middle_proxy/pool_status/runtime_snapshot.rs b/src/transport/middle_proxy/pool_status/runtime_snapshot.rs new file mode 100644 index 0000000..9291d85 --- /dev/null +++ b/src/transport/middle_proxy/pool_status/runtime_snapshot.rs @@ -0,0 +1,250 @@ +use super::*; + +impl MePool { + #[allow(dead_code)] + pub(crate) async fn api_runtime_snapshot(&self) -> MeApiRuntimeSnapshot { + let reinit = self.reinit.status.load_full(); + self.api_runtime_snapshot_for_reinit(reinit.as_ref()).await + } + + async fn api_runtime_snapshot_for_reinit( + &self, + reinit: &ReinitStatusSnapshot, + ) -> MeApiRuntimeSnapshot { + let now = Instant::now(); + let now_epoch_secs = Self::now_epoch_secs(); + let pending_started_at = reinit.pending_hardswap_started_at_epoch_secs; + let pending_hardswap_age_secs = + (pending_started_at > 0).then_some(now_epoch_secs.saturating_sub(pending_started_at)); + + let mut quarantined_endpoints = Vec::::new(); + { + let guard = self.endpoint_quarantine.lock().await; + for (endpoint, expires_at) in guard.iter() { + if *expires_at <= now { + continue; + } + let remaining_ms = expires_at.duration_since(now).as_millis() as u64; + quarantined_endpoints.push(MeApiQuarantinedEndpointSnapshot { + endpoint: *endpoint, + remaining_ms, + }); + } + } + quarantined_endpoints.sort_by_key(|entry| entry.endpoint); + + let mut network_path = Vec::::new(); + if let Some(upstream) = &self.upstream { + for dc in 1..=5 { + let dc_idx = dc as i16; + let ip_preference = upstream + .get_dc_ip_preference(dc_idx) + .await + .map(ip_preference_label); + let selected_addr_v4 = upstream.get_dc_addr(dc_idx, false).await; + let selected_addr_v6 = upstream.get_dc_addr(dc_idx, true).await; + network_path.push(MeApiDcPathSnapshot { + dc: dc_idx, + ip_preference, + selected_addr_v4, + selected_addr_v6, + }); + } + } + + MeApiRuntimeSnapshot { + active_generation: reinit.active_generation, + warm_generation: reinit.warm_generations.last().copied().unwrap_or(0), + warm_generations: reinit.warm_generations.clone(), + pending_hardswap_generation: reinit.pending_hardswap_generation, + pending_hardswap_age_secs, + reinit_inflight: reinit.inflight, + reinit_max_concurrency_effective: self + .reinit + .max_concurrency_effective + .load(Ordering::Acquire), + hardswap_enabled: self.reinit.hardswap.load(Ordering::Relaxed), + floor_mode: floor_mode_label(self.floor_mode()), + adaptive_floor_idle_secs: self + .floor_runtime + .me_adaptive_floor_idle_secs + .load(Ordering::Relaxed), + adaptive_floor_min_writers_single_endpoint: self + .floor_runtime + .me_adaptive_floor_min_writers_single_endpoint + .load(Ordering::Relaxed), + adaptive_floor_min_writers_multi_endpoint: self + .floor_runtime + .me_adaptive_floor_min_writers_multi_endpoint + .load(Ordering::Relaxed), + adaptive_floor_recover_grace_secs: self + .floor_runtime + .me_adaptive_floor_recover_grace_secs + .load(Ordering::Relaxed), + adaptive_floor_writers_per_core_total: self + .floor_runtime + .me_adaptive_floor_writers_per_core_total + .load(Ordering::Relaxed) as u16, + adaptive_floor_cpu_cores_override: self + .floor_runtime + .me_adaptive_floor_cpu_cores_override + .load(Ordering::Relaxed) as u16, + adaptive_floor_max_extra_writers_single_per_core: self + .floor_runtime + .me_adaptive_floor_max_extra_writers_single_per_core + .load(Ordering::Relaxed) + as u16, + adaptive_floor_max_extra_writers_multi_per_core: self + .floor_runtime + .me_adaptive_floor_max_extra_writers_multi_per_core + .load(Ordering::Relaxed) + as u16, + adaptive_floor_max_active_writers_per_core: self + .floor_runtime + .me_adaptive_floor_max_active_writers_per_core + .load(Ordering::Relaxed) + as u16, + adaptive_floor_max_warm_writers_per_core: self + .floor_runtime + .me_adaptive_floor_max_warm_writers_per_core + .load(Ordering::Relaxed) + as u16, + adaptive_floor_max_active_writers_global: self + .floor_runtime + .me_adaptive_floor_max_active_writers_global + .load(Ordering::Relaxed), + adaptive_floor_max_warm_writers_global: self + .floor_runtime + .me_adaptive_floor_max_warm_writers_global + .load(Ordering::Relaxed), + adaptive_floor_cpu_cores_detected: self + .floor_runtime + .me_adaptive_floor_cpu_cores_detected + .load(Ordering::Relaxed), + adaptive_floor_cpu_cores_effective: self + .floor_runtime + .me_adaptive_floor_cpu_cores_effective + .load(Ordering::Relaxed), + adaptive_floor_global_cap_raw: self + .floor_runtime + .me_adaptive_floor_global_cap_raw + .load(Ordering::Relaxed), + adaptive_floor_global_cap_effective: self + .floor_runtime + .me_adaptive_floor_global_cap_effective + .load(Ordering::Relaxed), + adaptive_floor_target_writers_total: self + .floor_runtime + .me_adaptive_floor_target_writers_total + .load(Ordering::Relaxed), + adaptive_floor_active_cap_configured: self + .floor_runtime + .me_adaptive_floor_active_cap_configured + .load(Ordering::Relaxed), + adaptive_floor_active_cap_effective: self + .floor_runtime + .me_adaptive_floor_active_cap_effective + .load(Ordering::Relaxed), + adaptive_floor_warm_cap_configured: self + .floor_runtime + .me_adaptive_floor_warm_cap_configured + .load(Ordering::Relaxed), + adaptive_floor_warm_cap_effective: self + .floor_runtime + .me_adaptive_floor_warm_cap_effective + .load(Ordering::Relaxed), + adaptive_floor_active_writers_current: self + .floor_runtime + .me_adaptive_floor_active_writers_current + .load(Ordering::Relaxed), + adaptive_floor_warm_writers_current: self + .floor_runtime + .me_adaptive_floor_warm_writers_current + .load(Ordering::Relaxed), + me_keepalive_enabled: self.writer_lifecycle.me_keepalive_enabled, + me_keepalive_interval_secs: self.writer_lifecycle.me_keepalive_interval.as_secs(), + me_keepalive_jitter_secs: self.writer_lifecycle.me_keepalive_jitter.as_secs(), + me_keepalive_payload_random: self.writer_lifecycle.me_keepalive_payload_random, + rpc_proxy_req_every_secs: self + .writer_lifecycle + .rpc_proxy_req_every_secs + .load(Ordering::Relaxed), + me_reconnect_max_concurrent_per_dc: self + .reconnect_runtime + .me_reconnect_max_concurrent_per_dc, + me_reconnect_backoff_base_ms: self + .reconnect_runtime + .me_reconnect_backoff_base + .as_millis() as u64, + me_reconnect_backoff_cap_ms: self.reconnect_runtime.me_reconnect_backoff_cap.as_millis() + as u64, + me_reconnect_fast_retry_count: self.reconnect_runtime.me_reconnect_fast_retry_count, + me_pool_drain_ttl_secs: self + .drain_runtime + .me_pool_drain_ttl_secs + .load(Ordering::Relaxed), + me_pool_force_close_secs: self + .drain_runtime + .me_pool_force_close_secs + .load(Ordering::Relaxed), + me_pool_min_fresh_ratio: Self::permille_to_ratio( + self.drain_runtime + .me_pool_min_fresh_ratio_permille + .load(Ordering::Relaxed), + ), + me_bind_stale_mode: bind_stale_mode_label(self.bind_stale_mode()), + me_bind_stale_ttl_secs: self + .binding_policy + .me_bind_stale_ttl_secs + .load(Ordering::Relaxed), + me_single_endpoint_shadow_writers: self + .single_endpoint_runtime + .me_single_endpoint_shadow_writers + .load(Ordering::Relaxed), + me_single_endpoint_outage_mode_enabled: self + .single_endpoint_runtime + .me_single_endpoint_outage_mode_enabled + .load(Ordering::Relaxed), + me_single_endpoint_outage_disable_quarantine: self + .single_endpoint_runtime + .me_single_endpoint_outage_disable_quarantine + .load(Ordering::Relaxed), + me_single_endpoint_outage_backoff_min_ms: self + .single_endpoint_runtime + .me_single_endpoint_outage_backoff_min_ms + .load(Ordering::Relaxed), + me_single_endpoint_outage_backoff_max_ms: self + .single_endpoint_runtime + .me_single_endpoint_outage_backoff_max_ms + .load(Ordering::Relaxed), + me_single_endpoint_shadow_rotate_every_secs: self + .single_endpoint_runtime + .me_single_endpoint_shadow_rotate_every_secs + .load(Ordering::Relaxed), + me_deterministic_writer_sort: self + .writer_selection_policy + .me_deterministic_writer_sort + .load(Ordering::Relaxed), + me_writer_pick_mode: writer_pick_mode_label(self.writer_pick_mode()), + me_writer_pick_sample_size: self.writer_pick_sample_size() as u8, + me_socks_kdf_policy: socks_kdf_policy_label(self.socks_kdf_policy()), + quarantined_endpoints, + network_path, + } + } + + pub(crate) async fn api_coherent_snapshots( + &self, + ) -> (MeApiStatusSnapshot, MeApiRuntimeSnapshot) { + let mut attempts = 0usize; + loop { + let reinit = self.reinit.status.load_full(); + let status = self.api_status_snapshot_for_reinit(reinit.as_ref()).await; + let runtime = self.api_runtime_snapshot_for_reinit(reinit.as_ref()).await; + attempts += 1; + if Arc::ptr_eq(&reinit, &self.reinit.status.load_full()) || attempts >= 3 { + return (status, runtime); + } + } + } +} diff --git a/src/transport/middle_proxy/pool_status/status_snapshot.rs b/src/transport/middle_proxy/pool_status/status_snapshot.rs new file mode 100644 index 0000000..660db57 --- /dev/null +++ b/src/transport/middle_proxy/pool_status/status_snapshot.rs @@ -0,0 +1,308 @@ +use super::*; + +impl MePool { + pub(crate) async fn admission_ready_conditional_cast(&self) -> bool { + let mut endpoints_by_dc = BTreeMap::>::new(); + if self.decision.ipv4_me { + let map = self.proxy_map_v4.read().await.clone(); + extend_signed_endpoints(&mut endpoints_by_dc, map); + } + if self.decision.ipv6_me { + let map = self.proxy_map_v6.read().await.clone(); + extend_signed_endpoints(&mut endpoints_by_dc, map); + } + + if endpoints_by_dc.is_empty() { + return false; + } + + let writers = self.writers.read().await.clone(); + let mut live_writers_by_dc = HashMap::::new(); + for writer in writers.iter() { + if writer.draining.load(Ordering::Relaxed) { + continue; + } + if let Ok(dc) = i16::try_from(writer.writer_dc) { + *live_writers_by_dc.entry(dc).or_insert(0) += 1; + } + } + + for dc in endpoints_by_dc.keys() { + let alive = live_writers_by_dc.get(dc).copied().unwrap_or(0); + if alive == 0 { + return false; + } + } + + true + } + + #[allow(dead_code)] + pub(crate) async fn admission_ready_full_floor(&self) -> bool { + let mut endpoints_by_dc = BTreeMap::>::new(); + if self.decision.ipv4_me { + let map = self.proxy_map_v4.read().await.clone(); + extend_signed_endpoints(&mut endpoints_by_dc, map); + } + if self.decision.ipv6_me { + let map = self.proxy_map_v6.read().await.clone(); + extend_signed_endpoints(&mut endpoints_by_dc, map); + } + + if endpoints_by_dc.is_empty() { + return false; + } + + let writers = self.writers.read().await.clone(); + let mut live_writers_by_dc = HashMap::::new(); + for writer in writers.iter() { + if writer.draining.load(Ordering::Relaxed) { + continue; + } + if let Ok(dc) = i16::try_from(writer.writer_dc) { + *live_writers_by_dc.entry(dc).or_insert(0) += 1; + } + } + + for (dc, endpoints) in endpoints_by_dc { + let endpoint_count = endpoints.len(); + if endpoint_count == 0 { + return false; + } + let required = self.required_writers_for_dc_with_floor_mode(endpoint_count, false); + let alive = live_writers_by_dc.get(&dc).copied().unwrap_or(0); + if alive < required { + return false; + } + } + + true + } + + pub(crate) async fn api_status_snapshot(&self) -> MeApiStatusSnapshot { + let reinit = self.reinit.status.load_full(); + self.api_status_snapshot_for_reinit(reinit.as_ref()).await + } + + pub(super) async fn api_status_snapshot_for_reinit( + &self, + reinit: &ReinitStatusSnapshot, + ) -> MeApiStatusSnapshot { + let now_epoch_secs = Self::now_epoch_secs(); + let active_generation = reinit.active_generation; + let drain_ttl_secs = self + .drain_runtime + .me_pool_drain_ttl_secs + .load(Ordering::Relaxed); + + let mut endpoints_by_dc = BTreeMap::>::new(); + if self.decision.ipv4_me { + let map = self.proxy_map_v4.read().await.clone(); + extend_signed_endpoints(&mut endpoints_by_dc, map); + } + if self.decision.ipv6_me { + let map = self.proxy_map_v6.read().await.clone(); + extend_signed_endpoints(&mut endpoints_by_dc, map); + } + + let configured_dc_groups = endpoints_by_dc.len(); + let configured_endpoints = endpoints_by_dc.values().map(BTreeSet::len).sum(); + + let required_writers = endpoints_by_dc + .values() + .map(|endpoints| self.required_writers_for_dc_with_floor_mode(endpoints.len(), false)) + .sum(); + + let idle_since = self.registry.writer_idle_since_snapshot().await; + let activity = self.registry.writer_activity_snapshot().await; + let rtt = self.rtt_stats.lock().await.clone(); + let writers = self.writers.read().await.clone(); + + let mut live_writers_by_dc_endpoint = HashMap::<(i16, SocketAddr), usize>::new(); + let mut live_writers_by_dc = HashMap::::new(); + let mut fresh_writers_by_dc = HashMap::::new(); + let mut dc_rtt_agg = HashMap::::new(); + let mut writer_rows = Vec::::with_capacity(writers.len()); + + for writer in writers.iter() { + let endpoint = writer.addr; + let dc = i16::try_from(writer.writer_dc).ok(); + let draining = writer.draining.load(Ordering::Relaxed); + let degraded = writer.degraded.load(Ordering::Relaxed); + let matches_active_generation = writer.generation == active_generation; + let in_desired_map = dc + .and_then(|dc_idx| endpoints_by_dc.get(&dc_idx)) + .is_some_and(|endpoints| endpoints.contains(&endpoint)); + let bound_clients = activity + .bound_clients_by_writer + .get(&writer.id) + .copied() + .unwrap_or(0); + let idle_for_secs = idle_since + .get(&writer.id) + .map(|idle_ts| now_epoch_secs.saturating_sub(*idle_ts)); + let rtt_ema_ms = rtt.get(&writer.id).map(|(_, ema)| *ema); + let allow_drain_fallback = writer.allow_drain_fallback.load(Ordering::Relaxed); + let drain_started_at_epoch_secs = writer + .draining_started_at_epoch_secs + .load(Ordering::Relaxed); + let drain_deadline_epoch_secs = + writer.drain_deadline_epoch_secs.load(Ordering::Relaxed); + let drain_started_at_epoch_secs = + (drain_started_at_epoch_secs != 0).then_some(drain_started_at_epoch_secs); + let drain_deadline_epoch_secs = + (drain_deadline_epoch_secs != 0).then_some(drain_deadline_epoch_secs); + let drain_over_ttl = draining + && drain_ttl_secs > 0 + && drain_started_at_epoch_secs + .is_some_and(|started| now_epoch_secs.saturating_sub(started) > drain_ttl_secs); + let state = match WriterContour::from_u8(writer.contour.load(Ordering::Relaxed)) { + WriterContour::Warm => "warm", + WriterContour::Active => "active", + WriterContour::Draining => "draining", + }; + + if !draining && let Some(dc_idx) = dc { + *live_writers_by_dc_endpoint + .entry((dc_idx, endpoint)) + .or_insert(0) += 1; + *live_writers_by_dc.entry(dc_idx).or_insert(0) += 1; + if let Some(ema_ms) = rtt_ema_ms { + let entry = dc_rtt_agg.entry(dc_idx).or_insert((0.0, 0)); + entry.0 += ema_ms; + entry.1 += 1; + } + if matches_active_generation && in_desired_map { + *fresh_writers_by_dc.entry(dc_idx).or_insert(0) += 1; + } + } + + writer_rows.push(MeApiWriterStatusSnapshot { + writer_id: writer.id, + dc, + endpoint, + generation: writer.generation, + state, + draining, + degraded, + bound_clients, + idle_for_secs, + rtt_ema_ms, + matches_active_generation, + in_desired_map, + allow_drain_fallback, + drain_started_at_epoch_secs, + drain_deadline_epoch_secs, + drain_over_ttl, + }); + } + + writer_rows.sort_by_key(|row| (row.dc.unwrap_or(i16::MAX), row.endpoint, row.writer_id)); + + let mut dcs = Vec::::with_capacity(endpoints_by_dc.len()); + let mut available_endpoints = 0usize; + let mut alive_writers = 0usize; + let mut fresh_alive_writers = 0usize; + let floor_mode = self.floor_mode(); + let adaptive_cpu_cores = (self + .floor_runtime + .me_adaptive_floor_cpu_cores_effective + .load(Ordering::Relaxed) as usize) + .max(1); + for (dc, endpoints) in endpoints_by_dc { + let endpoint_count = endpoints.len(); + let dc_available_endpoints = endpoints + .iter() + .filter(|endpoint| live_writers_by_dc_endpoint.contains_key(&(dc, **endpoint))) + .count(); + let base_required = self.required_writers_for_dc(endpoint_count); + let dc_required_writers = + self.required_writers_for_dc_with_floor_mode(endpoint_count, false); + let floor_min = if endpoint_count <= 1 { + (self + .floor_runtime + .me_adaptive_floor_min_writers_single_endpoint + .load(Ordering::Relaxed) as usize) + .max(1) + .min(base_required.max(1)) + } else { + (self + .floor_runtime + .me_adaptive_floor_min_writers_multi_endpoint + .load(Ordering::Relaxed) as usize) + .max(1) + .min(base_required.max(1)) + }; + let extra_per_core = if endpoint_count <= 1 { + self.floor_runtime + .me_adaptive_floor_max_extra_writers_single_per_core + .load(Ordering::Relaxed) as usize + } else { + self.floor_runtime + .me_adaptive_floor_max_extra_writers_multi_per_core + .load(Ordering::Relaxed) as usize + }; + let floor_max = + base_required.saturating_add(adaptive_cpu_cores.saturating_mul(extra_per_core)); + let floor_capped = + matches!(floor_mode, MeFloorMode::Adaptive) && dc_required_writers < base_required; + let dc_alive_writers = live_writers_by_dc.get(&dc).copied().unwrap_or(0); + let dc_fresh_alive_writers = fresh_writers_by_dc.get(&dc).copied().unwrap_or(0); + let dc_load = activity + .active_sessions_by_target_dc + .get(&dc) + .copied() + .unwrap_or(0); + let dc_rtt_ms = dc_rtt_agg + .get(&dc) + .and_then(|(sum, count)| (*count > 0).then_some(*sum / (*count as f64))); + + available_endpoints += dc_available_endpoints; + alive_writers += dc_alive_writers; + fresh_alive_writers += dc_fresh_alive_writers; + + dcs.push(MeApiDcStatusSnapshot { + dc, + endpoint_writers: endpoints + .iter() + .map(|endpoint| MeApiDcEndpointWriterSnapshot { + endpoint: *endpoint, + active_writers: live_writers_by_dc_endpoint + .get(&(dc, *endpoint)) + .copied() + .unwrap_or(0), + }) + .collect(), + endpoints: endpoints.into_iter().collect(), + available_endpoints: dc_available_endpoints, + available_pct: ratio_pct(dc_available_endpoints, endpoint_count), + required_writers: dc_required_writers, + floor_min, + floor_target: dc_required_writers, + floor_max, + floor_capped, + alive_writers: dc_alive_writers, + coverage_pct: ratio_pct(dc_alive_writers, dc_required_writers), + fresh_alive_writers: dc_fresh_alive_writers, + fresh_coverage_pct: ratio_pct(dc_fresh_alive_writers, dc_required_writers), + rtt_ms: dc_rtt_ms, + load: dc_load, + }); + } + + MeApiStatusSnapshot { + generated_at_epoch_secs: now_epoch_secs, + configured_dc_groups, + configured_endpoints, + available_endpoints, + available_pct: ratio_pct(available_endpoints, configured_endpoints), + required_writers, + alive_writers, + coverage_pct: ratio_pct(alive_writers, required_writers), + fresh_alive_writers, + fresh_coverage_pct: ratio_pct(fresh_alive_writers, required_writers), + writers: writer_rows, + dcs, + } + } +} diff --git a/src/transport/middle_proxy/pool_writer.rs b/src/transport/middle_proxy/pool_writer.rs index 9c54c07..6841a29 100644 --- a/src/transport/middle_proxy/pool_writer.rs +++ b/src/transport/middle_proxy/pool_writer.rs @@ -22,6 +22,9 @@ use super::pool::{MePool, MeWriter, WriterContour}; use super::reader::reader_loop; use super::wire::build_proxy_req_payload; +// Writer admission, teardown, and drain-state transitions. +mod runtime; + const ME_ACTIVE_PING_SECS: u64 = 25; const ME_ACTIVE_PING_JITTER_SECS: i64 = 5; const ME_IDLE_KEEPALIVE_MAX_SECS: u64 = 5; @@ -327,465 +330,3 @@ async fn rpc_proxy_req_signal_loop( conn_lease.unregister().await; } } - -impl MePool { - pub(crate) async fn prune_closed_writers(self: &Arc) { - let closed_writer_ids: Vec = { - let ws = self.writers.read().await; - ws.iter() - .filter(|w| w.tx.is_closed()) - .map(|w| w.id) - .collect() - }; - if closed_writer_ids.is_empty() { - return; - } - - for writer_id in closed_writer_ids { - let _ = self.remove_writer_and_close_clients(writer_id).await; - } - } - - pub(crate) async fn connect_one_for_dc( - self: &Arc, - addr: SocketAddr, - writer_dc: i32, - rng: &SecureRandom, - ) -> Result<()> { - self.connect_one_with_generation_contour( - addr, - rng, - self.current_generation(), - WriterContour::Active, - writer_dc, - ) - .await - } - - pub(super) async fn connect_one_with_generation_contour( - self: &Arc, - addr: SocketAddr, - rng: &SecureRandom, - generation: u64, - contour: WriterContour, - writer_dc: i32, - ) -> Result<()> { - self.connect_one_with_generation_contour_for_dc(addr, rng, generation, contour, writer_dc) - .await - } - - pub(super) async fn connect_one_with_generation_contour_for_dc( - self: &Arc, - addr: SocketAddr, - rng: &SecureRandom, - generation: u64, - contour: WriterContour, - writer_dc: i32, - ) -> Result<()> { - self.connect_one_with_generation_contour_for_dc_with_cap_policy( - addr, rng, generation, contour, writer_dc, false, - ) - .await - } - - pub(super) async fn connect_one_with_generation_contour_for_dc_with_cap_policy( - self: &Arc, - addr: SocketAddr, - rng: &SecureRandom, - generation: u64, - contour: WriterContour, - writer_dc: i32, - allow_coverage_override: bool, - ) -> Result<()> { - let Some(_writer_open_reservation) = self - .reserve_writer_open(contour, allow_coverage_override, writer_dc) - .await - else { - return Err(ProxyError::Proxy(format!( - "ME {contour:?} writer cap reached" - ))); - }; - - let secret_len = self.proxy_secret.read().await.secret.len(); - if secret_len < 32 { - return Err(ProxyError::Proxy( - "proxy-secret too short for ME auth".into(), - )); - } - - let dc_idx = i16::try_from(writer_dc).ok(); - let (stream, _connect_ms, upstream_egress) = self.connect_tcp(addr, dc_idx).await?; - let hs = self - .handshake_only(stream, addr, upstream_egress, rng) - .await?; - let Some(task_registration) = self.lifecycle.try_register() else { - return Err(ProxyError::Proxy("ME pool lifecycle closed".into())); - }; - - let writer_id = self.next_writer_id.fetch_add(1, Ordering::Relaxed); - let contour = Arc::new(AtomicU8::new(contour.as_u8())); - let cancel = CancellationToken::new(); - let degraded = Arc::new(AtomicBool::new(false)); - let rtt_ema_ms_x10 = Arc::new(AtomicU32::new(0)); - let draining = Arc::new(AtomicBool::new(false)); - let draining_started_at_epoch_secs = Arc::new(AtomicU64::new(0)); - let drain_deadline_epoch_secs = Arc::new(AtomicU64::new(0)); - let allow_drain_fallback = Arc::new(AtomicBool::new(false)); - let byte_budget = self.new_writer_byte_budget(); - let (tx, rx) = - mpsc::channel::(self.writer_lifecycle.writer_cmd_channel_capacity); - let rpc_writer = RpcWriter { - writer: hs.wr, - key: hs.write_key, - iv: hs.write_iv, - seq_no: 0, - crc_mode: hs.crc_mode, - frame_buf: Vec::new(), - }; - let writer = MeWriter { - id: writer_id, - addr, - source_ip: hs.source_ip, - writer_dc, - generation, - contour: contour.clone(), - created_at: Instant::now(), - tx: tx.clone(), - byte_budget: byte_budget.clone(), - cancel: cancel.clone(), - degraded: degraded.clone(), - rtt_ema_ms_x10: rtt_ema_ms_x10.clone(), - draining: draining.clone(), - draining_started_at_epoch_secs: draining_started_at_epoch_secs.clone(), - drain_deadline_epoch_secs: drain_deadline_epoch_secs.clone(), - allow_drain_fallback: allow_drain_fallback.clone(), - }; - self.writers - .update(|writers| writers.push(writer.clone())) - .await; - self.registry - .register_writer(writer_id, tx.clone(), byte_budget) - .await; - self.registry.mark_writer_idle(writer_id).await; - self.conn_count.fetch_add(1, Ordering::Relaxed); - self.notify_writer_epoch(); - - let reg = self.registry.clone(); - let writers_arc = self.writers_arc(); - let ping_tracker = Arc::new(tokio::sync::Mutex::new(HashMap::::new())); - let ping_tracker_reader = ping_tracker.clone(); - let ping_tracker_ping = ping_tracker.clone(); - let rtt_stats = self.rtt_stats.clone(); - let stats_reader = self.stats.clone(); - let stats_reader_close = self.stats.clone(); - let stats_ping = self.stats.clone(); - let stats_signal = self.stats.clone(); - let pool_lifecycle = Arc::downgrade(self); - let pool_ping = Arc::downgrade(self); - let pool_signal = Arc::downgrade(self); - let tx_reader = tx.clone(); - let tx_ping = tx.clone(); - let tx_signal = tx.clone(); - let keepalive_enabled = self.writer_lifecycle.me_keepalive_enabled; - let keepalive_interval = self.writer_lifecycle.me_keepalive_interval; - let keepalive_jitter = self.writer_lifecycle.me_keepalive_jitter; - let keepalive_jitter_signal = self.writer_lifecycle.me_keepalive_jitter; - let rpc_proxy_req_every_secs = self - .writer_lifecycle - .rpc_proxy_req_every_secs - .load(Ordering::Relaxed); - let cancel_reader = cancel.clone(); - let cancel_writer = cancel.clone(); - let cancel_ping = cancel.clone(); - let cancel_signal = cancel.clone(); - let cancel_select = cancel.clone(); - let cancel_cleanup = cancel.clone(); - let route_backpressure_enabled = - self.transport_policy.me_route_backpressure_enabled.clone(); - let route_fairshare_enabled = self.transport_policy.me_route_fairshare_enabled.clone(); - let reader_route_data_wait_ms = self.transport_policy.me_reader_route_data_wait_ms.clone(); - - self.lifecycle.spawn_registered_writer(task_registration, async move { - // Reader MUST be the first branch in biased select! to avoid read starvation. - let exit = tokio::select! { - biased; - - reader_res = reader_loop( - hs.rd, - hs.read_key, - hs.read_iv, - hs.crc_mode, - reg.clone(), - BytesMut::new(), - BytesMut::new(), - tx_reader, - ping_tracker_reader, - rtt_stats, - stats_reader, - writer_id, - degraded, - rtt_ema_ms_x10, - route_backpressure_enabled, - route_fairshare_enabled, - reader_route_data_wait_ms, - cancel_reader, - ) => WriterLifecycleExit::Reader(reader_res), - writer_res = writer_command_loop(rx, rpc_writer, cancel_writer) => { - WriterLifecycleExit::Writer(writer_res) - } - _ = ping_loop( - pool_ping, - writer_id, - tx_ping, - ping_tracker_ping, - stats_ping, - keepalive_enabled, - keepalive_interval, - keepalive_jitter, - cancel_ping, - ) => WriterLifecycleExit::Ping, - _ = rpc_proxy_req_signal_loop( - pool_signal, - writer_id, - tx_signal, - stats_signal, - cancel_signal, - keepalive_jitter_signal, - rpc_proxy_req_every_secs, - ) => WriterLifecycleExit::Signal, - _ = cancel_select.cancelled() => WriterLifecycleExit::Cancelled, - }; - - match exit { - WriterLifecycleExit::Reader(res) => { - let idle_close_by_peer = if let Err(e) = res.as_ref() { - is_me_peer_closed_error(e) && reg.is_writer_empty(writer_id).await - } else { - false - }; - if idle_close_by_peer { - stats_reader_close.increment_me_idle_close_by_peer_total(); - info!(writer_id, "ME socket closed by peer on idle writer"); - } - if let Err(e) = res - && !idle_close_by_peer - { - warn!(error = %e, "ME reader ended"); - } - } - WriterLifecycleExit::Writer(res) => { - if let Err(e) = res { - warn!(error = %e, "ME writer command loop ended"); - } - } - WriterLifecycleExit::Ping => { - debug!(writer_id, "ME ping loop finished"); - } - WriterLifecycleExit::Signal => { - debug!(writer_id, "ME rpc_proxy_req signal loop finished"); - } - WriterLifecycleExit::Cancelled => {} - } - - if let Some(pool) = pool_lifecycle.upgrade() { - pool.remove_writer_and_close_clients(writer_id).await; - } else { - // Fallback for shutdown races: make lifecycle exit observable by prune. - cancel_cleanup.cancel(); - } - - let remaining = writers_arc.read().await.len(); - debug!(writer_id, remaining, "ME writer lifecycle task finished"); - }); - - Ok(()) - } - - pub(crate) async fn remove_writer_and_close_clients(self: &Arc, writer_id: u64) { - // Full client cleanup now happens inside `registry.writer_lost` to keep - // writer reap/remove paths strictly non-blocking per connection. - let _ = self - .remove_writer_with_mode(writer_id, WriterTeardownMode::Any) - .await; - } - - pub(super) async fn remove_draining_writer_hard_detach( - self: &Arc, - writer_id: u64, - ) -> bool { - self.remove_writer_with_mode(writer_id, WriterTeardownMode::DrainingOnly) - .await - } - - #[allow(dead_code)] - async fn remove_writer_only(self: &Arc, writer_id: u64) -> bool { - self.remove_writer_with_mode(writer_id, WriterTeardownMode::Any) - .await - } - - // Authoritative teardown primitive shared by normal cleanup and watchdog path. - // Lock-order invariant: - // 1) mutate `writers` under pool write lock, - // 2) release pool lock, - // 3) run registry/metrics/refill side effects. - // `registry.writer_lost` must never run while `writers` lock is held. - async fn remove_writer_with_mode( - self: &Arc, - writer_id: u64, - mode: WriterTeardownMode, - ) -> bool { - let mut close_tx: Option> = None; - let mut removed_addr: Option = None; - let mut removed_dc: Option = None; - let mut removed_uptime: Option = None; - let mut trigger_refill = false; - let mut removed = false; - { - let mut ws = self.writers.write().await; - if let Some(pos) = ws.iter().position(|w| w.id == writer_id) { - if matches!(mode, WriterTeardownMode::DrainingOnly) - && !ws[pos].draining.load(Ordering::Relaxed) - { - return false; - } - let w = ws.remove(pos); - let was_draining = w.draining.load(Ordering::Relaxed); - if was_draining { - self.stats.decrement_pool_drain_active(); - self.decrement_draining_active_runtime(); - } - self.stats.increment_me_writer_removed_total(); - w.cancel.cancel(); - removed_addr = Some(w.addr); - removed_dc = Some(w.writer_dc); - removed_uptime = Some(w.created_at.elapsed()); - trigger_refill = !was_draining; - if trigger_refill { - self.stats.increment_me_writer_removed_unexpected_total(); - } - close_tx = Some(w.tx.clone()); - self.conn_count.fetch_sub(1, Ordering::Relaxed); - removed = true; - } - } - // State invariant: - // - writer is removed from `self.writers` (pool visibility), - // - writer is removed from registry routing/binding maps via `writer_lost`. - // The close command below is only a best-effort accelerator for task shutdown. - // Cleanup progress must never depend on command-channel availability. - let _ = self.registry.writer_lost(writer_id).await; - self.rtt_stats.lock().await.remove(&writer_id); - if let Some(tx) = close_tx { - // Keep teardown critical path non-blocking: close is best-effort only. - let _ = tx.try_send(WriterCommand::Close); - } - if let Some(addr) = removed_addr { - if let Some(uptime) = removed_uptime { - // Quarantine contract: only unexpected removals are considered endpoint flap. - if trigger_refill { - self.stats - .increment_me_endpoint_quarantine_unexpected_total(); - self.maybe_quarantine_flapping_endpoint(addr, uptime, "unexpected") - .await; - } else { - self.stats - .increment_me_endpoint_quarantine_draining_suppressed_total(); - debug!( - %addr, - uptime_ms = uptime.as_millis(), - "Skipping endpoint quarantine for draining writer removal" - ); - } - } - if trigger_refill && let Some(writer_dc) = removed_dc { - self.trigger_immediate_refill_for_dc(addr, writer_dc); - } - } - if removed { - self.notify_writer_epoch(); - } - removed - } - - pub(crate) async fn mark_writer_draining_with_timeout( - self: &Arc, - writer_id: u64, - timeout: Option, - allow_drain_fallback: bool, - ) { - let timeout = timeout.filter(|d| !d.is_zero()); - let found = { - let mut ws = self.writers.write().await; - if let Some(w) = ws.iter_mut().find(|w| w.id == writer_id) { - let already_draining = w.draining.swap(true, Ordering::Relaxed); - w.allow_drain_fallback - .store(allow_drain_fallback, Ordering::Relaxed); - let now_epoch_secs = Self::now_epoch_secs(); - w.draining_started_at_epoch_secs - .store(now_epoch_secs, Ordering::Relaxed); - let drain_deadline_epoch_secs = timeout - .map(|duration| now_epoch_secs.saturating_add(duration.as_secs())) - .unwrap_or(0); - w.drain_deadline_epoch_secs - .store(drain_deadline_epoch_secs, Ordering::Relaxed); - if !already_draining { - self.stats.increment_pool_drain_active(); - self.increment_draining_active_runtime(); - } - w.contour - .store(WriterContour::Draining.as_u8(), Ordering::Relaxed); - w.draining.store(true, Ordering::Relaxed); - true - } else { - false - } - }; - - if !found { - return; - } - - let timeout_secs = timeout.map(|d| d.as_secs()).unwrap_or(0); - debug!( - writer_id, - timeout_secs, allow_drain_fallback, "ME writer marked draining" - ); - } - - pub(crate) async fn mark_writer_draining(self: &Arc, writer_id: u64) { - self.mark_writer_draining_with_timeout(writer_id, Some(Duration::from_secs(300)), false) - .await; - } - - pub(super) fn writer_accepts_new_binding(&self, writer: &MeWriter) -> bool { - if !writer.draining.load(Ordering::Relaxed) { - return true; - } - if !writer.allow_drain_fallback.load(Ordering::Relaxed) { - return false; - } - - match self.bind_stale_mode() { - MeBindStaleMode::Never => false, - MeBindStaleMode::Always => true, - MeBindStaleMode::Ttl => { - let ttl_secs = self - .binding_policy - .me_bind_stale_ttl_secs - .load(Ordering::Relaxed); - if ttl_secs == 0 { - return true; - } - - let started = writer - .draining_started_at_epoch_secs - .load(Ordering::Relaxed); - if started == 0 { - return false; - } - - Self::now_epoch_secs().saturating_sub(started) <= ttl_secs - } - } - } -} diff --git a/src/transport/middle_proxy/pool_writer/runtime.rs b/src/transport/middle_proxy/pool_writer/runtime.rs new file mode 100644 index 0000000..d59218d --- /dev/null +++ b/src/transport/middle_proxy/pool_writer/runtime.rs @@ -0,0 +1,467 @@ +use super::*; + +impl MePool { + pub(crate) async fn prune_closed_writers(self: &Arc) { + let closed_writer_ids: Vec = { + let ws = self.writers.read().await; + ws.iter() + .filter(|w| w.tx.is_closed()) + .map(|w| w.id) + .collect() + }; + if closed_writer_ids.is_empty() { + return; + } + + for writer_id in closed_writer_ids { + let _ = self.remove_writer_and_close_clients(writer_id).await; + } + } + + pub(crate) async fn connect_one_for_dc( + self: &Arc, + addr: SocketAddr, + writer_dc: i32, + rng: &SecureRandom, + ) -> Result<()> { + self.connect_one_with_generation_contour( + addr, + rng, + self.current_generation(), + WriterContour::Active, + writer_dc, + ) + .await + } + + pub(in crate::transport::middle_proxy) async fn connect_one_with_generation_contour( + self: &Arc, + addr: SocketAddr, + rng: &SecureRandom, + generation: u64, + contour: WriterContour, + writer_dc: i32, + ) -> Result<()> { + self.connect_one_with_generation_contour_for_dc(addr, rng, generation, contour, writer_dc) + .await + } + + pub(in crate::transport::middle_proxy) async fn connect_one_with_generation_contour_for_dc( + self: &Arc, + addr: SocketAddr, + rng: &SecureRandom, + generation: u64, + contour: WriterContour, + writer_dc: i32, + ) -> Result<()> { + self.connect_one_with_generation_contour_for_dc_with_cap_policy( + addr, rng, generation, contour, writer_dc, false, + ) + .await + } + + pub(in crate::transport::middle_proxy) async fn connect_one_with_generation_contour_for_dc_with_cap_policy( + self: &Arc, + addr: SocketAddr, + rng: &SecureRandom, + generation: u64, + contour: WriterContour, + writer_dc: i32, + allow_coverage_override: bool, + ) -> Result<()> { + let Some(_writer_open_reservation) = self + .reserve_writer_open(contour, allow_coverage_override, writer_dc) + .await + else { + return Err(ProxyError::Proxy(format!( + "ME {contour:?} writer cap reached" + ))); + }; + + let secret_len = self.proxy_secret.read().await.secret.len(); + if secret_len < 32 { + return Err(ProxyError::Proxy( + "proxy-secret too short for ME auth".into(), + )); + } + + let dc_idx = i16::try_from(writer_dc).ok(); + let (stream, _connect_ms, upstream_egress) = self.connect_tcp(addr, dc_idx).await?; + let hs = self + .handshake_only(stream, addr, upstream_egress, rng) + .await?; + let Some(task_registration) = self.lifecycle.try_register() else { + return Err(ProxyError::Proxy("ME pool lifecycle closed".into())); + }; + + let writer_id = self.next_writer_id.fetch_add(1, Ordering::Relaxed); + let contour = Arc::new(AtomicU8::new(contour.as_u8())); + let cancel = CancellationToken::new(); + let degraded = Arc::new(AtomicBool::new(false)); + let rtt_ema_ms_x10 = Arc::new(AtomicU32::new(0)); + let draining = Arc::new(AtomicBool::new(false)); + let draining_started_at_epoch_secs = Arc::new(AtomicU64::new(0)); + let drain_deadline_epoch_secs = Arc::new(AtomicU64::new(0)); + let allow_drain_fallback = Arc::new(AtomicBool::new(false)); + let byte_budget = self.new_writer_byte_budget(); + let (tx, rx) = + mpsc::channel::(self.writer_lifecycle.writer_cmd_channel_capacity); + let rpc_writer = RpcWriter { + writer: hs.wr, + key: hs.write_key, + iv: hs.write_iv, + seq_no: 0, + crc_mode: hs.crc_mode, + frame_buf: Vec::new(), + }; + let writer = MeWriter { + id: writer_id, + addr, + source_ip: hs.source_ip, + writer_dc, + generation, + contour: contour.clone(), + created_at: Instant::now(), + tx: tx.clone(), + byte_budget: byte_budget.clone(), + cancel: cancel.clone(), + degraded: degraded.clone(), + rtt_ema_ms_x10: rtt_ema_ms_x10.clone(), + draining: draining.clone(), + draining_started_at_epoch_secs: draining_started_at_epoch_secs.clone(), + drain_deadline_epoch_secs: drain_deadline_epoch_secs.clone(), + allow_drain_fallback: allow_drain_fallback.clone(), + }; + self.writers + .update(|writers| writers.push(writer.clone())) + .await; + self.registry + .register_writer(writer_id, tx.clone(), byte_budget) + .await; + self.registry.mark_writer_idle(writer_id).await; + self.conn_count.fetch_add(1, Ordering::Relaxed); + self.notify_writer_epoch(); + + let reg = self.registry.clone(); + let writers_arc = self.writers_arc(); + let ping_tracker = Arc::new(tokio::sync::Mutex::new(HashMap::::new())); + let ping_tracker_reader = ping_tracker.clone(); + let ping_tracker_ping = ping_tracker.clone(); + let rtt_stats = self.rtt_stats.clone(); + let stats_reader = self.stats.clone(); + let stats_reader_close = self.stats.clone(); + let stats_ping = self.stats.clone(); + let stats_signal = self.stats.clone(); + let pool_lifecycle = Arc::downgrade(self); + let pool_ping = Arc::downgrade(self); + let pool_signal = Arc::downgrade(self); + let tx_reader = tx.clone(); + let tx_ping = tx.clone(); + let tx_signal = tx.clone(); + let keepalive_enabled = self.writer_lifecycle.me_keepalive_enabled; + let keepalive_interval = self.writer_lifecycle.me_keepalive_interval; + let keepalive_jitter = self.writer_lifecycle.me_keepalive_jitter; + let keepalive_jitter_signal = self.writer_lifecycle.me_keepalive_jitter; + let rpc_proxy_req_every_secs = self + .writer_lifecycle + .rpc_proxy_req_every_secs + .load(Ordering::Relaxed); + let cancel_reader = cancel.clone(); + let cancel_writer = cancel.clone(); + let cancel_ping = cancel.clone(); + let cancel_signal = cancel.clone(); + let cancel_select = cancel.clone(); + let cancel_cleanup = cancel.clone(); + let route_backpressure_enabled = + self.transport_policy.me_route_backpressure_enabled.clone(); + let route_fairshare_enabled = self.transport_policy.me_route_fairshare_enabled.clone(); + let reader_route_data_wait_ms = self.transport_policy.me_reader_route_data_wait_ms.clone(); + + self.lifecycle + .spawn_registered_writer(task_registration, async move { + // Reader MUST be the first branch in biased select! to avoid read starvation. + let exit = tokio::select! { + biased; + + reader_res = reader_loop( + hs.rd, + hs.read_key, + hs.read_iv, + hs.crc_mode, + reg.clone(), + BytesMut::new(), + BytesMut::new(), + tx_reader, + ping_tracker_reader, + rtt_stats, + stats_reader, + writer_id, + degraded, + rtt_ema_ms_x10, + route_backpressure_enabled, + route_fairshare_enabled, + reader_route_data_wait_ms, + cancel_reader, + ) => WriterLifecycleExit::Reader(reader_res), + writer_res = writer_command_loop(rx, rpc_writer, cancel_writer) => { + WriterLifecycleExit::Writer(writer_res) + } + _ = ping_loop( + pool_ping, + writer_id, + tx_ping, + ping_tracker_ping, + stats_ping, + keepalive_enabled, + keepalive_interval, + keepalive_jitter, + cancel_ping, + ) => WriterLifecycleExit::Ping, + _ = rpc_proxy_req_signal_loop( + pool_signal, + writer_id, + tx_signal, + stats_signal, + cancel_signal, + keepalive_jitter_signal, + rpc_proxy_req_every_secs, + ) => WriterLifecycleExit::Signal, + _ = cancel_select.cancelled() => WriterLifecycleExit::Cancelled, + }; + + match exit { + WriterLifecycleExit::Reader(res) => { + let idle_close_by_peer = if let Err(e) = res.as_ref() { + is_me_peer_closed_error(e) && reg.is_writer_empty(writer_id).await + } else { + false + }; + if idle_close_by_peer { + stats_reader_close.increment_me_idle_close_by_peer_total(); + info!(writer_id, "ME socket closed by peer on idle writer"); + } + if let Err(e) = res + && !idle_close_by_peer + { + warn!(error = %e, "ME reader ended"); + } + } + WriterLifecycleExit::Writer(res) => { + if let Err(e) = res { + warn!(error = %e, "ME writer command loop ended"); + } + } + WriterLifecycleExit::Ping => { + debug!(writer_id, "ME ping loop finished"); + } + WriterLifecycleExit::Signal => { + debug!(writer_id, "ME rpc_proxy_req signal loop finished"); + } + WriterLifecycleExit::Cancelled => {} + } + + if let Some(pool) = pool_lifecycle.upgrade() { + pool.remove_writer_and_close_clients(writer_id).await; + } else { + // Fallback for shutdown races: make lifecycle exit observable by prune. + cancel_cleanup.cancel(); + } + + let remaining = writers_arc.read().await.len(); + debug!(writer_id, remaining, "ME writer lifecycle task finished"); + }); + + Ok(()) + } + + pub(crate) async fn remove_writer_and_close_clients(self: &Arc, writer_id: u64) { + // Full client cleanup now happens inside `registry.writer_lost` to keep + // writer reap/remove paths strictly non-blocking per connection. + let _ = self + .remove_writer_with_mode(writer_id, WriterTeardownMode::Any) + .await; + } + + pub(in crate::transport::middle_proxy) async fn remove_draining_writer_hard_detach( + self: &Arc, + writer_id: u64, + ) -> bool { + self.remove_writer_with_mode(writer_id, WriterTeardownMode::DrainingOnly) + .await + } + + #[allow(dead_code)] + async fn remove_writer_only(self: &Arc, writer_id: u64) -> bool { + self.remove_writer_with_mode(writer_id, WriterTeardownMode::Any) + .await + } + + // Authoritative teardown primitive shared by normal cleanup and watchdog path. + // Lock-order invariant: + // 1) mutate `writers` under pool write lock, + // 2) release pool lock, + // 3) run registry/metrics/refill side effects. + // `registry.writer_lost` must never run while `writers` lock is held. + async fn remove_writer_with_mode( + self: &Arc, + writer_id: u64, + mode: WriterTeardownMode, + ) -> bool { + let mut close_tx: Option> = None; + let mut removed_addr: Option = None; + let mut removed_dc: Option = None; + let mut removed_uptime: Option = None; + let mut trigger_refill = false; + let mut removed = false; + { + let mut ws = self.writers.write().await; + if let Some(pos) = ws.iter().position(|w| w.id == writer_id) { + if matches!(mode, WriterTeardownMode::DrainingOnly) + && !ws[pos].draining.load(Ordering::Relaxed) + { + return false; + } + let w = ws.remove(pos); + let was_draining = w.draining.load(Ordering::Relaxed); + if was_draining { + self.stats.decrement_pool_drain_active(); + self.decrement_draining_active_runtime(); + } + self.stats.increment_me_writer_removed_total(); + w.cancel.cancel(); + removed_addr = Some(w.addr); + removed_dc = Some(w.writer_dc); + removed_uptime = Some(w.created_at.elapsed()); + trigger_refill = !was_draining; + if trigger_refill { + self.stats.increment_me_writer_removed_unexpected_total(); + } + close_tx = Some(w.tx.clone()); + self.conn_count.fetch_sub(1, Ordering::Relaxed); + removed = true; + } + } + // State invariant: + // - writer is removed from `self.writers` (pool visibility), + // - writer is removed from registry routing/binding maps via `writer_lost`. + // The close command below is only a best-effort accelerator for task shutdown. + // Cleanup progress must never depend on command-channel availability. + let _ = self.registry.writer_lost(writer_id).await; + self.rtt_stats.lock().await.remove(&writer_id); + if let Some(tx) = close_tx { + // Keep teardown critical path non-blocking: close is best-effort only. + let _ = tx.try_send(WriterCommand::Close); + } + if let Some(addr) = removed_addr { + if let Some(uptime) = removed_uptime { + // Quarantine contract: only unexpected removals are considered endpoint flap. + if trigger_refill { + self.stats + .increment_me_endpoint_quarantine_unexpected_total(); + self.maybe_quarantine_flapping_endpoint(addr, uptime, "unexpected") + .await; + } else { + self.stats + .increment_me_endpoint_quarantine_draining_suppressed_total(); + debug!( + %addr, + uptime_ms = uptime.as_millis(), + "Skipping endpoint quarantine for draining writer removal" + ); + } + } + if trigger_refill && let Some(writer_dc) = removed_dc { + self.trigger_immediate_refill_for_dc(addr, writer_dc); + } + } + if removed { + self.notify_writer_epoch(); + } + removed + } + + pub(crate) async fn mark_writer_draining_with_timeout( + self: &Arc, + writer_id: u64, + timeout: Option, + allow_drain_fallback: bool, + ) { + let timeout = timeout.filter(|d| !d.is_zero()); + let found = { + let mut ws = self.writers.write().await; + if let Some(w) = ws.iter_mut().find(|w| w.id == writer_id) { + let already_draining = w.draining.swap(true, Ordering::Relaxed); + w.allow_drain_fallback + .store(allow_drain_fallback, Ordering::Relaxed); + let now_epoch_secs = Self::now_epoch_secs(); + w.draining_started_at_epoch_secs + .store(now_epoch_secs, Ordering::Relaxed); + let drain_deadline_epoch_secs = timeout + .map(|duration| now_epoch_secs.saturating_add(duration.as_secs())) + .unwrap_or(0); + w.drain_deadline_epoch_secs + .store(drain_deadline_epoch_secs, Ordering::Relaxed); + if !already_draining { + self.stats.increment_pool_drain_active(); + self.increment_draining_active_runtime(); + } + w.contour + .store(WriterContour::Draining.as_u8(), Ordering::Relaxed); + w.draining.store(true, Ordering::Relaxed); + true + } else { + false + } + }; + + if !found { + return; + } + + let timeout_secs = timeout.map(|d| d.as_secs()).unwrap_or(0); + debug!( + writer_id, + timeout_secs, allow_drain_fallback, "ME writer marked draining" + ); + } + + pub(crate) async fn mark_writer_draining(self: &Arc, writer_id: u64) { + self.mark_writer_draining_with_timeout(writer_id, Some(Duration::from_secs(300)), false) + .await; + } + + pub(in crate::transport::middle_proxy) fn writer_accepts_new_binding( + &self, + writer: &MeWriter, + ) -> bool { + if !writer.draining.load(Ordering::Relaxed) { + return true; + } + if !writer.allow_drain_fallback.load(Ordering::Relaxed) { + return false; + } + + match self.bind_stale_mode() { + MeBindStaleMode::Never => false, + MeBindStaleMode::Always => true, + MeBindStaleMode::Ttl => { + let ttl_secs = self + .binding_policy + .me_bind_stale_ttl_secs + .load(Ordering::Relaxed); + if ttl_secs == 0 { + return true; + } + + let started = writer + .draining_started_at_epoch_secs + .load(Ordering::Relaxed); + if started == 0 { + return false; + } + + Self::now_epoch_secs().saturating_sub(started) <= ttl_secs + } + } + } +} diff --git a/src/transport/middle_proxy/tests/health_adversarial_tests.rs b/src/transport/middle_proxy/tests/health_adversarial_tests.rs index e42151d..cafff0f 100644 --- a/src/transport/middle_proxy/tests/health_adversarial_tests.rs +++ b/src/transport/middle_proxy/tests/health_adversarial_tests.rs @@ -234,383 +234,9 @@ async fn set_writer_runtime_state( } } -#[tokio::test] -async fn reap_draining_writers_clears_warn_state_when_pool_empty() { - let (pool, _rng) = make_pool(128, 1, 1).await; - let mut warn_next_allowed = HashMap::new(); - warn_next_allowed.insert(11, Instant::now() + Duration::from_secs(5)); - warn_next_allowed.insert(22, Instant::now() + Duration::from_secs(5)); - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - - assert!(warn_next_allowed.is_empty()); -} - -#[tokio::test] -async fn reap_draining_writers_respects_threshold_across_multiple_overflow_cycles() { - let threshold = 3u64; - let (pool, _rng) = make_pool(threshold, 1, 1).await; - let now_epoch_secs = MePool::now_epoch_secs(); - - for writer_id in 1..=60u64 { - insert_draining_writer(&pool, writer_id, now_epoch_secs.saturating_sub(20), 1, 0).await; - } - - let mut warn_next_allowed = HashMap::new(); - for _ in 0..64 { - reap_draining_writers(&pool, &mut warn_next_allowed).await; - if writer_count(&pool).await <= threshold as usize { - break; - } - } - - assert_eq!(writer_count(&pool).await, threshold as usize); - assert_eq!(sorted_writer_ids(&pool).await, vec![1, 2, 3]); -} - -#[tokio::test] -async fn reap_draining_writers_handles_large_empty_writer_population() { - let (pool, _rng) = make_pool(128, 1, 1).await; - let now_epoch_secs = MePool::now_epoch_secs(); - let total = health_drain_close_budget() - .saturating_mul(3) - .saturating_add(27); - - for writer_id in 1..=total as u64 { - insert_draining_writer(&pool, writer_id, now_epoch_secs.saturating_sub(120), 0, 0).await; - } - - let mut warn_next_allowed = HashMap::new(); - for _ in 0..24 { - if writer_count(&pool).await == 0 { - break; - } - reap_draining_writers(&pool, &mut warn_next_allowed).await; - } - - assert_eq!(writer_count(&pool).await, 0); -} - -#[tokio::test] -async fn reap_draining_writers_processes_mass_deadline_expiry_without_unbounded_growth() { - let (pool, _rng) = make_pool(128, 1, 1).await; - let now_epoch_secs = MePool::now_epoch_secs(); - let total = health_drain_close_budget() - .saturating_mul(4) - .saturating_add(31); - - for writer_id in 1..=total as u64 { - insert_draining_writer( - &pool, - writer_id, - now_epoch_secs.saturating_sub(180), - 1, - now_epoch_secs.saturating_sub(1), - ) - .await; - } - - let mut warn_next_allowed = HashMap::new(); - for _ in 0..40 { - if writer_count(&pool).await == 0 { - break; - } - reap_draining_writers(&pool, &mut warn_next_allowed).await; - } - - assert_eq!(writer_count(&pool).await, 0); -} - -#[tokio::test] -async fn reap_draining_writers_maintains_warn_state_subset_property_under_bulk_churn() { - let (pool, _rng) = make_pool(128, 1, 1).await; - let now_epoch_secs = MePool::now_epoch_secs(); - let mut warn_next_allowed = HashMap::new(); - - for wave in 0..40u64 { - for offset in 0..8u64 { - insert_draining_writer( - &pool, - wave * 100 + offset, - now_epoch_secs.saturating_sub(400 + offset), - 1, - 0, - ) - .await; - } - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - assert!(warn_next_allowed.len() <= writer_count(&pool).await); - - let ids = sorted_writer_ids(&pool).await; - for writer_id in ids.into_iter().take(3) { - let _ = pool.remove_writer_and_close_clients(writer_id).await; - } - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - assert!(warn_next_allowed.len() <= writer_count(&pool).await); - } -} - -#[tokio::test] -async fn reap_draining_writers_budgeted_cleanup_never_increases_pool_size() { - let (pool, _rng) = make_pool(5, 1, 1).await; - let now_epoch_secs = MePool::now_epoch_secs(); - - for writer_id in 1..=200u64 { - insert_draining_writer( - &pool, - writer_id, - now_epoch_secs.saturating_sub(240).saturating_add(writer_id), - 1, - 0, - ) - .await; - } - - let mut warn_next_allowed = HashMap::new(); - let mut previous = writer_count(&pool).await; - for _ in 0..32 { - reap_draining_writers(&pool, &mut warn_next_allowed).await; - let current = writer_count(&pool).await; - assert!(current <= previous); - previous = current; - } -} - -#[tokio::test] -async fn me_health_monitor_converges_to_threshold_under_live_injection_churn() { - let threshold = 7u64; - let (pool, rng) = make_pool(threshold, 1, 1).await; - let now_epoch_secs = MePool::now_epoch_secs(); - - for writer_id in 1..=40u64 { - insert_draining_writer( - &pool, - writer_id, - now_epoch_secs.saturating_sub(300).saturating_add(writer_id), - 1, - 0, - ) - .await; - } - - let monitor = tokio::spawn(me_health_monitor(pool.clone(), rng, 0)); - - for wave in 0..8u64 { - for offset in 0..10u64 { - insert_draining_writer( - &pool, - 1000 + wave * 100 + offset, - now_epoch_secs.saturating_sub(120).saturating_add(offset), - 1, - 0, - ) - .await; - } - tokio::time::sleep(Duration::from_millis(5)).await; - } - - tokio::time::sleep(Duration::from_millis(120)).await; - monitor.abort(); - let _ = monitor.await; - - assert!(writer_count(&pool).await <= threshold as usize); -} - -#[tokio::test] -async fn me_health_monitor_drains_deadline_storm_with_budgeted_progress() { - let (pool, rng) = make_pool(128, 1, 1).await; - let now_epoch_secs = MePool::now_epoch_secs(); - - for writer_id in 1..=220u64 { - insert_draining_writer( - &pool, - writer_id, - now_epoch_secs.saturating_sub(120), - 1, - now_epoch_secs.saturating_sub(1), - ) - .await; - } - - let monitor = tokio::spawn(me_health_monitor(pool.clone(), rng, 0)); - tokio::time::sleep(Duration::from_millis(120)).await; - monitor.abort(); - let _ = monitor.await; - - assert_eq!(writer_count(&pool).await, 0); -} - -#[tokio::test] -async fn me_health_monitor_eliminates_mixed_empty_and_deadline_backlog() { - let threshold = 12u64; - let (pool, rng) = make_pool(threshold, 1, 1).await; - let now_epoch_secs = MePool::now_epoch_secs(); - - for writer_id in 1..=180u64 { - let bound_clients = if writer_id % 3 == 0 { 0 } else { 1 }; - let deadline = if writer_id % 2 == 0 { - now_epoch_secs.saturating_sub(1) - } else { - 0 - }; - insert_draining_writer( - &pool, - writer_id, - now_epoch_secs.saturating_sub(250).saturating_add(writer_id), - bound_clients, - deadline, - ) - .await; - } - - let monitor = tokio::spawn(me_health_monitor(pool.clone(), rng, 0)); - tokio::time::sleep(Duration::from_millis(140)).await; - monitor.abort(); - let _ = monitor.await; - - assert!(writer_count(&pool).await <= threshold as usize); -} - -#[tokio::test] -async fn reap_draining_writers_deterministic_mixed_state_churn_preserves_invariants() { - let threshold = 9u64; - let (pool, _rng) = make_pool(threshold, 1, 1).await; - let mut warn_next_allowed = HashMap::new(); - let mut seed = 0x9E37_79B9_7F4A_7C15u64; - let mut next_writer_id = 20_000u64; - let now_epoch_secs = MePool::now_epoch_secs(); - - for writer_id in 1..=72u64 { - let bound_clients = if writer_id % 4 == 0 { 0 } else { 1 }; - let deadline = if writer_id % 5 == 0 { - now_epoch_secs.saturating_sub(1) - } else { - 0 - }; - insert_draining_writer( - &pool, - writer_id, - now_epoch_secs.saturating_sub(500).saturating_add(writer_id), - bound_clients, - deadline, - ) - .await; - } - - for _round in 0..90 { - reap_draining_writers(&pool, &mut warn_next_allowed).await; - - let draining_ids = draining_writer_ids(&pool).await; - assert!( - warn_next_allowed.keys().all(|id| draining_ids.contains(id)), - "warn-state keys must always be a subset of live draining writers" - ); - - let writer_ids = sorted_writer_ids(&pool).await; - if writer_ids.is_empty() { - continue; - } - - let remove_n = (lcg_next(&mut seed) % 3) as usize; - for writer_id in writer_ids.iter().copied().take(remove_n) { - let _ = pool.remove_writer_and_close_clients(writer_id).await; - } - - let survivors = sorted_writer_ids(&pool).await; - if !survivors.is_empty() { - let idx = (lcg_next(&mut seed) as usize) % survivors.len(); - let target = survivors[idx]; - set_writer_runtime_state(&pool, target, false, 0, 0).await; - } - - let survivors = sorted_writer_ids(&pool).await; - if survivors.len() > 1 { - let idx = (lcg_next(&mut seed) as usize) % survivors.len(); - let target = survivors[idx]; - let expired_deadline = if lcg_next(&mut seed) & 1 == 0 { - now_epoch_secs.saturating_sub(1) - } else { - 0 - }; - set_writer_runtime_state( - &pool, - target, - true, - now_epoch_secs.saturating_sub(120), - expired_deadline, - ) - .await; - } - - let inject_n = (lcg_next(&mut seed) % 4) as usize; - for _ in 0..inject_n { - let bound_clients = if lcg_next(&mut seed) & 1 == 0 { 0 } else { 1 }; - let deadline = if lcg_next(&mut seed) & 1 == 0 { - now_epoch_secs.saturating_sub(1) - } else { - 0 - }; - insert_draining_writer( - &pool, - next_writer_id, - now_epoch_secs.saturating_sub(240), - bound_clients, - deadline, - ) - .await; - next_writer_id = next_writer_id.saturating_add(1); - } - } - - for _ in 0..64 { - reap_draining_writers(&pool, &mut warn_next_allowed).await; - if writer_count(&pool).await <= threshold as usize { - break; - } - } - - assert!(writer_count(&pool).await <= threshold as usize); - let draining_ids = draining_writer_ids(&pool).await; - assert!(warn_next_allowed.keys().all(|id| draining_ids.contains(id))); -} - -#[tokio::test] -async fn reap_draining_writers_repeated_draining_flips_never_leave_stale_warn_state() { - let (pool, _rng) = make_pool(64, 1, 1).await; - let now_epoch_secs = MePool::now_epoch_secs(); - - for writer_id in 1..=24u64 { - insert_draining_writer(&pool, writer_id, now_epoch_secs.saturating_sub(240), 1, 0).await; - } - - let mut warn_next_allowed = HashMap::new(); - for _round in 0..48u64 { - for writer_id in 1..=24u64 { - let draining = (writer_id + _round) % 3 != 0; - set_writer_runtime_state( - &pool, - writer_id, - draining, - now_epoch_secs.saturating_sub(120), - 0, - ) - .await; - } - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - - let draining_ids = draining_writer_ids(&pool).await; - assert!( - warn_next_allowed.keys().all(|id| draining_ids.contains(id)), - "warn-state map must not retain entries for writers outside draining set" - ); - } -} - -#[test] -fn health_drain_close_budget_is_within_expected_bounds() { - let budget = health_drain_close_budget(); - assert!((16..=256).contains(&budget)); -} +// Bulk and monitor-driven drain adversarial cases. +#[path = "health_adversarial_tests/bulk.rs"] +mod bulk; +// Deterministic churn and stale-state adversarial cases. +#[path = "health_adversarial_tests/churn.rs"] +mod churn; diff --git a/src/transport/middle_proxy/tests/health_adversarial_tests/bulk.rs b/src/transport/middle_proxy/tests/health_adversarial_tests/bulk.rs new file mode 100644 index 0000000..b27179b --- /dev/null +++ b/src/transport/middle_proxy/tests/health_adversarial_tests/bulk.rs @@ -0,0 +1,240 @@ +use super::*; + +#[tokio::test] +async fn reap_draining_writers_clears_warn_state_when_pool_empty() { + let (pool, _rng) = make_pool(128, 1, 1).await; + let mut warn_next_allowed = HashMap::new(); + warn_next_allowed.insert(11, Instant::now() + Duration::from_secs(5)); + warn_next_allowed.insert(22, Instant::now() + Duration::from_secs(5)); + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + + assert!(warn_next_allowed.is_empty()); +} + +#[tokio::test] +async fn reap_draining_writers_respects_threshold_across_multiple_overflow_cycles() { + let threshold = 3u64; + let (pool, _rng) = make_pool(threshold, 1, 1).await; + let now_epoch_secs = MePool::now_epoch_secs(); + + for writer_id in 1..=60u64 { + insert_draining_writer(&pool, writer_id, now_epoch_secs.saturating_sub(20), 1, 0).await; + } + + let mut warn_next_allowed = HashMap::new(); + for _ in 0..64 { + reap_draining_writers(&pool, &mut warn_next_allowed).await; + if writer_count(&pool).await <= threshold as usize { + break; + } + } + + assert_eq!(writer_count(&pool).await, threshold as usize); + assert_eq!(sorted_writer_ids(&pool).await, vec![1, 2, 3]); +} + +#[tokio::test] +async fn reap_draining_writers_handles_large_empty_writer_population() { + let (pool, _rng) = make_pool(128, 1, 1).await; + let now_epoch_secs = MePool::now_epoch_secs(); + let total = health_drain_close_budget() + .saturating_mul(3) + .saturating_add(27); + + for writer_id in 1..=total as u64 { + insert_draining_writer(&pool, writer_id, now_epoch_secs.saturating_sub(120), 0, 0).await; + } + + let mut warn_next_allowed = HashMap::new(); + for _ in 0..24 { + if writer_count(&pool).await == 0 { + break; + } + reap_draining_writers(&pool, &mut warn_next_allowed).await; + } + + assert_eq!(writer_count(&pool).await, 0); +} + +#[tokio::test] +async fn reap_draining_writers_processes_mass_deadline_expiry_without_unbounded_growth() { + let (pool, _rng) = make_pool(128, 1, 1).await; + let now_epoch_secs = MePool::now_epoch_secs(); + let total = health_drain_close_budget() + .saturating_mul(4) + .saturating_add(31); + + for writer_id in 1..=total as u64 { + insert_draining_writer( + &pool, + writer_id, + now_epoch_secs.saturating_sub(180), + 1, + now_epoch_secs.saturating_sub(1), + ) + .await; + } + + let mut warn_next_allowed = HashMap::new(); + for _ in 0..40 { + if writer_count(&pool).await == 0 { + break; + } + reap_draining_writers(&pool, &mut warn_next_allowed).await; + } + + assert_eq!(writer_count(&pool).await, 0); +} + +#[tokio::test] +async fn reap_draining_writers_maintains_warn_state_subset_property_under_bulk_churn() { + let (pool, _rng) = make_pool(128, 1, 1).await; + let now_epoch_secs = MePool::now_epoch_secs(); + let mut warn_next_allowed = HashMap::new(); + + for wave in 0..40u64 { + for offset in 0..8u64 { + insert_draining_writer( + &pool, + wave * 100 + offset, + now_epoch_secs.saturating_sub(400 + offset), + 1, + 0, + ) + .await; + } + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + assert!(warn_next_allowed.len() <= writer_count(&pool).await); + + let ids = sorted_writer_ids(&pool).await; + for writer_id in ids.into_iter().take(3) { + let _ = pool.remove_writer_and_close_clients(writer_id).await; + } + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + assert!(warn_next_allowed.len() <= writer_count(&pool).await); + } +} + +#[tokio::test] +async fn reap_draining_writers_budgeted_cleanup_never_increases_pool_size() { + let (pool, _rng) = make_pool(5, 1, 1).await; + let now_epoch_secs = MePool::now_epoch_secs(); + + for writer_id in 1..=200u64 { + insert_draining_writer( + &pool, + writer_id, + now_epoch_secs.saturating_sub(240).saturating_add(writer_id), + 1, + 0, + ) + .await; + } + + let mut warn_next_allowed = HashMap::new(); + let mut previous = writer_count(&pool).await; + for _ in 0..32 { + reap_draining_writers(&pool, &mut warn_next_allowed).await; + let current = writer_count(&pool).await; + assert!(current <= previous); + previous = current; + } +} + +#[tokio::test] +async fn me_health_monitor_converges_to_threshold_under_live_injection_churn() { + let threshold = 7u64; + let (pool, rng) = make_pool(threshold, 1, 1).await; + let now_epoch_secs = MePool::now_epoch_secs(); + + for writer_id in 1..=40u64 { + insert_draining_writer( + &pool, + writer_id, + now_epoch_secs.saturating_sub(300).saturating_add(writer_id), + 1, + 0, + ) + .await; + } + + let monitor = tokio::spawn(me_health_monitor(pool.clone(), rng, 0)); + + for wave in 0..8u64 { + for offset in 0..10u64 { + insert_draining_writer( + &pool, + 1000 + wave * 100 + offset, + now_epoch_secs.saturating_sub(120).saturating_add(offset), + 1, + 0, + ) + .await; + } + tokio::time::sleep(Duration::from_millis(5)).await; + } + + tokio::time::sleep(Duration::from_millis(120)).await; + monitor.abort(); + let _ = monitor.await; + + assert!(writer_count(&pool).await <= threshold as usize); +} + +#[tokio::test] +async fn me_health_monitor_drains_deadline_storm_with_budgeted_progress() { + let (pool, rng) = make_pool(128, 1, 1).await; + let now_epoch_secs = MePool::now_epoch_secs(); + + for writer_id in 1..=220u64 { + insert_draining_writer( + &pool, + writer_id, + now_epoch_secs.saturating_sub(120), + 1, + now_epoch_secs.saturating_sub(1), + ) + .await; + } + + let monitor = tokio::spawn(me_health_monitor(pool.clone(), rng, 0)); + tokio::time::sleep(Duration::from_millis(120)).await; + monitor.abort(); + let _ = monitor.await; + + assert_eq!(writer_count(&pool).await, 0); +} + +#[tokio::test] +async fn me_health_monitor_eliminates_mixed_empty_and_deadline_backlog() { + let threshold = 12u64; + let (pool, rng) = make_pool(threshold, 1, 1).await; + let now_epoch_secs = MePool::now_epoch_secs(); + + for writer_id in 1..=180u64 { + let bound_clients = if writer_id % 3 == 0 { 0 } else { 1 }; + let deadline = if writer_id % 2 == 0 { + now_epoch_secs.saturating_sub(1) + } else { + 0 + }; + insert_draining_writer( + &pool, + writer_id, + now_epoch_secs.saturating_sub(250).saturating_add(writer_id), + bound_clients, + deadline, + ) + .await; + } + + let monitor = tokio::spawn(me_health_monitor(pool.clone(), rng, 0)); + tokio::time::sleep(Duration::from_millis(140)).await; + monitor.abort(); + let _ = monitor.await; + + assert!(writer_count(&pool).await <= threshold as usize); +} diff --git a/src/transport/middle_proxy/tests/health_adversarial_tests/churn.rs b/src/transport/middle_proxy/tests/health_adversarial_tests/churn.rs new file mode 100644 index 0000000..e263b51 --- /dev/null +++ b/src/transport/middle_proxy/tests/health_adversarial_tests/churn.rs @@ -0,0 +1,143 @@ +use super::*; + +#[tokio::test] +async fn reap_draining_writers_deterministic_mixed_state_churn_preserves_invariants() { + let threshold = 9u64; + let (pool, _rng) = make_pool(threshold, 1, 1).await; + let mut warn_next_allowed = HashMap::new(); + let mut seed = 0x9E37_79B9_7F4A_7C15u64; + let mut next_writer_id = 20_000u64; + let now_epoch_secs = MePool::now_epoch_secs(); + + for writer_id in 1..=72u64 { + let bound_clients = if writer_id % 4 == 0 { 0 } else { 1 }; + let deadline = if writer_id % 5 == 0 { + now_epoch_secs.saturating_sub(1) + } else { + 0 + }; + insert_draining_writer( + &pool, + writer_id, + now_epoch_secs.saturating_sub(500).saturating_add(writer_id), + bound_clients, + deadline, + ) + .await; + } + + for _round in 0..90 { + reap_draining_writers(&pool, &mut warn_next_allowed).await; + + let draining_ids = draining_writer_ids(&pool).await; + assert!( + warn_next_allowed.keys().all(|id| draining_ids.contains(id)), + "warn-state keys must always be a subset of live draining writers" + ); + + let writer_ids = sorted_writer_ids(&pool).await; + if writer_ids.is_empty() { + continue; + } + + let remove_n = (lcg_next(&mut seed) % 3) as usize; + for writer_id in writer_ids.iter().copied().take(remove_n) { + let _ = pool.remove_writer_and_close_clients(writer_id).await; + } + + let survivors = sorted_writer_ids(&pool).await; + if !survivors.is_empty() { + let idx = (lcg_next(&mut seed) as usize) % survivors.len(); + let target = survivors[idx]; + set_writer_runtime_state(&pool, target, false, 0, 0).await; + } + + let survivors = sorted_writer_ids(&pool).await; + if survivors.len() > 1 { + let idx = (lcg_next(&mut seed) as usize) % survivors.len(); + let target = survivors[idx]; + let expired_deadline = if lcg_next(&mut seed) & 1 == 0 { + now_epoch_secs.saturating_sub(1) + } else { + 0 + }; + set_writer_runtime_state( + &pool, + target, + true, + now_epoch_secs.saturating_sub(120), + expired_deadline, + ) + .await; + } + + let inject_n = (lcg_next(&mut seed) % 4) as usize; + for _ in 0..inject_n { + let bound_clients = if lcg_next(&mut seed) & 1 == 0 { 0 } else { 1 }; + let deadline = if lcg_next(&mut seed) & 1 == 0 { + now_epoch_secs.saturating_sub(1) + } else { + 0 + }; + insert_draining_writer( + &pool, + next_writer_id, + now_epoch_secs.saturating_sub(240), + bound_clients, + deadline, + ) + .await; + next_writer_id = next_writer_id.saturating_add(1); + } + } + + for _ in 0..64 { + reap_draining_writers(&pool, &mut warn_next_allowed).await; + if writer_count(&pool).await <= threshold as usize { + break; + } + } + + assert!(writer_count(&pool).await <= threshold as usize); + let draining_ids = draining_writer_ids(&pool).await; + assert!(warn_next_allowed.keys().all(|id| draining_ids.contains(id))); +} + +#[tokio::test] +async fn reap_draining_writers_repeated_draining_flips_never_leave_stale_warn_state() { + let (pool, _rng) = make_pool(64, 1, 1).await; + let now_epoch_secs = MePool::now_epoch_secs(); + + for writer_id in 1..=24u64 { + insert_draining_writer(&pool, writer_id, now_epoch_secs.saturating_sub(240), 1, 0).await; + } + + let mut warn_next_allowed = HashMap::new(); + for _round in 0..48u64 { + for writer_id in 1..=24u64 { + let draining = (writer_id + _round) % 3 != 0; + set_writer_runtime_state( + &pool, + writer_id, + draining, + now_epoch_secs.saturating_sub(120), + 0, + ) + .await; + } + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + + let draining_ids = draining_writer_ids(&pool).await; + assert!( + warn_next_allowed.keys().all(|id| draining_ids.contains(id)), + "warn-state map must not retain entries for writers outside draining set" + ); + } +} + +#[test] +fn health_drain_close_budget_is_within_expected_bounds() { + let budget = health_drain_close_budget(); + assert!((16..=256).contains(&budget)); +} diff --git a/src/transport/middle_proxy/tests/health_regression_tests.rs b/src/transport/middle_proxy/tests/health_regression_tests.rs index 04c7773..b06e721 100644 --- a/src/transport/middle_proxy/tests/health_regression_tests.rs +++ b/src/transport/middle_proxy/tests/health_regression_tests.rs @@ -203,422 +203,9 @@ async fn set_writer_draining(pool: &Arc, writer_id: u64, draining: bool) } } -#[tokio::test] -async fn reap_draining_writers_drops_warn_state_for_removed_writer() { - let pool = make_pool(128).await; - let now_epoch_secs = MePool::now_epoch_secs(); - let conn_ids = insert_draining_writer( - &pool, - 7, - now_epoch_secs.saturating_sub(180), - 1, - now_epoch_secs.saturating_add(3_600), - ) - .await; - let mut warn_next_allowed = HashMap::new(); - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - assert!(warn_next_allowed.contains_key(&7)); - - let _ = pool.remove_writer_and_close_clients(7).await; - assert!(pool.registry.get_writer(conn_ids[0]).await.is_none()); - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - assert!(!warn_next_allowed.contains_key(&7)); -} - -#[tokio::test] -async fn reap_draining_writers_removes_empty_draining_writers() { - let pool = make_pool(128).await; - let now_epoch_secs = MePool::now_epoch_secs(); - insert_draining_writer(&pool, 1, now_epoch_secs.saturating_sub(40), 0, 0).await; - insert_draining_writer(&pool, 2, now_epoch_secs.saturating_sub(30), 0, 0).await; - insert_draining_writer(&pool, 3, now_epoch_secs.saturating_sub(20), 1, 0).await; - let mut warn_next_allowed = HashMap::new(); - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - - assert_eq!(current_writer_ids(&pool).await, vec![3]); -} - -#[tokio::test] -async fn reap_draining_writers_overflow_closes_oldest_non_empty_writers() { - let pool = make_pool(2).await; - let now_epoch_secs = MePool::now_epoch_secs(); - insert_draining_writer(&pool, 11, now_epoch_secs.saturating_sub(40), 1, 0).await; - insert_draining_writer(&pool, 22, now_epoch_secs.saturating_sub(30), 1, 0).await; - insert_draining_writer(&pool, 33, now_epoch_secs.saturating_sub(20), 1, 0).await; - insert_draining_writer(&pool, 44, now_epoch_secs.saturating_sub(10), 1, 0).await; - let mut warn_next_allowed = HashMap::new(); - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - - assert_eq!(current_writer_ids(&pool).await, vec![33, 44]); -} - -#[tokio::test] -async fn reap_draining_writers_deadline_force_close_applies_under_threshold() { - let pool = make_pool(128).await; - let now_epoch_secs = MePool::now_epoch_secs(); - insert_draining_writer( - &pool, - 50, - now_epoch_secs.saturating_sub(15), - 1, - now_epoch_secs.saturating_sub(1), - ) - .await; - let mut warn_next_allowed = HashMap::new(); - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - - assert!(current_writer_ids(&pool).await.is_empty()); -} - -#[tokio::test] -async fn reap_draining_writers_limits_closes_per_health_tick() { - let pool = make_pool(1).await; - let now_epoch_secs = MePool::now_epoch_secs(); - let close_budget = health_drain_close_budget(); - let writer_total = close_budget.saturating_add(20); - for writer_id in 1..=writer_total as u64 { - insert_draining_writer(&pool, writer_id, now_epoch_secs.saturating_sub(20), 1, 0).await; - } - let mut warn_next_allowed = HashMap::new(); - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - - assert_eq!(pool.writers.read().await.len(), writer_total - close_budget); -} - -#[tokio::test] -async fn reap_draining_writers_keeps_warn_state_for_deadline_backlog_writers() { - let pool = make_pool(0).await; - let now_epoch_secs = MePool::now_epoch_secs(); - let close_budget = health_drain_close_budget(); - let writer_total = close_budget.saturating_add(5); - for writer_id in 1..=writer_total as u64 { - insert_draining_writer( - &pool, - writer_id, - now_epoch_secs.saturating_sub(60), - 1, - now_epoch_secs.saturating_sub(1), - ) - .await; - } - let target_writer_id = writer_total as u64; - let mut warn_next_allowed = HashMap::new(); - warn_next_allowed.insert(target_writer_id, Instant::now() + Duration::from_secs(300)); - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - - assert!(writer_exists(&pool, target_writer_id).await); - assert!(warn_next_allowed.contains_key(&target_writer_id)); -} - -#[tokio::test] -async fn reap_draining_writers_keeps_warn_state_for_overflow_backlog_writers() { - let pool = make_pool(1).await; - let now_epoch_secs = MePool::now_epoch_secs(); - let close_budget = health_drain_close_budget(); - let writer_total = close_budget.saturating_add(6); - for writer_id in 1..=writer_total as u64 { - insert_draining_writer( - &pool, - writer_id, - now_epoch_secs.saturating_sub(300).saturating_add(writer_id), - 1, - 0, - ) - .await; - } - let target_writer_id = writer_total.saturating_sub(1) as u64; - let mut warn_next_allowed = HashMap::new(); - warn_next_allowed.insert(target_writer_id, Instant::now() + Duration::from_secs(300)); - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - - assert!(writer_exists(&pool, target_writer_id).await); - assert!(warn_next_allowed.contains_key(&target_writer_id)); -} - -#[tokio::test] -async fn reap_draining_writers_drops_warn_state_when_writer_exits_draining_state() { - let pool = make_pool(128).await; - let now_epoch_secs = MePool::now_epoch_secs(); - insert_draining_writer(&pool, 71, now_epoch_secs.saturating_sub(60), 1, 0).await; - - let mut warn_next_allowed = HashMap::new(); - warn_next_allowed.insert(71, Instant::now() + Duration::from_secs(300)); - - set_writer_draining(&pool, 71, false).await; - reap_draining_writers(&pool, &mut warn_next_allowed).await; - - assert!(writer_exists(&pool, 71).await); - assert!( - !warn_next_allowed.contains_key(&71), - "warn cooldown state must be dropped after writer leaves draining state" - ); -} - -#[tokio::test] -async fn reap_draining_writers_preserves_warn_state_across_multiple_budget_deferrals() { - let pool = make_pool(0).await; - let now_epoch_secs = MePool::now_epoch_secs(); - let close_budget = health_drain_close_budget(); - let writer_total = close_budget.saturating_mul(2).saturating_add(1); - for writer_id in 1..=writer_total as u64 { - insert_draining_writer( - &pool, - writer_id, - now_epoch_secs.saturating_sub(120), - 1, - now_epoch_secs.saturating_sub(1), - ) - .await; - } - - let tail_writer_id = writer_total as u64; - let mut warn_next_allowed = HashMap::new(); - warn_next_allowed.insert(tail_writer_id, Instant::now() + Duration::from_secs(300)); - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - assert!(writer_exists(&pool, tail_writer_id).await); - assert!(warn_next_allowed.contains_key(&tail_writer_id)); - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - assert!(writer_exists(&pool, tail_writer_id).await); - assert!(warn_next_allowed.contains_key(&tail_writer_id)); - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - assert!(!writer_exists(&pool, tail_writer_id).await); - assert!( - !warn_next_allowed.contains_key(&tail_writer_id), - "warn cooldown state must clear once writer is actually removed" - ); -} - -#[tokio::test] -async fn reap_draining_writers_backlog_drains_across_ticks() { - let pool = make_pool(128).await; - let now_epoch_secs = MePool::now_epoch_secs(); - let close_budget = health_drain_close_budget(); - let writer_total = close_budget.saturating_mul(2).saturating_add(7); - for writer_id in 1..=writer_total as u64 { - insert_draining_writer(&pool, writer_id, now_epoch_secs.saturating_sub(20), 0, 0).await; - } - let mut warn_next_allowed = HashMap::new(); - - for _ in 0..8 { - if pool.writers.read().await.is_empty() { - break; - } - reap_draining_writers(&pool, &mut warn_next_allowed).await; - } - - assert!(pool.writers.read().await.is_empty()); -} - -#[tokio::test] -async fn reap_draining_writers_threshold_backlog_converges_to_threshold() { - let threshold = 5u64; - let pool = make_pool(threshold).await; - let now_epoch_secs = MePool::now_epoch_secs(); - let close_budget = health_drain_close_budget(); - let writer_total = threshold as usize + close_budget.saturating_add(12); - for writer_id in 1..=writer_total as u64 { - insert_draining_writer(&pool, writer_id, now_epoch_secs.saturating_sub(20), 1, 0).await; - } - let mut warn_next_allowed = HashMap::new(); - - for _ in 0..16 { - reap_draining_writers(&pool, &mut warn_next_allowed).await; - if pool.writers.read().await.len() <= threshold as usize { - break; - } - } - - assert_eq!(pool.writers.read().await.len(), threshold as usize); -} - -#[tokio::test] -async fn reap_draining_writers_threshold_zero_preserves_non_expired_non_empty_writers() { - let pool = make_pool(0).await; - let now_epoch_secs = MePool::now_epoch_secs(); - insert_draining_writer(&pool, 10, now_epoch_secs.saturating_sub(40), 1, 0).await; - insert_draining_writer(&pool, 20, now_epoch_secs.saturating_sub(30), 1, 0).await; - insert_draining_writer(&pool, 30, now_epoch_secs.saturating_sub(20), 1, 0).await; - let mut warn_next_allowed = HashMap::new(); - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - - assert_eq!(current_writer_ids(&pool).await, vec![10, 20, 30]); -} - -#[tokio::test] -async fn reap_draining_writers_prioritizes_force_close_before_empty_cleanup() { - let pool = make_pool(1).await; - let now_epoch_secs = MePool::now_epoch_secs(); - let close_budget = health_drain_close_budget(); - for writer_id in 1..=close_budget.saturating_add(1) as u64 { - insert_draining_writer(&pool, writer_id, now_epoch_secs.saturating_sub(20), 1, 0).await; - } - let empty_writer_id = close_budget.saturating_add(2) as u64; - insert_draining_writer( - &pool, - empty_writer_id, - now_epoch_secs.saturating_sub(20), - 0, - 0, - ) - .await; - let mut warn_next_allowed = HashMap::new(); - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - - assert_eq!(current_writer_ids(&pool).await, vec![1, empty_writer_id]); -} - -#[tokio::test] -async fn reap_draining_writers_empty_cleanup_does_not_increment_force_close_metric() { - let pool = make_pool(128).await; - let now_epoch_secs = MePool::now_epoch_secs(); - insert_draining_writer(&pool, 1, now_epoch_secs.saturating_sub(60), 0, 0).await; - insert_draining_writer(&pool, 2, now_epoch_secs.saturating_sub(50), 0, 0).await; - let mut warn_next_allowed = HashMap::new(); - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - - assert!(current_writer_ids(&pool).await.is_empty()); - assert_eq!(pool.stats.get_pool_force_close_total(), 0); -} - -#[tokio::test] -async fn reap_draining_writers_handles_duplicate_force_close_requests_for_same_writer() { - let pool = make_pool(1).await; - let now_epoch_secs = MePool::now_epoch_secs(); - insert_draining_writer( - &pool, - 10, - now_epoch_secs.saturating_sub(30), - 1, - now_epoch_secs.saturating_sub(1), - ) - .await; - insert_draining_writer( - &pool, - 20, - now_epoch_secs.saturating_sub(20), - 1, - now_epoch_secs.saturating_sub(1), - ) - .await; - let mut warn_next_allowed = HashMap::new(); - - reap_draining_writers(&pool, &mut warn_next_allowed).await; - - assert!(current_writer_ids(&pool).await.is_empty()); -} - -#[tokio::test] -async fn reap_draining_writers_warn_state_never_exceeds_live_draining_population_under_churn() { - let pool = make_pool(128).await; - let now_epoch_secs = MePool::now_epoch_secs(); - let mut warn_next_allowed = HashMap::new(); - - for wave in 0..12u64 { - for offset in 0..9u64 { - insert_draining_writer( - &pool, - wave * 100 + offset, - now_epoch_secs.saturating_sub(120 + offset), - 1, - 0, - ) - .await; - } - reap_draining_writers(&pool, &mut warn_next_allowed).await; - assert!(warn_next_allowed.len() <= pool.writers.read().await.len()); - - let existing_writer_ids = current_writer_ids(&pool).await; - for writer_id in existing_writer_ids.into_iter().take(4) { - let _ = pool.remove_writer_and_close_clients(writer_id).await; - } - reap_draining_writers(&pool, &mut warn_next_allowed).await; - assert!(warn_next_allowed.len() <= pool.writers.read().await.len()); - } -} - -#[tokio::test] -async fn reap_draining_writers_mixed_backlog_converges_without_leaking_warn_state() { - let pool = make_pool(6).await; - let now_epoch_secs = MePool::now_epoch_secs(); - let mut warn_next_allowed = HashMap::new(); - - for writer_id in 1..=18u64 { - let bound_clients = if writer_id % 3 == 0 { 0 } else { 1 }; - let deadline = if writer_id % 2 == 0 { - now_epoch_secs.saturating_sub(1) - } else { - 0 - }; - insert_draining_writer( - &pool, - writer_id, - now_epoch_secs.saturating_sub(300).saturating_add(writer_id), - bound_clients, - deadline, - ) - .await; - } - - for _ in 0..16 { - reap_draining_writers(&pool, &mut warn_next_allowed).await; - if pool.writers.read().await.len() <= 6 { - break; - } - } - - assert!(pool.writers.read().await.len() <= 6); - assert!(warn_next_allowed.len() <= pool.writers.read().await.len()); -} - -#[test] -fn general_config_default_drain_threshold_remains_enabled() { - assert_eq!(GeneralConfig::default().me_pool_drain_threshold, 32); - assert!(GeneralConfig::default().me_pool_drain_soft_evict_enabled); - assert_eq!( - GeneralConfig::default().me_pool_drain_soft_evict_grace_secs, - 10 - ); - assert_eq!( - GeneralConfig::default().me_pool_drain_soft_evict_per_writer, - 2 - ); - assert_eq!( - GeneralConfig::default().me_pool_drain_soft_evict_budget_per_core, - 16 - ); - assert_eq!( - GeneralConfig::default().me_pool_drain_soft_evict_cooldown_ms, - 1000 - ); - assert_eq!( - GeneralConfig::default().me_bind_stale_mode, - MeBindStaleMode::Never - ); -} - -#[tokio::test] -async fn prune_closed_writers_closes_bound_clients_when_writer_is_non_empty() { - let pool = make_pool(128).await; - let now_epoch_secs = MePool::now_epoch_secs(); - let conn_ids = - insert_draining_writer(&pool, 910, now_epoch_secs.saturating_sub(60), 1, 0).await; - - pool.prune_closed_writers().await; - - assert!(!writer_exists(&pool, 910).await); - assert!(pool.registry.get_writer(conn_ids[0]).await.is_none()); -} +// Budgeted drain cleanup regression cases. +#[path = "health_regression_tests/budgeted.rs"] +mod budgeted; +// Backlog convergence and pruning regression cases. +#[path = "health_regression_tests/convergence.rs"] +mod convergence; diff --git a/src/transport/middle_proxy/tests/health_regression_tests/budgeted.rs b/src/transport/middle_proxy/tests/health_regression_tests/budgeted.rs new file mode 100644 index 0000000..93f0732 --- /dev/null +++ b/src/transport/middle_proxy/tests/health_regression_tests/budgeted.rs @@ -0,0 +1,240 @@ +use super::*; + +#[tokio::test] +async fn reap_draining_writers_drops_warn_state_for_removed_writer() { + let pool = make_pool(128).await; + let now_epoch_secs = MePool::now_epoch_secs(); + let conn_ids = insert_draining_writer( + &pool, + 7, + now_epoch_secs.saturating_sub(180), + 1, + now_epoch_secs.saturating_add(3_600), + ) + .await; + let mut warn_next_allowed = HashMap::new(); + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + assert!(warn_next_allowed.contains_key(&7)); + + let _ = pool.remove_writer_and_close_clients(7).await; + assert!(pool.registry.get_writer(conn_ids[0]).await.is_none()); + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + assert!(!warn_next_allowed.contains_key(&7)); +} + +#[tokio::test] +async fn reap_draining_writers_removes_empty_draining_writers() { + let pool = make_pool(128).await; + let now_epoch_secs = MePool::now_epoch_secs(); + insert_draining_writer(&pool, 1, now_epoch_secs.saturating_sub(40), 0, 0).await; + insert_draining_writer(&pool, 2, now_epoch_secs.saturating_sub(30), 0, 0).await; + insert_draining_writer(&pool, 3, now_epoch_secs.saturating_sub(20), 1, 0).await; + let mut warn_next_allowed = HashMap::new(); + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + + assert_eq!(current_writer_ids(&pool).await, vec![3]); +} + +#[tokio::test] +async fn reap_draining_writers_overflow_closes_oldest_non_empty_writers() { + let pool = make_pool(2).await; + let now_epoch_secs = MePool::now_epoch_secs(); + insert_draining_writer(&pool, 11, now_epoch_secs.saturating_sub(40), 1, 0).await; + insert_draining_writer(&pool, 22, now_epoch_secs.saturating_sub(30), 1, 0).await; + insert_draining_writer(&pool, 33, now_epoch_secs.saturating_sub(20), 1, 0).await; + insert_draining_writer(&pool, 44, now_epoch_secs.saturating_sub(10), 1, 0).await; + let mut warn_next_allowed = HashMap::new(); + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + + assert_eq!(current_writer_ids(&pool).await, vec![33, 44]); +} + +#[tokio::test] +async fn reap_draining_writers_deadline_force_close_applies_under_threshold() { + let pool = make_pool(128).await; + let now_epoch_secs = MePool::now_epoch_secs(); + insert_draining_writer( + &pool, + 50, + now_epoch_secs.saturating_sub(15), + 1, + now_epoch_secs.saturating_sub(1), + ) + .await; + let mut warn_next_allowed = HashMap::new(); + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + + assert!(current_writer_ids(&pool).await.is_empty()); +} + +#[tokio::test] +async fn reap_draining_writers_limits_closes_per_health_tick() { + let pool = make_pool(1).await; + let now_epoch_secs = MePool::now_epoch_secs(); + let close_budget = health_drain_close_budget(); + let writer_total = close_budget.saturating_add(20); + for writer_id in 1..=writer_total as u64 { + insert_draining_writer(&pool, writer_id, now_epoch_secs.saturating_sub(20), 1, 0).await; + } + let mut warn_next_allowed = HashMap::new(); + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + + assert_eq!(pool.writers.read().await.len(), writer_total - close_budget); +} + +#[tokio::test] +async fn reap_draining_writers_keeps_warn_state_for_deadline_backlog_writers() { + let pool = make_pool(0).await; + let now_epoch_secs = MePool::now_epoch_secs(); + let close_budget = health_drain_close_budget(); + let writer_total = close_budget.saturating_add(5); + for writer_id in 1..=writer_total as u64 { + insert_draining_writer( + &pool, + writer_id, + now_epoch_secs.saturating_sub(60), + 1, + now_epoch_secs.saturating_sub(1), + ) + .await; + } + let target_writer_id = writer_total as u64; + let mut warn_next_allowed = HashMap::new(); + warn_next_allowed.insert(target_writer_id, Instant::now() + Duration::from_secs(300)); + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + + assert!(writer_exists(&pool, target_writer_id).await); + assert!(warn_next_allowed.contains_key(&target_writer_id)); +} + +#[tokio::test] +async fn reap_draining_writers_keeps_warn_state_for_overflow_backlog_writers() { + let pool = make_pool(1).await; + let now_epoch_secs = MePool::now_epoch_secs(); + let close_budget = health_drain_close_budget(); + let writer_total = close_budget.saturating_add(6); + for writer_id in 1..=writer_total as u64 { + insert_draining_writer( + &pool, + writer_id, + now_epoch_secs.saturating_sub(300).saturating_add(writer_id), + 1, + 0, + ) + .await; + } + let target_writer_id = writer_total.saturating_sub(1) as u64; + let mut warn_next_allowed = HashMap::new(); + warn_next_allowed.insert(target_writer_id, Instant::now() + Duration::from_secs(300)); + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + + assert!(writer_exists(&pool, target_writer_id).await); + assert!(warn_next_allowed.contains_key(&target_writer_id)); +} + +#[tokio::test] +async fn reap_draining_writers_drops_warn_state_when_writer_exits_draining_state() { + let pool = make_pool(128).await; + let now_epoch_secs = MePool::now_epoch_secs(); + insert_draining_writer(&pool, 71, now_epoch_secs.saturating_sub(60), 1, 0).await; + + let mut warn_next_allowed = HashMap::new(); + warn_next_allowed.insert(71, Instant::now() + Duration::from_secs(300)); + + set_writer_draining(&pool, 71, false).await; + reap_draining_writers(&pool, &mut warn_next_allowed).await; + + assert!(writer_exists(&pool, 71).await); + assert!( + !warn_next_allowed.contains_key(&71), + "warn cooldown state must be dropped after writer leaves draining state" + ); +} + +#[tokio::test] +async fn reap_draining_writers_preserves_warn_state_across_multiple_budget_deferrals() { + let pool = make_pool(0).await; + let now_epoch_secs = MePool::now_epoch_secs(); + let close_budget = health_drain_close_budget(); + let writer_total = close_budget.saturating_mul(2).saturating_add(1); + for writer_id in 1..=writer_total as u64 { + insert_draining_writer( + &pool, + writer_id, + now_epoch_secs.saturating_sub(120), + 1, + now_epoch_secs.saturating_sub(1), + ) + .await; + } + + let tail_writer_id = writer_total as u64; + let mut warn_next_allowed = HashMap::new(); + warn_next_allowed.insert(tail_writer_id, Instant::now() + Duration::from_secs(300)); + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + assert!(writer_exists(&pool, tail_writer_id).await); + assert!(warn_next_allowed.contains_key(&tail_writer_id)); + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + assert!(writer_exists(&pool, tail_writer_id).await); + assert!(warn_next_allowed.contains_key(&tail_writer_id)); + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + assert!(!writer_exists(&pool, tail_writer_id).await); + assert!( + !warn_next_allowed.contains_key(&tail_writer_id), + "warn cooldown state must clear once writer is actually removed" + ); +} + +#[tokio::test] +async fn reap_draining_writers_backlog_drains_across_ticks() { + let pool = make_pool(128).await; + let now_epoch_secs = MePool::now_epoch_secs(); + let close_budget = health_drain_close_budget(); + let writer_total = close_budget.saturating_mul(2).saturating_add(7); + for writer_id in 1..=writer_total as u64 { + insert_draining_writer(&pool, writer_id, now_epoch_secs.saturating_sub(20), 0, 0).await; + } + let mut warn_next_allowed = HashMap::new(); + + for _ in 0..8 { + if pool.writers.read().await.is_empty() { + break; + } + reap_draining_writers(&pool, &mut warn_next_allowed).await; + } + + assert!(pool.writers.read().await.is_empty()); +} + +#[tokio::test] +async fn reap_draining_writers_threshold_backlog_converges_to_threshold() { + let threshold = 5u64; + let pool = make_pool(threshold).await; + let now_epoch_secs = MePool::now_epoch_secs(); + let close_budget = health_drain_close_budget(); + let writer_total = threshold as usize + close_budget.saturating_add(12); + for writer_id in 1..=writer_total as u64 { + insert_draining_writer(&pool, writer_id, now_epoch_secs.saturating_sub(20), 1, 0).await; + } + let mut warn_next_allowed = HashMap::new(); + + for _ in 0..16 { + reap_draining_writers(&pool, &mut warn_next_allowed).await; + if pool.writers.read().await.len() <= threshold as usize { + break; + } + } + + assert_eq!(pool.writers.read().await.len(), threshold as usize); +} diff --git a/src/transport/middle_proxy/tests/health_regression_tests/convergence.rs b/src/transport/middle_proxy/tests/health_regression_tests/convergence.rs new file mode 100644 index 0000000..513d958 --- /dev/null +++ b/src/transport/middle_proxy/tests/health_regression_tests/convergence.rs @@ -0,0 +1,182 @@ +use super::*; + +#[tokio::test] +async fn reap_draining_writers_threshold_zero_preserves_non_expired_non_empty_writers() { + let pool = make_pool(0).await; + let now_epoch_secs = MePool::now_epoch_secs(); + insert_draining_writer(&pool, 10, now_epoch_secs.saturating_sub(40), 1, 0).await; + insert_draining_writer(&pool, 20, now_epoch_secs.saturating_sub(30), 1, 0).await; + insert_draining_writer(&pool, 30, now_epoch_secs.saturating_sub(20), 1, 0).await; + let mut warn_next_allowed = HashMap::new(); + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + + assert_eq!(current_writer_ids(&pool).await, vec![10, 20, 30]); +} + +#[tokio::test] +async fn reap_draining_writers_prioritizes_force_close_before_empty_cleanup() { + let pool = make_pool(1).await; + let now_epoch_secs = MePool::now_epoch_secs(); + let close_budget = health_drain_close_budget(); + for writer_id in 1..=close_budget.saturating_add(1) as u64 { + insert_draining_writer(&pool, writer_id, now_epoch_secs.saturating_sub(20), 1, 0).await; + } + let empty_writer_id = close_budget.saturating_add(2) as u64; + insert_draining_writer( + &pool, + empty_writer_id, + now_epoch_secs.saturating_sub(20), + 0, + 0, + ) + .await; + let mut warn_next_allowed = HashMap::new(); + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + + assert_eq!(current_writer_ids(&pool).await, vec![1, empty_writer_id]); +} + +#[tokio::test] +async fn reap_draining_writers_empty_cleanup_does_not_increment_force_close_metric() { + let pool = make_pool(128).await; + let now_epoch_secs = MePool::now_epoch_secs(); + insert_draining_writer(&pool, 1, now_epoch_secs.saturating_sub(60), 0, 0).await; + insert_draining_writer(&pool, 2, now_epoch_secs.saturating_sub(50), 0, 0).await; + let mut warn_next_allowed = HashMap::new(); + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + + assert!(current_writer_ids(&pool).await.is_empty()); + assert_eq!(pool.stats.get_pool_force_close_total(), 0); +} + +#[tokio::test] +async fn reap_draining_writers_handles_duplicate_force_close_requests_for_same_writer() { + let pool = make_pool(1).await; + let now_epoch_secs = MePool::now_epoch_secs(); + insert_draining_writer( + &pool, + 10, + now_epoch_secs.saturating_sub(30), + 1, + now_epoch_secs.saturating_sub(1), + ) + .await; + insert_draining_writer( + &pool, + 20, + now_epoch_secs.saturating_sub(20), + 1, + now_epoch_secs.saturating_sub(1), + ) + .await; + let mut warn_next_allowed = HashMap::new(); + + reap_draining_writers(&pool, &mut warn_next_allowed).await; + + assert!(current_writer_ids(&pool).await.is_empty()); +} + +#[tokio::test] +async fn reap_draining_writers_warn_state_never_exceeds_live_draining_population_under_churn() { + let pool = make_pool(128).await; + let now_epoch_secs = MePool::now_epoch_secs(); + let mut warn_next_allowed = HashMap::new(); + + for wave in 0..12u64 { + for offset in 0..9u64 { + insert_draining_writer( + &pool, + wave * 100 + offset, + now_epoch_secs.saturating_sub(120 + offset), + 1, + 0, + ) + .await; + } + reap_draining_writers(&pool, &mut warn_next_allowed).await; + assert!(warn_next_allowed.len() <= pool.writers.read().await.len()); + + let existing_writer_ids = current_writer_ids(&pool).await; + for writer_id in existing_writer_ids.into_iter().take(4) { + let _ = pool.remove_writer_and_close_clients(writer_id).await; + } + reap_draining_writers(&pool, &mut warn_next_allowed).await; + assert!(warn_next_allowed.len() <= pool.writers.read().await.len()); + } +} + +#[tokio::test] +async fn reap_draining_writers_mixed_backlog_converges_without_leaking_warn_state() { + let pool = make_pool(6).await; + let now_epoch_secs = MePool::now_epoch_secs(); + let mut warn_next_allowed = HashMap::new(); + + for writer_id in 1..=18u64 { + let bound_clients = if writer_id % 3 == 0 { 0 } else { 1 }; + let deadline = if writer_id % 2 == 0 { + now_epoch_secs.saturating_sub(1) + } else { + 0 + }; + insert_draining_writer( + &pool, + writer_id, + now_epoch_secs.saturating_sub(300).saturating_add(writer_id), + bound_clients, + deadline, + ) + .await; + } + + for _ in 0..16 { + reap_draining_writers(&pool, &mut warn_next_allowed).await; + if pool.writers.read().await.len() <= 6 { + break; + } + } + + assert!(pool.writers.read().await.len() <= 6); + assert!(warn_next_allowed.len() <= pool.writers.read().await.len()); +} + +#[test] +fn general_config_default_drain_threshold_remains_enabled() { + assert_eq!(GeneralConfig::default().me_pool_drain_threshold, 32); + assert!(GeneralConfig::default().me_pool_drain_soft_evict_enabled); + assert_eq!( + GeneralConfig::default().me_pool_drain_soft_evict_grace_secs, + 10 + ); + assert_eq!( + GeneralConfig::default().me_pool_drain_soft_evict_per_writer, + 2 + ); + assert_eq!( + GeneralConfig::default().me_pool_drain_soft_evict_budget_per_core, + 16 + ); + assert_eq!( + GeneralConfig::default().me_pool_drain_soft_evict_cooldown_ms, + 1000 + ); + assert_eq!( + GeneralConfig::default().me_bind_stale_mode, + MeBindStaleMode::Never + ); +} + +#[tokio::test] +async fn prune_closed_writers_closes_bound_clients_when_writer_is_non_empty() { + let pool = make_pool(128).await; + let now_epoch_secs = MePool::now_epoch_secs(); + let conn_ids = + insert_draining_writer(&pool, 910, now_epoch_secs.saturating_sub(60), 1, 0).await; + + pool.prune_closed_writers().await; + + assert!(!writer_exists(&pool, 910).await); + assert!(pool.registry.get_writer(conn_ids[0]).await.is_none()); +} diff --git a/src/transport/upstream.rs b/src/transport/upstream.rs index f899b7e..49141b8 100644 --- a/src/transport/upstream.rs +++ b/src/transport/upstream.rs @@ -337,1955 +337,15 @@ pub struct UpstreamManager { dns_resolver: Arc, } -impl UpstreamManager { - fn is_unscoped_upstream(upstream: &UpstreamConfig) -> bool { - upstream.scopes.is_empty() - } - - fn should_check_in_default_dc_connectivity( - has_unscoped: bool, - upstream: &UpstreamConfig, - ) -> bool { - !has_unscoped || Self::is_unscoped_upstream(upstream) - } - - pub fn new( - configs: Vec, - connect_retry_attempts: u32, - connect_retry_backoff_ms: u64, - connect_budget_ms: u64, - tg_connect_timeout_secs: u64, - unhealthy_fail_threshold: u32, - connect_failfast_hard_errors: bool, - stats: Arc, - ) -> Self { - let states = configs - .into_iter() - .filter(|c| c.enabled) - .map(UpstreamState::new) - .collect(); - - Self { - upstreams: Arc::new(RwLock::new(states)), - connect_retry_attempts: connect_retry_attempts.max(1), - connect_retry_backoff: Duration::from_millis(connect_retry_backoff_ms), - connect_budget: Duration::from_millis(connect_budget_ms.max(1)), - tg_connect_timeout_secs: tg_connect_timeout_secs.max(1), - unhealthy_fail_threshold: unhealthy_fail_threshold.max(1), - connect_failfast_hard_errors, - no_upstreams_warn_epoch_ms: Arc::new(AtomicU64::new(0)), - no_healthy_warn_epoch_ms: Arc::new(AtomicU64::new(0)), - stats, - dns_resolver: Arc::new(GenerationDnsResolver::default()), - } - } - - pub(crate) fn with_dns_overrides(self, entries: &[String]) -> Result { - self.dns_resolver.apply_entries(entries)?; - Ok(self) - } - - pub(crate) fn update_dns_overrides(&self, entries: &[String]) -> Result<()> { - self.dns_resolver.apply_entries(entries) - } - - pub(crate) fn dns_resolver(&self) -> Arc { - Arc::clone(&self.dns_resolver) - } - - pub(crate) async fn resolve_all(&self, host: &str, port: u16) -> Result> { - if let Some(addr) = self.dns_resolver.resolve_socket_addr(host, port) { - return Ok(vec![addr]); - } - let addrs = tokio::net::lookup_host((host, port)) - .await - .map_err(ProxyError::Io)? - .take(DNS_RESULT_MAX_ADDRESSES) - .collect::>(); - if addrs.is_empty() { - return Err(ProxyError::Proxy(format!( - "DNS returned no addresses for {host}:{port}" - ))); - } - Ok(addrs) - } - - pub(crate) async fn resolve_hostname(&self, host: &str, port: u16) -> Result { - let addrs = self.resolve_all(host, port).await?; - if let Some(addr) = addrs.iter().copied().find(SocketAddr::is_ipv4) { - return Ok(addr); - } - addrs.first().copied().ok_or_else(|| { - ProxyError::Proxy(format!("DNS returned no addresses for {host}:{port}")) - }) - } - - fn now_epoch_ms() -> u64 { - std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_millis() as u64 - } - - fn should_emit_warn(last_epoch_ms: &AtomicU64, cooldown_ms: u64) -> bool { - let now_epoch_ms = Self::now_epoch_ms(); - let previous_epoch_ms = last_epoch_ms.load(Ordering::Relaxed); - if now_epoch_ms.saturating_sub(previous_epoch_ms) < cooldown_ms { - return false; - } - last_epoch_ms - .compare_exchange( - previous_epoch_ms, - now_epoch_ms, - Ordering::AcqRel, - Ordering::Relaxed, - ) - .is_ok() - } - - pub fn try_api_snapshot(&self) -> Option { - let guard = self.upstreams.try_read().ok()?; - let now = std::time::Instant::now(); - - let mut summary = UpstreamApiSummarySnapshot { - configured_total: guard.len(), - ..UpstreamApiSummarySnapshot::default() - }; - let mut upstreams = Vec::with_capacity(guard.len()); - - for (idx, upstream) in guard.iter().enumerate() { - if upstream.healthy { - summary.healthy_total += 1; - } else { - summary.unhealthy_total += 1; - } - - let (route_kind, address) = Self::describe_upstream(&upstream.config.upstream_type); - match route_kind { - UpstreamRouteKind::Direct => summary.direct_total += 1, - UpstreamRouteKind::Socks4 => summary.socks4_total += 1, - UpstreamRouteKind::Socks5 => summary.socks5_total += 1, - UpstreamRouteKind::Shadowsocks => summary.shadowsocks_total += 1, - } - - let mut dc = Vec::with_capacity(NUM_DCS); - for dc_idx in 0..NUM_DCS { - dc.push(UpstreamApiDcSnapshot { - dc: (dc_idx + 1) as i16, - latency_ema_ms: upstream.dc_latency[dc_idx].get(), - ip_preference: upstream.dc_ip_pref[dc_idx], - }); - } - - upstreams.push(UpstreamApiItemSnapshot { - upstream_id: idx, - route_kind, - address, - weight: upstream.config.weight, - scopes: upstream.config.scopes.clone(), - healthy: upstream.healthy, - fails: upstream.fails, - last_check_age_secs: now.saturating_duration_since(upstream.last_check).as_secs(), - effective_latency_ms: upstream.effective_latency(None), - dc, - }); - } - - Some(UpstreamApiSnapshot { summary, upstreams }) - } - - pub async fn api_health_summary(&self) -> UpstreamApiHealthSummary { - let guard = self.upstreams.read().await; - let mut summary = UpstreamApiHealthSummary { - configured_total: guard.len(), - healthy_total: 0, - }; - for upstream in guard.iter() { - if upstream.healthy { - summary.healthy_total += 1; - } - } - summary - } - - fn describe_upstream(upstream_type: &UpstreamType) -> (UpstreamRouteKind, String) { - match upstream_type { - UpstreamType::Direct { .. } => (UpstreamRouteKind::Direct, "direct".to_string()), - UpstreamType::Socks4 { address, .. } => (UpstreamRouteKind::Socks4, address.clone()), - UpstreamType::Socks5 { address, .. } => (UpstreamRouteKind::Socks5, address.clone()), - UpstreamType::Shadowsocks { url, .. } => ( - UpstreamRouteKind::Shadowsocks, - sanitize_shadowsocks_url(url).unwrap_or_else(|_| "invalid".to_string()), - ), - } - } - - pub fn api_policy_snapshot(&self) -> UpstreamApiPolicySnapshot { - UpstreamApiPolicySnapshot { - connect_retry_attempts: self.connect_retry_attempts, - connect_retry_backoff_ms: self.connect_retry_backoff.as_millis() as u64, - connect_budget_ms: self.connect_budget.as_millis() as u64, - unhealthy_fail_threshold: self.unhealthy_fail_threshold, - connect_failfast_hard_errors: self.connect_failfast_hard_errors, - } - } - - fn resolve_probe_dc_families( - upstream: &UpstreamConfig, - ipv4_available: bool, - ipv6_available: bool, - ) -> (bool, bool) { - ( - upstream.ipv4.unwrap_or(ipv4_available), - upstream.ipv6.unwrap_or(ipv6_available), - ) - } - - fn resolve_runtime_dc_families( - upstream: &UpstreamConfig, - dc_preference: IpPreference, - ) -> (bool, bool) { - let (auto_ipv4, auto_ipv6) = match dc_preference { - IpPreference::PreferV4 => (true, false), - IpPreference::PreferV6 => (false, true), - IpPreference::BothWork | IpPreference::Unknown | IpPreference::Unavailable => { - (true, true) - } - }; - - ( - upstream.ipv4.unwrap_or(auto_ipv4), - upstream.ipv6.unwrap_or(auto_ipv6), - ) - } - - fn dc_table_addr(dc_idx: i16, ipv6: bool, port: u16) -> Option { - let arr_idx = UpstreamState::dc_array_idx(dc_idx)?; - let ip = if ipv6 { - TG_DATACENTERS_V6[arr_idx] - } else { - TG_DATACENTERS_V4[arr_idx] - }; - Some(SocketAddr::new(ip, port)) - } - - fn resolve_runtime_dc_target( - target: SocketAddr, - dc_idx: Option, - upstream: &UpstreamConfig, - dc_preference: IpPreference, - ) -> Result { - let (allow_ipv4, allow_ipv6) = Self::resolve_runtime_dc_families(upstream, dc_preference); - let preferred_ipv6 = match dc_preference { - IpPreference::PreferV6 => Some(true), - IpPreference::PreferV4 => Some(false), - IpPreference::BothWork | IpPreference::Unknown | IpPreference::Unavailable => { - upstream.prefer.map(|prefer| prefer == 6) - } - }; - if let Some(preferred_ipv6) = preferred_ipv6 - && target.is_ipv6() != preferred_ipv6 - { - let preferred_allowed = if preferred_ipv6 { - allow_ipv6 - } else { - allow_ipv4 - }; - if preferred_allowed { - if let Some(dc_idx) = dc_idx - && let Some(remapped) = - Self::dc_table_addr(dc_idx, preferred_ipv6, target.port()) - { - return Ok(remapped); - } - } - } - - if (target.is_ipv4() && allow_ipv4) || (target.is_ipv6() && allow_ipv6) { - return Ok(target); - } - - if !allow_ipv4 && !allow_ipv6 { - return Err(ProxyError::Config(format!( - "Upstream DC family policy blocks all families for target {target}" - ))); - } - - let Some(dc_idx) = dc_idx else { - return Err(ProxyError::Config(format!( - "Upstream DC family policy cannot remap target {target} without dc_idx" - ))); - }; - - let remapped = if target.is_ipv4() { - if allow_ipv6 { - Self::dc_table_addr(dc_idx, true, target.port()) - } else { - None - } - } else if allow_ipv4 { - Self::dc_table_addr(dc_idx, false, target.port()) - } else { - None - }; - - remapped.ok_or_else(|| { - ProxyError::Config(format!( - "Upstream DC family policy rejected target {target} (dc_idx={dc_idx})" - )) - }) - } - - #[cfg(unix)] - fn resolve_interface_addrs(name: &str, want_ipv6: bool) -> Vec { - use nix::ifaddrs::getifaddrs; - - let mut out = Vec::new(); - if let Ok(addrs) = getifaddrs() { - for iface in addrs { - if iface.interface_name != name { - continue; - } - if let Some(address) = iface.address { - if let Some(v4) = address.as_sockaddr_in() { - if !want_ipv6 { - out.push(IpAddr::V4(v4.ip())); - } - } else if let Some(v6) = address.as_sockaddr_in6() - && want_ipv6 - { - out.push(IpAddr::V6(v6.ip())); - } - } - } - } - out.sort_unstable(); - out.dedup(); - out - } - - pub(crate) fn resolve_bind_address( - interface: &Option, - bind_addresses: &Option>, - target: SocketAddr, - rr: Option<&AtomicUsize>, - validate_ip_on_interface: bool, - ) -> Option { - let want_ipv6 = target.is_ipv6(); - - if let Some(addrs) = bind_addresses.as_ref().filter(|v| !v.is_empty()) { - let mut candidates: Vec = addrs - .iter() - .filter_map(|s| s.parse::().ok()) - .filter(|ip| ip.is_ipv6() == want_ipv6) - .collect(); - - // Explicit bind IP has strict priority over interface auto-selection. - if validate_ip_on_interface - && let Some(iface) = interface - && iface.parse::().is_err() - { - #[cfg(unix)] - { - let iface_addrs = Self::resolve_interface_addrs(iface, want_ipv6); - if !iface_addrs.is_empty() { - candidates.retain(|ip| { - let ok = iface_addrs.contains(ip); - if !ok { - warn!( - interface = %iface, - bind_ip = %ip, - target = %target, - "Configured bind address is not assigned to interface" - ); - } - ok - }); - } else if !candidates.is_empty() { - warn!( - interface = %iface, - target = %target, - "Configured interface has no addresses for target family" - ); - candidates.clear(); - } - } - } - - if !candidates.is_empty() { - if let Some(counter) = rr { - let idx = counter.fetch_add(1, Ordering::Relaxed) % candidates.len(); - return Some(candidates[idx]); - } - return candidates.first().copied(); - } - - if validate_ip_on_interface - && interface - .as_ref() - .is_some_and(|iface| iface.parse::().is_err()) - { - warn!( - interface = interface.as_deref().unwrap_or(""), - target = %target, - "No valid bind_addresses left for interface" - ); - } - - return None; - } - - if let Some(iface) = interface { - if let Ok(ip) = iface.parse::() { - if ip.is_ipv6() == want_ipv6 { - return Some(ip); - } - } else { - #[cfg(unix)] - if let Some(ip) = resolve_interface_ip(iface, want_ipv6) { - return Some(ip); - } - } - } - - None - } - - async fn connect_hostname_with_dns_override( - &self, - address: &str, - connect_timeout: Duration, - ) -> Result { - if let Some((host, port)) = split_host_port(address) - && let Some(addr) = self.dns_resolver.resolve_socket_addr(&host, port) - { - return match tokio::time::timeout(connect_timeout, TcpStream::connect(addr)).await { - Ok(Ok(stream)) => Ok(stream), - Ok(Err(e)) => Err(ProxyError::Io(e)), - Err(_) => Err(ProxyError::ConnectionTimeout { - addr: addr.to_string(), - }), - }; - } - - match tokio::time::timeout(connect_timeout, TcpStream::connect(address)).await { - Ok(Ok(stream)) => Ok(stream), - Ok(Err(e)) => Err(ProxyError::Io(e)), - Err(_) => Err(ProxyError::ConnectionTimeout { - addr: address.to_string(), - }), - } - } - - fn retry_backoff_with_jitter(&self) -> Duration { - if self.connect_retry_backoff.is_zero() { - return Duration::ZERO; - } - let base_ms = self.connect_retry_backoff.as_millis() as u64; - if base_ms == 0 { - return self.connect_retry_backoff; - } - let jitter_cap_ms = (base_ms / 2).max(1); - let jitter_ms = rand::rng().random_range(0..=jitter_cap_ms); - Duration::from_millis(base_ms.saturating_add(jitter_ms)) - } - - fn is_hard_connect_error(error: &ProxyError) -> bool { - match error { - ProxyError::Config(_) | ProxyError::ConnectionRefused { .. } => true, - ProxyError::Io(ioe) => matches!( - ioe.kind(), - std::io::ErrorKind::ConnectionRefused - | std::io::ErrorKind::AddrInUse - | std::io::ErrorKind::AddrNotAvailable - | std::io::ErrorKind::InvalidInput - | std::io::ErrorKind::Unsupported - ), - _ => false, - } - } - - /// Select upstream using latency-weighted random selection. - async fn select_upstream(&self, dc_idx: Option, scope: Option<&str>) -> Option { - let upstreams = self.upstreams.read().await; - if upstreams.is_empty() { - return None; - } - // Scope filter: - // If scope is set: only scoped and matched items - // If scope is not set: only unscoped items - let filtered_upstreams: Vec = upstreams - .iter() - .enumerate() - .filter(|(_, u)| { - scope.map_or(u.config.scopes.is_empty(), |req_scope| { - u.config - .scopes - .split(',') - .map(str::trim) - .any(|s| s == req_scope) - }) - }) - .map(|(i, _)| i) - .collect(); - - // Healthy filter - let healthy: Vec = filtered_upstreams - .iter() - .filter(|&&i| upstreams[i].healthy) - .copied() - .collect(); - - if filtered_upstreams.is_empty() { - if Self::should_emit_warn(self.no_upstreams_warn_epoch_ms.as_ref(), 5_000) { - warn!( - scope = scope, - "No upstreams available! Using first (direct?)" - ); - } - return None; - } - - if healthy.is_empty() { - if Self::should_emit_warn(self.no_healthy_warn_epoch_ms.as_ref(), 5_000) { - warn!( - scope = scope, - "No healthy upstreams available! Using random." - ); - } - return Some(filtered_upstreams[rand::rng().random_range(0..filtered_upstreams.len())]); - } - - if healthy.len() == 1 { - return Some(healthy[0]); - } - - let weights: Vec<(usize, f64)> = healthy - .iter() - .map(|&i| { - let base = upstreams[i].config.weight as f64; - let latency_factor = upstreams[i] - .effective_latency(dc_idx) - .map(|ms| if ms > 1.0 { 1000.0 / ms } else { 1000.0 }) - .unwrap_or(1.0); - - (i, base * latency_factor) - }) - .collect(); - - let total: f64 = weights.iter().map(|(_, w)| w).sum(); - - if total <= 0.0 { - return Some(healthy[rand::rng().random_range(0..healthy.len())]); - } - - let mut choice: f64 = rand::rng().random_range(0.0..total); - - for &(idx, weight) in &weights { - if choice < weight { - trace!( - upstream = idx, - dc = ?dc_idx, - weight = format!("{:.2}", weight), - total = format!("{:.2}", total), - "Upstream selected" - ); - return Some(idx); - } - choice -= weight; - } - - Some(healthy[0]) - } - - /// Connect to target through a selected upstream. - pub async fn connect( - &self, - target: SocketAddr, - dc_idx: Option, - scope: Option<&str>, - ) -> Result { - let idx = self - .select_upstream(dc_idx, scope) - .await - .ok_or_else(|| ProxyError::Config("No upstreams available".to_string()))?; - - let (mut upstream, bind_rr, dc_preference) = { - let guard = self.upstreams.read().await; - let state = &guard[idx]; - let dc_preference = dc_idx - .and_then(UpstreamState::dc_array_idx) - .map(|dc_array_idx| state.dc_ip_pref[dc_array_idx]) - .unwrap_or(IpPreference::Unknown); - ( - state.config.clone(), - Some(state.bind_rr.clone()), - dc_preference, - ) - }; - - if let Some(s) = scope { - upstream.selected_scope = s.to_string(); - } - - let target = if dc_idx.is_some() { - Self::resolve_runtime_dc_target(target, dc_idx, &upstream, dc_preference)? - } else { - target - }; - - let (stream, _) = self - .connect_selected_upstream(idx, upstream, target, dc_idx, bind_rr) - .await?; - Ok(stream) - } - - /// Connect to target through a selected upstream and return egress details. - pub async fn connect_with_details( - &self, - target: SocketAddr, - dc_idx: Option, - scope: Option<&str>, - ) -> Result<(TcpStream, UpstreamEgressInfo)> { - let idx = self - .select_upstream(dc_idx, scope) - .await - .ok_or_else(|| ProxyError::Config("No upstreams available".to_string()))?; - - let (mut upstream, bind_rr, dc_preference) = { - let guard = self.upstreams.read().await; - let state = &guard[idx]; - let dc_preference = dc_idx - .and_then(UpstreamState::dc_array_idx) - .map(|dc_array_idx| state.dc_ip_pref[dc_array_idx]) - .unwrap_or(IpPreference::Unknown); - ( - state.config.clone(), - Some(state.bind_rr.clone()), - dc_preference, - ) - }; - - // Set scope for configuration copy - if let Some(s) = scope { - upstream.selected_scope = s.to_string(); - } - - let target = if dc_idx.is_some() { - Self::resolve_runtime_dc_target(target, dc_idx, &upstream, dc_preference)? - } else { - target - }; - - let (stream, egress) = self - .connect_selected_upstream(idx, upstream, target, dc_idx, bind_rr) - .await?; - Ok((stream.into_tcp()?, egress)) - } - - async fn connect_selected_upstream( - &self, - idx: usize, - upstream: UpstreamConfig, - target: SocketAddr, - dc_idx: Option, - bind_rr: Option>, - ) -> Result<(UpstreamStream, UpstreamEgressInfo)> { - let connect_started_at = Instant::now(); - let mut last_error: Option = None; - let mut attempts_used = 0u32; - for attempt in 1..=self.connect_retry_attempts { - let elapsed = connect_started_at.elapsed(); - if elapsed >= self.connect_budget { - last_error = Some(ProxyError::ConnectionTimeout { - addr: target.to_string(), - }); - break; - } - let remaining_budget = self.connect_budget.saturating_sub(elapsed); - let attempt_timeout = - Duration::from_secs(self.tg_connect_timeout_secs).min(remaining_budget); - if attempt_timeout.is_zero() { - last_error = Some(ProxyError::ConnectionTimeout { - addr: target.to_string(), - }); - break; - } - attempts_used = attempt; - self.stats.increment_upstream_connect_attempt_total(); - let start = Instant::now(); - match self - .connect_via_upstream(idx, &upstream, target, bind_rr.clone(), attempt_timeout) - .await - { - Ok((stream, egress)) => { - let rtt_ms = start.elapsed().as_secs_f64() * 1000.0; - self.stats.increment_upstream_connect_success_total(); - self.stats - .observe_upstream_connect_attempts_per_request(attempts_used); - self.stats.observe_upstream_connect_duration_ms( - connect_started_at.elapsed().as_millis() as u64, - true, - ); - let mut guard = self.upstreams.write().await; - if let Some(u) = guard.get_mut(idx) { - if !u.healthy { - debug!(rtt_ms = format!("{:.1}", rtt_ms), "Upstream recovered"); - } - if attempt > 1 { - debug!( - attempt, - attempts = self.connect_retry_attempts, - rtt_ms = format!("{:.1}", rtt_ms), - "Upstream connect recovered after retry" - ); - } - u.healthy = true; - u.fails = 0; - - if let Some(di) = dc_idx.and_then(UpstreamState::dc_array_idx) { - u.dc_latency[di].update(rtt_ms); - } - } - return Ok((stream, egress)); - } - Err(e) => { - let hard_error = - self.connect_failfast_hard_errors && Self::is_hard_connect_error(&e); - if hard_error { - self.stats - .increment_upstream_connect_failfast_hard_error_total(); - } - if attempt < self.connect_retry_attempts && !hard_error { - debug!( - attempt, - attempts = self.connect_retry_attempts, - target = %target, - error = %e, - "Upstream connect attempt failed, retrying" - ); - let backoff = self.retry_backoff_with_jitter(); - if !backoff.is_zero() { - tokio::time::sleep(backoff).await; - } - } else if hard_error { - debug!( - attempt, - attempts = self.connect_retry_attempts, - target = %target, - error = %e, - "Upstream connect failed with hard error, failfast is active" - ); - } - last_error = Some(e); - if hard_error { - break; - } - } - } - } - - self.stats.increment_upstream_connect_fail_total(); - self.stats - .observe_upstream_connect_attempts_per_request(attempts_used); - self.stats.observe_upstream_connect_duration_ms( - connect_started_at.elapsed().as_millis() as u64, - false, - ); - - let error = last_error.unwrap_or_else(|| { - ProxyError::Config("Upstream connect attempts exhausted".to_string()) - }); - - let mut guard = self.upstreams.write().await; - if let Some(u) = guard.get_mut(idx) { - // Intermediate attempts are intentionally ignored here. - // Health state is degraded only when the entire connect cycle fails. - u.fails += 1; - warn!( - fails = u.fails, - attempts = self.connect_retry_attempts, - "Upstream failed after retries: {}", - error - ); - if u.fails >= self.unhealthy_fail_threshold { - u.healthy = false; - warn!( - fails = u.fails, - threshold = self.unhealthy_fail_threshold, - "Upstream marked unhealthy" - ); - } - } - Err(error) - } - - async fn connect_via_upstream( - &self, - upstream_id: usize, - config: &UpstreamConfig, - target: SocketAddr, - bind_rr: Option>, - connect_timeout: Duration, - ) -> Result<(UpstreamStream, UpstreamEgressInfo)> { - match &config.upstream_type { - UpstreamType::Direct { - interface, - bind_addresses, - bindtodevice, - } => { - let bind_ip = Self::resolve_bind_address( - interface, - bind_addresses, - target, - bind_rr.as_deref(), - true, - ); - if bind_ip.is_none() && bind_addresses.as_ref().is_some_and(|v| !v.is_empty()) { - return Err(ProxyError::Config(format!( - "No valid bind_addresses for target family {target}" - ))); - } - - let socket = create_outgoing_socket_bound(target, bind_ip)?; - if let Some(device) = bindtodevice.as_deref().filter(|value| !value.is_empty()) { - bind_outgoing_socket_to_device(&socket, device).map_err(ProxyError::Io)?; - debug!(bindtodevice = %device, target = %target, "Pinned socket to interface"); - } - if let Some(ip) = bind_ip { - debug!(bind = %ip, target = %target, "Bound outgoing socket"); - } else if interface.is_some() || bind_addresses.is_some() { - debug!(target = %target, "No matching bind address for target family"); - } - - socket.set_nonblocking(true)?; - match socket.connect(&target.into()) { - Ok(()) => {} - Err(err) - if err.raw_os_error() == Some(libc::EINPROGRESS) - || err.kind() == std::io::ErrorKind::WouldBlock => {} - Err(err) => return Err(ProxyError::Io(err)), - } - - let std_stream: std::net::TcpStream = socket.into(); - let stream = TcpStream::from_std(std_stream)?; - - match tokio::time::timeout(connect_timeout, stream.writable()).await { - Ok(Ok(())) => {} - Ok(Err(e)) => return Err(ProxyError::Io(e)), - Err(_) => { - return Err(ProxyError::ConnectionTimeout { - addr: target.to_string(), - }); - } - } - if let Some(e) = stream.take_error()? { - return Err(ProxyError::Io(e)); - } - - let local_addr = stream.local_addr().ok(); - Ok(( - UpstreamStream::Tcp(stream), - UpstreamEgressInfo { - upstream_id, - route_kind: UpstreamRouteKind::Direct, - local_addr, - direct_bind_ip: bind_ip, - socks_bound_addr: None, - socks_proxy_addr: None, - }, - )) - } - UpstreamType::Socks4 { - address, - interface, - user_id, - } => { - // Try to parse as SocketAddr first (IP:port), otherwise treat as hostname:port - let mut stream = if let Ok(proxy_addr) = address.parse::() { - // IP:port format - use socket with optional interface binding - let bind_ip = Self::resolve_bind_address( - interface, - &None, - proxy_addr, - bind_rr.as_deref(), - false, - ); - - let socket = create_outgoing_socket_bound(proxy_addr, bind_ip)?; - - socket.set_nonblocking(true)?; - match socket.connect(&proxy_addr.into()) { - Ok(()) => {} - Err(err) - if err.raw_os_error() == Some(libc::EINPROGRESS) - || err.kind() == std::io::ErrorKind::WouldBlock => {} - Err(err) => return Err(ProxyError::Io(err)), - } - - let std_stream: std::net::TcpStream = socket.into(); - let stream = TcpStream::from_std(std_stream)?; - - match tokio::time::timeout(connect_timeout, stream.writable()).await { - Ok(Ok(())) => {} - Ok(Err(e)) => return Err(ProxyError::Io(e)), - Err(_) => { - return Err(ProxyError::ConnectionTimeout { - addr: proxy_addr.to_string(), - }); - } - } - if let Some(e) = stream.take_error()? { - return Err(ProxyError::Io(e)); - } - stream - } else { - // Hostname:port format - use tokio DNS resolution - // Note: interface binding is not supported for hostnames - if interface.is_some() { - warn!( - "SOCKS4 interface binding is not supported for hostname addresses, ignoring" - ); - } - self.connect_hostname_with_dns_override(address, connect_timeout) - .await? - }; - - // replace socks user_id with config.selected_scope, if set - let scope: Option<&str> = - Some(config.selected_scope.as_str()).filter(|s| !s.is_empty()); - let _user_id: Option<&str> = scope.or(user_id.as_deref()); - - let bound = match tokio::time::timeout( - connect_timeout, - connect_socks4(&mut stream, target, _user_id), - ) - .await - { - Ok(Ok(bound)) => bound, - Ok(Err(e)) => return Err(e), - Err(_) => { - return Err(ProxyError::ConnectionTimeout { - addr: target.to_string(), - }); - } - }; - let local_addr = stream.local_addr().ok(); - let socks_proxy_addr = stream.peer_addr().ok(); - Ok(( - UpstreamStream::Tcp(stream), - UpstreamEgressInfo { - upstream_id, - route_kind: UpstreamRouteKind::Socks4, - local_addr, - direct_bind_ip: None, - socks_bound_addr: Some(bound.addr), - socks_proxy_addr, - }, - )) - } - UpstreamType::Socks5 { - address, - interface, - username, - password, - } => { - // Try to parse as SocketAddr first (IP:port), otherwise treat as hostname:port - let mut stream = if let Ok(proxy_addr) = address.parse::() { - // IP:port format - use socket with optional interface binding - let bind_ip = Self::resolve_bind_address( - interface, - &None, - proxy_addr, - bind_rr.as_deref(), - false, - ); - - let socket = create_outgoing_socket_bound(proxy_addr, bind_ip)?; - - socket.set_nonblocking(true)?; - match socket.connect(&proxy_addr.into()) { - Ok(()) => {} - Err(err) - if err.raw_os_error() == Some(libc::EINPROGRESS) - || err.kind() == std::io::ErrorKind::WouldBlock => {} - Err(err) => return Err(ProxyError::Io(err)), - } - - let std_stream: std::net::TcpStream = socket.into(); - let stream = TcpStream::from_std(std_stream)?; - - match tokio::time::timeout(connect_timeout, stream.writable()).await { - Ok(Ok(())) => {} - Ok(Err(e)) => return Err(ProxyError::Io(e)), - Err(_) => { - return Err(ProxyError::ConnectionTimeout { - addr: proxy_addr.to_string(), - }); - } - } - if let Some(e) = stream.take_error()? { - return Err(ProxyError::Io(e)); - } - stream - } else { - // Hostname:port format - use tokio DNS resolution - // Note: interface binding is not supported for hostnames - if interface.is_some() { - warn!( - "SOCKS5 interface binding is not supported for hostname addresses, ignoring" - ); - } - self.connect_hostname_with_dns_override(address, connect_timeout) - .await? - }; - - debug!(config = ?config, "Socks5 connection"); - // replace socks user:pass with config.selected_scope, if set - let scope: Option<&str> = - Some(config.selected_scope.as_str()).filter(|s| !s.is_empty()); - let _username: Option<&str> = scope.or(username.as_deref()); - let _password: Option<&str> = scope.or(password.as_deref()); - - let bound = match tokio::time::timeout( - connect_timeout, - connect_socks5(&mut stream, target, _username, _password), - ) - .await - { - Ok(Ok(bound)) => bound, - Ok(Err(e)) => return Err(e), - Err(_) => { - return Err(ProxyError::ConnectionTimeout { - addr: target.to_string(), - }); - } - }; - let local_addr = stream.local_addr().ok(); - let socks_proxy_addr = stream.peer_addr().ok(); - Ok(( - UpstreamStream::Tcp(stream), - UpstreamEgressInfo { - upstream_id, - route_kind: UpstreamRouteKind::Socks5, - local_addr, - direct_bind_ip: None, - socks_bound_addr: Some(bound.addr), - socks_proxy_addr, - }, - )) - } - UpstreamType::Shadowsocks { url, interface } => { - let stream = connect_shadowsocks(url, interface, target, connect_timeout).await?; - let local_addr = stream.get_ref().local_addr().ok(); - Ok(( - UpstreamStream::Shadowsocks(Box::new(stream)), - UpstreamEgressInfo { - upstream_id, - route_kind: UpstreamRouteKind::Shadowsocks, - local_addr, - direct_bind_ip: None, - socks_bound_addr: None, - socks_proxy_addr: None, - }, - )) - } - } - } - - // ============= Startup Ping (test both IPv6 and IPv4) ============= - - /// Ping all Telegram DCs through all upstreams. - /// Tests BOTH IPv6 and IPv4, returns separate results for each. - pub async fn ping_all_dcs( - &self, - prefer_ipv6: bool, - dc_overrides: &HashMap>, - ipv4_enabled: bool, - ipv6_enabled: bool, - ) -> Vec { - let upstreams: Vec<(usize, UpstreamConfig, Arc)> = { - let guard = self.upstreams.read().await; - guard - .iter() - .enumerate() - .map(|(i, u)| (i, u.config.clone(), u.bind_rr.clone())) - .collect() - }; - let has_unscoped = upstreams - .iter() - .any(|(_, cfg, _)| Self::is_unscoped_upstream(cfg)); - - let mut all_results = Vec::new(); - - for (upstream_idx, upstream_config, bind_rr) in &upstreams { - // DC connectivity checks should follow the default routing path. - // Scoped upstreams are included only when no unscoped upstream exists. - if !Self::should_check_in_default_dc_connectivity(has_unscoped, upstream_config) { - continue; - } - - let (upstream_ipv4_enabled, upstream_ipv6_enabled) = - Self::resolve_probe_dc_families(upstream_config, ipv4_enabled, ipv6_enabled); - let upstream_prefer_ipv6 = upstream_config.prefer_ipv6(prefer_ipv6); - let upstream_name = match &upstream_config.upstream_type { - UpstreamType::Direct { - interface, - bind_addresses, - bindtodevice, - } => { - let mut direct_parts = Vec::new(); - if let Some(dev) = interface.as_deref().filter(|v| !v.is_empty()) { - direct_parts.push(format!("dev={dev}")); - } - if let Some(src) = bind_addresses.as_ref().filter(|v| !v.is_empty()) { - direct_parts.push(format!("src={}", src.join(","))); - } - if let Some(device) = bindtodevice.as_deref().filter(|v| !v.is_empty()) { - direct_parts.push(format!("bindtodevice={device}")); - } - if direct_parts.is_empty() { - "direct".to_string() - } else { - format!("direct {}", direct_parts.join(" ")) - } - } - UpstreamType::Socks4 { address, .. } => format!("socks4://{}", address), - UpstreamType::Socks5 { address, .. } => format!("socks5://{}", address), - UpstreamType::Shadowsocks { url, .. } => { - let address = - sanitize_shadowsocks_url(url).unwrap_or_else(|_| "invalid".to_string()); - format!("shadowsocks://{address}") - } - }; - - let mut v6_results = Vec::with_capacity(NUM_DCS); - if upstream_ipv6_enabled { - for dc_zero_idx in 0..NUM_DCS { - let dc_v6 = TG_DATACENTERS_V6[dc_zero_idx]; - let addr_v6 = SocketAddr::new(dc_v6, TG_DATACENTER_PORT); - - let result = tokio::time::timeout( - Duration::from_secs(DC_PING_TIMEOUT_SECS), - self.ping_single_dc( - *upstream_idx, - upstream_config, - Some(bind_rr.clone()), - addr_v6, - ), - ) - .await; - - let ping_result = match result { - Ok(Ok(rtt_ms)) => { - let mut guard = self.upstreams.write().await; - if let Some(u) = guard.get_mut(*upstream_idx) { - u.dc_latency[dc_zero_idx].update(rtt_ms); - } - DcPingResult { - dc_idx: dc_zero_idx + 1, - dc_addr: addr_v6, - rtt_ms: Some(rtt_ms), - error: None, - } - } - Ok(Err(e)) => DcPingResult { - dc_idx: dc_zero_idx + 1, - dc_addr: addr_v6, - rtt_ms: None, - error: Some(e.to_string()), - }, - Err(_) => DcPingResult { - dc_idx: dc_zero_idx + 1, - dc_addr: addr_v6, - rtt_ms: None, - error: Some("timeout".to_string()), - }, - }; - v6_results.push(ping_result); - } - } else { - for dc_zero_idx in 0..NUM_DCS { - let dc_v6 = TG_DATACENTERS_V6[dc_zero_idx]; - v6_results.push(DcPingResult { - dc_idx: dc_zero_idx + 1, - dc_addr: SocketAddr::new(dc_v6, TG_DATACENTER_PORT), - rtt_ms: None, - error: Some(if ipv6_enabled { - "ipv6 disabled by upstream policy".to_string() - } else { - "ipv6 disabled".to_string() - }), - }); - } - } - - let mut v4_results = Vec::with_capacity(NUM_DCS); - if upstream_ipv4_enabled { - for dc_zero_idx in 0..NUM_DCS { - let dc_v4 = TG_DATACENTERS_V4[dc_zero_idx]; - let addr_v4 = SocketAddr::new(dc_v4, TG_DATACENTER_PORT); - - let result = tokio::time::timeout( - Duration::from_secs(DC_PING_TIMEOUT_SECS), - self.ping_single_dc( - *upstream_idx, - upstream_config, - Some(bind_rr.clone()), - addr_v4, - ), - ) - .await; - - let ping_result = match result { - Ok(Ok(rtt_ms)) => { - let mut guard = self.upstreams.write().await; - if let Some(u) = guard.get_mut(*upstream_idx) { - u.dc_latency[dc_zero_idx].update(rtt_ms); - } - DcPingResult { - dc_idx: dc_zero_idx + 1, - dc_addr: addr_v4, - rtt_ms: Some(rtt_ms), - error: None, - } - } - Ok(Err(e)) => DcPingResult { - dc_idx: dc_zero_idx + 1, - dc_addr: addr_v4, - rtt_ms: None, - error: Some(e.to_string()), - }, - Err(_) => DcPingResult { - dc_idx: dc_zero_idx + 1, - dc_addr: addr_v4, - rtt_ms: None, - error: Some("timeout".to_string()), - }, - }; - v4_results.push(ping_result); - } - } else { - for dc_zero_idx in 0..NUM_DCS { - let dc_v4 = TG_DATACENTERS_V4[dc_zero_idx]; - v4_results.push(DcPingResult { - dc_idx: dc_zero_idx + 1, - dc_addr: SocketAddr::new(dc_v4, TG_DATACENTER_PORT), - rtt_ms: None, - error: Some(if ipv4_enabled { - "ipv4 disabled by upstream policy".to_string() - } else { - "ipv4 disabled".to_string() - }), - }); - } - } - - // === Ping DC overrides (v4/v6) === - for (dc_key, addrs) in dc_overrides { - let dc_num: i16 = match dc_key.parse::() { - Ok(v) if v > 0 => v, - Err(_) => { - warn!(dc = %dc_key, "Invalid dc_overrides key, skipping"); - continue; - } - _ => continue, - }; - let dc_idx = dc_num as usize; - for addr_str in addrs { - match addr_str.parse::() { - Ok(addr) => { - let is_v6 = addr.is_ipv6(); - if (is_v6 && !upstream_ipv6_enabled) - || (!is_v6 && !upstream_ipv4_enabled) - { - continue; - } - let result = tokio::time::timeout( - Duration::from_secs(DC_PING_TIMEOUT_SECS), - self.ping_single_dc( - *upstream_idx, - upstream_config, - Some(bind_rr.clone()), - addr, - ), - ) - .await; - - let ping_result = match result { - Ok(Ok(rtt_ms)) => DcPingResult { - dc_idx, - dc_addr: addr, - rtt_ms: Some(rtt_ms), - error: None, - }, - Ok(Err(e)) => DcPingResult { - dc_idx, - dc_addr: addr, - rtt_ms: None, - error: Some(e.to_string()), - }, - Err(_) => DcPingResult { - dc_idx, - dc_addr: addr, - rtt_ms: None, - error: Some("timeout".to_string()), - }, - }; - - if is_v6 { - v6_results.push(ping_result); - } else { - v4_results.push(ping_result); - } - } - Err(_) => { - warn!(dc = %dc_idx, addr = %addr_str, "Invalid dc_overrides address, skipping") - } - } - } - } - - // Check if both IP versions have at least one working DC - let v6_has_working = v6_results.iter().any(|r| r.rtt_ms.is_some()); - let v4_has_working = v4_results.iter().any(|r| r.rtt_ms.is_some()); - let both_available = v6_has_working && v4_has_working; - - // Update IP preference for each DC - { - let mut guard = self.upstreams.write().await; - if let Some(u) = guard.get_mut(*upstream_idx) { - for dc_zero_idx in 0..NUM_DCS { - let v6_ok = v6_results[dc_zero_idx].rtt_ms.is_some(); - let v4_ok = v4_results[dc_zero_idx].rtt_ms.is_some(); - - u.dc_ip_pref[dc_zero_idx] = match (v6_ok, v4_ok) { - (true, true) => IpPreference::BothWork, - (true, false) => IpPreference::PreferV6, - (false, true) => IpPreference::PreferV4, - (false, false) => IpPreference::Unavailable, - }; - } - } - } - - all_results.push(StartupPingResult { - v6_results, - v4_results, - upstream_name, - prefer_ipv6: upstream_prefer_ipv6, - both_available, - }); - } - - all_results - } - - async fn ping_single_dc( - &self, - upstream_id: usize, - config: &UpstreamConfig, - bind_rr: Option>, - target: SocketAddr, - ) -> Result { - let start = Instant::now(); - let _ = self - .connect_via_upstream( - upstream_id, - config, - target, - bind_rr, - Duration::from_secs(DC_PING_TIMEOUT_SECS), - ) - .await?; - Ok(start.elapsed().as_secs_f64() * 1000.0) - } - - fn required_healthy_group_count(total_groups: usize) -> usize { - if total_groups == 0 { - 0 - } else { - total_groups.min(MIN_HEALTHY_DC_GROUPS) - } - } - - fn build_health_check_groups( - ipv4_enabled: bool, - ipv6_enabled: bool, - dc_overrides: &HashMap>, - ) -> Vec { - let mut v4_by_dc: HashMap> = HashMap::new(); - let mut v6_by_dc: HashMap> = HashMap::new(); - - if ipv4_enabled { - for (idx, dc_ip) in TG_DATACENTERS_V4.iter().enumerate() { - let dc_idx = (idx + 1) as i16; - v4_by_dc - .entry(dc_idx) - .or_default() - .push(SocketAddr::new(*dc_ip, TG_DATACENTER_PORT)); - } - } - - if ipv6_enabled { - for (idx, dc_ip) in TG_DATACENTERS_V6.iter().enumerate() { - let dc_idx = (idx + 1) as i16; - v6_by_dc - .entry(dc_idx) - .or_default() - .push(SocketAddr::new(*dc_ip, TG_DATACENTER_PORT)); - } - } - - for (dc_key, addrs) in dc_overrides { - let dc_idx = match dc_key.parse::() { - Ok(v) if v > 0 => v, - _ => { - warn!(dc = %dc_key, "Invalid dc_overrides key for health-check, skipping"); - continue; - } - }; - - for addr_str in addrs { - match addr_str.parse::() { - Ok(addr) if addr.is_ipv6() => { - if ipv6_enabled { - v6_by_dc.entry(dc_idx).or_default().push(addr); - } - } - Ok(addr) => { - if ipv4_enabled { - v4_by_dc.entry(dc_idx).or_default().push(addr); - } - } - Err(_) => { - warn!( - dc = %dc_idx, - addr = %addr_str, - "Invalid dc_overrides address for health-check, skipping" - ); - } - } - } - } - - for addrs in v4_by_dc.values_mut() { - addrs.sort_unstable(); - addrs.dedup(); - } - for addrs in v6_by_dc.values_mut() { - addrs.sort_unstable(); - addrs.dedup(); - } - - let mut all_dcs = BTreeSet::new(); - all_dcs.extend(v4_by_dc.keys().copied()); - all_dcs.extend(v6_by_dc.keys().copied()); - - let mut groups = Vec::with_capacity(all_dcs.len()); - for dc_idx in all_dcs { - let v4_endpoints = v4_by_dc.remove(&dc_idx).unwrap_or_default(); - let v6_endpoints = v6_by_dc.remove(&dc_idx).unwrap_or_default(); - - if v4_endpoints.is_empty() && v6_endpoints.is_empty() { - continue; - } - - groups.push(HealthCheckGroup { - dc_idx, - v4_endpoints, - v6_endpoints, - }); - } - - groups - } - - fn health_check_endpoint_order( - group: &HealthCheckGroup, - prefer_ipv6: bool, - ) -> [(bool, &[SocketAddr]); 2] { - if prefer_ipv6 { - [(true, &group.v6_endpoints), (false, &group.v4_endpoints)] - } else { - [(true, &group.v4_endpoints), (false, &group.v6_endpoints)] - } - } - - // ============= Health Checks ============= - - /// Background health check based on reachable DC groups through each upstream. - /// Upstream stays healthy while at least `MIN_HEALTHY_DC_GROUPS` groups are reachable. - pub async fn run_health_checks( - &self, - prefer_ipv6: bool, - ipv4_enabled: bool, - ipv6_enabled: bool, - dc_overrides: HashMap>, - ) { - let (health_ipv4_enabled, health_ipv6_enabled) = { - let guard = self.upstreams.read().await; - ( - ipv4_enabled - || guard - .iter() - .any(|upstream| upstream.config.ipv4 == Some(true)), - ipv6_enabled - || guard - .iter() - .any(|upstream| upstream.config.ipv6 == Some(true)), - ) - }; - let groups = Self::build_health_check_groups( - health_ipv4_enabled, - health_ipv6_enabled, - &dc_overrides, - ); - let required_healthy_groups = Self::required_healthy_group_count(groups.len()); - let mut endpoint_rotation: HashMap<(usize, i16, bool), usize> = HashMap::new(); - - if groups.is_empty() { - warn!("No DC groups available for upstream health-checks"); - } - - loop { - tokio::time::sleep(Duration::from_secs(HEALTH_CHECK_INTERVAL_SECS)).await; - - if groups.is_empty() || required_healthy_groups == 0 { - continue; - } - - let target_upstreams: Vec = { - let guard = self.upstreams.read().await; - let has_unscoped = guard - .iter() - .any(|upstream| Self::is_unscoped_upstream(&upstream.config)); - guard - .iter() - .enumerate() - .filter(|(_, upstream)| { - Self::should_check_in_default_dc_connectivity( - has_unscoped, - &upstream.config, - ) - }) - .map(|(idx, _)| idx) - .collect() - }; - - for i in target_upstreams { - let (config, bind_rr) = { - let guard = self.upstreams.read().await; - let u = &guard[i]; - (u.config.clone(), u.bind_rr.clone()) - }; - let (upstream_ipv4_enabled, upstream_ipv6_enabled) = - Self::resolve_probe_dc_families(&config, ipv4_enabled, ipv6_enabled); - let upstream_prefer_ipv6 = config.prefer_ipv6(prefer_ipv6); - - let mut healthy_groups = 0usize; - let mut latency_updates: Vec<(usize, f64)> = Vec::new(); - - for group in &groups { - let mut group_ok = false; - let mut group_rtt_ms = None; - - for (is_primary, endpoints) in - Self::health_check_endpoint_order(group, upstream_prefer_ipv6) - { - if endpoints.is_empty() { - continue; - } - - let filtered_endpoints: Vec = endpoints - .iter() - .copied() - .filter(|endpoint| { - if endpoint.is_ipv4() { - upstream_ipv4_enabled - } else { - upstream_ipv6_enabled - } - }) - .collect(); - - if filtered_endpoints.is_empty() { - continue; - } - - let rotation_key = (i, group.dc_idx, is_primary); - let start_idx = *endpoint_rotation.entry(rotation_key).or_insert(0) - % filtered_endpoints.len(); - let mut next_idx = (start_idx + 1) % filtered_endpoints.len(); - - for step in 0..filtered_endpoints.len() { - let endpoint_idx = (start_idx + step) % filtered_endpoints.len(); - let endpoint = filtered_endpoints[endpoint_idx]; - - let start = Instant::now(); - let result = tokio::time::timeout( - Duration::from_secs(HEALTH_CHECK_CONNECT_TIMEOUT_SECS), - self.connect_via_upstream( - i, - &config, - endpoint, - Some(bind_rr.clone()), - Duration::from_secs(HEALTH_CHECK_CONNECT_TIMEOUT_SECS), - ), - ) - .await; - - match result { - Ok(Ok(_stream)) => { - group_ok = true; - group_rtt_ms = Some(start.elapsed().as_secs_f64() * 1000.0); - next_idx = (endpoint_idx + 1) % filtered_endpoints.len(); - break; - } - Ok(Err(e)) => { - debug!( - upstream = i, - dc = group.dc_idx, - endpoint = %endpoint, - primary = is_primary, - error = %e, - "Health-check endpoint failed" - ); - } - Err(_) => { - debug!( - upstream = i, - dc = group.dc_idx, - endpoint = %endpoint, - primary = is_primary, - "Health-check endpoint timed out" - ); - } - } - } - - endpoint_rotation.insert(rotation_key, next_idx); - - if group_ok { - break; - } - } - - if group_ok { - healthy_groups += 1; - if let (Some(dc_array_idx), Some(rtt_ms)) = - (UpstreamState::dc_array_idx(group.dc_idx), group_rtt_ms) - { - latency_updates.push((dc_array_idx, rtt_ms)); - } - } - } - - let mut guard = self.upstreams.write().await; - let u = &mut guard[i]; - - for (dc_array_idx, rtt_ms) in latency_updates { - u.dc_latency[dc_array_idx].update(rtt_ms); - } - - if healthy_groups >= required_healthy_groups { - if !u.healthy { - info!( - upstream = i, - healthy_groups, - total_groups = groups.len(), - required_groups = required_healthy_groups, - "Upstream recovered by DC-group health threshold" - ); - } - u.healthy = true; - u.fails = 0; - } else { - u.fails += 1; - debug!( - upstream = i, - healthy_groups, - total_groups = groups.len(), - required_groups = required_healthy_groups, - fails = u.fails, - "Upstream health-check below DC-group threshold" - ); - if u.fails >= self.unhealthy_fail_threshold { - u.healthy = false; - warn!( - upstream = i, - healthy_groups, - total_groups = groups.len(), - required_groups = required_healthy_groups, - fails = u.fails, - threshold = self.unhealthy_fail_threshold, - "Upstream unhealthy (insufficient reachable DC groups)" - ); - } - } - - u.last_check = std::time::Instant::now(); - } - } - } - - /// Get the preferred IP for a DC (for use by other components) - #[allow(dead_code)] - pub async fn get_dc_ip_preference(&self, dc_idx: i16) -> Option { - let guard = self.upstreams.read().await; - if guard.is_empty() { - return None; - } - - UpstreamState::dc_array_idx(dc_idx).map(|idx| guard[0].dc_ip_pref[idx]) - } - - /// Get preferred DC address based on config preference - #[allow(dead_code)] - pub async fn get_dc_addr(&self, dc_idx: i16, prefer_ipv6: bool) -> Option { - let arr_idx = UpstreamState::dc_array_idx(dc_idx)?; - - let ip = if prefer_ipv6 { - TG_DATACENTERS_V6[arr_idx] - } else { - TG_DATACENTERS_V4[arr_idx] - }; - - Some(SocketAddr::new(ip, TG_DATACENTER_PORT)) - } -} - +// Upstream manager configuration, DNS resolution, and API snapshots. +mod manager_config; +// Latency- and scope-aware upstream selection. +mod selection; +// Direct and proxied connection establishment. +mod connect; +// Startup Telegram DC connectivity probes. +mod startup_probe; +// Periodic upstream health checks. +mod health_checks; #[cfg(test)] -mod tests { - use super::*; - use std::sync::Arc; - - use crate::stats::Stats; - - const TEST_SHADOWSOCKS_URL: &str = - "ss://2022-blake3-aes-256-gcm:MDEyMzQ1Njc4OTAxMjM0NTY3ODkwMTIzNDU2Nzg5MDE=@127.0.0.1:8388"; - - fn manager_with_dns(entries: &[String]) -> UpstreamManager { - UpstreamManager::new(Vec::new(), 1, 1, 1, 1, 1, false, Arc::new(Stats::new())) - .with_dns_overrides(entries) - .unwrap() - } - - #[tokio::test] - async fn generation_local_dns_overrides_are_isolated_and_case_insensitive() { - let active = manager_with_dns(&["Front.Example:443:192.0.2.10".to_string()]); - let candidate = manager_with_dns(&["front.example:443:[2001:db8::10]".to_string()]); - - assert_eq!( - active.resolve_hostname("front.example", 443).await.unwrap(), - "192.0.2.10:443".parse::().unwrap() - ); - assert_eq!( - candidate - .resolve_hostname("FRONT.EXAMPLE", 443) - .await - .unwrap(), - "[2001:db8::10]:443".parse::().unwrap() - ); - - candidate - .update_dns_overrides(&["front.example:443:192.0.2.20".to_string()]) - .unwrap(); - assert_eq!( - active.resolve_hostname("FRONT.EXAMPLE", 443).await.unwrap(), - "192.0.2.10:443".parse::().unwrap() - ); - assert_eq!( - candidate - .resolve_hostname("front.example", 443) - .await - .unwrap(), - "192.0.2.20:443".parse::().unwrap() - ); - } - - #[test] - fn required_healthy_group_count_applies_three_group_threshold() { - assert_eq!(UpstreamManager::required_healthy_group_count(0), 0); - assert_eq!(UpstreamManager::required_healthy_group_count(1), 1); - assert_eq!(UpstreamManager::required_healthy_group_count(2), 2); - assert_eq!(UpstreamManager::required_healthy_group_count(3), 3); - assert_eq!(UpstreamManager::required_healthy_group_count(5), 3); - } - - #[test] - fn build_health_check_groups_merges_family_endpoints_with_preference() { - let mut overrides = HashMap::new(); - overrides.insert( - "2".to_string(), - vec![ - "203.0.113.10:443".to_string(), - "203.0.113.11:443".to_string(), - "[2001:db8::10]:443".to_string(), - ], - ); - - let groups = UpstreamManager::build_health_check_groups(true, true, &overrides); - let dc2 = groups - .iter() - .find(|g| g.dc_idx == 2) - .expect("dc2 must be present"); - - assert!(dc2.v6_endpoints.iter().all(|addr| addr.is_ipv6())); - assert!(dc2.v4_endpoints.iter().all(|addr| addr.is_ipv4())); - assert!( - dc2.v6_endpoints - .contains(&"[2001:db8::10]:443".parse::().unwrap()) - ); - assert!( - dc2.v4_endpoints - .contains(&"203.0.113.10:443".parse::().unwrap()) - ); - assert!( - dc2.v4_endpoints - .contains(&"203.0.113.11:443".parse::().unwrap()) - ); - - let ordered = UpstreamManager::health_check_endpoint_order(dc2, true); - assert!(ordered[0].1.iter().all(|addr| addr.is_ipv6())); - assert!(ordered[1].1.iter().all(|addr| addr.is_ipv4())); - } - - #[test] - fn build_health_check_groups_keeps_multiple_endpoints_per_group() { - let mut overrides = HashMap::new(); - overrides.insert( - "9".to_string(), - vec![ - "198.51.100.1:443".to_string(), - "198.51.100.2:443".to_string(), - "198.51.100.1:443".to_string(), - ], - ); - - let groups = UpstreamManager::build_health_check_groups(true, false, &overrides); - let dc9 = groups - .iter() - .find(|g| g.dc_idx == 9) - .expect("override-only dc group must be present"); - - assert_eq!(dc9.v4_endpoints.len(), 2); - assert!( - dc9.v4_endpoints - .contains(&"198.51.100.1:443".parse::().unwrap()) - ); - assert!( - dc9.v4_endpoints - .contains(&"198.51.100.2:443".parse::().unwrap()) - ); - assert!(dc9.v6_endpoints.is_empty()); - } - - #[test] - fn hard_connect_error_classification_detects_connection_refused() { - let error = ProxyError::ConnectionRefused { - addr: "127.0.0.1:443".to_string(), - }; - assert!(UpstreamManager::is_hard_connect_error(&error)); - } - - #[test] - fn hard_connect_error_classification_skips_timeouts() { - let error = ProxyError::ConnectionTimeout { - addr: "127.0.0.1:443".to_string(), - }; - assert!(!UpstreamManager::is_hard_connect_error(&error)); - } - - #[test] - fn unscoped_selection_detects_default_route_upstream() { - let mut upstream = UpstreamConfig { - upstream_type: UpstreamType::Direct { - interface: None, - bind_addresses: None, - bindtodevice: None, - }, - weight: 1, - enabled: true, - scopes: String::new(), - selected_scope: String::new(), - ipv4: None, - ipv6: None, - prefer: None, - }; - - assert!(UpstreamManager::is_unscoped_upstream(&upstream)); - upstream.scopes = "local".to_string(); - assert!(!UpstreamManager::is_unscoped_upstream(&upstream)); - assert!(!UpstreamManager::should_check_in_default_dc_connectivity( - true, &upstream - )); - assert!(UpstreamManager::should_check_in_default_dc_connectivity( - false, &upstream - )); - } - - #[test] - fn resolve_bind_address_prefers_explicit_bind_ip() { - let target = "203.0.113.10:443".parse::().unwrap(); - let bind = UpstreamManager::resolve_bind_address( - &Some("198.51.100.20".to_string()), - &Some(vec!["198.51.100.10".to_string()]), - target, - None, - true, - ); - - assert_eq!(bind, Some("198.51.100.10".parse::().unwrap())); - } - - #[test] - fn resolve_bind_address_does_not_fallback_to_interface_when_bind_addresses_present() { - let target = "203.0.113.10:443".parse::().unwrap(); - let bind = UpstreamManager::resolve_bind_address( - &Some("198.51.100.20".to_string()), - &Some(vec!["2001:db8::10".to_string()]), - target, - None, - true, - ); - - assert_eq!(bind, None); - } - - #[test] - fn api_snapshot_reports_shadowsocks_as_sanitized_route() { - let manager = UpstreamManager::new( - vec![UpstreamConfig { - upstream_type: UpstreamType::Shadowsocks { - url: TEST_SHADOWSOCKS_URL.to_string(), - interface: None, - }, - weight: 2, - enabled: true, - scopes: String::new(), - selected_scope: String::new(), - ipv4: None, - ipv6: None, - prefer: None, - }], - 1, - 100, - 1000, - 10, - 1, - false, - Arc::new(Stats::new()), - ); - - let snapshot = manager.try_api_snapshot().expect("snapshot"); - assert_eq!(snapshot.summary.configured_total, 1); - assert_eq!(snapshot.summary.shadowsocks_total, 1); - assert_eq!(snapshot.upstreams.len(), 1); - assert_eq!( - snapshot.upstreams[0].route_kind, - UpstreamRouteKind::Shadowsocks - ); - assert_eq!(snapshot.upstreams[0].address, "127.0.0.1:8388"); - } -} +mod tests; diff --git a/src/transport/upstream/connect.rs b/src/transport/upstream/connect.rs new file mode 100644 index 0000000..af63bc8 --- /dev/null +++ b/src/transport/upstream/connect.rs @@ -0,0 +1,413 @@ +use super::*; + +impl UpstreamManager { + pub(super) async fn connect_selected_upstream( + &self, + idx: usize, + upstream: UpstreamConfig, + target: SocketAddr, + dc_idx: Option, + bind_rr: Option>, + ) -> Result<(UpstreamStream, UpstreamEgressInfo)> { + let connect_started_at = Instant::now(); + let mut last_error: Option = None; + let mut attempts_used = 0u32; + for attempt in 1..=self.connect_retry_attempts { + let elapsed = connect_started_at.elapsed(); + if elapsed >= self.connect_budget { + last_error = Some(ProxyError::ConnectionTimeout { + addr: target.to_string(), + }); + break; + } + let remaining_budget = self.connect_budget.saturating_sub(elapsed); + let attempt_timeout = + Duration::from_secs(self.tg_connect_timeout_secs).min(remaining_budget); + if attempt_timeout.is_zero() { + last_error = Some(ProxyError::ConnectionTimeout { + addr: target.to_string(), + }); + break; + } + attempts_used = attempt; + self.stats.increment_upstream_connect_attempt_total(); + let start = Instant::now(); + match self + .connect_via_upstream(idx, &upstream, target, bind_rr.clone(), attempt_timeout) + .await + { + Ok((stream, egress)) => { + let rtt_ms = start.elapsed().as_secs_f64() * 1000.0; + self.stats.increment_upstream_connect_success_total(); + self.stats + .observe_upstream_connect_attempts_per_request(attempts_used); + self.stats.observe_upstream_connect_duration_ms( + connect_started_at.elapsed().as_millis() as u64, + true, + ); + let mut guard = self.upstreams.write().await; + if let Some(u) = guard.get_mut(idx) { + if !u.healthy { + debug!(rtt_ms = format!("{:.1}", rtt_ms), "Upstream recovered"); + } + if attempt > 1 { + debug!( + attempt, + attempts = self.connect_retry_attempts, + rtt_ms = format!("{:.1}", rtt_ms), + "Upstream connect recovered after retry" + ); + } + u.healthy = true; + u.fails = 0; + + if let Some(di) = dc_idx.and_then(UpstreamState::dc_array_idx) { + u.dc_latency[di].update(rtt_ms); + } + } + return Ok((stream, egress)); + } + Err(e) => { + let hard_error = + self.connect_failfast_hard_errors && Self::is_hard_connect_error(&e); + if hard_error { + self.stats + .increment_upstream_connect_failfast_hard_error_total(); + } + if attempt < self.connect_retry_attempts && !hard_error { + debug!( + attempt, + attempts = self.connect_retry_attempts, + target = %target, + error = %e, + "Upstream connect attempt failed, retrying" + ); + let backoff = self.retry_backoff_with_jitter(); + if !backoff.is_zero() { + tokio::time::sleep(backoff).await; + } + } else if hard_error { + debug!( + attempt, + attempts = self.connect_retry_attempts, + target = %target, + error = %e, + "Upstream connect failed with hard error, failfast is active" + ); + } + last_error = Some(e); + if hard_error { + break; + } + } + } + } + + self.stats.increment_upstream_connect_fail_total(); + self.stats + .observe_upstream_connect_attempts_per_request(attempts_used); + self.stats.observe_upstream_connect_duration_ms( + connect_started_at.elapsed().as_millis() as u64, + false, + ); + + let error = last_error.unwrap_or_else(|| { + ProxyError::Config("Upstream connect attempts exhausted".to_string()) + }); + + let mut guard = self.upstreams.write().await; + if let Some(u) = guard.get_mut(idx) { + // Intermediate attempts are intentionally ignored here. + // Health state is degraded only when the entire connect cycle fails. + u.fails += 1; + warn!( + fails = u.fails, + attempts = self.connect_retry_attempts, + "Upstream failed after retries: {}", + error + ); + if u.fails >= self.unhealthy_fail_threshold { + u.healthy = false; + warn!( + fails = u.fails, + threshold = self.unhealthy_fail_threshold, + "Upstream marked unhealthy" + ); + } + } + Err(error) + } + + pub(super) async fn connect_via_upstream( + &self, + upstream_id: usize, + config: &UpstreamConfig, + target: SocketAddr, + bind_rr: Option>, + connect_timeout: Duration, + ) -> Result<(UpstreamStream, UpstreamEgressInfo)> { + match &config.upstream_type { + UpstreamType::Direct { + interface, + bind_addresses, + bindtodevice, + } => { + let bind_ip = Self::resolve_bind_address( + interface, + bind_addresses, + target, + bind_rr.as_deref(), + true, + ); + if bind_ip.is_none() && bind_addresses.as_ref().is_some_and(|v| !v.is_empty()) { + return Err(ProxyError::Config(format!( + "No valid bind_addresses for target family {target}" + ))); + } + + let socket = create_outgoing_socket_bound(target, bind_ip)?; + if let Some(device) = bindtodevice.as_deref().filter(|value| !value.is_empty()) { + bind_outgoing_socket_to_device(&socket, device).map_err(ProxyError::Io)?; + debug!(bindtodevice = %device, target = %target, "Pinned socket to interface"); + } + if let Some(ip) = bind_ip { + debug!(bind = %ip, target = %target, "Bound outgoing socket"); + } else if interface.is_some() || bind_addresses.is_some() { + debug!(target = %target, "No matching bind address for target family"); + } + + socket.set_nonblocking(true)?; + match socket.connect(&target.into()) { + Ok(()) => {} + Err(err) + if err.raw_os_error() == Some(libc::EINPROGRESS) + || err.kind() == std::io::ErrorKind::WouldBlock => {} + Err(err) => return Err(ProxyError::Io(err)), + } + + let std_stream: std::net::TcpStream = socket.into(); + let stream = TcpStream::from_std(std_stream)?; + + match tokio::time::timeout(connect_timeout, stream.writable()).await { + Ok(Ok(())) => {} + Ok(Err(e)) => return Err(ProxyError::Io(e)), + Err(_) => { + return Err(ProxyError::ConnectionTimeout { + addr: target.to_string(), + }); + } + } + if let Some(e) = stream.take_error()? { + return Err(ProxyError::Io(e)); + } + + let local_addr = stream.local_addr().ok(); + Ok(( + UpstreamStream::Tcp(stream), + UpstreamEgressInfo { + upstream_id, + route_kind: UpstreamRouteKind::Direct, + local_addr, + direct_bind_ip: bind_ip, + socks_bound_addr: None, + socks_proxy_addr: None, + }, + )) + } + UpstreamType::Socks4 { + address, + interface, + user_id, + } => { + // Try to parse as SocketAddr first (IP:port), otherwise treat as hostname:port + let mut stream = if let Ok(proxy_addr) = address.parse::() { + // IP:port format - use socket with optional interface binding + let bind_ip = Self::resolve_bind_address( + interface, + &None, + proxy_addr, + bind_rr.as_deref(), + false, + ); + + let socket = create_outgoing_socket_bound(proxy_addr, bind_ip)?; + + socket.set_nonblocking(true)?; + match socket.connect(&proxy_addr.into()) { + Ok(()) => {} + Err(err) + if err.raw_os_error() == Some(libc::EINPROGRESS) + || err.kind() == std::io::ErrorKind::WouldBlock => {} + Err(err) => return Err(ProxyError::Io(err)), + } + + let std_stream: std::net::TcpStream = socket.into(); + let stream = TcpStream::from_std(std_stream)?; + + match tokio::time::timeout(connect_timeout, stream.writable()).await { + Ok(Ok(())) => {} + Ok(Err(e)) => return Err(ProxyError::Io(e)), + Err(_) => { + return Err(ProxyError::ConnectionTimeout { + addr: proxy_addr.to_string(), + }); + } + } + if let Some(e) = stream.take_error()? { + return Err(ProxyError::Io(e)); + } + stream + } else { + // Hostname:port format - use tokio DNS resolution + // Note: interface binding is not supported for hostnames + if interface.is_some() { + warn!( + "SOCKS4 interface binding is not supported for hostname addresses, ignoring" + ); + } + self.connect_hostname_with_dns_override(address, connect_timeout) + .await? + }; + + // replace socks user_id with config.selected_scope, if set + let scope: Option<&str> = + Some(config.selected_scope.as_str()).filter(|s| !s.is_empty()); + let _user_id: Option<&str> = scope.or(user_id.as_deref()); + + let bound = match tokio::time::timeout( + connect_timeout, + connect_socks4(&mut stream, target, _user_id), + ) + .await + { + Ok(Ok(bound)) => bound, + Ok(Err(e)) => return Err(e), + Err(_) => { + return Err(ProxyError::ConnectionTimeout { + addr: target.to_string(), + }); + } + }; + let local_addr = stream.local_addr().ok(); + let socks_proxy_addr = stream.peer_addr().ok(); + Ok(( + UpstreamStream::Tcp(stream), + UpstreamEgressInfo { + upstream_id, + route_kind: UpstreamRouteKind::Socks4, + local_addr, + direct_bind_ip: None, + socks_bound_addr: Some(bound.addr), + socks_proxy_addr, + }, + )) + } + UpstreamType::Socks5 { + address, + interface, + username, + password, + } => { + // Try to parse as SocketAddr first (IP:port), otherwise treat as hostname:port + let mut stream = if let Ok(proxy_addr) = address.parse::() { + // IP:port format - use socket with optional interface binding + let bind_ip = Self::resolve_bind_address( + interface, + &None, + proxy_addr, + bind_rr.as_deref(), + false, + ); + + let socket = create_outgoing_socket_bound(proxy_addr, bind_ip)?; + + socket.set_nonblocking(true)?; + match socket.connect(&proxy_addr.into()) { + Ok(()) => {} + Err(err) + if err.raw_os_error() == Some(libc::EINPROGRESS) + || err.kind() == std::io::ErrorKind::WouldBlock => {} + Err(err) => return Err(ProxyError::Io(err)), + } + + let std_stream: std::net::TcpStream = socket.into(); + let stream = TcpStream::from_std(std_stream)?; + + match tokio::time::timeout(connect_timeout, stream.writable()).await { + Ok(Ok(())) => {} + Ok(Err(e)) => return Err(ProxyError::Io(e)), + Err(_) => { + return Err(ProxyError::ConnectionTimeout { + addr: proxy_addr.to_string(), + }); + } + } + if let Some(e) = stream.take_error()? { + return Err(ProxyError::Io(e)); + } + stream + } else { + // Hostname:port format - use tokio DNS resolution + // Note: interface binding is not supported for hostnames + if interface.is_some() { + warn!( + "SOCKS5 interface binding is not supported for hostname addresses, ignoring" + ); + } + self.connect_hostname_with_dns_override(address, connect_timeout) + .await? + }; + + debug!(config = ?config, "Socks5 connection"); + // replace socks user:pass with config.selected_scope, if set + let scope: Option<&str> = + Some(config.selected_scope.as_str()).filter(|s| !s.is_empty()); + let _username: Option<&str> = scope.or(username.as_deref()); + let _password: Option<&str> = scope.or(password.as_deref()); + + let bound = match tokio::time::timeout( + connect_timeout, + connect_socks5(&mut stream, target, _username, _password), + ) + .await + { + Ok(Ok(bound)) => bound, + Ok(Err(e)) => return Err(e), + Err(_) => { + return Err(ProxyError::ConnectionTimeout { + addr: target.to_string(), + }); + } + }; + let local_addr = stream.local_addr().ok(); + let socks_proxy_addr = stream.peer_addr().ok(); + Ok(( + UpstreamStream::Tcp(stream), + UpstreamEgressInfo { + upstream_id, + route_kind: UpstreamRouteKind::Socks5, + local_addr, + direct_bind_ip: None, + socks_bound_addr: Some(bound.addr), + socks_proxy_addr, + }, + )) + } + UpstreamType::Shadowsocks { url, interface } => { + let stream = connect_shadowsocks(url, interface, target, connect_timeout).await?; + let local_addr = stream.get_ref().local_addr().ok(); + Ok(( + UpstreamStream::Shadowsocks(Box::new(stream)), + UpstreamEgressInfo { + upstream_id, + route_kind: UpstreamRouteKind::Shadowsocks, + local_addr, + direct_bind_ip: None, + socks_bound_addr: None, + socks_proxy_addr: None, + }, + )) + } + } + } +} diff --git a/src/transport/upstream/health_checks.rs b/src/transport/upstream/health_checks.rs new file mode 100644 index 0000000..fe1f43f --- /dev/null +++ b/src/transport/upstream/health_checks.rs @@ -0,0 +1,243 @@ +use super::*; + +impl UpstreamManager { + /// Background health check based on reachable DC groups through each upstream. + /// Upstream stays healthy while at least `MIN_HEALTHY_DC_GROUPS` groups are reachable. + pub async fn run_health_checks( + &self, + prefer_ipv6: bool, + ipv4_enabled: bool, + ipv6_enabled: bool, + dc_overrides: HashMap>, + ) { + let (health_ipv4_enabled, health_ipv6_enabled) = { + let guard = self.upstreams.read().await; + ( + ipv4_enabled + || guard + .iter() + .any(|upstream| upstream.config.ipv4 == Some(true)), + ipv6_enabled + || guard + .iter() + .any(|upstream| upstream.config.ipv6 == Some(true)), + ) + }; + let groups = Self::build_health_check_groups( + health_ipv4_enabled, + health_ipv6_enabled, + &dc_overrides, + ); + let required_healthy_groups = Self::required_healthy_group_count(groups.len()); + let mut endpoint_rotation: HashMap<(usize, i16, bool), usize> = HashMap::new(); + + if groups.is_empty() { + warn!("No DC groups available for upstream health-checks"); + } + + loop { + tokio::time::sleep(Duration::from_secs(HEALTH_CHECK_INTERVAL_SECS)).await; + + if groups.is_empty() || required_healthy_groups == 0 { + continue; + } + + let target_upstreams: Vec = { + let guard = self.upstreams.read().await; + let has_unscoped = guard + .iter() + .any(|upstream| Self::is_unscoped_upstream(&upstream.config)); + guard + .iter() + .enumerate() + .filter(|(_, upstream)| { + Self::should_check_in_default_dc_connectivity( + has_unscoped, + &upstream.config, + ) + }) + .map(|(idx, _)| idx) + .collect() + }; + + for i in target_upstreams { + let (config, bind_rr) = { + let guard = self.upstreams.read().await; + let u = &guard[i]; + (u.config.clone(), u.bind_rr.clone()) + }; + let (upstream_ipv4_enabled, upstream_ipv6_enabled) = + Self::resolve_probe_dc_families(&config, ipv4_enabled, ipv6_enabled); + let upstream_prefer_ipv6 = config.prefer_ipv6(prefer_ipv6); + + let mut healthy_groups = 0usize; + let mut latency_updates: Vec<(usize, f64)> = Vec::new(); + + for group in &groups { + let mut group_ok = false; + let mut group_rtt_ms = None; + + for (is_primary, endpoints) in + Self::health_check_endpoint_order(group, upstream_prefer_ipv6) + { + if endpoints.is_empty() { + continue; + } + + let filtered_endpoints: Vec = endpoints + .iter() + .copied() + .filter(|endpoint| { + if endpoint.is_ipv4() { + upstream_ipv4_enabled + } else { + upstream_ipv6_enabled + } + }) + .collect(); + + if filtered_endpoints.is_empty() { + continue; + } + + let rotation_key = (i, group.dc_idx, is_primary); + let start_idx = *endpoint_rotation.entry(rotation_key).or_insert(0) + % filtered_endpoints.len(); + let mut next_idx = (start_idx + 1) % filtered_endpoints.len(); + + for step in 0..filtered_endpoints.len() { + let endpoint_idx = (start_idx + step) % filtered_endpoints.len(); + let endpoint = filtered_endpoints[endpoint_idx]; + + let start = Instant::now(); + let result = tokio::time::timeout( + Duration::from_secs(HEALTH_CHECK_CONNECT_TIMEOUT_SECS), + self.connect_via_upstream( + i, + &config, + endpoint, + Some(bind_rr.clone()), + Duration::from_secs(HEALTH_CHECK_CONNECT_TIMEOUT_SECS), + ), + ) + .await; + + match result { + Ok(Ok(_stream)) => { + group_ok = true; + group_rtt_ms = Some(start.elapsed().as_secs_f64() * 1000.0); + next_idx = (endpoint_idx + 1) % filtered_endpoints.len(); + break; + } + Ok(Err(e)) => { + debug!( + upstream = i, + dc = group.dc_idx, + endpoint = %endpoint, + primary = is_primary, + error = %e, + "Health-check endpoint failed" + ); + } + Err(_) => { + debug!( + upstream = i, + dc = group.dc_idx, + endpoint = %endpoint, + primary = is_primary, + "Health-check endpoint timed out" + ); + } + } + } + + endpoint_rotation.insert(rotation_key, next_idx); + + if group_ok { + break; + } + } + + if group_ok { + healthy_groups += 1; + if let (Some(dc_array_idx), Some(rtt_ms)) = + (UpstreamState::dc_array_idx(group.dc_idx), group_rtt_ms) + { + latency_updates.push((dc_array_idx, rtt_ms)); + } + } + } + + let mut guard = self.upstreams.write().await; + let u = &mut guard[i]; + + for (dc_array_idx, rtt_ms) in latency_updates { + u.dc_latency[dc_array_idx].update(rtt_ms); + } + + if healthy_groups >= required_healthy_groups { + if !u.healthy { + info!( + upstream = i, + healthy_groups, + total_groups = groups.len(), + required_groups = required_healthy_groups, + "Upstream recovered by DC-group health threshold" + ); + } + u.healthy = true; + u.fails = 0; + } else { + u.fails += 1; + debug!( + upstream = i, + healthy_groups, + total_groups = groups.len(), + required_groups = required_healthy_groups, + fails = u.fails, + "Upstream health-check below DC-group threshold" + ); + if u.fails >= self.unhealthy_fail_threshold { + u.healthy = false; + warn!( + upstream = i, + healthy_groups, + total_groups = groups.len(), + required_groups = required_healthy_groups, + fails = u.fails, + threshold = self.unhealthy_fail_threshold, + "Upstream unhealthy (insufficient reachable DC groups)" + ); + } + } + + u.last_check = std::time::Instant::now(); + } + } + } + + /// Get the preferred IP for a DC (for use by other components) + #[allow(dead_code)] + pub async fn get_dc_ip_preference(&self, dc_idx: i16) -> Option { + let guard = self.upstreams.read().await; + if guard.is_empty() { + return None; + } + + UpstreamState::dc_array_idx(dc_idx).map(|idx| guard[0].dc_ip_pref[idx]) + } + + /// Get preferred DC address based on config preference + #[allow(dead_code)] + pub async fn get_dc_addr(&self, dc_idx: i16, prefer_ipv6: bool) -> Option { + let arr_idx = UpstreamState::dc_array_idx(dc_idx)?; + + let ip = if prefer_ipv6 { + TG_DATACENTERS_V6[arr_idx] + } else { + TG_DATACENTERS_V4[arr_idx] + }; + + Some(SocketAddr::new(ip, TG_DATACENTER_PORT)) + } +} diff --git a/src/transport/upstream/manager_config.rs b/src/transport/upstream/manager_config.rs new file mode 100644 index 0000000..bb47662 --- /dev/null +++ b/src/transport/upstream/manager_config.rs @@ -0,0 +1,470 @@ +use super::*; + +impl UpstreamManager { + pub(super) fn is_unscoped_upstream(upstream: &UpstreamConfig) -> bool { + upstream.scopes.is_empty() + } + + pub(super) fn should_check_in_default_dc_connectivity( + has_unscoped: bool, + upstream: &UpstreamConfig, + ) -> bool { + !has_unscoped || Self::is_unscoped_upstream(upstream) + } + + pub fn new( + configs: Vec, + connect_retry_attempts: u32, + connect_retry_backoff_ms: u64, + connect_budget_ms: u64, + tg_connect_timeout_secs: u64, + unhealthy_fail_threshold: u32, + connect_failfast_hard_errors: bool, + stats: Arc, + ) -> Self { + let states = configs + .into_iter() + .filter(|c| c.enabled) + .map(UpstreamState::new) + .collect(); + + Self { + upstreams: Arc::new(RwLock::new(states)), + connect_retry_attempts: connect_retry_attempts.max(1), + connect_retry_backoff: Duration::from_millis(connect_retry_backoff_ms), + connect_budget: Duration::from_millis(connect_budget_ms.max(1)), + tg_connect_timeout_secs: tg_connect_timeout_secs.max(1), + unhealthy_fail_threshold: unhealthy_fail_threshold.max(1), + connect_failfast_hard_errors, + no_upstreams_warn_epoch_ms: Arc::new(AtomicU64::new(0)), + no_healthy_warn_epoch_ms: Arc::new(AtomicU64::new(0)), + stats, + dns_resolver: Arc::new(GenerationDnsResolver::default()), + } + } + + pub(crate) fn with_dns_overrides(self, entries: &[String]) -> Result { + self.dns_resolver.apply_entries(entries)?; + Ok(self) + } + + pub(crate) fn update_dns_overrides(&self, entries: &[String]) -> Result<()> { + self.dns_resolver.apply_entries(entries) + } + + pub(crate) fn dns_resolver(&self) -> Arc { + Arc::clone(&self.dns_resolver) + } + + pub(crate) async fn resolve_all(&self, host: &str, port: u16) -> Result> { + if let Some(addr) = self.dns_resolver.resolve_socket_addr(host, port) { + return Ok(vec![addr]); + } + let addrs = tokio::net::lookup_host((host, port)) + .await + .map_err(ProxyError::Io)? + .take(DNS_RESULT_MAX_ADDRESSES) + .collect::>(); + if addrs.is_empty() { + return Err(ProxyError::Proxy(format!( + "DNS returned no addresses for {host}:{port}" + ))); + } + Ok(addrs) + } + + pub(crate) async fn resolve_hostname(&self, host: &str, port: u16) -> Result { + let addrs = self.resolve_all(host, port).await?; + if let Some(addr) = addrs.iter().copied().find(SocketAddr::is_ipv4) { + return Ok(addr); + } + addrs.first().copied().ok_or_else(|| { + ProxyError::Proxy(format!("DNS returned no addresses for {host}:{port}")) + }) + } + + pub(super) fn now_epoch_ms() -> u64 { + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_millis() as u64 + } + + pub(super) fn should_emit_warn(last_epoch_ms: &AtomicU64, cooldown_ms: u64) -> bool { + let now_epoch_ms = Self::now_epoch_ms(); + let previous_epoch_ms = last_epoch_ms.load(Ordering::Relaxed); + if now_epoch_ms.saturating_sub(previous_epoch_ms) < cooldown_ms { + return false; + } + last_epoch_ms + .compare_exchange( + previous_epoch_ms, + now_epoch_ms, + Ordering::AcqRel, + Ordering::Relaxed, + ) + .is_ok() + } + + pub fn try_api_snapshot(&self) -> Option { + let guard = self.upstreams.try_read().ok()?; + let now = std::time::Instant::now(); + + let mut summary = UpstreamApiSummarySnapshot { + configured_total: guard.len(), + ..UpstreamApiSummarySnapshot::default() + }; + let mut upstreams = Vec::with_capacity(guard.len()); + + for (idx, upstream) in guard.iter().enumerate() { + if upstream.healthy { + summary.healthy_total += 1; + } else { + summary.unhealthy_total += 1; + } + + let (route_kind, address) = Self::describe_upstream(&upstream.config.upstream_type); + match route_kind { + UpstreamRouteKind::Direct => summary.direct_total += 1, + UpstreamRouteKind::Socks4 => summary.socks4_total += 1, + UpstreamRouteKind::Socks5 => summary.socks5_total += 1, + UpstreamRouteKind::Shadowsocks => summary.shadowsocks_total += 1, + } + + let mut dc = Vec::with_capacity(NUM_DCS); + for dc_idx in 0..NUM_DCS { + dc.push(UpstreamApiDcSnapshot { + dc: (dc_idx + 1) as i16, + latency_ema_ms: upstream.dc_latency[dc_idx].get(), + ip_preference: upstream.dc_ip_pref[dc_idx], + }); + } + + upstreams.push(UpstreamApiItemSnapshot { + upstream_id: idx, + route_kind, + address, + weight: upstream.config.weight, + scopes: upstream.config.scopes.clone(), + healthy: upstream.healthy, + fails: upstream.fails, + last_check_age_secs: now.saturating_duration_since(upstream.last_check).as_secs(), + effective_latency_ms: upstream.effective_latency(None), + dc, + }); + } + + Some(UpstreamApiSnapshot { summary, upstreams }) + } + + pub async fn api_health_summary(&self) -> UpstreamApiHealthSummary { + let guard = self.upstreams.read().await; + let mut summary = UpstreamApiHealthSummary { + configured_total: guard.len(), + healthy_total: 0, + }; + for upstream in guard.iter() { + if upstream.healthy { + summary.healthy_total += 1; + } + } + summary + } + + pub(super) fn describe_upstream(upstream_type: &UpstreamType) -> (UpstreamRouteKind, String) { + match upstream_type { + UpstreamType::Direct { .. } => (UpstreamRouteKind::Direct, "direct".to_string()), + UpstreamType::Socks4 { address, .. } => (UpstreamRouteKind::Socks4, address.clone()), + UpstreamType::Socks5 { address, .. } => (UpstreamRouteKind::Socks5, address.clone()), + UpstreamType::Shadowsocks { url, .. } => ( + UpstreamRouteKind::Shadowsocks, + sanitize_shadowsocks_url(url).unwrap_or_else(|_| "invalid".to_string()), + ), + } + } + + pub fn api_policy_snapshot(&self) -> UpstreamApiPolicySnapshot { + UpstreamApiPolicySnapshot { + connect_retry_attempts: self.connect_retry_attempts, + connect_retry_backoff_ms: self.connect_retry_backoff.as_millis() as u64, + connect_budget_ms: self.connect_budget.as_millis() as u64, + unhealthy_fail_threshold: self.unhealthy_fail_threshold, + connect_failfast_hard_errors: self.connect_failfast_hard_errors, + } + } + + pub(super) fn resolve_probe_dc_families( + upstream: &UpstreamConfig, + ipv4_available: bool, + ipv6_available: bool, + ) -> (bool, bool) { + ( + upstream.ipv4.unwrap_or(ipv4_available), + upstream.ipv6.unwrap_or(ipv6_available), + ) + } + + pub(super) fn resolve_runtime_dc_families( + upstream: &UpstreamConfig, + dc_preference: IpPreference, + ) -> (bool, bool) { + let (auto_ipv4, auto_ipv6) = match dc_preference { + IpPreference::PreferV4 => (true, false), + IpPreference::PreferV6 => (false, true), + IpPreference::BothWork | IpPreference::Unknown | IpPreference::Unavailable => { + (true, true) + } + }; + + ( + upstream.ipv4.unwrap_or(auto_ipv4), + upstream.ipv6.unwrap_or(auto_ipv6), + ) + } + + pub(super) fn dc_table_addr(dc_idx: i16, ipv6: bool, port: u16) -> Option { + let arr_idx = UpstreamState::dc_array_idx(dc_idx)?; + let ip = if ipv6 { + TG_DATACENTERS_V6[arr_idx] + } else { + TG_DATACENTERS_V4[arr_idx] + }; + Some(SocketAddr::new(ip, port)) + } + + pub(super) fn resolve_runtime_dc_target( + target: SocketAddr, + dc_idx: Option, + upstream: &UpstreamConfig, + dc_preference: IpPreference, + ) -> Result { + let (allow_ipv4, allow_ipv6) = Self::resolve_runtime_dc_families(upstream, dc_preference); + let preferred_ipv6 = match dc_preference { + IpPreference::PreferV6 => Some(true), + IpPreference::PreferV4 => Some(false), + IpPreference::BothWork | IpPreference::Unknown | IpPreference::Unavailable => { + upstream.prefer.map(|prefer| prefer == 6) + } + }; + if let Some(preferred_ipv6) = preferred_ipv6 + && target.is_ipv6() != preferred_ipv6 + { + let preferred_allowed = if preferred_ipv6 { + allow_ipv6 + } else { + allow_ipv4 + }; + if preferred_allowed { + if let Some(dc_idx) = dc_idx + && let Some(remapped) = + Self::dc_table_addr(dc_idx, preferred_ipv6, target.port()) + { + return Ok(remapped); + } + } + } + + if (target.is_ipv4() && allow_ipv4) || (target.is_ipv6() && allow_ipv6) { + return Ok(target); + } + + if !allow_ipv4 && !allow_ipv6 { + return Err(ProxyError::Config(format!( + "Upstream DC family policy blocks all families for target {target}" + ))); + } + + let Some(dc_idx) = dc_idx else { + return Err(ProxyError::Config(format!( + "Upstream DC family policy cannot remap target {target} without dc_idx" + ))); + }; + + let remapped = if target.is_ipv4() { + if allow_ipv6 { + Self::dc_table_addr(dc_idx, true, target.port()) + } else { + None + } + } else if allow_ipv4 { + Self::dc_table_addr(dc_idx, false, target.port()) + } else { + None + }; + + remapped.ok_or_else(|| { + ProxyError::Config(format!( + "Upstream DC family policy rejected target {target} (dc_idx={dc_idx})" + )) + }) + } + + #[cfg(unix)] + pub(super) fn resolve_interface_addrs(name: &str, want_ipv6: bool) -> Vec { + use nix::ifaddrs::getifaddrs; + + let mut out = Vec::new(); + if let Ok(addrs) = getifaddrs() { + for iface in addrs { + if iface.interface_name != name { + continue; + } + if let Some(address) = iface.address { + if let Some(v4) = address.as_sockaddr_in() { + if !want_ipv6 { + out.push(IpAddr::V4(v4.ip())); + } + } else if let Some(v6) = address.as_sockaddr_in6() + && want_ipv6 + { + out.push(IpAddr::V6(v6.ip())); + } + } + } + } + out.sort_unstable(); + out.dedup(); + out + } + + pub(crate) fn resolve_bind_address( + interface: &Option, + bind_addresses: &Option>, + target: SocketAddr, + rr: Option<&AtomicUsize>, + validate_ip_on_interface: bool, + ) -> Option { + let want_ipv6 = target.is_ipv6(); + + if let Some(addrs) = bind_addresses.as_ref().filter(|v| !v.is_empty()) { + let mut candidates: Vec = addrs + .iter() + .filter_map(|s| s.parse::().ok()) + .filter(|ip| ip.is_ipv6() == want_ipv6) + .collect(); + + // Explicit bind IP has strict priority over interface auto-selection. + if validate_ip_on_interface + && let Some(iface) = interface + && iface.parse::().is_err() + { + #[cfg(unix)] + { + let iface_addrs = Self::resolve_interface_addrs(iface, want_ipv6); + if !iface_addrs.is_empty() { + candidates.retain(|ip| { + let ok = iface_addrs.contains(ip); + if !ok { + warn!( + interface = %iface, + bind_ip = %ip, + target = %target, + "Configured bind address is not assigned to interface" + ); + } + ok + }); + } else if !candidates.is_empty() { + warn!( + interface = %iface, + target = %target, + "Configured interface has no addresses for target family" + ); + candidates.clear(); + } + } + } + + if !candidates.is_empty() { + if let Some(counter) = rr { + let idx = counter.fetch_add(1, Ordering::Relaxed) % candidates.len(); + return Some(candidates[idx]); + } + return candidates.first().copied(); + } + + if validate_ip_on_interface + && interface + .as_ref() + .is_some_and(|iface| iface.parse::().is_err()) + { + warn!( + interface = interface.as_deref().unwrap_or(""), + target = %target, + "No valid bind_addresses left for interface" + ); + } + + return None; + } + + if let Some(iface) = interface { + if let Ok(ip) = iface.parse::() { + if ip.is_ipv6() == want_ipv6 { + return Some(ip); + } + } else { + #[cfg(unix)] + if let Some(ip) = resolve_interface_ip(iface, want_ipv6) { + return Some(ip); + } + } + } + + None + } + + pub(super) async fn connect_hostname_with_dns_override( + &self, + address: &str, + connect_timeout: Duration, + ) -> Result { + if let Some((host, port)) = split_host_port(address) + && let Some(addr) = self.dns_resolver.resolve_socket_addr(&host, port) + { + return match tokio::time::timeout(connect_timeout, TcpStream::connect(addr)).await { + Ok(Ok(stream)) => Ok(stream), + Ok(Err(e)) => Err(ProxyError::Io(e)), + Err(_) => Err(ProxyError::ConnectionTimeout { + addr: addr.to_string(), + }), + }; + } + + match tokio::time::timeout(connect_timeout, TcpStream::connect(address)).await { + Ok(Ok(stream)) => Ok(stream), + Ok(Err(e)) => Err(ProxyError::Io(e)), + Err(_) => Err(ProxyError::ConnectionTimeout { + addr: address.to_string(), + }), + } + } + + pub(super) fn retry_backoff_with_jitter(&self) -> Duration { + if self.connect_retry_backoff.is_zero() { + return Duration::ZERO; + } + let base_ms = self.connect_retry_backoff.as_millis() as u64; + if base_ms == 0 { + return self.connect_retry_backoff; + } + let jitter_cap_ms = (base_ms / 2).max(1); + let jitter_ms = rand::rng().random_range(0..=jitter_cap_ms); + Duration::from_millis(base_ms.saturating_add(jitter_ms)) + } + + pub(super) fn is_hard_connect_error(error: &ProxyError) -> bool { + match error { + ProxyError::Config(_) | ProxyError::ConnectionRefused { .. } => true, + ProxyError::Io(ioe) => matches!( + ioe.kind(), + std::io::ErrorKind::ConnectionRefused + | std::io::ErrorKind::AddrInUse + | std::io::ErrorKind::AddrNotAvailable + | std::io::ErrorKind::InvalidInput + | std::io::ErrorKind::Unsupported + ), + _ => false, + } + } +} diff --git a/src/transport/upstream/selection.rs b/src/transport/upstream/selection.rs new file mode 100644 index 0000000..9f51ac3 --- /dev/null +++ b/src/transport/upstream/selection.rs @@ -0,0 +1,185 @@ +use super::*; + +impl UpstreamManager { + /// Select upstream using latency-weighted random selection. + pub(super) async fn select_upstream( + &self, + dc_idx: Option, + scope: Option<&str>, + ) -> Option { + let upstreams = self.upstreams.read().await; + if upstreams.is_empty() { + return None; + } + // Scope filter: + // If scope is set: only scoped and matched items + // If scope is not set: only unscoped items + let filtered_upstreams: Vec = upstreams + .iter() + .enumerate() + .filter(|(_, u)| { + scope.map_or(u.config.scopes.is_empty(), |req_scope| { + u.config + .scopes + .split(',') + .map(str::trim) + .any(|s| s == req_scope) + }) + }) + .map(|(i, _)| i) + .collect(); + + // Healthy filter + let healthy: Vec = filtered_upstreams + .iter() + .filter(|&&i| upstreams[i].healthy) + .copied() + .collect(); + + if filtered_upstreams.is_empty() { + if Self::should_emit_warn(self.no_upstreams_warn_epoch_ms.as_ref(), 5_000) { + warn!( + scope = scope, + "No upstreams available! Using first (direct?)" + ); + } + return None; + } + + if healthy.is_empty() { + if Self::should_emit_warn(self.no_healthy_warn_epoch_ms.as_ref(), 5_000) { + warn!( + scope = scope, + "No healthy upstreams available! Using random." + ); + } + return Some(filtered_upstreams[rand::rng().random_range(0..filtered_upstreams.len())]); + } + + if healthy.len() == 1 { + return Some(healthy[0]); + } + + let weights: Vec<(usize, f64)> = healthy + .iter() + .map(|&i| { + let base = upstreams[i].config.weight as f64; + let latency_factor = upstreams[i] + .effective_latency(dc_idx) + .map(|ms| if ms > 1.0 { 1000.0 / ms } else { 1000.0 }) + .unwrap_or(1.0); + + (i, base * latency_factor) + }) + .collect(); + + let total: f64 = weights.iter().map(|(_, w)| w).sum(); + + if total <= 0.0 { + return Some(healthy[rand::rng().random_range(0..healthy.len())]); + } + + let mut choice: f64 = rand::rng().random_range(0.0..total); + + for &(idx, weight) in &weights { + if choice < weight { + trace!( + upstream = idx, + dc = ?dc_idx, + weight = format!("{:.2}", weight), + total = format!("{:.2}", total), + "Upstream selected" + ); + return Some(idx); + } + choice -= weight; + } + + Some(healthy[0]) + } + + /// Connect to target through a selected upstream. + pub async fn connect( + &self, + target: SocketAddr, + dc_idx: Option, + scope: Option<&str>, + ) -> Result { + let idx = self + .select_upstream(dc_idx, scope) + .await + .ok_or_else(|| ProxyError::Config("No upstreams available".to_string()))?; + + let (mut upstream, bind_rr, dc_preference) = { + let guard = self.upstreams.read().await; + let state = &guard[idx]; + let dc_preference = dc_idx + .and_then(UpstreamState::dc_array_idx) + .map(|dc_array_idx| state.dc_ip_pref[dc_array_idx]) + .unwrap_or(IpPreference::Unknown); + ( + state.config.clone(), + Some(state.bind_rr.clone()), + dc_preference, + ) + }; + + if let Some(s) = scope { + upstream.selected_scope = s.to_string(); + } + + let target = if dc_idx.is_some() { + Self::resolve_runtime_dc_target(target, dc_idx, &upstream, dc_preference)? + } else { + target + }; + + let (stream, _) = self + .connect_selected_upstream(idx, upstream, target, dc_idx, bind_rr) + .await?; + Ok(stream) + } + + /// Connect to target through a selected upstream and return egress details. + pub async fn connect_with_details( + &self, + target: SocketAddr, + dc_idx: Option, + scope: Option<&str>, + ) -> Result<(TcpStream, UpstreamEgressInfo)> { + let idx = self + .select_upstream(dc_idx, scope) + .await + .ok_or_else(|| ProxyError::Config("No upstreams available".to_string()))?; + + let (mut upstream, bind_rr, dc_preference) = { + let guard = self.upstreams.read().await; + let state = &guard[idx]; + let dc_preference = dc_idx + .and_then(UpstreamState::dc_array_idx) + .map(|dc_array_idx| state.dc_ip_pref[dc_array_idx]) + .unwrap_or(IpPreference::Unknown); + ( + state.config.clone(), + Some(state.bind_rr.clone()), + dc_preference, + ) + }; + + // Set scope for configuration copy + if let Some(s) = scope { + upstream.selected_scope = s.to_string(); + } + + let target = if dc_idx.is_some() { + Self::resolve_runtime_dc_target(target, dc_idx, &upstream, dc_preference)? + } else { + target + }; + + let (stream, egress) = self + .connect_selected_upstream(idx, upstream, target, dc_idx, bind_rr) + .await?; + Ok((stream.into_tcp()?, egress)) + } +} diff --git a/src/transport/upstream/startup_probe.rs b/src/transport/upstream/startup_probe.rs new file mode 100644 index 0000000..515ebac --- /dev/null +++ b/src/transport/upstream/startup_probe.rs @@ -0,0 +1,420 @@ +use super::*; + +impl UpstreamManager { + /// Ping all Telegram DCs through all upstreams. + /// Tests BOTH IPv6 and IPv4, returns separate results for each. + pub async fn ping_all_dcs( + &self, + prefer_ipv6: bool, + dc_overrides: &HashMap>, + ipv4_enabled: bool, + ipv6_enabled: bool, + ) -> Vec { + let upstreams: Vec<(usize, UpstreamConfig, Arc)> = { + let guard = self.upstreams.read().await; + guard + .iter() + .enumerate() + .map(|(i, u)| (i, u.config.clone(), u.bind_rr.clone())) + .collect() + }; + let has_unscoped = upstreams + .iter() + .any(|(_, cfg, _)| Self::is_unscoped_upstream(cfg)); + + let mut all_results = Vec::new(); + + for (upstream_idx, upstream_config, bind_rr) in &upstreams { + // DC connectivity checks should follow the default routing path. + // Scoped upstreams are included only when no unscoped upstream exists. + if !Self::should_check_in_default_dc_connectivity(has_unscoped, upstream_config) { + continue; + } + + let (upstream_ipv4_enabled, upstream_ipv6_enabled) = + Self::resolve_probe_dc_families(upstream_config, ipv4_enabled, ipv6_enabled); + let upstream_prefer_ipv6 = upstream_config.prefer_ipv6(prefer_ipv6); + let upstream_name = match &upstream_config.upstream_type { + UpstreamType::Direct { + interface, + bind_addresses, + bindtodevice, + } => { + let mut direct_parts = Vec::new(); + if let Some(dev) = interface.as_deref().filter(|v| !v.is_empty()) { + direct_parts.push(format!("dev={dev}")); + } + if let Some(src) = bind_addresses.as_ref().filter(|v| !v.is_empty()) { + direct_parts.push(format!("src={}", src.join(","))); + } + if let Some(device) = bindtodevice.as_deref().filter(|v| !v.is_empty()) { + direct_parts.push(format!("bindtodevice={device}")); + } + if direct_parts.is_empty() { + "direct".to_string() + } else { + format!("direct {}", direct_parts.join(" ")) + } + } + UpstreamType::Socks4 { address, .. } => format!("socks4://{}", address), + UpstreamType::Socks5 { address, .. } => format!("socks5://{}", address), + UpstreamType::Shadowsocks { url, .. } => { + let address = + sanitize_shadowsocks_url(url).unwrap_or_else(|_| "invalid".to_string()); + format!("shadowsocks://{address}") + } + }; + + let mut v6_results = Vec::with_capacity(NUM_DCS); + if upstream_ipv6_enabled { + for dc_zero_idx in 0..NUM_DCS { + let dc_v6 = TG_DATACENTERS_V6[dc_zero_idx]; + let addr_v6 = SocketAddr::new(dc_v6, TG_DATACENTER_PORT); + + let result = tokio::time::timeout( + Duration::from_secs(DC_PING_TIMEOUT_SECS), + self.ping_single_dc( + *upstream_idx, + upstream_config, + Some(bind_rr.clone()), + addr_v6, + ), + ) + .await; + + let ping_result = match result { + Ok(Ok(rtt_ms)) => { + let mut guard = self.upstreams.write().await; + if let Some(u) = guard.get_mut(*upstream_idx) { + u.dc_latency[dc_zero_idx].update(rtt_ms); + } + DcPingResult { + dc_idx: dc_zero_idx + 1, + dc_addr: addr_v6, + rtt_ms: Some(rtt_ms), + error: None, + } + } + Ok(Err(e)) => DcPingResult { + dc_idx: dc_zero_idx + 1, + dc_addr: addr_v6, + rtt_ms: None, + error: Some(e.to_string()), + }, + Err(_) => DcPingResult { + dc_idx: dc_zero_idx + 1, + dc_addr: addr_v6, + rtt_ms: None, + error: Some("timeout".to_string()), + }, + }; + v6_results.push(ping_result); + } + } else { + for dc_zero_idx in 0..NUM_DCS { + let dc_v6 = TG_DATACENTERS_V6[dc_zero_idx]; + v6_results.push(DcPingResult { + dc_idx: dc_zero_idx + 1, + dc_addr: SocketAddr::new(dc_v6, TG_DATACENTER_PORT), + rtt_ms: None, + error: Some(if ipv6_enabled { + "ipv6 disabled by upstream policy".to_string() + } else { + "ipv6 disabled".to_string() + }), + }); + } + } + + let mut v4_results = Vec::with_capacity(NUM_DCS); + if upstream_ipv4_enabled { + for dc_zero_idx in 0..NUM_DCS { + let dc_v4 = TG_DATACENTERS_V4[dc_zero_idx]; + let addr_v4 = SocketAddr::new(dc_v4, TG_DATACENTER_PORT); + + let result = tokio::time::timeout( + Duration::from_secs(DC_PING_TIMEOUT_SECS), + self.ping_single_dc( + *upstream_idx, + upstream_config, + Some(bind_rr.clone()), + addr_v4, + ), + ) + .await; + + let ping_result = match result { + Ok(Ok(rtt_ms)) => { + let mut guard = self.upstreams.write().await; + if let Some(u) = guard.get_mut(*upstream_idx) { + u.dc_latency[dc_zero_idx].update(rtt_ms); + } + DcPingResult { + dc_idx: dc_zero_idx + 1, + dc_addr: addr_v4, + rtt_ms: Some(rtt_ms), + error: None, + } + } + Ok(Err(e)) => DcPingResult { + dc_idx: dc_zero_idx + 1, + dc_addr: addr_v4, + rtt_ms: None, + error: Some(e.to_string()), + }, + Err(_) => DcPingResult { + dc_idx: dc_zero_idx + 1, + dc_addr: addr_v4, + rtt_ms: None, + error: Some("timeout".to_string()), + }, + }; + v4_results.push(ping_result); + } + } else { + for dc_zero_idx in 0..NUM_DCS { + let dc_v4 = TG_DATACENTERS_V4[dc_zero_idx]; + v4_results.push(DcPingResult { + dc_idx: dc_zero_idx + 1, + dc_addr: SocketAddr::new(dc_v4, TG_DATACENTER_PORT), + rtt_ms: None, + error: Some(if ipv4_enabled { + "ipv4 disabled by upstream policy".to_string() + } else { + "ipv4 disabled".to_string() + }), + }); + } + } + + // === Ping DC overrides (v4/v6) === + for (dc_key, addrs) in dc_overrides { + let dc_num: i16 = match dc_key.parse::() { + Ok(v) if v > 0 => v, + Err(_) => { + warn!(dc = %dc_key, "Invalid dc_overrides key, skipping"); + continue; + } + _ => continue, + }; + let dc_idx = dc_num as usize; + for addr_str in addrs { + match addr_str.parse::() { + Ok(addr) => { + let is_v6 = addr.is_ipv6(); + if (is_v6 && !upstream_ipv6_enabled) + || (!is_v6 && !upstream_ipv4_enabled) + { + continue; + } + let result = tokio::time::timeout( + Duration::from_secs(DC_PING_TIMEOUT_SECS), + self.ping_single_dc( + *upstream_idx, + upstream_config, + Some(bind_rr.clone()), + addr, + ), + ) + .await; + + let ping_result = match result { + Ok(Ok(rtt_ms)) => DcPingResult { + dc_idx, + dc_addr: addr, + rtt_ms: Some(rtt_ms), + error: None, + }, + Ok(Err(e)) => DcPingResult { + dc_idx, + dc_addr: addr, + rtt_ms: None, + error: Some(e.to_string()), + }, + Err(_) => DcPingResult { + dc_idx, + dc_addr: addr, + rtt_ms: None, + error: Some("timeout".to_string()), + }, + }; + + if is_v6 { + v6_results.push(ping_result); + } else { + v4_results.push(ping_result); + } + } + Err(_) => { + warn!(dc = %dc_idx, addr = %addr_str, "Invalid dc_overrides address, skipping") + } + } + } + } + + // Check if both IP versions have at least one working DC + let v6_has_working = v6_results.iter().any(|r| r.rtt_ms.is_some()); + let v4_has_working = v4_results.iter().any(|r| r.rtt_ms.is_some()); + let both_available = v6_has_working && v4_has_working; + + // Update IP preference for each DC + { + let mut guard = self.upstreams.write().await; + if let Some(u) = guard.get_mut(*upstream_idx) { + for dc_zero_idx in 0..NUM_DCS { + let v6_ok = v6_results[dc_zero_idx].rtt_ms.is_some(); + let v4_ok = v4_results[dc_zero_idx].rtt_ms.is_some(); + + u.dc_ip_pref[dc_zero_idx] = match (v6_ok, v4_ok) { + (true, true) => IpPreference::BothWork, + (true, false) => IpPreference::PreferV6, + (false, true) => IpPreference::PreferV4, + (false, false) => IpPreference::Unavailable, + }; + } + } + } + + all_results.push(StartupPingResult { + v6_results, + v4_results, + upstream_name, + prefer_ipv6: upstream_prefer_ipv6, + both_available, + }); + } + + all_results + } + + pub(super) async fn ping_single_dc( + &self, + upstream_id: usize, + config: &UpstreamConfig, + bind_rr: Option>, + target: SocketAddr, + ) -> Result { + let start = Instant::now(); + let _ = self + .connect_via_upstream( + upstream_id, + config, + target, + bind_rr, + Duration::from_secs(DC_PING_TIMEOUT_SECS), + ) + .await?; + Ok(start.elapsed().as_secs_f64() * 1000.0) + } + + pub(super) fn required_healthy_group_count(total_groups: usize) -> usize { + if total_groups == 0 { + 0 + } else { + total_groups.min(MIN_HEALTHY_DC_GROUPS) + } + } + + pub(super) fn build_health_check_groups( + ipv4_enabled: bool, + ipv6_enabled: bool, + dc_overrides: &HashMap>, + ) -> Vec { + let mut v4_by_dc: HashMap> = HashMap::new(); + let mut v6_by_dc: HashMap> = HashMap::new(); + + if ipv4_enabled { + for (idx, dc_ip) in TG_DATACENTERS_V4.iter().enumerate() { + let dc_idx = (idx + 1) as i16; + v4_by_dc + .entry(dc_idx) + .or_default() + .push(SocketAddr::new(*dc_ip, TG_DATACENTER_PORT)); + } + } + + if ipv6_enabled { + for (idx, dc_ip) in TG_DATACENTERS_V6.iter().enumerate() { + let dc_idx = (idx + 1) as i16; + v6_by_dc + .entry(dc_idx) + .or_default() + .push(SocketAddr::new(*dc_ip, TG_DATACENTER_PORT)); + } + } + + for (dc_key, addrs) in dc_overrides { + let dc_idx = match dc_key.parse::() { + Ok(v) if v > 0 => v, + _ => { + warn!(dc = %dc_key, "Invalid dc_overrides key for health-check, skipping"); + continue; + } + }; + + for addr_str in addrs { + match addr_str.parse::() { + Ok(addr) if addr.is_ipv6() => { + if ipv6_enabled { + v6_by_dc.entry(dc_idx).or_default().push(addr); + } + } + Ok(addr) => { + if ipv4_enabled { + v4_by_dc.entry(dc_idx).or_default().push(addr); + } + } + Err(_) => { + warn!( + dc = %dc_idx, + addr = %addr_str, + "Invalid dc_overrides address for health-check, skipping" + ); + } + } + } + } + + for addrs in v4_by_dc.values_mut() { + addrs.sort_unstable(); + addrs.dedup(); + } + for addrs in v6_by_dc.values_mut() { + addrs.sort_unstable(); + addrs.dedup(); + } + + let mut all_dcs = BTreeSet::new(); + all_dcs.extend(v4_by_dc.keys().copied()); + all_dcs.extend(v6_by_dc.keys().copied()); + + let mut groups = Vec::with_capacity(all_dcs.len()); + for dc_idx in all_dcs { + let v4_endpoints = v4_by_dc.remove(&dc_idx).unwrap_or_default(); + let v6_endpoints = v6_by_dc.remove(&dc_idx).unwrap_or_default(); + + if v4_endpoints.is_empty() && v6_endpoints.is_empty() { + continue; + } + + groups.push(HealthCheckGroup { + dc_idx, + v4_endpoints, + v6_endpoints, + }); + } + + groups + } + + pub(super) fn health_check_endpoint_order( + group: &HealthCheckGroup, + prefer_ipv6: bool, + ) -> [(bool, &[SocketAddr]); 2] { + if prefer_ipv6 { + [(true, &group.v6_endpoints), (false, &group.v4_endpoints)] + } else { + [(true, &group.v4_endpoints), (false, &group.v6_endpoints)] + } + } +} diff --git a/src/transport/upstream/tests.rs b/src/transport/upstream/tests.rs new file mode 100644 index 0000000..d1968f4 --- /dev/null +++ b/src/transport/upstream/tests.rs @@ -0,0 +1,231 @@ +use super::*; +use std::sync::Arc; + +use crate::stats::Stats; + +const TEST_SHADOWSOCKS_URL: &str = + "ss://2022-blake3-aes-256-gcm:MDEyMzQ1Njc4OTAxMjM0NTY3ODkwMTIzNDU2Nzg5MDE=@127.0.0.1:8388"; + +fn manager_with_dns(entries: &[String]) -> UpstreamManager { + UpstreamManager::new(Vec::new(), 1, 1, 1, 1, 1, false, Arc::new(Stats::new())) + .with_dns_overrides(entries) + .unwrap() +} + +#[tokio::test] +async fn generation_local_dns_overrides_are_isolated_and_case_insensitive() { + let active = manager_with_dns(&["Front.Example:443:192.0.2.10".to_string()]); + let candidate = manager_with_dns(&["front.example:443:[2001:db8::10]".to_string()]); + + assert_eq!( + active.resolve_hostname("front.example", 443).await.unwrap(), + "192.0.2.10:443".parse::().unwrap() + ); + assert_eq!( + candidate + .resolve_hostname("FRONT.EXAMPLE", 443) + .await + .unwrap(), + "[2001:db8::10]:443".parse::().unwrap() + ); + + candidate + .update_dns_overrides(&["front.example:443:192.0.2.20".to_string()]) + .unwrap(); + assert_eq!( + active.resolve_hostname("FRONT.EXAMPLE", 443).await.unwrap(), + "192.0.2.10:443".parse::().unwrap() + ); + assert_eq!( + candidate + .resolve_hostname("front.example", 443) + .await + .unwrap(), + "192.0.2.20:443".parse::().unwrap() + ); +} + +#[test] +fn required_healthy_group_count_applies_three_group_threshold() { + assert_eq!(UpstreamManager::required_healthy_group_count(0), 0); + assert_eq!(UpstreamManager::required_healthy_group_count(1), 1); + assert_eq!(UpstreamManager::required_healthy_group_count(2), 2); + assert_eq!(UpstreamManager::required_healthy_group_count(3), 3); + assert_eq!(UpstreamManager::required_healthy_group_count(5), 3); +} + +#[test] +fn build_health_check_groups_merges_family_endpoints_with_preference() { + let mut overrides = HashMap::new(); + overrides.insert( + "2".to_string(), + vec![ + "203.0.113.10:443".to_string(), + "203.0.113.11:443".to_string(), + "[2001:db8::10]:443".to_string(), + ], + ); + + let groups = UpstreamManager::build_health_check_groups(true, true, &overrides); + let dc2 = groups + .iter() + .find(|g| g.dc_idx == 2) + .expect("dc2 must be present"); + + assert!(dc2.v6_endpoints.iter().all(|addr| addr.is_ipv6())); + assert!(dc2.v4_endpoints.iter().all(|addr| addr.is_ipv4())); + assert!( + dc2.v6_endpoints + .contains(&"[2001:db8::10]:443".parse::().unwrap()) + ); + assert!( + dc2.v4_endpoints + .contains(&"203.0.113.10:443".parse::().unwrap()) + ); + assert!( + dc2.v4_endpoints + .contains(&"203.0.113.11:443".parse::().unwrap()) + ); + + let ordered = UpstreamManager::health_check_endpoint_order(dc2, true); + assert!(ordered[0].1.iter().all(|addr| addr.is_ipv6())); + assert!(ordered[1].1.iter().all(|addr| addr.is_ipv4())); +} + +#[test] +fn build_health_check_groups_keeps_multiple_endpoints_per_group() { + let mut overrides = HashMap::new(); + overrides.insert( + "9".to_string(), + vec![ + "198.51.100.1:443".to_string(), + "198.51.100.2:443".to_string(), + "198.51.100.1:443".to_string(), + ], + ); + + let groups = UpstreamManager::build_health_check_groups(true, false, &overrides); + let dc9 = groups + .iter() + .find(|g| g.dc_idx == 9) + .expect("override-only dc group must be present"); + + assert_eq!(dc9.v4_endpoints.len(), 2); + assert!( + dc9.v4_endpoints + .contains(&"198.51.100.1:443".parse::().unwrap()) + ); + assert!( + dc9.v4_endpoints + .contains(&"198.51.100.2:443".parse::().unwrap()) + ); + assert!(dc9.v6_endpoints.is_empty()); +} + +#[test] +fn hard_connect_error_classification_detects_connection_refused() { + let error = ProxyError::ConnectionRefused { + addr: "127.0.0.1:443".to_string(), + }; + assert!(UpstreamManager::is_hard_connect_error(&error)); +} + +#[test] +fn hard_connect_error_classification_skips_timeouts() { + let error = ProxyError::ConnectionTimeout { + addr: "127.0.0.1:443".to_string(), + }; + assert!(!UpstreamManager::is_hard_connect_error(&error)); +} + +#[test] +fn unscoped_selection_detects_default_route_upstream() { + let mut upstream = UpstreamConfig { + upstream_type: UpstreamType::Direct { + interface: None, + bind_addresses: None, + bindtodevice: None, + }, + weight: 1, + enabled: true, + scopes: String::new(), + selected_scope: String::new(), + ipv4: None, + ipv6: None, + prefer: None, + }; + + assert!(UpstreamManager::is_unscoped_upstream(&upstream)); + upstream.scopes = "local".to_string(); + assert!(!UpstreamManager::is_unscoped_upstream(&upstream)); + assert!(!UpstreamManager::should_check_in_default_dc_connectivity( + true, &upstream + )); + assert!(UpstreamManager::should_check_in_default_dc_connectivity( + false, &upstream + )); +} + +#[test] +fn resolve_bind_address_prefers_explicit_bind_ip() { + let target = "203.0.113.10:443".parse::().unwrap(); + let bind = UpstreamManager::resolve_bind_address( + &Some("198.51.100.20".to_string()), + &Some(vec!["198.51.100.10".to_string()]), + target, + None, + true, + ); + + assert_eq!(bind, Some("198.51.100.10".parse::().unwrap())); +} + +#[test] +fn resolve_bind_address_does_not_fallback_to_interface_when_bind_addresses_present() { + let target = "203.0.113.10:443".parse::().unwrap(); + let bind = UpstreamManager::resolve_bind_address( + &Some("198.51.100.20".to_string()), + &Some(vec!["2001:db8::10".to_string()]), + target, + None, + true, + ); + + assert_eq!(bind, None); +} + +#[test] +fn api_snapshot_reports_shadowsocks_as_sanitized_route() { + let manager = UpstreamManager::new( + vec![UpstreamConfig { + upstream_type: UpstreamType::Shadowsocks { + url: TEST_SHADOWSOCKS_URL.to_string(), + interface: None, + }, + weight: 2, + enabled: true, + scopes: String::new(), + selected_scope: String::new(), + ipv4: None, + ipv6: None, + prefer: None, + }], + 1, + 100, + 1000, + 10, + 1, + false, + Arc::new(Stats::new()), + ); + + let snapshot = manager.try_api_snapshot().expect("snapshot"); + assert_eq!(snapshot.summary.configured_total, 1); + assert_eq!(snapshot.summary.shadowsocks_total, 1); + assert_eq!(snapshot.upstreams.len(), 1); + assert_eq!( + snapshot.upstreams[0].route_kind, + UpstreamRouteKind::Shadowsocks + ); + assert_eq!(snapshot.upstreams[0].address, "127.0.0.1:8388"); +} diff --git a/src/web/manager/operator_lifecycle.rs b/src/web/manager/operator_lifecycle.rs index 06df139..fed6171 100644 --- a/src/web/manager/operator_lifecycle.rs +++ b/src/web/manager/operator_lifecycle.rs @@ -18,8 +18,8 @@ pub(crate) use status::{ }; // Admission fencing and registration-drain synchronization. mod admission; -use admission::{OperatorAdmission, OperatorAdmissionRejection}; pub(super) use admission::OperatorRegistration; +use admission::{OperatorAdmission, OperatorAdmissionRejection}; // Mutable and published lifecycle state storage. mod state; use state::{ActiveDrain, OperatorLifecycleInner, OperatorSnapshot, WorkCounts}; @@ -246,9 +246,9 @@ impl WebProcessRuntime { self.operator_lifecycle .transition_locked(&mut inner, OperatorLifecycleState::Paused); } - self.operator_lifecycle.admission.close( - OperatorAdmissionRejection::for_state(inner.state), - ); + self.operator_lifecycle + .admission + .close(OperatorAdmissionRejection::for_state(inner.state)); self.operator_lifecycle.publish_locked(&inner); } self.operator_lifecycle diff --git a/src/web/manager/operator_lifecycle/admission.rs b/src/web/manager/operator_lifecycle/admission.rs index 5462241..741163f 100644 --- a/src/web/manager/operator_lifecycle/admission.rs +++ b/src/web/manager/operator_lifecycle/admission.rs @@ -53,9 +53,7 @@ impl OperatorAdmissionRejection { match self { Self::Paused => crate::web::telemetry::WebRejectionReason::OperatorPaused, Self::Draining => crate::web::telemetry::WebRejectionReason::OperatorDraining, - Self::ForceClosing => { - crate::web::telemetry::WebRejectionReason::OperatorForceClosing - } + Self::ForceClosing => crate::web::telemetry::WebRejectionReason::OperatorForceClosing, Self::Drained => crate::web::telemetry::WebRejectionReason::OperatorDrained, Self::RuntimeClosed => crate::web::telemetry::WebRejectionReason::RuntimeClosed, } @@ -90,7 +88,7 @@ impl OperatorAdmission { loop { if state & OPERATOR_ADMISSION_CLOSED != 0 { return Err( - OperatorAdmissionRejection::from_admission_state(state).telemetry_reason(), + OperatorAdmissionRejection::from_admission_state(state).telemetry_reason() ); } if state & OPERATOR_REGISTRATION_COUNT == OPERATOR_REGISTRATION_COUNT { @@ -113,15 +111,11 @@ impl OperatorAdmission { let reason = (reason as usize) << OPERATOR_REJECTION_SHIFT; let mut state = self.state.load(Ordering::Acquire); loop { - let next = (state & OPERATOR_REGISTRATION_COUNT) - | OPERATOR_ADMISSION_CLOSED - | reason; - match self.state.compare_exchange_weak( - state, - next, - Ordering::AcqRel, - Ordering::Acquire, - ) { + let next = (state & OPERATOR_REGISTRATION_COUNT) | OPERATOR_ADMISSION_CLOSED | reason; + match self + .state + .compare_exchange_weak(state, next, Ordering::AcqRel, Ordering::Acquire) + { Ok(_) => return, Err(observed) => state = observed, } diff --git a/src/web/manager/operator_lifecycle/tests.rs b/src/web/manager/operator_lifecycle/tests.rs index 2ff32b7..a72a705 100644 --- a/src/web/manager/operator_lifecycle/tests.rs +++ b/src/web/manager/operator_lifecycle/tests.rs @@ -138,11 +138,8 @@ async fn drain_request_returns_after_registering_its_worker() { let (runtime, generation) = test_runtime(); let registration = runtime.try_operator_admission().unwrap(); let drain_runtime = Arc::clone(&runtime); - let drain = tokio::spawn(async move { - drain_runtime - .drain_operator(Duration::from_secs(30)) - .await - }); + let drain = + tokio::spawn(async move { drain_runtime.drain_operator(Duration::from_secs(30)).await }); tokio::time::timeout(Duration::from_secs(1), async { while runtime.operator_lifecycle_status().state != OperatorLifecycleState::Draining { diff --git a/src/web/manager/status.rs b/src/web/manager/status.rs index 3bdf96d..825542c 100644 --- a/src/web/manager/status.rs +++ b/src/web/manager/status.rs @@ -515,53 +515,8 @@ impl WebProcessRuntime { } } -/// Opaque session-reference validation failure. -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub(crate) enum SessionRefError { - /// The reference does not use the canonical versioned shape. - Invalid, - /// The reference belongs to another process runtime. - StaleInstance, -} - -/// Tests immutable candidate fields before any optional state-lock read. -pub(super) fn immutable_matches( - session: &crate::web::session::WebSession, - index: &super::state::LiveSessionIndex, - filter: &SessionFilter, -) -> bool { - filter - .trace_session_id - .is_none_or(|value| value == session.trace_session_id()) - && filter - .client_ip - .is_none_or(|value| value == session.client_ip()) - && filter - .host - .as_deref() - .is_none_or(|value| value == session.profile_host()) - && filter - .user - .as_deref() - .is_none_or(|value| value == session.profile_user()) - && filter - .key_id - .as_deref() - .is_none_or(|value| session.key_id() == value) - && filter - .carrier - .is_none_or(|value| value == session.carrier()) - && filter - .user_agent_id - .is_none_or(|value| index.user_agent_id == Some(value)) -} - -fn permits(semaphore: &Arc, capacity: usize) -> PermitStatus { - let available = semaphore.available_permits().min(capacity); - PermitStatus { - used: capacity.saturating_sub(available), - available, - capacity, - closed: semaphore.is_closed(), - } -} +// Opaque session references and immutable filter matching. +mod reference; +pub(crate) use reference::SessionRefError; +pub(super) use reference::immutable_matches; +use reference::permits; diff --git a/src/web/manager/status/reference.rs b/src/web/manager/status/reference.rs new file mode 100644 index 0000000..034ec3d --- /dev/null +++ b/src/web/manager/status/reference.rs @@ -0,0 +1,52 @@ +use super::*; + +/// Opaque session-reference validation failure. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum SessionRefError { + /// The reference does not use the canonical versioned shape. + Invalid, + /// The reference belongs to another process runtime. + StaleInstance, +} + +/// Tests immutable candidate fields before any optional state-lock read. +pub(in crate::web::manager) fn immutable_matches( + session: &crate::web::session::WebSession, + index: &crate::web::manager::state::LiveSessionIndex, + filter: &SessionFilter, +) -> bool { + filter + .trace_session_id + .is_none_or(|value| value == session.trace_session_id()) + && filter + .client_ip + .is_none_or(|value| value == session.client_ip()) + && filter + .host + .as_deref() + .is_none_or(|value| value == session.profile_host()) + && filter + .user + .as_deref() + .is_none_or(|value| value == session.profile_user()) + && filter + .key_id + .as_deref() + .is_none_or(|value| session.key_id() == value) + && filter + .carrier + .is_none_or(|value| value == session.carrier()) + && filter + .user_agent_id + .is_none_or(|value| index.user_agent_id == Some(value)) +} + +pub(super) fn permits(semaphore: &Arc, capacity: usize) -> PermitStatus { + let available = semaphore.available_permits().min(capacity); + PermitStatus { + used: capacity.saturating_sub(available), + available, + capacity, + closed: semaphore.is_closed(), + } +}