diff --git a/src/client/actor/download_loop.rs b/src/client/actor/download_loop.rs index 22d3fad..5296674 100644 --- a/src/client/actor/download_loop.rs +++ b/src/client/actor/download_loop.rs @@ -34,7 +34,7 @@ pub struct DownloadLoopActor { write_half: tokio::net::tcp::OwnedWriteHalf, cipher: Option>, encoding: EncodingType, - cookie_val: String, + stream_id: String, http_client: Arc, state: Arc, max_bytes: Option, @@ -46,7 +46,7 @@ impl DownloadLoopActor { initial_response: wreq::Response, write_half: tokio::net::tcp::OwnedWriteHalf, cipher: Option>, - cookie_val: String, + stream_id: String, http_client: Arc, state: Arc, ) -> Self { @@ -58,7 +58,7 @@ impl DownloadLoopActor { let use_prefetch = prefetch_at.is_some_and(|at| at > 0); let (prefetch_trigger, prefetch_rx) = if rotate_enabled && use_prefetch { - let (tx, rx) = spawn_prefetch_continuation(&http_client, &state, &cookie_val); + let (tx, rx) = spawn_prefetch_continuation(&http_client, &state, &stream_id); (Some(tx), Some(rx)) } else { (None, None) @@ -68,7 +68,7 @@ impl DownloadLoopActor { write_half, cipher: cipher_dyn, encoding, - cookie_val, + stream_id, http_client, state, max_bytes, @@ -151,15 +151,15 @@ impl DownloadLoopActor { match tokio::time::timeout(PREFETCH_ROTATE_TIMEOUT, rx).await { Ok(Ok(Ok(resp))) => resp, Ok(Ok(Err(_))) | Ok(Err(_)) => { - send_continue_request(&self.http_client, &self.state, &self.cookie_val).await? + send_continue_request(&self.http_client, &self.state, &self.stream_id).await? } Err(_elapsed) => { warn!("prefetch timed out, falling back to synchronous continue"); - send_continue_request(&self.http_client, &self.state, &self.cookie_val).await? + send_continue_request(&self.http_client, &self.state, &self.stream_id).await? } } } else { - send_continue_request(&self.http_client, &self.state, &self.cookie_val).await? + send_continue_request(&self.http_client, &self.state, &self.stream_id).await? }; let use_prefetch = self @@ -168,7 +168,7 @@ impl DownloadLoopActor { .is_some_and(|at| at > 0); let (prefetch_trigger, prefetch_rx) = if use_prefetch { let (tx, rx) = - spawn_prefetch_continuation(&self.http_client, &self.state, &self.cookie_val); + spawn_prefetch_continuation(&self.http_client, &self.state, &self.stream_id); (Some(tx), Some(rx)) } else { (None, None) @@ -238,10 +238,10 @@ async fn download_single_response( async fn send_continue_request( http_client: &wreq::Client, state: &SharedState, - cookie_val: &str, + stream_id: &str, ) -> Result { let mut cookie = String::new(); - utils::build_tunnel_cookie(&mut cookie, cookie_val); + utils::build_stream_cookie(&mut cookie, stream_id); let mut req = http_client .post(state.remote_str.as_str()) .header("Cookie", cookie); @@ -259,7 +259,7 @@ async fn send_continue_request( fn spawn_prefetch_continuation( http_client: &Arc, state: &Arc, - cookie_val: &str, + stream_id: &str, ) -> ( oneshot::Sender<()>, oneshot::Receiver>, @@ -268,13 +268,13 @@ fn spawn_prefetch_continuation( let (result_tx, result_rx) = oneshot::channel(); let pre_client = Arc::clone(http_client); let pre_state = Arc::clone(state); - let pre_cookie = cookie_val.to_owned(); + let pre_stream_id = stream_id.to_owned(); tokio::spawn( async move { if trigger_rx.await.is_err() { return; } - match send_continue_request(&pre_client, &pre_state, &pre_cookie).await { + match send_continue_request(&pre_client, &pre_state, &pre_stream_id).await { Ok(resp) => { let _ = result_tx.send(Ok(resp)); } diff --git a/src/client/actor/upload_loop.rs b/src/client/actor/upload_loop.rs index d61bd92..89450c8 100644 --- a/src/client/actor/upload_loop.rs +++ b/src/client/actor/upload_loop.rs @@ -35,7 +35,7 @@ enum Phase { pub struct UploadLoopActor { http_client: Arc, state: Arc, - session_cookie: String, + stream_id: String, shaped: ShaperStream, request_sem: Arc, bytes_sem: Arc, @@ -50,7 +50,7 @@ impl UploadLoopActor { initial_payload: Bytes, read_half: tokio::net::tcp::OwnedReadHalf, cipher: Option>, - session_cookie: String, + stream_id: String, start_seq: u64, ) -> Self { let reader = AsyncReadExt::chain(std::io::Cursor::new(initial_payload), read_half); @@ -65,7 +65,7 @@ impl UploadLoopActor { Self { http_client, state, - session_cookie, + stream_id, shaped, request_sem: Arc::new(Semaphore::new(UPLOAD_CONCURRENCY)), bytes_sem: Arc::new(Semaphore::new(MAX_IN_FLIGHT_BYTES)), @@ -181,12 +181,12 @@ impl UploadLoopActor { let body = batch_buf.freeze(); let http_client = Arc::clone(&self.http_client); let state_ref = Arc::clone(&self.state); - let session_val = self.session_cookie.clone(); + let stream_id = self.stream_id.clone(); self.tasks.spawn( async move { let _req_guard = req_permit; let _bytes = bytes_permits; - send_upload_post(&http_client, &state_ref, body, &session_val).await + send_upload_post(&http_client, &state_ref, body, &stream_id).await } .instrument(tracing::Span::current()), ); @@ -230,11 +230,11 @@ async fn send_upload_post( http_client: &wreq::Client, state: &SharedState, body: Bytes, - session_cookie_val: &str, + stream_id: &str, ) -> Result<()> { debug_assert!(!body.is_empty(), "empty upload body"); let mut cookie = String::new(); - utils::build_tunnel_cookie(&mut cookie, session_cookie_val); + utils::build_stream_cookie(&mut cookie, stream_id); let mut req = http_client .post(state.remote_str.as_str()) .header("Accept-Encoding", "identity") diff --git a/src/client/connection.rs b/src/client/connection.rs index 1451ce9..7006e13 100644 --- a/src/client/connection.rs +++ b/src/client/connection.rs @@ -21,7 +21,7 @@ pub(crate) async fn handle_plain_proxy( ) -> Result<()> { let stream_id = uuid::Uuid::new_v4().to_string(); let mut cookie = String::new(); - utils::build_tunnel_cookie(&mut cookie, &stream_id); + utils::build_stream_cookie(&mut cookie, &stream_id); let (early_data, remaining_payload, frames_sent) = utils::encode_initial_payload( &payload, @@ -30,7 +30,7 @@ pub(crate) async fn handle_plain_proxy( &state.traffic_config, )?; - info!(target = %target_host, "connection initiated"); + info!(stream_id = %stream_id, target = %target_host, "connection initiated"); let response = tokio::time::timeout( DOWNLOAD_CONNECT_TIMEOUT, @@ -64,7 +64,7 @@ pub(crate) async fn handle_plain_proxy( response, write_half, None, - stream_id.to_owned(), + stream_id, Arc::clone(&http_client), Arc::clone(&state), ); diff --git a/src/client/handshake.rs b/src/client/handshake.rs index 340d6e4..2b4f46b 100644 --- a/src/client/handshake.rs +++ b/src/client/handshake.rs @@ -49,22 +49,23 @@ pub async fn try_pq_connect( let session_id = &ticket.session_id; info!(session_id = %session_id, target = %target_host, "session resumption: attempting to reuse session"); - let conn_nonce: [u8; 16] = rand::rng().random(); + let stream_id = uuid::Uuid::new_v4(); + let stream_id_bytes: [u8; 16] = *stream_id.as_bytes(); let (upload_key, download_key, target_key) = - crypto::derive_connection_keys(master, &conn_nonce); + crypto::derive_connection_keys(master, &stream_id_bytes); let upload_cipher = Arc::new(AesFrameCipher::new(&upload_key)); let download_cipher = Arc::new(AesFrameCipher::new(&download_key)); let enc_target = crypto::encrypt_bytes(&target_key, target_host.as_bytes())?; - let cookie_nonce_key = crypto::derive_cookie_nonce_key(master); - let enc_conn_nonce = crypto::encrypt_bytes(&cookie_nonce_key, &conn_nonce)?; + let cookie_stream_key = crypto::derive_cookie_stream_key(master); + let enc_stream_id = crypto::encrypt_bytes(&cookie_stream_key, &stream_id_bytes)?; let cookie_val = format!( "{}:{}:{}", session_id, base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(&enc_target), - base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(&enc_conn_nonce) + base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(&enc_stream_id) ); let (early_data, remaining_payload, frames_sent) = utils::encode_initial_payload( @@ -108,7 +109,7 @@ pub async fn try_pq_connect( let upload_client = Arc::clone(http_client); let upload_state = Arc::clone(state); let upload_cipher_clone = Arc::clone(&upload_cipher); - let session_cookie_val = cookie_val.clone(); + let stream_id_str = stream_id.to_string(); let upload_actor = UploadLoopActor::new( upload_client.clone(), @@ -116,7 +117,7 @@ pub async fn try_pq_connect( remaining_payload, read_half, Some(upload_cipher_clone), - session_cookie_val, + stream_id_str.clone(), frames_sent, ); let upload_task = @@ -126,7 +127,7 @@ pub async fn try_pq_connect( response, write_half, Some(download_cipher), - cookie_val.clone(), + stream_id_str, Arc::clone(http_client), Arc::clone(state), ); diff --git a/src/client/utils.rs b/src/client/utils.rs index 8d680ec..2974682 100644 --- a/src/client/utils.rs +++ b/src/client/utils.rs @@ -9,17 +9,28 @@ use crate::client::constants::{MIN_PADDING, PADDING_POOL}; use crate::shaper::{self, FrameCipher}; #[inline] -pub fn build_tunnel_cookie(buf: &mut String, session_val: &str) { +fn build_cookie_into(buf: &mut String, name: &str, value: &str) { buf.clear(); - let cap = 8 + session_val.len() + MIN_PADDING + PADDING_POOL.len(); + let cap = name.len() + 1 + value.len() + 2 + MIN_PADDING + PADDING_POOL.len(); buf.reserve(cap); - buf.push_str("session="); - buf.push_str(session_val); + buf.push_str(name); + buf.push('='); + buf.push_str(value); buf.push_str("; "); let padding_len = rand::rng().random_range(MIN_PADDING..PADDING_POOL.len()); buf.push_str(std::str::from_utf8(&PADDING_POOL[..padding_len]).expect("Invalid UTF-8")) } +#[inline] +pub fn build_tunnel_cookie(buf: &mut String, session_val: &str) { + build_cookie_into(buf, "session", session_val) +} + +#[inline] +pub fn build_stream_cookie(buf: &mut String, stream_id: &str) { + build_cookie_into(buf, "stream", stream_id) +} + pub fn encode_initial_payload( initial_payload: &[u8], max_bytes: usize, diff --git a/src/crypto/handshake.rs b/src/crypto/handshake.rs index c7462b9..d093d4e 100644 --- a/src/crypto/handshake.rs +++ b/src/crypto/handshake.rs @@ -21,18 +21,18 @@ pub fn derive_initial_master(mlkem_ss: &[u8], x25519_ss: &[u8; 32]) -> Zeroizing master } -pub fn derive_cookie_nonce_key(master: &[u8; 32]) -> Zeroizing<[u8; 32]> { +pub fn derive_cookie_stream_key(master: &[u8; 32]) -> Zeroizing<[u8; 32]> { let hkdf = Hkdf::::new(None, master); let mut key = Zeroizing::new([0u8; 32]); - hkdf.expand(b"cookie_nonce_key", &mut *key) + hkdf.expand(b"cookie_stream_key", &mut *key) .expect("32 bytes is valid for HKDF"); key } -pub fn derive_connection_keys(master: &[u8; 32], conn_nonce: &[u8; 16]) -> super::ConnectionKeys { +pub fn derive_connection_keys(master: &[u8; 32], stream_id: &[u8; 16]) -> super::ConnectionKeys { let hkdf = Hkdf::::new(None, master); let mut info = Vec::with_capacity(16 + 15); - info.extend_from_slice(conn_nonce); + info.extend_from_slice(stream_id); info.extend_from_slice(b"connection_keys"); let mut buf = Zeroizing::new([0u8; 96]); hkdf.expand(&info, &mut *buf) diff --git a/src/crypto/mod.rs b/src/crypto/mod.rs index aeac139..475ef08 100644 --- a/src/crypto/mod.rs +++ b/src/crypto/mod.rs @@ -4,7 +4,7 @@ mod keys; pub use cipher::{AesFrameCipher, decrypt_bytes, encrypt_bytes}; pub use handshake::{ - derive_connection_keys, derive_cookie_nonce_key, derive_handshake_key, derive_initial_master, + derive_connection_keys, derive_cookie_stream_key, derive_handshake_key, derive_initial_master, }; pub use keys::{ b64_to_private_key, b64_to_public_key, bytes_to_encapsulation_key, diffie_hellman, diff --git a/src/server/actor/tunnel.rs b/src/server/actor/tunnel.rs index 9653459..9c894a2 100644 --- a/src/server/actor/tunnel.rs +++ b/src/server/actor/tunnel.rs @@ -15,7 +15,7 @@ use crate::server::constants::{ DOWNLOAD_CHANNEL_CAPACITY, ROTATION_STALENESS, STREAM_IDLE_TIMEOUT_SECS, UPLOAD_CMD_CHANNEL_CAPACITY, UPLOAD_DONE_TIMEOUT, }; -use crate::server::nonce_registry::NonceRegistry; +use crate::server::stream_registry::StreamRegistry; use crate::shaper::{FrameCipher, TrafficConfig, TrafficShaper}; pub enum TunnelCmd { @@ -54,9 +54,8 @@ pub struct TunnelActor { pending_continue: Vec>>>>, segment_done_tx: mpsc::Sender>, segment_done_rx: mpsc::Receiver>, - session_id: String, - conn_nonce: Option<[u8; 16]>, - nonce_registry: Arc, + stream_id: String, + stream_registry: Arc, shutdown_signal: Arc, max_download_bytes: Option, pending_write_half: Option, @@ -69,9 +68,8 @@ impl TunnelActor { pub fn new( rx: mpsc::Receiver, download_tx: mpsc::Sender>, - session_id: String, - conn_nonce: Option<[u8; 16]>, - nonce_registry: Arc, + stream_id: String, + stream_registry: Arc, max_download_bytes: Option, ) -> Self { let (seg_tx, seg_rx) = mpsc::channel::>(2); @@ -86,9 +84,8 @@ impl TunnelActor { pending_continue: Vec::new(), segment_done_tx: seg_tx, segment_done_rx: seg_rx, - session_id, - conn_nonce, - nonce_registry, + stream_id, + stream_registry, shutdown_signal: Arc::new(Notify::new()), max_download_bytes, pending_write_half: None, @@ -237,7 +234,7 @@ impl TunnelActor { } } _ = &mut rotation_timeout, if matches!(self.phase, Phase::Rotating) => { - warn!(session_id = %self.session_id, "rotation staleness timeout"); + warn!(stream_id = %self.stream_id, "rotation staleness timeout"); break; } _ = self.shutdown_signal.notified() => { @@ -248,7 +245,7 @@ impl TunnelActor { .saturating_sub(self.last_activity.load(Ordering::Relaxed)); if idle_secs >= STREAM_IDLE_TIMEOUT_SECS { warn!( - session_id = %self.session_id, + stream_id = %self.stream_id, "stream idle timeout, closing tunnel" ); break; @@ -332,7 +329,7 @@ impl TunnelActor { async fn cleanup(&mut self) { self.phase = Phase::Closed; - self.consume_nonce(); + self.consume_stream(); if let Some(handle) = self.download_handle.take() { handle.abort(); } @@ -342,21 +339,19 @@ impl TunnelActor { } self.upload_tx = None; if let Err(join_err) = handle.await { - warn!(session_id = %self.session_id, error = %join_err, "upload actor panicked"); + warn!(stream_id = %self.stream_id, error = %join_err, "upload actor panicked"); } } - info!(session_id = %self.session_id, "tunnel actor closed"); + info!(stream_id = %self.stream_id, "tunnel actor closed"); } - fn consume_nonce(&mut self) { - if let Some(nonce) = self.conn_nonce.take() { - self.nonce_registry.mark_consumed(&self.session_id, &nonce); - } + fn consume_stream(&mut self) { + self.stream_registry.mark_consumed(&self.stream_id); } } impl Drop for TunnelActor { fn drop(&mut self) { - self.consume_nonce(); + self.consume_stream(); } } diff --git a/src/server/constants.rs b/src/server/constants.rs index a064ef5..aeeb7c0 100644 --- a/src/server/constants.rs +++ b/src/server/constants.rs @@ -14,7 +14,7 @@ pub const ROTATION_STALENESS: Duration = Duration::from_secs(10); pub const UPLOAD_DONE_TIMEOUT: Duration = Duration::from_secs(30); pub const JANITOR_INTERVAL: Duration = Duration::from_secs(30); -pub const NONCE_CLEANUP_INTERVAL: Duration = Duration::from_secs(60); +pub const MASTER_CLEANUP_INTERVAL: Duration = Duration::from_secs(60); pub const CONNECT_TIMEOUT: Duration = Duration::from_secs(10); pub const WRITE_TIMEOUT: Duration = Duration::from_secs(30); diff --git a/src/server/handlers.rs b/src/server/handlers.rs index 1e5d47f..bd8f38d 100644 --- a/src/server/handlers.rs +++ b/src/server/handlers.rs @@ -18,11 +18,38 @@ use crate::server::constants::{ TUNNEL_CMD_CHANNEL_CAPACITY, }; use crate::server::stream::FrameDecoder; +use crate::server::stream_registry::StreamQueryResult; use crate::server::{SessionHandle, utils}; use crate::shaper::{self, FrameCipher}; use tokio::sync::oneshot; use super::AppState; +use crate::server::stream_registry::StreamRegistry; + +struct StreamGuard { + registry: Arc, + stream_id: String, + armed: bool, +} +impl Drop for StreamGuard { + fn drop(&mut self) { + if self.armed { + self.registry.mark_consumed(&self.stream_id); + } + } +} +impl StreamGuard { + fn new(registry: Arc, stream_id: String) -> Self { + Self { + registry, + stream_id, + armed: true, + } + } + fn disarm(&mut self) { + self.armed = false; + } +} struct ActorGuard { actors: Arc>, @@ -73,11 +100,9 @@ async fn setup_tunnel_response( body: Body, host: &str, port: u16, - session_id: &str, - conn_nonce: Option<[u8; 16]>, + stream_id: &str, upload_cipher: Option>, download_cipher: Option>, - actor_key: &str, ) -> Result { let encoding = state.traffic_config.encoding_type; @@ -130,9 +155,8 @@ async fn setup_tunnel_response( let mut actor = crate::server::actor::tunnel::TunnelActor::new( actor_rx, download_tx, - session_id.to_owned(), - conn_nonce, - Arc::clone(&state.nonce_registry), + stream_id.to_owned(), + Arc::clone(&state.stream_registry), state.traffic_config.max_download_bytes, ); actor.set_upload_channel(upload_tx_for_actor, upload_handle); @@ -149,10 +173,10 @@ async fn setup_tunnel_response( upload_cipher, encoding, }; - state.actors.insert(actor_key.to_owned(), handle); - let mut early_guard = ActorGuard::new(Arc::clone(&state.actors), actor_key.to_owned()); + state.actors.insert(stream_id.to_owned(), handle); + let mut early_guard = ActorGuard::new(Arc::clone(&state.actors), stream_id.to_owned()); - let key = actor_key.to_owned(); + let key = stream_id.to_owned(); let actors_ref2 = Arc::clone(&state.actors); let actor_handle = tokio::spawn( async move { @@ -179,19 +203,21 @@ async fn setup_tunnel_response( Ok(response) } -async fn spawn_tunnel_actor( +async fn spawn_encrypted_tunnel( state: Arc, - cookie_val: &str, + session_cookie_val: &str, body: Body, span: tracing::Span, ) -> Result { - let parts: Vec<&str> = cookie_val.splitn(3, ':').collect(); + let parts: Vec<&str> = session_cookie_val.splitn(3, ':').collect(); if parts.len() != 3 { return Err(ServerError::precondition_required( "invalid session cookie format", )); } - let (session_id, enc_target_b64, enc_nonce_b64) = (parts[0], parts[1], parts[2]); + let (session_id, enc_target_b64, enc_stream_id_b64) = (parts[0], parts[1], parts[2]); + + utils::validate_uuid(session_id)?; let entry = state .master_store @@ -211,28 +237,32 @@ async fn spawn_tunnel_actor( span.record("user", &username); drop(entry); - let cookie_nonce_key = crypto::derive_cookie_nonce_key(&master); - let enc_nonce = base64::engine::general_purpose::URL_SAFE_NO_PAD - .decode(enc_nonce_b64) + let cookie_stream_key = crypto::derive_cookie_stream_key(&master); + let enc_stream_id = base64::engine::general_purpose::URL_SAFE_NO_PAD + .decode(enc_stream_id_b64) .map_err(|_| ServerError::precondition_required("invalid cookie encoding"))?; - let conn_nonce_bytes = crypto::decrypt_bytes(&cookie_nonce_key, &enc_nonce) - .map_err(|_| ServerError::precondition_required("failed to decrypt conn_nonce"))?; - let conn_nonce: [u8; 16] = conn_nonce_bytes + let stream_id_bytes = crypto::decrypt_bytes(&cookie_stream_key, &enc_stream_id) + .map_err(|_| ServerError::precondition_required("failed to decrypt stream ID"))?; + let stream_id_bytes_arr: [u8; 16] = stream_id_bytes .try_into() - .map_err(|_| ServerError::precondition_required("invalid conn_nonce length"))?; - - match state.nonce_registry.try_claim(session_id, &conn_nonce) { - Ok(true) => {} - Ok(false) => { - return Err(ServerError::precondition_required("nonce already claimed")); - } - Err(_) => { - return Err(ServerError::precondition_required("nonce already consumed")); - } + .map_err(|_| ServerError::precondition_required("invalid stream ID length"))?; + let stream_uuid = uuid::Uuid::from_bytes(stream_id_bytes_arr); + let stream_id = stream_uuid.to_string(); + + utils::validate_uuid(&stream_id)?; + + if !state + .stream_registry + .register(&stream_id, crate::now_secs()) + { + warn!(stream_id = %stream_id, session_id = %session_id, "duplicate stream registration attempt"); + return Err(ServerError::precondition_required("stream already active")); } + let mut stream_guard = StreamGuard::new(Arc::clone(&state.stream_registry), stream_id.clone()); + let (upload_key, download_key, target_key) = - crypto::derive_connection_keys(&master, &conn_nonce); + crypto::derive_connection_keys(&master, stream_uuid.as_bytes()); let enc_target = base64::engine::general_purpose::URL_SAFE_NO_PAD .decode(enc_target_b64) .map_err(|_| ServerError::precondition_required("invalid cookie encoding"))?; @@ -251,21 +281,20 @@ async fn spawn_tunnel_actor( let download_cipher: Arc = Arc::new(AesFrameCipher::new(&download_key)); let upload_cipher = Arc::new(AesFrameCipher::new(&upload_key)); - let log_target = target.clone(); let response = setup_tunnel_response( &state, body, host, port, - session_id, - Some(conn_nonce), + &stream_id, Some(upload_cipher as Arc), Some(download_cipher), - cookie_val, ) .await?; - info!(session_id = %session_id, user = %username, target = %log_target, "tunnel actor spawned"); + stream_guard.disarm(); + + info!(stream_id = %stream_id, session_id = %session_id, "tunnel established"); Ok(response) } @@ -356,33 +385,83 @@ pub async fn dispatch( body: Body, ) -> Result { let span = tracing::Span::current(); - let session_cookie = utils::extract_cookie_value(&headers, "session"); + let has_target = headers.get("X-Target").is_some(); - let has_encrypted_cookie = session_cookie.is_some_and(|c| c.contains(':')); - if headers.get("X-Target").is_some() && !has_encrypted_cookie { - return handle_plaintext_download(state, headers, body, span).await; - } + let stream_cookie = utils::extract_cookie_value(&headers, "stream").map(|s| s.to_owned()); + if let Some(ref stream_id) = stream_cookie { + utils::validate_uuid(stream_id)?; - if let Some(cookie_val) = session_cookie { - if let Some(handle) = state.actors.get(cookie_val).map(|r| r.value().clone()) { + if let Some(handle) = state.actors.get(stream_id).map(|r| r.value().clone()) { return dispatch_to_actor(handle, body).await; } - let session_id = cookie_val.split(':').next().unwrap_or(cookie_val); - if state.master_store.get(session_id).is_some() { - return spawn_tunnel_actor(state, cookie_val, body, span).await; + if has_target { + return handle_plaintext_tunnel(state, headers, body, span, Some(stream_id.clone())) + .await; } - return Err(ServerError::precondition_required("session not found")); + + return handle_stream_not_found(&state, stream_id); + } + + let session_cookie = utils::extract_cookie_value(&headers, "session").map(|s| s.to_owned()); + if let Some(ref session_val) = session_cookie { + if session_val.contains(':') { + return spawn_encrypted_tunnel(state, session_val, body, span).await; + } + if has_target { + utils::validate_uuid(session_val)?; + return handle_plaintext_tunnel(state, headers, body, span, Some(session_val.clone())) + .await; + } + return Err(ServerError::precondition_required( + "invalid session cookie — missing target or encrypted payload", + )); } - handle_fresh_handshake(state, headers, body, span).await + if has_target { + return handle_plaintext_tunnel(state, headers, body, span, None).await; + } + + let has_handshake_cookie = utils::extract_cookie_value(&headers, "eph_pk_a").is_some(); + if has_handshake_cookie { + return handle_fresh_handshake(state, headers, body, span).await; + } + + Err(ServerError::bad_request( + "missing required cookies or headers", + )) +} + +fn handle_stream_not_found( + state: &Arc, + stream_id: &str, +) -> Result { + match state.stream_registry.check(stream_id) { + StreamQueryResult::Consumed => { + warn!(stream_id = %stream_id, "replay attempt on consumed stream"); + Err(ServerError::precondition_required( + "stream already consumed", + )) + } + StreamQueryResult::Active => { + warn!(stream_id = %stream_id, "stream registered but initial tunnel setup did not complete"); + Err(ServerError::precondition_required( + "stream setup incomplete", + )) + } + StreamQueryResult::Fresh => { + warn!(stream_id = %stream_id, "upload for unknown stream"); + Err(ServerError::precondition_required("stream not found")) + } + } } -async fn handle_plaintext_download( +async fn handle_plaintext_tunnel( state: Arc, headers: HeaderMap, body: Body, span: tracing::Span, + stream_id_opt: Option, ) -> Result { let user = validate_jwt_if_needed(&headers, &state.decoding_key, &state.jwt_validation)?; span.record("user", &user); @@ -392,7 +471,14 @@ async fn handle_plaintext_download( .and_then(|v| v.to_str().ok()) .ok_or_else(|| ServerError::bad_request("missing X-Target header"))?; span.record("target", target); - info!(user = %user, target = %target, "connection initiated"); + + let stream_id = match stream_id_opt { + Some(id) => { + utils::validate_uuid(&id)?; + id + } + None => uuid::Uuid::new_v4().to_string(), + }; let (host, port_str) = target .rsplit_once(':') @@ -401,24 +487,21 @@ async fn handle_plaintext_download( .parse() .map_err(|_| ServerError::bad_request("invalid port"))?; - let session_id = utils::extract_cookie_value(&headers, "session") - .map(|s| s.to_owned()) - .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()); + if !state + .stream_registry + .register(&stream_id, crate::now_secs()) + { + warn!(stream_id = %stream_id, "duplicate stream registration for plaintext tunnel"); + return Err(ServerError::precondition_required("stream already active")); + } - let response = setup_tunnel_response( - &state, - body, - host, - port, - &session_id, - None, - None, - None, - &session_id, - ) - .await?; + let mut stream_guard = StreamGuard::new(Arc::clone(&state.stream_registry), stream_id.clone()); + + let response = setup_tunnel_response(&state, body, host, port, &stream_id, None, None).await?; + + stream_guard.disarm(); - info!(stream_id = %session_id, user = %user, target = %target, "stream established"); + info!(stream_id = %stream_id, "tunnel established"); Ok(response) } @@ -431,7 +514,7 @@ async fn handle_fresh_handshake( let user = validate_jwt_if_needed(&headers, &state.decoding_key, &state.jwt_validation)?; span.record("user", &user); - info!(user = %user, "handshake: received ClientHello"); + info!("handshake: received ClientHello"); let eph_pk_a_b64 = utils::extract_cookie_value(&headers, "eph_pk_a") .ok_or_else(|| ServerError::bad_request("missing eph_pk_a cookie"))?; @@ -514,7 +597,7 @@ async fn handle_fresh_handshake( (user.clone(), master, crate::now_secs()), ); - info!(session_id = %session_id, user = %user, "handshake: master key derived"); + info!(session_id = %session_id, "handshake: master key derived"); let ct_bytes: &[u8] = &ct; let ct_bytes = ct_bytes.to_vec(); diff --git a/src/server/janitor.rs b/src/server/janitor.rs index 732b829..2a1ddb3 100644 --- a/src/server/janitor.rs +++ b/src/server/janitor.rs @@ -2,36 +2,41 @@ use dashmap::DashMap; use std::sync::Arc; use crate::server::SessionHandle; -use crate::server::constants::{JANITOR_INTERVAL, MASTER_EXPIRY, NONCE_CLEANUP_INTERVAL, now_secs}; -use crate::server::nonce_registry::NonceRegistry; +use crate::server::constants::{ + JANITOR_INTERVAL, MASTER_CLEANUP_INTERVAL, MASTER_EXPIRY, now_secs, +}; +use crate::server::stream_registry::StreamRegistry; -pub async fn master_and_nonce_janitor( +pub async fn master_and_stream_janitor( master_store: Arc>, - nonce_registry: Arc, + stream_registry: Arc, ) { - let mut interval = tokio::time::interval(NONCE_CLEANUP_INTERVAL); + let mut interval = tokio::time::interval(MASTER_CLEANUP_INTERVAL); interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); loop { interval.tick().await; - let mut has_deleted = false; let now = now_secs(); let expiry_limit = MASTER_EXPIRY.as_secs(); master_store.retain(|session_id, (_, _, created)| { if now.saturating_sub(*created) >= expiry_limit { - nonce_registry.remove_session(session_id); - has_deleted = true; + tracing::info!( + session_id = %session_id, + "master key expired, removing from store" + ); false } else { true } }); - if has_deleted { - master_store.shrink_to_fit(); - nonce_registry.shrink_to_fit(); + let cutoff = now.saturating_sub(expiry_limit); + let pruned = stream_registry.remove_consumed_before(cutoff); + if pruned > 0 { + tracing::debug!(pruned, "pruned consumed stream registry entries"); + stream_registry.shrink_to_fit(); } } } @@ -42,19 +47,19 @@ pub async fn stream_janitor(actors: Arc>) { loop { interval.tick().await; - let mut has_deleted = false; - actors.retain(|_, handle| { + actors.retain(|stream_id, handle| { if handle.cmd_tx.is_closed() { - has_deleted = true; + tracing::info!( + stream_id = %stream_id, + "stream actor channel closed, removing from actor map" + ); false } else { true } }); - if has_deleted { - actors.shrink_to_fit(); - } + actors.shrink_to_fit(); } } diff --git a/src/server/mod.rs b/src/server/mod.rs index 8c40852..d71e505 100644 --- a/src/server/mod.rs +++ b/src/server/mod.rs @@ -3,8 +3,8 @@ pub mod connection; pub mod constants; pub mod handlers; pub mod janitor; -pub mod nonce_registry; pub mod stream; +pub mod stream_registry; pub mod utils; use crate::config::ServerTopConfig; @@ -17,7 +17,6 @@ use axum::serve::ListenerExt; use axum::{Router, body::Body, routing::post}; use dashmap::DashMap; use jsonwebtoken::{Algorithm, DecodingKey, Validation}; -use nonce_registry::NonceRegistry; use serde::{Deserialize, Serialize}; use std::{ net::{IpAddr, SocketAddr}, @@ -26,6 +25,7 @@ use std::{ atomic::{AtomicU64, Ordering}, }, }; +use stream_registry::StreamRegistry; use tokio::sync::mpsc; use tower::ServiceBuilder; use tower_http::trace::TraceLayer; @@ -59,7 +59,7 @@ pub struct AppState { pub traffic_config: Arc, pub private_key: Option, pub master_store: Arc>, - pub nonce_registry: Arc, + pub stream_registry: Arc, pub actors: Arc>, pub stream_id_counter: Arc, } @@ -103,7 +103,7 @@ pub async fn build_state(config: &mut ServerTopConfig) -> anyhow::Result anyhow::Result<()> { pub fn spawn_janitors( state: &Arc, ) -> (tokio::task::JoinHandle<()>, tokio::task::JoinHandle<()>) { - let master_jh = tokio::spawn(janitor::master_and_nonce_janitor( + let master_jh = tokio::spawn(janitor::master_and_stream_janitor( Arc::clone(&state.master_store), - Arc::clone(&state.nonce_registry), + Arc::clone(&state.stream_registry), )); let stream_jh = tokio::spawn(janitor::stream_janitor(Arc::clone(&state.actors))); (master_jh, stream_jh) diff --git a/src/server/nonce_registry.rs b/src/server/nonce_registry.rs deleted file mode 100644 index e97712a..0000000 --- a/src/server/nonce_registry.rs +++ /dev/null @@ -1,240 +0,0 @@ -use dashmap::DashMap; -use std::sync::atomic::{AtomicU8, Ordering}; - -#[repr(u8)] -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum NonceState { - Active = 0, - Consumed = 1, -} - -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum NonceQueryResult { - Fresh, - Active, - Consumed, -} - -#[derive(Debug)] -pub struct NonceConsumedError; - -pub struct NonceRegistry { - sessions: DashMap>, -} - -impl NonceRegistry { - pub fn new() -> Self { - Self { - sessions: DashMap::new(), - } - } - - pub fn try_claim( - &self, - session_id: &str, - conn_nonce: &[u8; 16], - ) -> Result { - let per_session = self.sessions.entry(session_id.to_owned()).or_default(); - - match per_session.entry(*conn_nonce) { - dashmap::Entry::Occupied(entry) => { - let existing = entry.get().load(Ordering::Acquire); - if existing == NonceState::Active as u8 { - Ok(false) - } else { - Err(NonceConsumedError) - } - } - dashmap::Entry::Vacant(entry) => { - entry.insert(AtomicU8::new(NonceState::Active as u8)); - Ok(true) - } - } - } - - pub fn mark_consumed(&self, session_id: &str, conn_nonce: &[u8; 16]) { - if let Some(per_session) = self.sessions.get(session_id) - && let Some(entry) = per_session.get(conn_nonce) - { - entry.store(NonceState::Consumed as u8, Ordering::Release); - } - } - - pub fn check_nonce(&self, session_id: &str, conn_nonce: &[u8; 16]) -> NonceQueryResult { - match self.sessions.get(session_id) { - None => NonceQueryResult::Fresh, - Some(per_session) => match per_session.get(conn_nonce) { - None => NonceQueryResult::Fresh, - Some(entry) => match entry.load(Ordering::Acquire) { - s if s == NonceState::Active as u8 => NonceQueryResult::Active, - _ => NonceQueryResult::Consumed, - }, - }, - } - } - - pub fn remove_session(&self, session_id: &str) { - self.sessions.remove(session_id); - } - - pub fn shrink_to_fit(&self) { - self.sessions.shrink_to_fit(); - } -} - -impl Default for NonceRegistry { - fn default() -> Self { - Self::new() - } -} - -#[cfg(test)] -mod tests { - use super::*; - use std::sync::Arc; - - #[test] - fn fresh_claim_succeeds() { - let reg = NonceRegistry::new(); - let nonce = [0xAA; 16]; - assert!(matches!(reg.try_claim("s1", &nonce), Ok(true))); - } - - #[test] - fn active_nonce_reports_active() { - let reg = NonceRegistry::new(); - let nonce = [0xBB; 16]; - reg.try_claim("s1", &nonce).unwrap(); - assert_eq!(reg.check_nonce("s1", &nonce), NonceQueryResult::Active); - } - - #[test] - fn consumed_nonce_reports_consumed() { - let reg = NonceRegistry::new(); - let nonce = [0xCC; 16]; - reg.try_claim("s1", &nonce).unwrap(); - reg.mark_consumed("s1", &nonce); - assert_eq!(reg.check_nonce("s1", &nonce), NonceQueryResult::Consumed); - } - - #[test] - fn consumed_nonce_rejects_replay() { - let reg = NonceRegistry::new(); - let nonce = [0xDD; 16]; - reg.try_claim("s1", &nonce).unwrap(); - reg.mark_consumed("s1", &nonce); - assert!(reg.try_claim("s1", &nonce).is_err()); - } - - #[test] - fn active_nonce_allows_duplicate_claim() { - let reg = NonceRegistry::new(); - let nonce = [0xEE; 16]; - assert!(reg.try_claim("s1", &nonce).unwrap()); - assert!(matches!(reg.try_claim("s1", &nonce), Ok(false))); - assert_eq!(reg.check_nonce("s1", &nonce), NonceQueryResult::Active); - } - - #[test] - fn remove_session_clears_all_nonces() { - let reg = NonceRegistry::new(); - reg.try_claim("s1", &[1u8; 16]).unwrap(); - reg.try_claim("s1", &[2u8; 16]).unwrap(); - reg.remove_session("s1"); - assert_eq!(reg.check_nonce("s1", &[1u8; 16]), NonceQueryResult::Fresh); - assert_eq!(reg.check_nonce("s1", &[2u8; 16]), NonceQueryResult::Fresh); - } - - #[test] - fn concurrent_sessions_independent() { - let reg = NonceRegistry::new(); - let n1 = [1u8; 16]; - let n2 = [2u8; 16]; - reg.try_claim("s1", &n1).unwrap(); - reg.try_claim("s2", &n2).unwrap(); - reg.mark_consumed("s1", &n1); - assert_eq!(reg.check_nonce("s1", &n1), NonceQueryResult::Consumed); - assert_eq!(reg.check_nonce("s2", &n2), NonceQueryResult::Active); - } - - #[test] - fn fresh_nonce_returns_fresh() { - let reg = NonceRegistry::new(); - assert_eq!( - reg.check_nonce("nonexistent", &[0xFF; 16]), - NonceQueryResult::Fresh - ); - } - - #[tokio::test] - async fn concurrent_claims_different_nonces() { - let reg = Arc::new(NonceRegistry::new()); - let mut handles = Vec::new(); - - for i in 0..16u8 { - let reg = Arc::clone(®); - handles.push(tokio::spawn(async move { - let nonce = [i; 16]; - reg.try_claim("concurrent-session", &nonce) - })); - } - - let mut success = 0; - for handle in handles { - if let Ok(true) = handle.await.unwrap() { - success += 1 - } - } - assert_eq!(success, 16, "all 16 concurrent claims should succeed"); - } - - #[tokio::test] - async fn concurrent_claim_and_consume_race() { - let reg = Arc::new(NonceRegistry::new()); - let nonce = [0x42u8; 16]; - - assert!(matches!(reg.try_claim("race-session", &nonce), Ok(true))); - - let reg_clone = Arc::clone(®); - let h1 = tokio::spawn(async move { reg_clone.try_claim("race-session", &nonce) }); - - let reg_clone2 = Arc::clone(®); - let h2 = tokio::spawn(async move { - reg_clone2.mark_consumed("race-session", &nonce); - }); - - let (r1, _) = tokio::join!(h1, h2); - match r1.expect("claim task should not panic") { - Ok(false) => {} - Err(NonceConsumedError) => {} - Ok(true) => panic!("should not claim an already-Active nonce"), - } - } - - #[tokio::test] - async fn concurrent_remove_and_claim_race() { - let reg = Arc::new(NonceRegistry::new()); - let nonce = [0x99u8; 16]; - - assert!(matches!(reg.try_claim("remove-session", &nonce), Ok(true))); - - let reg_clone = Arc::clone(®); - let h1 = tokio::spawn(async move { - reg_clone.remove_session("remove-session"); - }); - - let other_nonce = [0x88u8; 16]; - let _ = reg.try_claim("remove-session", &other_nonce); - - h1.await.unwrap(); - - assert_eq!( - reg.check_nonce("remove-session", &nonce), - NonceQueryResult::Fresh - ); - assert_eq!( - reg.check_nonce("remove-session", &other_nonce), - NonceQueryResult::Fresh - ); - } -} diff --git a/src/server/stream_registry.rs b/src/server/stream_registry.rs new file mode 100644 index 0000000..e9554c9 --- /dev/null +++ b/src/server/stream_registry.rs @@ -0,0 +1,196 @@ +use dashmap::DashMap; +use std::sync::atomic::{AtomicU8, Ordering}; + +#[repr(u8)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum StreamState { + Active = 0, + Consumed = 1, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum StreamQueryResult { + Fresh, + Active, + Consumed, +} + +#[derive(Debug)] +pub struct StreamConsumedError; + +pub struct StreamRegistry { + streams: DashMap, +} + +impl StreamRegistry { + pub fn new() -> Self { + Self { + streams: DashMap::new(), + } + } + + pub fn register(&self, stream_id: &str, now_secs: u64) -> bool { + match self.streams.entry(stream_id.to_owned()) { + dashmap::Entry::Occupied(_) => false, + dashmap::Entry::Vacant(entry) => { + entry.insert((AtomicU8::new(StreamState::Active as u8), now_secs)); + true + } + } + } + + pub fn mark_consumed(&self, stream_id: &str) { + if let Some(entry) = self.streams.get(stream_id) { + entry + .0 + .store(StreamState::Consumed as u8, Ordering::Release); + } + } + + pub fn check(&self, stream_id: &str) -> StreamQueryResult { + match self.streams.get(stream_id) { + None => StreamQueryResult::Fresh, + Some(entry) => match entry.0.load(Ordering::Acquire) { + s if s == StreamState::Active as u8 => StreamQueryResult::Active, + _ => StreamQueryResult::Consumed, + }, + } + } + + pub fn remove_consumed_before(&self, cutoff_secs: u64) -> usize { + let mut removed = 0; + self.streams.retain(|_id, (state, ts)| { + let keep = + state.load(Ordering::Acquire) == StreamState::Active as u8 || *ts >= cutoff_secs; + if !keep { + removed += 1; + } + keep + }); + removed + } + + pub fn shrink_to_fit(&self) { + self.streams.shrink_to_fit(); + } +} + +impl Default for StreamRegistry { + fn default() -> Self { + Self::new() + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::Arc; + + fn dummy_ts() -> u64 { + 1_700_000_000 + } + + #[test] + fn fresh_register_succeeds() { + let reg = StreamRegistry::new(); + assert!(reg.register("s1", dummy_ts())); + assert_eq!(reg.check("s1"), StreamQueryResult::Active); + } + + #[test] + fn duplicate_register_rejected() { + let reg = StreamRegistry::new(); + assert!(reg.register("s1", dummy_ts())); + assert!(!reg.register("s1", dummy_ts())); + } + + #[test] + fn consumed_stream_reports_consumed() { + let reg = StreamRegistry::new(); + reg.register("s1", dummy_ts()); + reg.mark_consumed("s1"); + assert_eq!(reg.check("s1"), StreamQueryResult::Consumed); + } + + #[test] + fn fresh_stream_returns_fresh() { + let reg = StreamRegistry::new(); + assert_eq!(reg.check("nonexistent"), StreamQueryResult::Fresh); + } + + #[test] + fn remove_consumed_before_removes_expired() { + let reg = StreamRegistry::new(); + reg.register("s1", 100); + reg.mark_consumed("s1"); + reg.register("s2", 200); + let removed = reg.remove_consumed_before(150); + assert!(removed >= 1); + assert_eq!(reg.check("s1"), StreamQueryResult::Fresh); + assert_eq!(reg.check("s2"), StreamQueryResult::Active); + } + + #[test] + fn remove_consumed_before_keeps_recent() { + let reg = StreamRegistry::new(); + reg.register("s1", 1000); + reg.mark_consumed("s1"); + let removed = reg.remove_consumed_before(900); + assert_eq!(removed, 0); + assert_eq!(reg.check("s1"), StreamQueryResult::Consumed); + } + + #[test] + fn concurrent_streams_independent() { + let reg = StreamRegistry::new(); + assert!(reg.register("s1", dummy_ts())); + assert!(reg.register("s2", dummy_ts())); + reg.mark_consumed("s1"); + assert_eq!(reg.check("s1"), StreamQueryResult::Consumed); + assert_eq!(reg.check("s2"), StreamQueryResult::Active); + } + + #[tokio::test] + async fn concurrent_registrations_different_ids() { + let reg = Arc::new(StreamRegistry::new()); + let mut handles = Vec::new(); + + for i in 0..16u8 { + let reg = Arc::clone(®); + handles.push(tokio::spawn(async move { + let id = format!("stream-{i}"); + reg.register(&id, dummy_ts()) + })); + } + + let mut success = 0; + for handle in handles { + if handle.await.unwrap() { + success += 1; + } + } + assert_eq!( + success, 16, + "all 16 concurrent registrations should succeed" + ); + } + + #[tokio::test] + async fn concurrent_register_and_consume_race() { + let reg = Arc::new(StreamRegistry::new()); + let id = "race-stream"; + + assert!(reg.register(id, dummy_ts())); + + let reg_clone = Arc::clone(®); + let h1 = tokio::spawn(async move { reg_clone.register(id, dummy_ts()) }); + + let reg_clone2 = Arc::clone(®); + let h2 = tokio::spawn(async move { + reg_clone2.mark_consumed(id); + }); + + let (r1, _) = tokio::join!(h1, h2); + assert!(!r1.expect("register task should not panic")); + } +} diff --git a/src/server/utils.rs b/src/server/utils.rs index 0088f91..4b73910 100644 --- a/src/server/utils.rs +++ b/src/server/utils.rs @@ -1,5 +1,7 @@ use rand::RngExt; +use uuid::Uuid; +use crate::error::ServerError; use crate::server::constants::PADDING_POOL; #[inline] @@ -41,6 +43,11 @@ pub fn extract_cookie_value<'a>(headers: &'a axum::http::HeaderMap, key: &str) - None } +#[inline] +pub fn validate_uuid(s: &str) -> Result { + Uuid::parse_str(s).map_err(|_| ServerError::bad_request("invalid UUID format")) +} + #[inline] pub fn random_padding() -> &'static [u8] { let padding_len = rand::rng().random_range(30..=PADDING_POOL.len());