pub mod control;
use std::collections::HashMap;
use std::path::Path;
use toride_ssh_core::SshPaths;
use toride_ssh_core::{Error, Result};
pub use control::{ControlSession, ForwardType, PortForward};
pub struct ForwardService<'a> {
paths: &'a SshPaths,
}
impl<'a> ForwardService<'a> {
#[must_use]
pub fn new(paths: &'a SshPaths) -> Self {
Self { paths }
}
pub async fn list(&self) -> Result<Vec<(ControlSession, Vec<PortForward>)>> {
let sessions = self.list_sessions().await?;
let listings =
control::list_forwards_bounded(sessions.iter().map(|s| s.control_path.clone())).await;
Ok(sessions
.into_iter()
.zip(listings)
.map(|(session, result)| match result {
Ok(forwards) => (session, forwards),
Err(e) => {
tracing::warn!(
"failed to list forwards for {}: {e}",
session.control_path.display()
);
(session, Vec::new())
}
})
.collect())
}
pub async fn list_sessions(&self) -> Result<Vec<ControlSession>> {
control::list_sessions(self.paths.ssh_dir()).await
}
pub async fn cancel(&self, control_path: &Path, local_port: u16) -> Result<()> {
control::cancel_forward(control_path, local_port).await
}
pub async fn list_forwards(&self, control_path: &Path) -> Result<Vec<PortForward>> {
control::list_forwards(control_path).await
}
pub async fn cancel_known(&self, control_path: &Path, forward: &PortForward) -> Result<()> {
control::cancel_known_forward(control_path, forward).await
}
pub async fn exit_session(&self, control_path: &Path) -> Result<()> {
control::exit_session(control_path).await
}
pub async fn conflicting_local_ports(&self) -> Result<HashMap<u16, Vec<std::path::PathBuf>>> {
let sessions = self.list_sessions().await?;
let listings =
control::list_forwards_bounded(sessions.iter().map(|s| s.control_path.clone())).await;
let mut port_owners: HashMap<u16, Vec<std::path::PathBuf>> = HashMap::new();
for (session, result) in sessions.iter().zip(listings) {
match result {
Ok(forwards) => {
for fwd in forwards {
port_owners
.entry(fwd.local_port)
.or_default()
.push(session.control_path.clone());
}
}
Err(e) => {
tracing::warn!(
"failed to list forwards for {}: {e}",
session.control_path.display()
);
}
}
}
port_owners.retain(|_, owners| owners.len() > 1);
Ok(port_owners)
}
const TEST_CONNECTIVITY_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(2);
pub async fn test_connectivity_with_timeout(
&self,
local_port: u16,
timeout: std::time::Duration,
) -> Result<()> {
let addr = format!("127.0.0.1:{local_port}");
tokio::time::timeout(timeout, tokio::net::TcpStream::connect(&addr))
.await
.map_err(|_| {
Error::ForwardFailed(format!(
"connection to {addr} timed out after {} seconds",
timeout.as_secs()
))
})?
.map_err(|e| Error::ForwardFailed(format!("cannot connect to {addr}: {e}")))?;
tracing::debug!("successfully connected to forwarded port {local_port}");
Ok(())
}
pub async fn test_connectivity(&self, local_port: u16) -> Result<()> {
self.test_connectivity_with_timeout(local_port, Self::TEST_CONNECTIVITY_TIMEOUT)
.await
}
}