use crate::metrics::ReverseMetrics;
use crate::{client_auth_handshake, 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(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,
pub tls: Option<crate::tls::ReverseClientTlsConfig>,
}
impl std::fmt::Debug for ReverseClientConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ReverseClientConfig")
.field("server_addr", &self.server_addr)
.field("auth_username", &self.auth_username)
.field(
"auth_password",
&self.auth_password.as_deref().map(|_| "****"),
)
.field("reconnect_initial_ms", &self.reconnect_initial_ms)
.field("reconnect_max_ms", &self.reconnect_max_ms)
.field("default_target_host", &self.default_target_host)
.field("default_target_port", &self.default_target_port)
.field("read_timeout_ms", &self.read_timeout_ms)
.field("drain_grace_ms", &self.drain_grace_ms)
.field("target_connect_timeout_ms", &self.target_connect_timeout_ms)
.field("tls", &self.tls)
.finish()
}
}
impl ReverseClientConfig {
pub fn validate(&self) -> Result<(), ProtocolError> {
if let Some(ref tls) = self.tls {
tls.validate()?;
}
Ok(())
}
}
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,
tls: None,
}
}
}
#[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> {
self.config.validate()?;
let tls_client_config: Option<(Arc<rustls::ClientConfig>, String)> =
match self.config.tls.as_ref() {
Some(tls) => {
let cfg = tls.build_client_config()?;
Some((cfg, tls.server_name.clone()))
}
None => None,
};
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(tls_client_config.as_ref()).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,
tls: Option<&(Arc<rustls::ClientConfig>, String)>,
) -> Result<(), ProtocolError> {
let connecting_start = Instant::now();
let tcp = tokio::select! {
result = TcpStream::connect(&self.config.server_addr) => {
result?
}
_ = self.cancel.cancelled() => {
return Err(ProtocolError::ConnectionClosed);
}
};
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 mut boxed: eggress_core::BoxStream = if let Some((cfg, server_name)) = tls {
let tcp_boxed: eggress_core::BoxStream = Box::new(tcp);
let handshake = eggress_transport_tls::tls_connect(tcp_boxed, cfg.clone(), server_name);
tokio::select! {
result = handshake => {
result.map_err(|e| {
let msg = format!("reverse control TLS handshake failed: {e}");
if let Some(m) = self.metrics.as_ref() {
m.record_error(&msg);
}
ProtocolError::Tls(msg)
})?
}
_ = self.cancel.cancelled() => {
return Err(ProtocolError::ConnectionClosed);
}
}
} else {
Box::new(tcp)
};
let authenticating_start = Instant::now();
if let (Some(ref username), Some(ref password)) =
(&self.config.auth_username, &self.config.auth_password)
{
client_auth_handshake(&mut boxed, 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"
);
} else {
crate::read_handshake(&mut boxed).await?;
if let Some(ref m) = self.metrics {
m.record_state_duration(
ControlState::Authenticating,
authenticating_start.elapsed().as_millis() as u64,
);
}
};
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"
);
let target_boxed: eggress_core::BoxStream = Box::new(target_stream);
crate::relay_bidirectional_boxed(
boxed,
target_boxed,
(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:?}"),
}
}
}