use std::net::SocketAddr;
use std::sync::Arc;
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::Semaphore;
use tracing::{debug, error, info, instrument, warn};
use crate::errors::{Error, Result};
use crate::mux::{SmuxConfig, SmuxSession};
use crate::protocol::SessionType;
use crate::session::Session;
#[derive(Debug, Clone)]
pub struct PortForwardConfig {
pub local_addr: SocketAddr,
pub max_connections: usize,
}
impl Default for PortForwardConfig {
fn default() -> Self {
Self {
local_addr: "127.0.0.1:0"
.parse()
.expect("127.0.0.1:0 is a valid socket address"),
max_connections: 100,
}
}
}
#[must_use = "dropping a PortForwarder stops port forwarding before any connections are accepted; call forward() to start forwarding"]
pub struct PortForwarder {
config: PortForwardConfig,
listener: TcpListener,
local_addr: SocketAddr,
}
impl PortForwarder {
#[instrument(skip(config), fields(local_addr = %config.local_addr))]
pub async fn bind(config: PortForwardConfig) -> Result<Self> {
let listener = TcpListener::bind(config.local_addr)
.await
.map_err(Error::Io)?;
let local_addr = listener.local_addr().map_err(Error::Io)?;
info!(local_addr = %local_addr, "Port forwarding listener bound");
Ok(Self {
config,
listener,
local_addr,
})
}
pub fn local_addr(&self) -> SocketAddr {
self.local_addr
}
#[instrument(skip(self, session, shutdown), fields(local_addr = %self.local_addr))]
pub async fn forward(
self,
session: Arc<Session>,
shutdown: crate::shutdown::ShutdownSignal,
) -> Result<()> {
if session.config().session_type != SessionType::Port {
return Err(Error::Config(format!(
"PortForwarder requires a port-forwarding session (SessionType::Port), \
got {:?}; use SessionBuilder::document(PortForwardingSession::new(port))",
session.config().session_type,
)));
}
let Self {
config, listener, ..
} = self;
let connect_timeout = session.config().connect_timeout;
let ready = tokio::select! {
ready = session.wait_for_ready(connect_timeout) => ready,
_ = shutdown.cancelled() => {
info!("Shutting down before session handshake completed");
return Ok(());
}
_ = session.wait_terminated() => {
info!("SSM session terminated before becoming ready");
return Ok(());
}
};
if !ready {
return Err(Error::InvalidState(
"SSM session not ready: handshake timed out".to_string(),
));
}
let mux = Arc::new(SmuxSession::new(
Arc::clone(&session),
SmuxConfig::default(),
));
let semaphore = Arc::new(Semaphore::new(config.max_connections));
let max_connections = config.max_connections;
let remote_port = session
.config()
.parameters
.get("portNumber")
.and_then(|v| v.first())
.and_then(|s| s.parse::<u16>().ok());
info!(
max_connections,
remote_port, "Accepting port forwarding connections"
);
loop {
tokio::select! {
biased;
_ = shutdown.cancelled() => {
info!("Port forwarder shutdown requested");
return Ok(());
}
_ = session.wait_terminated() => {
info!("SSM session terminated, stopping port forwarder");
shutdown.shutdown();
return Ok(());
}
result = listener.accept() => {
match result {
Ok((stream, peer_addr)) => {
debug!(peer_addr = %peer_addr, "Accepted connection");
let permit = match Arc::clone(&semaphore).try_acquire_owned() {
Ok(p) => p,
Err(_) => {
warn!(
max_connections,
"Connection limit reached, rejecting connection"
);
drop(stream);
continue;
}
};
let mux = Arc::clone(&mux);
let handler_shutdown = shutdown.clone();
tokio::spawn(async move {
let _permit = permit; if let Err(e) = Self::handle_connection(stream, mux, handler_shutdown).await {
error!(peer_addr = %peer_addr, error = ?e, "Connection handler error");
}
});
}
Err(e) => {
error!(error = ?e, "Failed to accept connection");
return Err(Error::Io(e));
}
}
}
}
}
}
#[instrument(skip(stream, mux, shutdown))]
async fn handle_connection(
stream: TcpStream,
mux: Arc<SmuxSession>,
shutdown: crate::shutdown::ShutdownSignal,
) -> Result<()> {
debug!("Starting connection handler");
let smux_stream = mux.open_stream()?;
let (mut tcp_rx, mut tcp_tx) = stream.into_split();
let (mut smux_rx, mut smux_tx) = tokio::io::split(smux_stream);
let local_to_remote = tokio::io::copy(&mut tcp_rx, &mut smux_tx);
let remote_to_local = tokio::io::copy(&mut smux_rx, &mut tcp_tx);
tokio::pin!(local_to_remote, remote_to_local);
tokio::select! {
biased;
_ = shutdown.cancelled() => {
debug!("Connection handler cancelled due to shutdown");
}
res = &mut local_to_remote => {
match res {
Ok(n) => info!(bytes = n, "Local\u{2192}remote copy completed"),
Err(e) => return Err(Error::Io(e)),
}
}
res = &mut remote_to_local => {
match res {
Ok(n) => info!(bytes = n, "Remote\u{2192}local copy completed"),
Err(e) => return Err(Error::Io(e)),
}
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_port_forward_config_default() {
let config = PortForwardConfig::default();
assert_eq!(config.max_connections, 100);
assert_eq!(config.local_addr.ip().to_string(), "127.0.0.1");
assert_eq!(config.local_addr.port(), 0);
}
#[test]
fn test_port_forward_config_custom() {
let config = PortForwardConfig {
local_addr: "127.0.0.1:8080".parse().unwrap(),
max_connections: 5,
};
assert_eq!(
config.local_addr,
"127.0.0.1:8080".parse::<std::net::SocketAddr>().unwrap()
);
assert_eq!(config.max_connections, 5);
}
#[tokio::test]
async fn test_bind_succeeds_with_default_config() {
let forwarder = PortForwarder::bind(PortForwardConfig::default())
.await
.expect("bind must succeed on loopback");
assert_ne!(
forwarder.local_addr().port(),
0,
"OS must assign a non-zero port"
);
}
#[tokio::test]
async fn test_bind_returns_correct_local_addr() {
let config = PortForwardConfig {
local_addr: "127.0.0.1:0".parse().unwrap(),
max_connections: 1,
};
let forwarder = PortForwarder::bind(config)
.await
.expect("bind must succeed on loopback");
assert_eq!(forwarder.local_addr().ip().to_string(), "127.0.0.1");
assert_ne!(forwarder.local_addr().port(), 0);
}
}