1use crate::metrics::ReverseMetrics;
2use crate::{
3 redact_auth, relay_bidirectional_with_timeout, server_auth_handshake, ControlState,
4 ProtocolError,
5};
6use std::net::{IpAddr, SocketAddr};
7use std::sync::atomic::{AtomicU32, Ordering};
8use std::sync::Arc;
9use std::time::{Duration, Instant};
10use tokio::net::{TcpListener, TcpStream};
11use tokio::sync::mpsc;
12use tokio::task::JoinSet;
13use tokio_util::sync::CancellationToken;
14use tracing::{debug, error, info, warn};
15
16const AUTH_FAILURE_DELAY: Duration = Duration::from_millis(100);
17
18#[derive(Debug, Clone)]
23pub struct ReverseServerConfig {
24 pub control_bind: SocketAddr,
26 pub external_bind: Option<SocketAddr>,
28 pub auth_username: Option<String>,
30 pub auth_password: Option<String>,
32 pub max_control_connections: u32,
34 pub read_timeout_ms: u64,
36 pub allow_bind: Option<Vec<SocketAddr>>,
40 pub max_listeners_per_client: u32,
44 pub max_streams_per_listener: u32,
46 pub max_pending_external: u32,
49}
50
51impl Default for ReverseServerConfig {
52 fn default() -> Self {
53 Self {
54 control_bind: "127.0.0.1:0".parse().unwrap(),
55 external_bind: None,
56 auth_username: None,
57 auth_password: None,
58 max_control_connections: 256,
59 read_timeout_ms: 300_000,
60 allow_bind: None,
61 max_listeners_per_client: 1,
62 max_streams_per_listener: 1024,
63 max_pending_external: 1024,
64 }
65 }
66}
67
68impl ReverseServerConfig {
69 pub fn is_bind_allowed(&self, addr: SocketAddr) -> bool {
73 match &self.allow_bind {
74 None => true,
75 Some(list) if list.is_empty() => true,
76 Some(list) => list.iter().any(|allowed| same_bind(allowed, &addr)),
77 }
78 }
79
80 pub fn is_loopback(addr: SocketAddr) -> bool {
82 match addr.ip() {
83 IpAddr::V4(v4) => v4.is_loopback(),
84 IpAddr::V6(v6) => v6.is_loopback(),
85 }
86 }
87
88 pub fn validate(&self) -> Result<(), ProtocolError> {
96 if let Some(external) = self.external_bind {
97 if !Self::is_loopback(external) {
102 let has_auth = self.auth_username.as_deref().is_some_and(|s| !s.is_empty())
103 && self.auth_password.as_deref().is_some_and(|s| !s.is_empty());
104 let has_allowlist = matches!(&self.allow_bind, Some(list) if !list.is_empty());
105 if !has_auth {
106 return Err(ProtocolError::ConfigInvalid(format!(
107 "reverse server external_bind={external} is non-loopback but no \
108 authentication is configured; set auth_username/auth_password or \
109 bind to loopback"
110 )));
111 }
112 if !has_allowlist {
113 return Err(ProtocolError::ConfigInvalid(format!(
114 "reverse server external_bind={external} is non-loopback but \
115 allow_bind is empty; configure an explicit allowlist"
116 )));
117 }
118 }
119 }
120 Ok(())
121 }
122}
123
124fn same_bind(a: &SocketAddr, b: &SocketAddr) -> bool {
125 a.port() == b.port()
126 && match (a.ip(), b.ip()) {
127 (IpAddr::V4(a4), IpAddr::V4(b4)) => a4 == b4,
128 (IpAddr::V6(a6), IpAddr::V6(b6)) => a6 == b6,
129 _ => false,
130 }
131}
132
133#[derive(Debug, Default)]
135pub struct ReverseServerState {
136 pub active_control: AtomicU32,
138 pub active_streams: AtomicU32,
140 pub pending_external: AtomicU32,
142 pub denied_bind: AtomicU32,
144 pub dropped_stream_limit: AtomicU32,
146 pub dropped_pending_limit: AtomicU32,
148}
149
150impl ReverseServerState {
151 pub fn snapshot(&self) -> ReverseServerStateSnapshot {
153 ReverseServerStateSnapshot {
154 active_control: self.active_control.load(Ordering::Relaxed),
155 active_streams: self.active_streams.load(Ordering::Relaxed),
156 pending_external: self.pending_external.load(Ordering::Relaxed),
157 denied_bind: self.denied_bind.load(Ordering::Relaxed),
158 dropped_stream_limit: self.dropped_stream_limit.load(Ordering::Relaxed),
159 dropped_pending_limit: self.dropped_pending_limit.load(Ordering::Relaxed),
160 }
161 }
162}
163
164#[derive(Debug, Clone, serde::Serialize)]
166pub struct ReverseServerStateSnapshot {
167 pub active_control: u32,
168 pub active_streams: u32,
169 pub pending_external: u32,
170 pub denied_bind: u32,
171 pub dropped_stream_limit: u32,
172 pub dropped_pending_limit: u32,
173}
174
175pub struct ReverseServer {
181 config: ReverseServerConfig,
182 cancel: CancellationToken,
183 metrics: Option<Arc<ReverseMetrics>>,
184 state: Arc<ReverseServerState>,
185}
186
187impl ReverseServer {
188 pub fn new(config: ReverseServerConfig) -> Self {
189 Self {
190 config,
191 cancel: CancellationToken::new(),
192 metrics: None,
193 state: Arc::new(ReverseServerState::default()),
194 }
195 }
196
197 pub fn set_metrics(&mut self, metrics: Arc<ReverseMetrics>) {
199 self.metrics = Some(metrics);
200 }
201
202 pub fn state_handle(&self) -> Arc<ReverseServerState> {
204 self.state.clone()
205 }
206
207 pub fn cancel_token(&self) -> CancellationToken {
209 self.cancel.clone()
210 }
211
212 async fn bind_external_listener(
215 config: &ReverseServerConfig,
216 state: &ReverseServerState,
217 ) -> Result<Option<TcpListener>, ProtocolError> {
218 let external_bind = match config.external_bind {
219 Some(addr) => addr,
220 None => return Ok(None),
221 };
222 if !config.is_bind_allowed(external_bind) {
223 state.denied_bind.fetch_add(1, Ordering::Relaxed);
224 return Err(ProtocolError::BindDenied(external_bind));
225 }
226 let listener = TcpListener::bind(external_bind).await?;
227 let addr = listener.local_addr()?;
228 info!(addr = %addr, "reverse server listening for external clients");
229 Ok(Some(listener))
230 }
231
232 pub async fn run(self) -> Result<(), ProtocolError> {
234 self.config.validate()?;
238
239 if self.config.auth_username.is_none() || self.config.auth_password.is_none() {
240 warn!(
241 control_bind = %self.config.control_bind,
242 "reverse server control channel has no authentication configured"
243 );
244 }
245
246 if let Some(external_bind) = self.config.external_bind {
248 if !self.config.is_bind_allowed(external_bind) {
249 self.state.denied_bind.fetch_add(1, Ordering::Relaxed);
250 return Err(ProtocolError::BindDenied(external_bind));
251 }
252 }
253
254 let control_listener = TcpListener::bind(&self.config.control_bind).await?;
255 let control_addr = control_listener.local_addr()?;
256 info!(addr = %control_addr, "reverse server listening for control connections");
257
258 let external_listener = Self::bind_external_listener(&self.config, &self.state).await?;
259
260 let config = Arc::new(self.config);
261 let cancel = self.cancel.clone();
262 let state = self.state.clone();
263 let metrics = self.metrics.clone();
264
265 let (control_tx, control_rx) = mpsc::channel::<ControlStream>(256);
267
268 let config_clone = config.clone();
270 let cancel_clone = cancel.clone();
271 let control_tx_clone = control_tx.clone();
272 let metrics_clone = metrics.clone();
273 let state_clone = state.clone();
274 let control_task = tokio::spawn(async move {
275 Self::accept_control_connections(
276 control_listener,
277 config_clone,
278 cancel_clone,
279 control_tx_clone,
280 metrics_clone,
281 state_clone,
282 )
283 .await;
284 });
285
286 let external_task = if let Some(external_listener) = external_listener {
288 let config_clone = config.clone();
289 let cancel_clone = cancel.clone();
290 let metrics_clone = metrics.clone();
291 let state_clone = state.clone();
292 Some(tokio::spawn(async move {
293 Self::accept_external_clients(
294 external_listener,
295 config_clone,
296 cancel_clone,
297 control_rx,
298 metrics_clone,
299 state_clone,
300 )
301 .await;
302 }))
303 } else {
304 let state_clone = state.clone();
309 let metrics_clone = metrics.clone();
310 let cancel_clone = cancel.clone();
311 Some(tokio::spawn(async move {
312 let mut control_rx = control_rx;
313 loop {
314 tokio::select! {
315 Some(ctrl) = control_rx.recv() => {
316 debug!(
317 control_peer = %ctrl.peer_addr,
318 "dropping control connection: no external listener"
319 );
320 drop(ctrl.stream);
321 state_clone.active_control.fetch_sub(1, Ordering::Relaxed);
322 if let Some(m) = metrics_clone.as_deref() {
323 m.record_control_closed();
324 }
325 }
326 _ = cancel_clone.cancelled() => break,
327 }
328 }
329 }))
330 };
331
332 cancel.cancelled().await;
334 let drain_start = Instant::now();
335 info!("reverse server shutting down, draining active streams");
336
337 let _ = control_task.await;
340 if let Some(task) = external_task {
341 let _ = task.await;
342 }
343 let drain_ms = drain_start.elapsed().as_millis() as u64;
344 if let Some(ref m) = metrics {
345 m.record_drain(drain_ms);
346 }
347 info!(drain_ms, "reverse server drain complete");
348 Ok(())
349 }
350
351 async fn accept_control_connections(
353 listener: TcpListener,
354 config: Arc<ReverseServerConfig>,
355 cancel: CancellationToken,
356 control_tx: mpsc::Sender<ControlStream>,
357 metrics: Option<Arc<ReverseMetrics>>,
358 state: Arc<ReverseServerState>,
359 ) {
360 loop {
361 tokio::select! {
362 result = listener.accept() => {
363 match result {
364 Ok((stream, peer_addr)) => {
365 let prev = state.active_control.fetch_add(1, Ordering::AcqRel);
368 if prev >= config.max_control_connections {
369 state.active_control.fetch_sub(1, Ordering::Relaxed);
370 warn!(
371 peer = %peer_addr,
372 max = config.max_control_connections,
373 "rejecting control connection: max reached"
374 );
375 if let Some(ref m) = metrics {
376 m.record_control_rejected(peer_addr, "max_control_connections");
377 }
378 drop(stream);
379 continue;
380 }
381
382 let config = config.clone();
383 let control_tx = control_tx.clone();
384 let metrics = metrics.clone();
385 let state = state.clone();
386 tokio::spawn(async move {
387 if let Err(e) = Self::handle_control_connection(
388 stream,
389 peer_addr,
390 config,
391 control_tx,
392 metrics.as_deref(),
393 state.clone(),
394 ).await {
395 state.active_control.fetch_sub(1, Ordering::Relaxed);
396 debug!(peer = %peer_addr, error = %e, "control connection handler error");
397 }
398 });
399 }
400 Err(e) => {
401 error!(error = %e, "failed to accept control connection");
402 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
403 }
404 }
405 }
406 _ = cancel.cancelled() => {
407 break;
408 }
409 }
410 }
411 }
412
413 async fn handle_control_connection(
415 mut stream: TcpStream,
416 peer_addr: SocketAddr,
417 config: Arc<ReverseServerConfig>,
418 control_tx: mpsc::Sender<ControlStream>,
419 metrics: Option<&ReverseMetrics>,
420 state: Arc<ReverseServerState>,
421 ) -> Result<(), ProtocolError> {
422 info!(peer = %peer_addr, state = ?ControlState::Connecting, "new control connection");
423
424 let redacted = if config.auth_username.is_some() && config.auth_password.is_some() {
426 let authenticating_start = Instant::now();
427 let result = server_auth_handshake(
428 &mut stream,
429 config.auth_username.as_deref(),
430 config.auth_password.as_deref(),
431 )
432 .await;
433 let elapsed = authenticating_start.elapsed().as_millis() as u64;
434
435 match result {
436 Ok(redacted) => {
437 info!(
438 peer = %peer_addr,
439 auth = %redacted,
440 duration_ms = elapsed,
441 state = ?ControlState::Authenticating,
442 "control connection authenticated"
443 );
444 if let Some(m) = metrics {
445 m.record_control_accepted(peer_addr);
446 m.record_state_duration(ControlState::Authenticating, elapsed);
447 }
448 Some(redacted)
449 }
450 Err(e) => {
451 warn!(
452 peer = %peer_addr,
453 error = %e,
454 duration_ms = elapsed,
455 state = ?ControlState::Authenticating,
456 "control connection auth failed"
457 );
458 if let Some(m) = metrics {
459 m.record_auth_failure(peer_addr, &e.to_string());
460 }
461 tokio::time::sleep(AUTH_FAILURE_DELAY).await;
462 return Err(e);
463 }
464 }
465 } else {
466 crate::write_handshake_accept(&mut stream).await?;
468 info!(
469 peer = %peer_addr,
470 state = ?ControlState::Authenticating,
471 "control connection accepted (no auth)"
472 );
473 if let Some(m) = metrics {
474 m.record_control_accepted(peer_addr);
475 }
476 None
477 };
478
479 let ctrl = ControlStream {
480 stream,
481 peer_addr,
482 redacted_auth: redacted,
483 };
484 if control_tx.try_send(ctrl).is_err() {
485 state.active_control.fetch_sub(1, Ordering::Relaxed);
486 if let Some(m) = metrics {
487 m.record_control_closed();
488 }
489 warn!(peer = %peer_addr, "control channel closed, cannot add to pool");
490 }
491
492 Ok(())
493 }
494
495 async fn accept_external_clients(
497 listener: TcpListener,
498 config: Arc<ReverseServerConfig>,
499 cancel: CancellationToken,
500 mut control_rx: mpsc::Receiver<ControlStream>,
501 metrics: Option<Arc<ReverseMetrics>>,
502 state: Arc<ReverseServerState>,
503 ) {
504 let mut relay_tasks = JoinSet::new();
505
506 loop {
507 tokio::select! {
508 result = listener.accept() => {
509 match result {
510 Ok((external_stream, peer_addr)) => {
511 match state.active_streams.fetch_update(
512 Ordering::AcqRel,
513 Ordering::Acquire,
514 |current| {
515 (current < config.max_streams_per_listener)
516 .then_some(current + 1)
517 },
518 ) {
519 Ok(_) => {}
520 Err(current) => {
521 warn!(
522 peer = %peer_addr,
523 active = current,
524 max = config.max_streams_per_listener,
525 "dropping external client: max_streams_per_listener reached"
526 );
527 state.dropped_stream_limit.fetch_add(1, Ordering::Relaxed);
528 drop(external_stream);
529 continue;
530 }
531 }
532
533 match state.pending_external.fetch_update(
534 Ordering::AcqRel,
535 Ordering::Acquire,
536 |current| {
537 (current < config.max_pending_external)
538 .then_some(current + 1)
539 },
540 ) {
541 Ok(_) => {}
542 Err(current) => {
543 state.active_streams.fetch_sub(1, Ordering::Release);
544 warn!(
545 peer = %peer_addr,
546 pending = current,
547 max = config.max_pending_external,
548 "dropping external client: max_pending_external reached"
549 );
550 state.dropped_pending_limit.fetch_add(1, Ordering::Relaxed);
551 drop(external_stream);
552 continue;
553 }
554 }
555 let control = tokio::select! {
557 control = control_rx.recv() => control,
558 _ = cancel.cancelled() => {
559 state.pending_external.fetch_sub(1, Ordering::Release);
560 state.active_streams.fetch_sub(1, Ordering::Release);
561 drop(external_stream);
562 break;
563 }
564 };
565 match control {
566 Some(control) => {
567 state.pending_external.fetch_sub(1, Ordering::Release);
568 let metrics = metrics.clone();
569 let state = state.clone();
570 let idle_timeout = (config.read_timeout_ms > 0).then(|| {
571 std::time::Duration::from_millis(config.read_timeout_ms)
572 });
573 state.active_control.fetch_sub(1, Ordering::Relaxed);
574 relay_tasks.spawn(async move {
575 info!(
576 peer = %peer_addr,
577 control_peer = %control.peer_addr,
578 "relaying external client through control connection"
579 );
580 if let Some(m) = metrics.as_deref() {
581 m.record_stream_opened();
582 m.record_state_duration(ControlState::Ready, 0);
583 }
584 let relay_result = relay_bidirectional_with_timeout(
585 external_stream,
586 control.stream,
587 idle_timeout,
588 )
589 .await;
590 match relay_result {
591 Ok(()) => {
592 debug!(peer = %peer_addr, "relay finished cleanly");
593 }
594 Err(e) => {
595 debug!(peer = %peer_addr, error = %e, "relay ended");
596 }
597 }
598 if let Some(m) = metrics.as_deref() {
599 m.record_stream_closed(0);
600 m.record_control_closed();
601 }
602 state.active_streams.fetch_sub(1, Ordering::Release);
603 debug!(peer = %peer_addr, "relay finished");
604 });
605 }
606 None => {
607 state.pending_external.fetch_sub(1, Ordering::Release);
608 state.active_streams.fetch_sub(1, Ordering::Release);
609 warn!(peer = %peer_addr, "no control connections available, rejecting external client");
610 drop(external_stream);
611 }
612 }
613 }
614 Err(e) => {
615 error!(error = %e, "failed to accept external client");
616 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
617 }
618 }
619 }
620 _ = cancel.cancelled() => {
621 break;
622 }
623 }
624 }
625
626 relay_tasks.abort_all();
627 while relay_tasks.join_next().await.is_some() {}
628 }
629
630 pub fn shutdown(&self) {
632 self.cancel.cancel();
633 }
634}
635
636pub struct ControlStream {
639 pub stream: TcpStream,
640 pub peer_addr: SocketAddr,
641 pub redacted_auth: Option<String>,
642}
643
644pub fn format_auth_redacted(auth: &str) -> String {
646 redact_auth(auth)
647}
648
649#[cfg(test)]
650mod tests {
651 use super::*;
652
653 #[test]
654 fn is_bind_allowed_with_none() {
655 let cfg = ReverseServerConfig {
656 allow_bind: None,
657 ..Default::default()
658 };
659 assert!(cfg.is_bind_allowed("127.0.0.1:8080".parse().unwrap()));
660 }
661
662 #[test]
663 fn is_bind_allowed_with_empty() {
664 let cfg = ReverseServerConfig {
665 allow_bind: Some(vec![]),
666 ..Default::default()
667 };
668 assert!(cfg.is_bind_allowed("127.0.0.1:8080".parse().unwrap()));
669 }
670
671 #[test]
672 fn is_bind_allowed_match() {
673 let cfg = ReverseServerConfig {
674 allow_bind: Some(vec!["127.0.0.1:8080".parse().unwrap()]),
675 ..Default::default()
676 };
677 assert!(cfg.is_bind_allowed("127.0.0.1:8080".parse().unwrap()));
678 }
679
680 #[test]
681 fn is_bind_allowed_mismatch() {
682 let cfg = ReverseServerConfig {
683 allow_bind: Some(vec!["127.0.0.1:8080".parse().unwrap()]),
684 ..Default::default()
685 };
686 assert!(!cfg.is_bind_allowed("0.0.0.0:8080".parse().unwrap()));
687 assert!(!cfg.is_bind_allowed("127.0.0.1:9090".parse().unwrap()));
688 }
689
690 #[test]
691 fn state_snapshot_round_trip() {
692 let s = ReverseServerState::default();
693 s.active_control.fetch_add(3, Ordering::Relaxed);
694 s.active_streams.fetch_add(2, Ordering::Relaxed);
695 s.pending_external.fetch_add(1, Ordering::Relaxed);
696 s.denied_bind.fetch_add(1, Ordering::Relaxed);
697 s.dropped_stream_limit.fetch_add(4, Ordering::Relaxed);
698 s.dropped_pending_limit.fetch_add(5, Ordering::Relaxed);
699 let snap = s.snapshot();
700 assert_eq!(snap.active_control, 3);
701 assert_eq!(snap.active_streams, 2);
702 assert_eq!(snap.pending_external, 1);
703 assert_eq!(snap.denied_bind, 1);
704 assert_eq!(snap.dropped_stream_limit, 4);
705 assert_eq!(snap.dropped_pending_limit, 5);
706 }
707
708 #[test]
709 fn format_auth_redacted_basic() {
710 assert_eq!(format_auth_redacted("user:pass"), "user:****");
711 }
712
713 #[test]
714 fn same_bind_v4() {
715 let a: SocketAddr = "127.0.0.1:8080".parse().unwrap();
716 let b: SocketAddr = "127.0.0.1:8080".parse().unwrap();
717 assert!(same_bind(&a, &b));
718 }
719
720 #[test]
721 fn same_bind_different_port() {
722 let a: SocketAddr = "127.0.0.1:8080".parse().unwrap();
723 let b: SocketAddr = "127.0.0.1:9090".parse().unwrap();
724 assert!(!same_bind(&a, &b));
725 }
726
727 #[test]
728 fn validate_loopback_ok() {
729 let cfg = ReverseServerConfig {
730 control_bind: "127.0.0.1:0".parse().unwrap(),
731 external_bind: Some("127.0.0.1:0".parse().unwrap()),
732 ..Default::default()
733 };
734 assert!(cfg.validate().is_ok());
735 }
736
737 #[test]
738 fn validate_no_external_bind_ok() {
739 let cfg = ReverseServerConfig {
740 control_bind: "127.0.0.1:0".parse().unwrap(),
741 external_bind: None,
742 ..Default::default()
743 };
744 assert!(cfg.validate().is_ok());
745 }
746
747 #[test]
748 fn validate_non_loopback_without_auth_rejected() {
749 let cfg = ReverseServerConfig {
750 control_bind: "127.0.0.1:0".parse().unwrap(),
751 external_bind: Some("0.0.0.0:9000".parse().unwrap()),
752 auth_username: None,
753 auth_password: None,
754 ..Default::default()
755 };
756 let err = cfg.validate().unwrap_err();
757 assert!(
758 matches!(err, ProtocolError::ConfigInvalid(_)),
759 "got: {err:?}"
760 );
761 }
762
763 #[test]
764 fn validate_non_loopback_with_auth_but_no_allowlist_rejected() {
765 let cfg = ReverseServerConfig {
766 control_bind: "127.0.0.1:0".parse().unwrap(),
767 external_bind: Some("0.0.0.0:9000".parse().unwrap()),
768 auth_username: Some("user".to_string()),
769 auth_password: Some("pass".to_string()),
770 allow_bind: None,
771 ..Default::default()
772 };
773 let err = cfg.validate().unwrap_err();
774 assert!(
775 matches!(err, ProtocolError::ConfigInvalid(_)),
776 "got: {err:?}"
777 );
778 }
779
780 #[test]
781 fn validate_non_loopback_with_auth_and_allowlist_ok() {
782 let cfg = ReverseServerConfig {
783 control_bind: "127.0.0.1:0".parse().unwrap(),
784 external_bind: Some("0.0.0.0:9000".parse().unwrap()),
785 auth_username: Some("user".to_string()),
786 auth_password: Some("pass".to_string()),
787 allow_bind: Some(vec!["0.0.0.0:9000".parse().unwrap()]),
788 ..Default::default()
789 };
790 assert!(cfg.validate().is_ok());
791 }
792
793 #[test]
794 fn validate_ipv6_loopback_ok() {
795 let cfg = ReverseServerConfig {
796 control_bind: "127.0.0.1:0".parse().unwrap(),
797 external_bind: Some("[::1]:9000".parse().unwrap()),
798 ..Default::default()
799 };
800 assert!(cfg.validate().is_ok());
801 }
802
803 #[test]
804 fn validate_ipv6_non_loopback_without_auth_rejected() {
805 let cfg = ReverseServerConfig {
806 control_bind: "127.0.0.1:0".parse().unwrap(),
807 external_bind: Some("[2001:db8::1]:9000".parse().unwrap()),
808 ..Default::default()
809 };
810 let err = cfg.validate().unwrap_err();
811 assert!(
812 matches!(err, ProtocolError::ConfigInvalid(_)),
813 "got: {err:?}"
814 );
815 }
816}