WEB Handshake OPEN -> DATA Misuse fixed

Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
This commit is contained in:
Alexey
2026-08-23 16:32:32 +03:00
parent b6b9a18f3d
commit 1bd85535f4
3 changed files with 367 additions and 15 deletions
+2 -2
View File
@@ -130,7 +130,7 @@ pub struct WebLimitsConfig {
/// Process-wide live logical-stream ceiling.
#[serde(default = "default_web_max_streams_global")]
pub max_streams_global: usize,
/// Process-wide concurrent inner MTProxy handshake ceiling.
/// Process-wide ceiling for inner MTProxy handshakes that received a first byte.
#[serde(default = "default_web_max_stream_handshakes")]
pub max_stream_handshakes: usize,
/// Closed stream identifiers retained by one session.
@@ -249,7 +249,7 @@ pub struct WebTimeoutsConfig {
/// Deadline for collecting one authenticated carrier request body.
#[serde(default = "default_web_body_timeout_secs")]
pub body_secs: u64,
/// Deadline for the inner MTProxy handshake on one logical stream.
/// Deadline from the first inner byte through MTProxy authentication.
#[serde(default = "default_web_stream_handshake_timeout_secs")]
pub stream_handshake_secs: u64,
/// Maximum wait for one empty downlink long poll.
+41 -13
View File
@@ -9,6 +9,10 @@ use crate::web::stream::WebLogicalStream;
use super::{WebSession, inbound_queue_cost};
#[cfg(test)]
#[path = "backend_tests.rs"]
mod tests;
impl WebSession {
/// Starts one owned inner handshake and relay task for an admitted stream.
pub(super) fn spawn_stream(self: &Arc<Self>, stream_id: u32, peer_port: u16) {
@@ -22,10 +26,6 @@ impl WebSession {
self.stream_finished(stream_id, peer_port);
return;
};
let Some(handshake_permit) = manager.try_stream_handshake() else {
self.stream_finished(stream_id, peer_port);
return;
};
let deps = generation.client_runtime_deps();
let replay_checker = Arc::clone(&generation.replay_checker);
let session = Arc::clone(self);
@@ -46,7 +46,6 @@ impl WebSession {
stream,
deps,
replay_checker,
handshake_permit,
peer_port,
) => {}
}
@@ -109,7 +108,6 @@ async fn run_stream(
stream: WebLogicalStream,
deps: crate::proxy::authenticated::ClientRuntimeDeps,
replay_checker: Arc<crate::stats::ReplayChecker>,
handshake_permit: tokio::sync::OwnedSemaphorePermit,
peer_port: u16,
) {
use tokio::io::AsyncReadExt;
@@ -122,10 +120,25 @@ async fn run_stream(
let mut handshake = [0u8; HANDSHAKE_LEN];
let peer = std::net::SocketAddr::new(session.client_ip, peer_port);
deps.stats.increment_connects_all();
// A carrier may publish OPEN before the local MTProto socket writes its
// first byte. Session and stream quotas bound this idle phase without
// consuming the process-wide active-handshake budget.
if reader.read_exact(&mut handshake[..1]).await.is_err() {
deps.stats
.increment_connects_bad_with_class("web_mtproto_handshake_io");
return;
}
let Some(manager) = session.manager.upgrade() else {
return;
};
let Some(handshake_permit) = manager.try_stream_handshake() else {
return;
};
let handshake_result = tokio::time::timeout(
Duration::from_secs(session.timeouts.stream_handshake_secs),
async {
reader.read_exact(&mut handshake).await?;
reader.read_exact(&mut handshake[1..]).await?;
Ok::<_, io::Error>(
handle_mtproto_handshake_for_web_user(
&handshake,
@@ -144,12 +157,27 @@ async fn run_stream(
)
.await;
drop(handshake_permit);
let Ok(Ok(crate::error::HandshakeResult::Success((reader, writer, success)))) =
handshake_result
else {
deps.stats
.increment_connects_bad_with_class("web_mtproto_bad_client");
return;
let (reader, writer, success) = match handshake_result {
Err(_) => {
deps.stats
.increment_connects_bad_with_class("web_mtproto_handshake_timeout");
deps.stats.increment_handshake_timeouts();
deps.stats.increment_handshake_failure_class("timeout");
return;
}
Ok(Err(_)) => {
deps.stats
.increment_connects_bad_with_class("web_mtproto_handshake_io");
return;
}
Ok(Ok(crate::error::HandshakeResult::Success((reader, writer, success)))) => {
(reader, writer, success)
}
Ok(Ok(_)) => {
deps.stats
.increment_connects_bad_with_class("web_mtproto_bad_client");
return;
}
};
let _ = run_authenticated(
reader,
+324
View File
@@ -0,0 +1,324 @@
use std::collections::BTreeMap;
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Duration;
use arc_swap::ArcSwap;
use tokio::net::TcpListener;
use super::*;
use crate::config::{
ProxyConfig, UpstreamConfig, UpstreamType, WebCarrier, WebRuntimeConfig, WebRuntimeProfile,
WebSecretMode,
};
use crate::crypto::{AesCtr, sha256};
use crate::maestro::generation::{RuntimeGeneration, test_runtime_generation};
use crate::protocol::constants::{
DC_IDX_POS, HANDSHAKE_LEN, IV_LEN, PREKEY_LEN, PROTO_TAG_POS, ProtoTag, SKIP_LEN,
};
use crate::web::frame;
use crate::web::manager::{ManagerError, WebProcessRuntime};
struct TestRuntime {
session: Arc<WebSession>,
manager: Arc<WebProcessRuntime>,
generation: Arc<RuntimeGeneration>,
}
impl TestRuntime {
fn process_frame(
&self,
stream_id: u32,
sequence: u64,
frame_type: FrameType,
payload: &[u8],
) -> Result<u64, ManagerError> {
let encoded = frame::encode(frame_type, stream_id, payload);
match self.session.carrier() {
WebCarrier::Https => self.session.process_up(sequence, &encoded),
WebCarrier::HttpsLanes => {
self.session
.process_up_lane(stream_id, sequence, &encoded)
}
}
}
async fn shutdown(self) {
self.session.close();
self.session.wait().await;
self.manager.shutdown().await;
self.generation.stop_sessions().await;
self.generation.stop_background_tasks().await;
}
}
fn test_runtime(carrier: WebCarrier, max_stream_handshakes: usize) -> TestRuntime {
test_runtime_with_dc(carrier, max_stream_handshakes, None)
}
fn test_runtime_with_dc(
carrier: WebCarrier,
max_stream_handshakes: usize,
dc_addr: Option<SocketAddr>,
) -> TestRuntime {
let profile = Arc::new(WebRuntimeProfile {
host: "proxy.example.com".to_string(),
public_addr: "203.0.113.10:443".parse().unwrap(),
user: "default".to_string(),
secret_mode: WebSecretMode::Plain,
carrier,
capability: [7; 32],
max_sessions: 4,
max_streams: 16,
max_streams_per_session: 4,
});
let mut config = ProxyConfig::default();
config.web.enabled = true;
config.web.carrier = carrier;
config.web.limits.max_stream_handshakes = max_stream_handshakes;
config.web.timeouts.stream_handshake_secs = 1;
config.web.timeouts.shutdown_secs = 1;
config.censorship.server_hello_delay_min_ms = 0;
config.censorship.server_hello_delay_max_ms = 0;
if let Some(dc_addr) = dc_addr {
config
.dc_overrides
.insert("2".to_string(), vec![dc_addr.to_string()]);
config.upstreams.push(UpstreamConfig {
upstream_type: UpstreamType::Direct {
interface: None,
bind_addresses: None,
bindtodevice: None,
},
weight: 1,
enabled: true,
scopes: String::new(),
selected_scope: String::new(),
ipv4: Some(true),
ipv6: Some(false),
prefer: Some(4),
});
}
config.web.runtime = Some(Arc::new(WebRuntimeConfig {
vhosts: BTreeMap::new(),
profiles: vec![Arc::clone(&profile)],
}));
config.rebuild_runtime_user_auth().unwrap();
let limits = config.web.limits.clone();
let timeouts = config.web.timeouts.clone();
let generation = test_runtime_generation(1, config);
let manager = WebProcessRuntime::start(Arc::new(ArcSwap::from(Arc::clone(&generation))));
let session = WebSession::new(
Arc::downgrade(&manager),
[8; 32],
"192.0.2.10".parse().unwrap(),
profile,
[7; 32],
limits,
timeouts,
);
TestRuntime {
session,
manager,
generation,
}
}
fn valid_plain_handshake() -> [u8; HANDSHAKE_LEN] {
let secret = [0u8; 16];
let mut handshake = [0x5a; HANDSHAKE_LEN];
for (index, byte) in handshake[SKIP_LEN..SKIP_LEN + PREKEY_LEN + IV_LEN]
.iter_mut()
.enumerate()
{
*byte = (index as u8).wrapping_add(1);
}
let dec_prekey = &handshake[SKIP_LEN..SKIP_LEN + PREKEY_LEN];
let dec_iv_bytes = &handshake[SKIP_LEN + PREKEY_LEN..SKIP_LEN + PREKEY_LEN + IV_LEN];
let mut dec_key_input = Vec::with_capacity(PREKEY_LEN + secret.len());
dec_key_input.extend_from_slice(dec_prekey);
dec_key_input.extend_from_slice(&secret);
let dec_key = sha256(&dec_key_input);
let mut dec_iv = [0u8; IV_LEN];
dec_iv.copy_from_slice(dec_iv_bytes);
let mut cipher = AesCtr::new(&dec_key, u128::from_be_bytes(dec_iv));
let keystream = cipher.encrypt(&[0u8; HANDSHAKE_LEN]);
let mut plaintext = [0u8; HANDSHAKE_LEN];
plaintext[PROTO_TAG_POS..PROTO_TAG_POS + 4]
.copy_from_slice(&ProtoTag::Intermediate.to_bytes());
plaintext[DC_IDX_POS..DC_IDX_POS + 2].copy_from_slice(&2i16.to_le_bytes());
for index in PROTO_TAG_POS..HANDSHAKE_LEN {
handshake[index] = plaintext[index] ^ keystream[index];
}
handshake
}
async fn settle_tasks() {
for _ in 0..4 {
tokio::task::yield_now().await;
}
}
fn bad_class(runtime: &TestRuntime, class: &str) -> u64 {
runtime
.generation
.stats
.get_connects_bad_class_counts()
.into_iter()
.find_map(|(name, total)| (name == class).then_some(total))
.unwrap_or(0)
}
#[tokio::test(start_paused = true)]
async fn open_without_data_does_not_start_the_inner_handshake_timeout() {
for carrier in [WebCarrier::Https, WebCarrier::HttpsLanes] {
let runtime = test_runtime(carrier, 1);
assert_eq!(runtime.process_frame(1, 1, FrameType::Open, &[]), Ok(1));
settle_tasks().await;
assert!(runtime.session.state.lock().streams.contains_key(&1));
tokio::time::advance(Duration::from_secs(2)).await;
settle_tasks().await;
assert!(
runtime.session.state.lock().streams.contains_key(&1),
"OPEN without DATA must remain live until session cancellation or idle expiry"
);
runtime.shutdown().await;
}
}
#[tokio::test(start_paused = true)]
async fn the_first_inner_byte_starts_the_handshake_timeout() {
let runtime = test_runtime(WebCarrier::Https, 1);
assert_eq!(runtime.process_frame(1, 1, FrameType::Open, &[]), Ok(1));
settle_tasks().await;
tokio::time::advance(Duration::from_secs(2)).await;
assert_eq!(runtime.process_frame(1, 2, FrameType::Data, &[0x5a]), Ok(2));
settle_tasks().await;
assert!(runtime.session.state.lock().streams.contains_key(&1));
tokio::time::advance(Duration::from_secs(2)).await;
settle_tasks().await;
assert!(!runtime.session.state.lock().streams.contains_key(&1));
assert_eq!(bad_class(&runtime, "web_mtproto_handshake_timeout"), 1);
assert_eq!(runtime.generation.stats.get_handshake_timeouts(), 1);
assert_eq!(
runtime
.generation
.stats
.get_handshake_failure_class_counts(),
vec![("timeout".to_string(), 1)]
);
runtime.shutdown().await;
}
#[tokio::test(start_paused = true)]
async fn delayed_complete_handshake_is_classified_after_data_arrives() {
let runtime = test_runtime(WebCarrier::Https, 1);
assert_eq!(runtime.process_frame(1, 1, FrameType::Open, &[]), Ok(1));
settle_tasks().await;
tokio::time::advance(Duration::from_secs(2)).await;
assert_eq!(
runtime.process_frame(1, 2, FrameType::Data, &[0; 64]),
Ok(2)
);
settle_tasks().await;
assert!(!runtime.session.state.lock().streams.contains_key(&1));
assert_eq!(bad_class(&runtime, "web_mtproto_bad_client"), 1);
assert_eq!(bad_class(&runtime, "web_mtproto_handshake_timeout"), 0);
assert_eq!(runtime.generation.stats.get_handshake_timeouts(), 0);
runtime.shutdown().await;
}
#[tokio::test(start_paused = true)]
async fn delayed_valid_handshake_reaches_the_authenticated_relay() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let runtime = test_runtime_with_dc(WebCarrier::Https, 1, Some(listener.local_addr().unwrap()));
assert_eq!(runtime.process_frame(1, 1, FrameType::Open, &[]), Ok(1));
settle_tasks().await;
tokio::time::advance(Duration::from_secs(2)).await;
assert_eq!(
runtime.process_frame(1, 2, FrameType::Data, &valid_plain_handshake()),
Ok(2)
);
let accepted = tokio::time::timeout(Duration::from_secs(1), listener.accept()).await;
assert!(
accepted.is_ok(),
"valid handshake did not reach the configured upstream; bad classes: {:?}",
runtime.generation.stats.get_connects_bad_class_counts()
);
let (upstream, _) = accepted.unwrap().unwrap();
assert!(runtime.session.state.lock().streams.contains_key(&1));
assert_eq!(runtime.generation.stats.get_connects_bad(), 0);
assert_eq!(runtime.generation.stats.get_handshake_timeouts(), 0);
runtime.shutdown().await;
drop(upstream);
}
#[tokio::test(start_paused = true)]
async fn silent_streams_do_not_consume_active_handshake_capacity() {
let runtime = test_runtime(WebCarrier::Https, 1);
assert_eq!(runtime.process_frame(1, 1, FrameType::Open, &[]), Ok(1));
assert_eq!(runtime.process_frame(2, 2, FrameType::Open, &[]), Ok(2));
settle_tasks().await;
tokio::time::advance(Duration::from_secs(2)).await;
settle_tasks().await;
{
let state = runtime.session.state.lock();
assert!(state.streams.contains_key(&1));
assert!(state.streams.contains_key(&2));
}
assert_eq!(runtime.process_frame(1, 3, FrameType::Data, &[1]), Ok(3));
settle_tasks().await;
assert_eq!(runtime.process_frame(2, 4, FrameType::Data, &[2]), Ok(4));
settle_tasks().await;
let state = runtime.session.state.lock();
assert!(state.streams.contains_key(&1));
assert!(!state.streams.contains_key(&2));
drop(state);
runtime.shutdown().await;
}
#[tokio::test(start_paused = true)]
async fn cancellation_while_waiting_for_data_releases_stream_ownership() {
let runtime = test_runtime(WebCarrier::Https, 1);
assert_eq!(runtime.process_frame(1, 1, FrameType::Open, &[]), Ok(1));
settle_tasks().await;
assert_eq!(runtime.session.tasks_live.load(Ordering::Acquire), 1);
runtime.session.close();
runtime.session.wait().await;
assert_eq!(runtime.session.tasks_live.load(Ordering::Acquire), 0);
assert!(runtime.session.state.lock().active_peer_ports.is_empty());
let peer_port = runtime
.manager
.try_acquire_stream(
runtime.session.profile_key,
runtime.session.profile.max_streams,
runtime.session.client_ip,
runtime.session.profile.public_addr,
)
.unwrap();
assert_eq!(peer_port, 1);
runtime.manager.release_stream(
runtime.session.profile_key,
runtime.session.client_ip,
runtime.session.profile.public_addr,
peer_port,
);
runtime.shutdown().await;
}