use crate::metrics::ReverseMetrics;
use crate::{client_auth_handshake, relay_bidirectional_with_timeout, ControlState, ProtocolError};
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::net::TcpStream;
use tokio_util::sync::CancellationToken;
use tracing::{debug, info, warn};
#[derive(Debug, Clone)]
pub struct ReverseClientConfig {
pub server_addr: SocketAddr,
pub auth_username: Option<String>,
pub auth_password: Option<String>,
pub reconnect_initial_ms: u64,
pub reconnect_max_ms: u64,
pub default_target_host: Option<String>,
pub default_target_port: Option<u16>,
pub read_timeout_ms: u64,
pub drain_grace_ms: u64,
pub target_connect_timeout_ms: u64,
}
impl Default for ReverseClientConfig {
fn default() -> Self {
Self {
server_addr: "127.0.0.1:0".parse().unwrap(),
auth_username: None,
auth_password: None,
reconnect_initial_ms: 1_000,
reconnect_max_ms: 30_000,
default_target_host: None,
default_target_port: None,
read_timeout_ms: 60_000,
drain_grace_ms: 5_000,
target_connect_timeout_ms: 10_000,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum TargetResolution {
Connect { host: String, port: u16 },
Reject { reason: String },
}
pub trait TargetResolver: Send + Sync {
fn resolve(&self) -> TargetResolution;
}
pub struct DefaultTargetResolver {
pub host: Option<String>,
pub port: Option<u16>,
}
impl DefaultTargetResolver {
pub fn new(host: Option<String>, port: Option<u16>) -> Self {
Self { host, port }
}
}
impl TargetResolver for DefaultTargetResolver {
fn resolve(&self) -> TargetResolution {
match (&self.host, self.port) {
(Some(h), Some(p)) => TargetResolution::Connect {
host: h.clone(),
port: p,
},
_ => TargetResolution::Reject {
reason: "no default target configured".to_string(),
},
}
}
}
pub struct ReverseClient {
config: ReverseClientConfig,
cancel: CancellationToken,
metrics: Option<Arc<ReverseMetrics>>,
resolver: Option<Arc<dyn TargetResolver>>,
}
impl ReverseClient {
pub fn new(config: ReverseClientConfig) -> Self {
let resolver: Arc<dyn TargetResolver> = Arc::new(DefaultTargetResolver::new(
config.default_target_host.clone(),
config.default_target_port,
));
Self {
config,
cancel: CancellationToken::new(),
metrics: None,
resolver: Some(resolver),
}
}
pub fn set_metrics(&mut self, metrics: Arc<ReverseMetrics>) {
self.metrics = Some(metrics);
}
pub fn set_resolver(&mut self, resolver: Arc<dyn TargetResolver>) {
self.resolver = Some(resolver);
}
pub fn cancel_token(&self) -> CancellationToken {
self.cancel.clone()
}
pub async fn run(&self) -> Result<(), ProtocolError> {
let mut backoff_ms = self.config.reconnect_initial_ms;
loop {
if self.cancel.is_cancelled() {
break;
}
let session_start = Instant::now();
match self.run_session().await {
Ok(()) => {
if let Some(ref m) = self.metrics {
m.record_state_duration(
ControlState::Ready,
session_start.elapsed().as_millis() as u64,
);
}
if self.cancel.is_cancelled() {
break;
}
backoff_ms = self.config.reconnect_initial_ms;
debug!("session ended, reconnecting immediately");
}
Err(e) => {
if self.cancel.is_cancelled() {
break;
}
if let Some(ref m) = self.metrics {
m.record_reconnect();
m.record_state_duration(
ControlState::Connecting,
session_start.elapsed().as_millis() as u64,
);
}
warn!(error = %e, backoff_ms, "session failed, reconnecting");
let sleep = tokio::time::sleep(Duration::from_millis(backoff_ms));
tokio::select! {
_ = sleep => {}
_ = self.cancel.cancelled() => break,
}
backoff_ms = (backoff_ms * 2).min(self.config.reconnect_max_ms);
}
}
}
let drain_start = Instant::now();
tokio::time::sleep(Duration::from_millis(50)).await;
if let Some(ref m) = self.metrics {
m.record_drain(drain_start.elapsed().as_millis() as u64);
}
info!("reverse client shut down");
Ok(())
}
async fn run_session(&self) -> Result<(), ProtocolError> {
let connecting_start = Instant::now();
let stream = TcpStream::connect(&self.config.server_addr).await?;
if let Some(ref m) = self.metrics {
m.record_state_duration(
ControlState::Connecting,
connecting_start.elapsed().as_millis() as u64,
);
}
info!(
server = %self.config.server_addr,
state = ?ControlState::Connecting,
"connected to reverse server"
);
let authenticating_start = Instant::now();
let stream = if let (Some(ref username), Some(ref password)) =
(&self.config.auth_username, &self.config.auth_password)
{
let mut s = stream;
client_auth_handshake(&mut s, username, password).await?;
if let Some(ref m) = self.metrics {
m.record_state_duration(
ControlState::Authenticating,
authenticating_start.elapsed().as_millis() as u64,
);
}
info!(
state = ?ControlState::Authenticating,
"authentication successful"
);
s
} else {
let mut s = stream;
crate::read_handshake(&mut s).await?;
if let Some(ref m) = self.metrics {
m.record_state_duration(
ControlState::Authenticating,
authenticating_start.elapsed().as_millis() as u64,
);
}
s
};
if let Some(ref m) = self.metrics {
m.record_stream_opened();
}
let ready_start = Instant::now();
let resolution =
self.resolver
.as_ref()
.map(|r| r.resolve())
.unwrap_or(TargetResolution::Reject {
reason: "no resolver configured".to_string(),
});
let session_result: Result<(), ProtocolError> = match resolution {
TargetResolution::Connect { host, port } => {
let target_addr = format!("{}:{}", host, port);
let connect_timeout = if self.config.target_connect_timeout_ms > 0 {
Duration::from_millis(self.config.target_connect_timeout_ms)
} else {
Duration::from_secs(30)
};
let connect_result =
tokio::time::timeout(connect_timeout, TcpStream::connect(&target_addr)).await;
match connect_result {
Ok(Ok(target_stream)) => {
info!(
target = %target_addr,
state = ?ControlState::Ready,
"connected to target, relaying"
);
relay_bidirectional_with_timeout(
stream,
target_stream,
(self.config.read_timeout_ms > 0)
.then(|| Duration::from_millis(self.config.read_timeout_ms)),
)
.await
}
Ok(Err(e)) => {
warn!(
target = %target_addr,
error = %e,
"failed to connect to target"
);
if let Some(ref m) = self.metrics {
m.record_error(&format!("target connect failed: {e}"));
}
Err(ProtocolError::Io(e))
}
Err(_elapsed) => {
let e = std::io::Error::new(
std::io::ErrorKind::TimedOut,
format!("target connect timed out: {target_addr}"),
);
warn!(
target = %target_addr,
"target connect timed out"
);
if let Some(ref m) = self.metrics {
m.record_error(&format!("target connect failed: {e}"));
}
Err(ProtocolError::Io(e))
}
}
}
TargetResolution::Reject { reason } => {
warn!(reason = %reason, "route resolution rejected, dropping control channel");
if let Some(ref m) = self.metrics {
m.record_error(&format!("route rejected: {reason}"));
}
Err(ProtocolError::AuthFailed)
}
};
if let Some(ref m) = self.metrics {
m.record_state_duration(
ControlState::Ready,
ready_start.elapsed().as_millis() as u64,
);
m.record_stream_closed(0);
}
session_result
}
pub fn shutdown(&self) {
self.cancel.cancel();
}
}
#[cfg(test)]
mod tests {
use super::*;
struct FixedResolver(TargetResolution);
impl TargetResolver for FixedResolver {
fn resolve(&self) -> TargetResolution {
self.0.clone()
}
}
#[test]
fn default_resolver_returns_configured_target() {
let r = DefaultTargetResolver::new(Some("127.0.0.1".to_string()), Some(8080));
assert_eq!(
r.resolve(),
TargetResolution::Connect {
host: "127.0.0.1".to_string(),
port: 8080,
}
);
}
#[test]
fn default_resolver_rejects_when_unset() {
let r = DefaultTargetResolver::new(None, None);
match r.resolve() {
TargetResolution::Reject { .. } => {}
other => panic!("expected Reject, got {other:?}"),
}
}
#[test]
fn default_resolver_rejects_partial() {
let r = DefaultTargetResolver::new(Some("127.0.0.1".to_string()), None);
match r.resolve() {
TargetResolution::Reject { .. } => {}
other => panic!("expected Reject, got {other:?}"),
}
}
#[test]
fn custom_resolver_can_reject() {
let r: Arc<dyn TargetResolver> = Arc::new(FixedResolver(TargetResolution::Reject {
reason: "policy".to_string(),
}));
match r.resolve() {
TargetResolution::Reject { reason } => assert_eq!(reason, "policy"),
other => panic!("expected Reject, got {other:?}"),
}
}
}