diff --git a/src/cli.rs b/src/cli.rs new file mode 100644 index 0000000..592eca4 --- /dev/null +++ b/src/cli.rs @@ -0,0 +1,55 @@ +use std::net::{IpAddr, Ipv6Addr}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct ListenCfg { + pub ip: Option, + pub port: u16, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct RelayCfg { + pub ip: Option, + pub min_port: u16, + pub max_port: u16, +} + +fn parse_ip(val: &str) -> Result<(Option, &str), &'static str> { + let val = val.trim(); + match val.as_bytes().first() { + Some(b'[') => { + let closing = val.find(']').ok_or("not a valid ipv6")?; + let v6 = val[1..closing] + .parse::() + .map_err(|_| "not a valid ipv6")?; + if val.get(closing + 1..closing + 2) == Some(":") { + Ok((Some(IpAddr::V6(v6)), &val[closing + 2..])) + } else { + Err("no port specified") + } + } + Some(b':') => Ok((None, &val[1..])), + _ => { + let (ip, port) = val + .split_once(':') + .ok_or("format must be ip:port or :port")?; + let ip = ip.parse::().map_err(|_| "not a valid ip")?; + Ok((Some(ip), port)) + } + } +} + +pub fn parse_listen(val: &str) -> Result { + let (ip, port) = parse_ip(val)?; + let port = port.parse().map_err(|_| "cannot parse port")?; + Ok(ListenCfg { ip, port }) +} + +pub fn parse_range(val: &str) -> Result { + let (ip, range) = parse_ip(val)?; + let (min, max) = range.split_once('-').ok_or("cannot parse as range")?; + Ok(RelayCfg { + ip, + min_port: min.parse().map_err(|_| "cannot parse min port")?, + max_port: max.parse().map_err(|_| "cannot parse max port")?, + }) +} diff --git a/src/main.rs b/src/main.rs index 9202633..df36b23 100644 --- a/src/main.rs +++ b/src/main.rs @@ -5,17 +5,20 @@ use std::sync::Arc; use std::time::Duration; use anyhow::{Context as _, Result}; +use clap::builder::ValueParser; use clap::{App, AppSettings, Arg}; use tokio::io::AsyncWriteExt; use tokio::net::{UdpSocket, UnixListener}; use turn::Error; use turn::auth::generate_long_term_credentials; use turn::auth::*; -use turn::relay::relay_static::RelayAddressGeneratorStatic; +use turn::relay::relay_range::RelayAddressGeneratorRanges; use turn::server::Server; use turn::server::config::{ConnConfig, ServerConfig}; use webrtc_util::vnet::net::Net; +mod cli; + fn listen_ips() -> BTreeSet { let mut ip_set = BTreeSet::new(); let interfaces = netdev::interface::get_interfaces(); @@ -41,6 +44,31 @@ fn is_link_local(ip: IpAddr) -> bool { } } +async fn create_conn_config( + listen_ip: IpAddr, + conn: Option>, + listen: &cli::ListenCfg, + relay: &cli::RelayCfg, +) -> Result { + println!("Listening on public IP: {listen_ip}"); + let conn = match conn { + Some(conn) => conn, // listener socket with user-specified host already created + None => Arc::new(UdpSocket::bind((listen_ip, listen.port)).await?), + }; + let relay_ip = relay.ip.unwrap_or(listen_ip); + Ok(ConnConfig { + conn, + relay_addr_generator: Box::new(RelayAddressGeneratorRanges { + relay_address: relay_ip, + address: relay_ip.to_string(), + min_port: relay.min_port, + max_port: relay.max_port, + max_retries: 0, // use the default + net: Arc::new(Net::new(None)), + }), + }) +} + /// Listens on the Unix socket, /// returning valid credentials to any connecting client. async fn socket_loop(path: &Path, shared_secret: &str) -> Result<()> { @@ -75,23 +103,39 @@ async fn main() -> Result<(), Error> { .setting(AppSettings::DeriveDisplayOrder) .setting(AppSettings::SubcommandsNegateReqs) .arg( - Arg::with_name("FULLHELP") + Arg::new("FULLHELP") .help("Prints more detailed help information") .long("fullhelp"), ) .arg( - Arg::with_name("realm") + Arg::new("realm") .default_value("webrtc.rs") .takes_value(true) .long("realm") .help("Realm (defaults to \"webrtc.rs\")"), ) .arg( - Arg::with_name("socket") + Arg::new("socket") .required(true) .takes_value(true) .long("socket") .help("Unix socket path"), + ) + .arg( + Arg::new("listen") + .default_value(":3478") + .takes_value(true) + .value_parser(ValueParser::new(cli::parse_listen)) + .long("listen") + .help("Address to bind TURN listener to: [ip]:"), + ) + .arg( + Arg::new("relayaddr") + .default_value(":49152-65535") + .takes_value(true) + .value_parser(ValueParser::new(cli::parse_range)) + .long("relay-addr") + .help("Host and port range available for TURN relay: [ip]:-"), ); let matches = app.clone().get_matches(); @@ -101,24 +145,27 @@ async fn main() -> Result<(), Error> { std::process::exit(0); } - let port = 3478; let realm = matches.value_of("realm").unwrap(); let socket_path = Path::new(matches.value_of("socket").unwrap()); - - let mut conn_configs = Vec::new(); - for listen_ip in listen_ips() { - println!("Listening on {listen_ip}"); - let conn = Arc::new(UdpSocket::bind((listen_ip, port)).await?); - let conn_config = ConnConfig { - conn, - relay_addr_generator: Box::new(RelayAddressGeneratorStatic { - relay_address: listen_ip, - address: listen_ip.to_string(), - net: Arc::new(Net::new(None)), - }), - }; - conn_configs.push(conn_config); - } + let listen: cli::ListenCfg = *matches.get_one("listen").unwrap(); + let relay: cli::RelayCfg = *matches.get_one("relayaddr").unwrap(); + + let conn = match listen.ip { + Some(ip) => Some(Arc::new(UdpSocket::bind((ip, listen.port)).await?)), + _ => None, + }; + + let conn_configs = if conn.is_some() && relay.ip.is_some() { + // do not iterate over available IPs + // when both hosts are explicitly specified + vec![create_conn_config(listen.ip.unwrap(), conn, &listen, &relay).await?] + } else { + let mut conn_configs = Vec::new(); + for listen_ip in listen_ips() { + conn_configs.push(create_conn_config(listen_ip, conn.clone(), &listen, &relay).await?); + } + conn_configs + }; let shared_secret = "north"; let auth_handler = LongTermAuthHandler::new(shared_secret.to_string());