diff --git a/examples/sdk_echo_server.rs b/examples/sdk_echo_server.rs index 11857b6..3350400 100644 --- a/examples/sdk_echo_server.rs +++ b/examples/sdk_echo_server.rs @@ -16,6 +16,9 @@ async fn main() -> Result<()> { client.user, client.session, client.conn_id ); } + DoshServerEvent::Disconnected(client) => { + eprintln!("disconnected conn={:?}", client.conn_id); + } DoshServerEvent::Session { conn_id, event: SessionEvent::Stream(TransportEvent::Open(open)), diff --git a/src/server.rs b/src/server.rs index afe564b..59e7b69 100644 --- a/src/server.rs +++ b/src/server.rs @@ -16,8 +16,8 @@ use crate::transport::{ use crate::udp::{is_transient_udp_error, is_transient_udp_send_error}; use anyhow::{Context, Result, anyhow, bail}; use ed25519_dalek::SigningKey; -use std::collections::{HashMap, HashSet}; -use std::net::SocketAddr; +use std::collections::{HashMap, HashSet, VecDeque}; +use std::net::{IpAddr, SocketAddr}; use std::path::PathBuf; use std::sync::Arc; use std::time::{Duration, Instant}; @@ -31,10 +31,13 @@ pub struct DoshServerConfig { pub transport: TransportConfig, pub require_current_user: bool, pub auth_timeout: Duration, + pub max_pending_auth: usize, + pub connection_timeout: Duration, } impl DoshServerConfig { pub fn new(server: ServerConfig) -> Self { + let connection_timeout = Duration::from_secs(server.client_timeout_secs.max(1)); Self { server, bind_addr: None, @@ -42,6 +45,8 @@ impl DoshServerConfig { transport: TransportConfig::default(), require_current_user: true, auth_timeout: Duration::from_secs(30), + max_pending_auth: 1024, + connection_timeout, } } @@ -71,6 +76,21 @@ impl DoshServerConfig { self.transport = transport; self } + + pub fn connection_timeout(mut self, timeout: Duration) -> Self { + self.connection_timeout = timeout.max(ADAPTIVE_RETRANSMIT_MIN); + self + } + + pub fn auth_timeout(mut self, timeout: Duration) -> Self { + self.auth_timeout = timeout.max(ADAPTIVE_RETRANSMIT_MIN); + self + } + + pub fn max_pending_auth(mut self, max_pending: usize) -> Self { + self.max_pending_auth = max_pending.max(1); + self + } } impl Default for DoshServerConfig { @@ -91,6 +111,7 @@ pub struct DoshAccepted { #[derive(Debug, Clone)] pub enum DoshServerEvent { Accepted(DoshAccepted), + Disconnected(DoshAccepted), Session { conn_id: [u8; 16], event: SessionEvent, @@ -107,13 +128,78 @@ struct PendingServerAuth { created_at: Instant, } +struct AuthRateLimiter { + per_minute: u32, + max_sources: usize, + buckets: HashMap, +} + +#[derive(Clone, Copy)] +struct AuthTokenBucket { + tokens: f64, + last_refill: Instant, +} + +impl AuthRateLimiter { + fn new(per_minute: u32, max_sources: usize) -> Self { + Self { + per_minute, + max_sources: max_sources.max(1), + buckets: HashMap::new(), + } + } + + fn check(&mut self, ip: IpAddr, now: Instant) -> Result { + if self.per_minute == 0 { + return Err(()); + } + self.evict_full(now); + if !self.buckets.contains_key(&ip) && self.buckets.len() >= self.max_sources { + return Err(()); + } + let capacity = self.per_minute as f64; + let refill_per_sec = capacity / 60.0; + let bucket = self.buckets.entry(ip).or_insert(AuthTokenBucket { + tokens: capacity, + last_refill: now, + }); + let elapsed = now + .saturating_duration_since(bucket.last_refill) + .as_secs_f64(); + bucket.tokens = (bucket.tokens + elapsed * refill_per_sec).min(capacity); + bucket.last_refill = now; + if bucket.tokens < 1.0 { + return Err(()); + } + bucket.tokens -= 1.0; + Ok(bucket.tokens as u32) + } + + fn evict_full(&mut self, now: Instant) { + if self.per_minute == 0 { + self.buckets.clear(); + return; + } + let capacity = self.per_minute as f64; + let refill_per_sec = capacity / 60.0; + self.buckets.retain(|_, bucket| { + let elapsed = now + .saturating_duration_since(bucket.last_refill) + .as_secs_f64(); + (bucket.tokens + elapsed * refill_per_sec) < capacity + }); + } +} + pub struct DoshServer { socket: Arc, config: DoshServerConfig, host_signing: SigningKey, + auth_limiter: AuthRateLimiter, pending: HashMap<[u8; 16], PendingServerAuth>, transports: HashMap<[u8; 16], DoshTransport>, accepted: HashMap<[u8; 16], DoshAccepted>, + disconnected: VecDeque, } impl DoshServer { @@ -139,13 +225,19 @@ impl DoshServer { }; let host_signing = load_or_create_host_key(&config.server)?; let socket = Arc::new(UdpSocket::bind(bind_addr).await?); + let auth_limiter = AuthRateLimiter::new( + config.server.native_auth_rate_limit_per_minute, + config.max_pending_auth, + ); Ok(Self { socket, config, host_signing, + auth_limiter, pending: HashMap::new(), transports: HashMap::new(), accepted: HashMap::new(), + disconnected: VecDeque::new(), }) } @@ -165,10 +257,19 @@ impl DoshServer { self.transports.get_mut(conn_id) } + pub fn remove_connection(&mut self, conn_id: &[u8; 16]) -> Option { + self.transports.remove(conn_id); + self.accepted.remove(conn_id) + } + pub async fn recv(&mut self) -> Result { let mut buf = vec![0u8; 65535]; loop { self.expire_pending(); + self.expire_connections(); + if let Some(connection) = self.disconnected.pop_front() { + return Ok(DoshServerEvent::Disconnected(connection)); + } let (n, peer) = match tokio::time::timeout( ADAPTIVE_RETRANSMIT_MIN, self.socket.recv_from(&mut buf), @@ -207,6 +308,9 @@ impl DoshServer { } if let Some(transport) = self.transports.get_mut(&packet.header.conn_id) { let event = transport.handle_datagram(datagram, peer).await?; + if let Some(accepted) = self.accepted.get_mut(&packet.header.conn_id) { + accepted.peer_addr = transport.peer_addr(); + } return Ok(Some(DoshServerEvent::Session { conn_id: packet.header.conn_id, event, @@ -270,7 +374,21 @@ impl DoshServer { self.send_reject(peer, [0u8; 16], &err.to_string()).await?; return Ok(()); } - let result = self.build_server_hello(req.hello, peer); + self.expire_pending(); + if self.pending.len() >= self.config.max_pending_auth { + self.send_reject(peer, [0u8; 16], "native auth server busy") + .await?; + return Ok(()); + } + let rate_limit_remaining = match self.auth_limiter.check(peer.ip(), Instant::now()) { + Ok(remaining) => remaining, + Err(()) => { + self.send_reject(peer, [0u8; 16], "native auth rate limit exceeded") + .await?; + return Ok(()); + } + }; + let result = self.build_server_hello(req.hello, peer, Some(rate_limit_remaining)); let (pending_id, hello) = match result { Ok(value) => value, Err(err) => { @@ -288,6 +406,7 @@ impl DoshServer { &mut self, client: native::NativeClientHello, peer: SocketAddr, + rate_limit_remaining: Option, ) -> Result<([u8; 16], NativeServerHello)> { if !self.config.server.native_auth { bail!("native auth disabled"); @@ -307,11 +426,14 @@ impl DoshServer { bail!("native auth requires a supported user key algorithm"); } if self.config.require_current_user { - let current_user = std::env::var("USER").unwrap_or_else(|_| "unknown".to_string()); + let current_user = local_username(); if client.requested_user != current_user { bail!("native auth user mismatch"); } } + if self.pending.len() >= self.config.max_pending_auth { + bail!("native auth server busy"); + } let (server_secret, server_public) = generate_native_ephemeral(); let mut server = NativeServerHello { @@ -322,7 +444,7 @@ impl DoshServer { chosen_aead: "chacha20poly1305".to_string(), server_key_epoch: 1, auth_challenge: crypto::random_32(), - rate_limit_remaining: None, + rate_limit_remaining, host_signature: Vec::new(), }; sign_server_hello(&self.host_signing, &client, &mut server)?; @@ -477,7 +599,37 @@ impl DoshServer { let timeout = self.config.auth_timeout; self.pending .retain(|_, pending| pending.created_at.elapsed() <= timeout); + self.auth_limiter.evict_full(Instant::now()); } + + fn expire_connections(&mut self) { + let timeout = self.config.connection_timeout; + let expired = self + .transports + .iter() + .filter_map(|(conn_id, transport)| { + (transport.stale_for() > timeout).then_some(*conn_id) + }) + .collect::>(); + for conn_id in expired { + self.transports.remove(&conn_id); + if let Some(connection) = self.accepted.remove(&conn_id) { + self.disconnected.push_back(connection); + } + } + } +} + +fn local_username() -> String { + local_username_from_env(|name| std::env::var(name).ok()) +} + +fn local_username_from_env(get: impl FnMut(&str) -> Option) -> String { + ["USER", "USERNAME"] + .into_iter() + .filter_map(get) + .find(|value| !value.is_empty()) + .unwrap_or_else(|| "unknown".to_string()) } async fn send_udp(socket: &UdpSocket, packet: &[u8], peer: SocketAddr) -> Result { @@ -530,6 +682,24 @@ mod tests { use crate::transport::TransportEvent; use ed25519_dalek::SigningKey; + fn native_client_hello(public: [u8; 32]) -> native::NativeClientHello { + native::NativeClientHello { + protocol_version: native::NATIVE_PROTOCOL_VERSION, + client_random: crypto::random_32(), + client_ephemeral_public: public, + requested_host: "127.0.0.1".to_string(), + requested_user: "sdk-user".to_string(), + requested_session: "test".to_string(), + requested_mode: "forward-only".to_string(), + terminal_size: (80, 24), + supported_aead: vec!["chacha20poly1305".to_string()], + supported_user_key_algorithms: vec!["ssh-ed25519".to_string()], + cached_host_key_fingerprint: None, + attach_ticket_envelope: None, + requested_env: Vec::new(), + } + } + #[tokio::test] async fn sdk_client_and_server_exchange_service_stream() { let dir = tempfile::tempdir().unwrap(); @@ -752,6 +922,209 @@ mod tests { assert_eq!(open.target_host, "@dosh-echo"); } + #[tokio::test] + async fn sdk_server_expires_disconnected_transport_and_reports_it() { + let dir = tempfile::tempdir().unwrap(); + let server_config = ServerConfig { + host_key: dir.path().join("host_key").to_string_lossy().to_string(), + ..ServerConfig::default() + }; + let server_config = DoshServerConfig::new(server_config) + .bind_addr("127.0.0.1:0".parse().unwrap()) + .require_current_user(false) + .connection_timeout(Duration::from_millis(30)); + let mut server = DoshServer::bind(server_config).await.unwrap(); + let peer_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap(); + let peer_addr = peer_socket.local_addr().unwrap(); + let conn_id = [81u8; 16]; + let transport = DoshTransport::new( + Arc::clone(&server.socket), + SessionTransportConfig { + role: SessionRole::Server, + conn_id, + session_key: [82u8; 32], + peer_addr, + initial_send_seq: 1, + initial_ack: 0, + stream: server.config.transport.clone(), + }, + ); + let accepted = DoshAccepted { + conn_id, + user: "sdk-user".to_string(), + session: "mobile".to_string(), + services: vec!["echo".to_string()], + peer_addr, + }; + server.transports.insert(conn_id, transport); + server.accepted.insert(conn_id, accepted.clone()); + + tokio::time::sleep(Duration::from_millis(40)).await; + let event = tokio::time::timeout(Duration::from_secs(1), server.recv()) + .await + .expect("server did not report expired connection") + .unwrap(); + match event { + DoshServerEvent::Disconnected(disconnected) => { + assert_eq!(disconnected.conn_id, conn_id); + assert_eq!(disconnected.session, "mobile"); + } + other => panic!("unexpected expiry event {other:?}"), + } + assert!(server.connection(&conn_id).is_none()); + assert!(server.transport(&conn_id).is_none()); + } + + #[tokio::test] + async fn sdk_server_updates_connection_metadata_after_authenticated_roam() { + let dir = tempfile::tempdir().unwrap(); + let server_config = ServerConfig { + host_key: dir.path().join("host_key").to_string_lossy().to_string(), + ..ServerConfig::default() + }; + let server_config = DoshServerConfig::new(server_config) + .bind_addr("127.0.0.1:0".parse().unwrap()) + .require_current_user(false); + let mut server = DoshServer::bind(server_config).await.unwrap(); + let original = UdpSocket::bind("127.0.0.1:0").await.unwrap(); + let roaming = UdpSocket::bind("127.0.0.1:0").await.unwrap(); + let original_addr = original.local_addr().unwrap(); + let roaming_addr = roaming.local_addr().unwrap(); + let conn_id = [83u8; 16]; + let session_key = [84u8; 32]; + let transport = DoshTransport::new( + Arc::clone(&server.socket), + SessionTransportConfig { + role: SessionRole::Server, + conn_id, + session_key, + peer_addr: original_addr, + initial_send_seq: 1, + initial_ack: 0, + stream: server.config.transport.clone(), + }, + ); + server.transports.insert(conn_id, transport); + server.accepted.insert( + conn_id, + DoshAccepted { + conn_id, + user: "sdk-user".to_string(), + session: "mobile".to_string(), + services: Vec::new(), + peer_addr: original_addr, + }, + ); + let ping = protocol::encode_encrypted( + PacketKind::Ping, + conn_id, + 1, + 0, + &session_key, + CLIENT_TO_SERVER, + b"", + ) + .unwrap(); + roaming + .send_to(&ping, server.local_addr().unwrap()) + .await + .unwrap(); + + let event = tokio::time::timeout(Duration::from_secs(1), server.recv()) + .await + .expect("server did not receive roaming ping") + .unwrap(); + assert!(matches!( + event, + DoshServerEvent::Session { + event: SessionEvent::Ping, + .. + } + )); + assert_eq!(server.connection(&conn_id).unwrap().peer_addr, roaming_addr); + } + + #[tokio::test] + async fn sdk_server_bounds_and_rate_limits_pending_authentication() { + let dir = tempfile::tempdir().unwrap(); + let server_config = ServerConfig { + host_key: dir.path().join("host_key").to_string_lossy().to_string(), + native_auth_rate_limit_per_minute: 1, + ..ServerConfig::default() + }; + let server_config = DoshServerConfig::new(server_config) + .bind_addr("127.0.0.1:0".parse().unwrap()) + .require_current_user(false) + .max_pending_auth(1); + let mut server = DoshServer::bind(server_config).await.unwrap(); + let client = UdpSocket::bind("127.0.0.1:0").await.unwrap(); + let peer = client.local_addr().unwrap(); + let (_, public) = native::generate_native_ephemeral(); + let body = protocol::to_body(&NativeClientHelloBody { + hello: native_client_hello(public), + }) + .unwrap(); + + server + .handle_client_hello(peer, body.clone()) + .await + .unwrap(); + let mut packet = [0u8; 65535]; + let (n, _) = client.recv_from(&mut packet).await.unwrap(); + let response = protocol::decode(&packet[..n]).unwrap(); + assert_eq!(response.header.kind, PacketKind::NativeServerHello); + let hello: NativeServerHelloBody = protocol::from_body(&response.body).unwrap(); + assert_eq!(hello.hello.rate_limit_remaining, Some(0)); + assert_eq!(server.pending.len(), 1); + + server + .handle_client_hello(peer, body.clone()) + .await + .unwrap(); + let (n, _) = client.recv_from(&mut packet).await.unwrap(); + let response = protocol::decode(&packet[..n]).unwrap(); + assert_eq!(response.header.kind, PacketKind::AttachReject); + let reject: AttachReject = protocol::from_body(&response.body).unwrap(); + assert_eq!(reject.reason, "native auth server busy"); + assert_eq!(server.pending.len(), 1); + + server.pending.clear(); + server.handle_client_hello(peer, body).await.unwrap(); + let (n, _) = client.recv_from(&mut packet).await.unwrap(); + let response = protocol::decode(&packet[..n]).unwrap(); + assert_eq!(response.header.kind, PacketKind::AttachReject); + let reject: AttachReject = protocol::from_body(&response.body).unwrap(); + assert_eq!(reject.reason, "native auth rate limit exceeded"); + assert!(server.pending.is_empty()); + } + + #[test] + fn sdk_auth_rate_limiter_bounds_source_tracking_and_refills() { + let now = Instant::now(); + let mut limiter = AuthRateLimiter::new(2, 1); + let first: IpAddr = "192.0.2.1".parse().unwrap(); + let second: IpAddr = "192.0.2.2".parse().unwrap(); + + assert_eq!(limiter.check(first, now), Ok(1)); + assert_eq!(limiter.check(first, now), Ok(0)); + assert_eq!(limiter.check(first, now), Err(())); + assert_eq!(limiter.check(second, now), Err(())); + assert_eq!(limiter.check(second, now + Duration::from_secs(60)), Ok(1)); + assert_eq!(limiter.buckets.len(), 1); + } + + #[test] + fn sdk_server_username_supports_windows_environment() { + assert_eq!( + local_username_from_env(|name| match name { + "USER" => None, + "USERNAME" => Some("palav-win".to_string()), + _ => None, + }), + "palav-win" + ); + } + #[tokio::test] async fn bad_native_auth_does_not_consume_pending_challenge() { let dir = tempfile::tempdir().unwrap(); @@ -778,22 +1151,10 @@ mod tests { let mut server = DoshServer::bind(server_config).await.unwrap(); let peer: SocketAddr = "127.0.0.1:9".parse().unwrap(); let (client_secret, client_public) = native::generate_native_ephemeral(); - let hello = native::NativeClientHello { - protocol_version: native::NATIVE_PROTOCOL_VERSION, - client_random: crypto::random_32(), - client_ephemeral_public: client_public, - requested_host: "127.0.0.1".to_string(), - requested_user: "sdk-user".to_string(), - requested_session: "test".to_string(), - requested_mode: "forward-only".to_string(), - terminal_size: (80, 24), - supported_aead: vec!["chacha20poly1305".to_string()], - supported_user_key_algorithms: vec!["ssh-ed25519".to_string()], - cached_host_key_fingerprint: None, - attach_ticket_envelope: None, - requested_env: Vec::new(), - }; - let (pending_id, server_hello) = server.build_server_hello(hello.clone(), peer).unwrap(); + let hello = native_client_hello(client_public); + let (pending_id, server_hello) = server + .build_server_hello(hello.clone(), peer, None) + .unwrap(); let session_key = native::derive_native_session_key( &client_secret, server_hello.server_ephemeral_public,