use std::net;
use crate::{
conn_expr::{Addr, ConnectionExpr, Port, PortSet, RouteExpr},
error::Error,
host_options::HostOptionsTable,
};
pub fn resolve_routes(
conn_expr: &ConnectionExpr,
host_opts_table: &HostOptionsTable,
) -> Result<Routes, Error> {
use ConnectionExpr::*;
let route_expr = match conn_expr {
RouteExpr(ref e) => e,
HostKey(ref k) => host_opts_table
.get(k)
.ok_or_else(|| Error::HostKeyNotFound(k.to_string()))?
.conn_expr
.try_as_route_expr()
.ok_or_else(|| Error::RecursiveHostKeysNotSupported(k.to_string()))?,
};
Ok(Routes {
inner: RoutesInner::try_from_route_expr(route_expr)?,
pos: 0,
})
}
#[derive(Clone, Debug)]
pub enum Route {
Direct(net::SocketAddr),
Tunneled(TunnelOptions),
}
#[derive(Clone, Debug)]
pub struct TunnelOptions {
pub ssh_user: Option<String>,
pub ssh_addr: Addr,
pub ssh_port: Option<Port>,
pub host_addr: Addr,
pub host_port: Port,
}
#[derive(Clone, Debug)]
pub struct Routes {
inner: RoutesInner,
pos: usize,
}
impl Iterator for Routes {
type Item = Route;
fn next(&mut self) -> Option<Self::Item> {
if self.pos < self.inner.len() {
let item = self.inner.produce(self.pos);
self.pos += 1;
Some(item)
} else {
None
}
}
}
#[derive(Clone, Debug)]
enum RoutesInner {
Direct {
ips: Vec<net::IpAddr>,
ports: PortSet,
},
Tunneled {
ssh_user: Option<String>,
ssh_addr: Addr,
ssh_ports: Option<PortSet>,
host_addr: Addr,
host_ports: PortSet,
},
}
impl RoutesInner {
fn try_from_route_expr(route_expr: &RouteExpr) -> Result<Self, Error> {
if let Some(ref tunnel) = route_expr.tunnel {
let host_addr = route_expr
.addr
.as_ref()
.expect("tunneling should guarantee final host address")
.clone();
Ok(RoutesInner::Tunneled {
ssh_user: tunnel.user.clone(),
ssh_addr: tunnel.addr.clone(),
ssh_ports: tunnel.ports.clone(),
host_addr,
host_ports: route_expr.ports.clone(),
})
} else {
let mut ips = match route_expr.addr {
None => dns_lookup::lookup_host("localhost")
.map_err(|_| Error::DomainNotFound("localhost".to_owned()))?,
Some(Addr::Domain(ref domain)) => dns_lookup::lookup_host(domain)
.map_err(|_| Error::DomainNotFound(domain.clone()))?,
Some(Addr::IP(ip)) => vec![ip],
};
ips.sort();
Ok(RoutesInner::Direct {
ips,
ports: route_expr.ports.clone(),
})
}
}
fn len(&self) -> usize {
match self {
RoutesInner::Direct { ips: addrs, ports } => {
addrs.len() * ports.as_slice().len()
}
RoutesInner::Tunneled {
ssh_ports: None,
host_ports,
..
} => host_ports.as_slice().len(),
RoutesInner::Tunneled {
ssh_ports: Some(ssh_ports),
host_ports,
..
} => ssh_ports.as_slice().len() * host_ports.as_slice().len(),
}
}
fn produce(&self, ix: usize) -> Route {
assert!(ix < self.len());
match self {
RoutesInner::Direct { ips: addrs, ports } => {
let ix_addr = ix % addrs.len();
let ix_port = ix / addrs.len();
Route::Direct(net::SocketAddr::new(
addrs[ix_addr],
ports.as_slice()[ix_port],
))
}
RoutesInner::Tunneled {
ssh_user,
ssh_addr,
ssh_ports: None,
host_addr,
host_ports,
} => Route::Tunneled(TunnelOptions {
ssh_user: ssh_user.clone(),
ssh_addr: ssh_addr.clone(),
ssh_port: None,
host_addr: host_addr.clone(),
host_port: host_ports.as_slice()[ix],
}),
RoutesInner::Tunneled {
ssh_user,
ssh_addr,
ssh_ports: Some(ssh_ports),
host_addr,
host_ports,
} => {
let ssh_ports = ssh_ports.as_slice();
let ix_ssh_port = ix % ssh_ports.len();
let ix_host_port = ix / ssh_ports.len();
Route::Tunneled(TunnelOptions {
ssh_user: ssh_user.clone(),
ssh_addr: ssh_addr.clone(),
ssh_port: Some(ssh_ports[ix_ssh_port]),
host_addr: host_addr.clone(),
host_port: host_ports.as_slice()[ix_host_port],
})
}
}
}
}