use anyhow::{Context, Result};
use std::collections::HashSet;
use std::sync::Arc;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpStream;
use tokio::sync::RwLock;
use tracing::{debug, error, info, warn};
use wispers_connect::P2pError;
use wispers_connect::{
IncomingConnections, NodeState, QuicConnection, ServingHandle, ServingSession, UdpConnection,
};
use crate::ipc;
#[derive(Clone, Debug)]
pub enum AllowedPorts {
All,
Whitelist(HashSet<u16>),
}
impl AllowedPorts {
pub fn parse(value: &str) -> Result<Self> {
if value.is_empty() {
return Ok(AllowedPorts::All);
}
let mut ports = HashSet::new();
for part in value.split(',') {
let part = part.trim();
if part.is_empty() {
continue;
}
let port: u16 = part
.parse()
.with_context(|| format!("invalid port number: {}", part))?;
ports.insert(port);
}
if ports.is_empty() {
Ok(AllowedPorts::All)
} else {
Ok(AllowedPorts::Whitelist(ports))
}
}
pub fn is_allowed(&self, port: u16) -> bool {
match self {
AllowedPorts::All => true,
AllowedPorts::Whitelist(ports) => ports.contains(&port),
}
}
}
pub async fn serve(
hub_override: Option<&str>,
profile: &str,
allowed_ports: Option<AllowedPorts>,
allow_egress: bool,
) -> Result<()> {
let storage = super::get_storage(hub_override, profile)?;
let node = super::load_node(&storage).await?;
if node.state() == NodeState::Pending {
anyhow::bail!("Not registered. Use 'wconnect register <token>' first.");
}
let cg_id = node.connectivity_group_id().unwrap().to_string();
let node_number = node.node_number().unwrap();
let ipc_server = ipc::Server::bind(&cg_id, node_number)
.await
.context("failed to start IPC server")?;
let allowed_ports = Arc::new(allowed_ports);
println!(
"Serving node {} in group {} (socket: {:?})",
node_number,
cg_id,
ipc_server.path()
);
if let Some(ref ports) = *allowed_ports {
match ports {
AllowedPorts::All => println!(" Port forwarding: all ports allowed"),
AllowedPorts::Whitelist(set) => {
let mut ports: Vec<_> = set.iter().collect();
ports.sort();
println!(" Port forwarding: allowed ports {:?}", ports);
}
}
} else {
println!(" Port forwarding: disabled");
}
if allow_egress {
println!(" Internet egress: enabled");
} else {
println!(" Internet egress: disabled");
}
let handle_state: Arc<RwLock<Option<ServingHandle>>> = Arc::new(RwLock::new(None));
let connect_handle_state = handle_state.clone();
let mut connect_task = tokio::spawn(async move {
let result: Result<(ServingHandle, ServingSession, IncomingConnections), anyhow::Error> =
node.start_serving()
.await
.context("failed to start serving");
if let Ok((handle, _session, _)) = &result {
*connect_handle_state.write().await = Some(handle.clone());
}
result
});
let mut session_task: Option<
tokio::task::JoinHandle<Result<(), wispers_connect::ServingError>>,
> = None;
let mut incoming_udp_rx: Option<tokio::sync::mpsc::Receiver<Result<UdpConnection, P2pError>>> =
None;
let mut incoming_quic_rx: Option<
tokio::sync::mpsc::Receiver<Result<QuicConnection, P2pError>>,
> = None;
loop {
tokio::select! {
result = &mut connect_task, if session_task.is_none() => {
match result {
Ok(Ok((handle, session, incoming))) => {
info!("Connected to hub");
*handle_state.write().await = Some(handle);
session_task = Some(tokio::spawn(async move { session.run().await }));
incoming_udp_rx = Some(incoming.udp);
incoming_quic_rx = Some(incoming.quic);
}
Ok(Err(e)) => {
return Err(e);
}
Err(e) => {
return Err(anyhow::anyhow!("Connect task panicked: {}", e));
}
}
}
result = async { session_task.as_mut().unwrap().await }, if session_task.is_some() => {
match result {
Ok(Ok(())) => {
info!("Session ended normally");
break;
}
Ok(Err(e)) => {
return Err(anyhow::anyhow!("Session error: {}", e));
}
Err(e) => {
return Err(anyhow::anyhow!("Session task panicked: {}", e));
}
}
}
Some(result) = async {
match incoming_udp_rx.as_mut() {
Some(rx) => rx.recv().await,
None => std::future::pending().await,
}
} => {
match result {
Ok(conn) => {
debug!(peer_node = conn.peer_node_number, "Incoming UDP P2P connection");
tokio::spawn(handle_udp_connection(conn));
}
Err(e) => {
warn!(error = %e, "UDP connection failed");
}
}
}
Some(result) = async {
match incoming_quic_rx.as_mut() {
Some(rx) => rx.recv().await,
None => std::future::pending().await,
}
} => {
match result {
Ok(conn) => {
debug!(peer_node = conn.peer_node_number, "Incoming QUIC P2P connection");
let allowed_ports = Arc::clone(&allowed_ports);
tokio::spawn(handle_quic_connection(conn, allowed_ports, allow_egress));
}
Err(e) => {
warn!(error = %e, "QUIC connection handshake failed");
}
}
}
result = ipc_server.accept() => {
match result {
Ok(stream) => {
let client_handle_state = handle_state.clone();
tokio::spawn(async move {
ipc::handle_client_with_optional_handle(stream, client_handle_state).await;
});
}
Err(e) => {
error!(error = %e, "Failed to accept IPC connection");
}
}
}
}
}
Ok(())
}
async fn handle_udp_connection(conn: UdpConnection) {
let peer = conn.peer_node_number;
debug!(peer, "UDP connected (connection already established)");
loop {
match conn.recv().await {
Ok(data) => {
if data == b"ping" {
info!(peer, "Received ping, sending pong");
if let Err(e) = conn.send(b"pong") {
warn!(peer, error = %e, "Failed to send pong");
break;
}
} else {
debug!(peer, bytes = data.len(), "Received data");
}
}
Err(e) => {
debug!(peer, error = %e, "UDP connection closed");
break;
}
}
}
}
async fn handle_quic_connection(
conn: QuicConnection,
allowed_ports: Arc<Option<AllowedPorts>>,
allow_egress: bool,
) {
let peer = conn.peer_node_number;
debug!(peer, "QUIC connected (connection already established)");
loop {
match conn.accept_stream().await {
Ok(stream) => {
let stream_id = stream.id();
debug!(peer, stream_id, "Accepted stream");
let allowed_ports = Arc::clone(&allowed_ports);
tokio::spawn(handle_quic_stream(
stream,
peer,
stream_id,
allowed_ports,
allow_egress,
));
}
Err(e) => {
debug!(peer, error = %e, "QUIC connection closed");
break;
}
}
}
}
async fn handle_quic_stream(
stream: wispers_connect::QuicStream,
_peer: i32,
stream_id: u64,
allowed_ports: Arc<Option<AllowedPorts>>,
allow_egress: bool,
) {
let mut buf = [0u8; 1024];
let n = match stream.read(&mut buf).await {
Ok(0) => {
debug!(stream_id, "Stream closed by peer before command");
return;
}
Ok(n) => n,
Err(e) => {
warn!(stream_id, error = %e, "Stream read error");
return;
}
};
let data = &buf[..n];
let line = match data.iter().position(|&b| b == b'\n') {
Some(pos) => &data[..pos],
None => data,
};
match line {
b"PING" => {
info!(stream_id, "Received PING, sending PONG");
if let Err(e) = stream.write_all(b"PONG\n").await {
warn!(stream_id, error = %e, "Failed to send PONG");
}
let _ = stream.finish().await;
}
cmd if cmd.starts_with(b"FORWARD ") => {
let port_str = String::from_utf8_lossy(&cmd[8..]);
match port_str.trim().parse::<u16>() {
Ok(port) => {
info!(stream_id, port, "Received FORWARD");
handle_forward_stream(stream, port, &allowed_ports).await;
}
Err(_) => {
let _ = stream.write_all(b"ERROR invalid port\n").await;
let _ = stream.finish().await;
}
}
}
cmd if cmd.starts_with(b"CONNECT ") => {
let target = String::from_utf8_lossy(&cmd[8..]).trim().to_string();
info!(stream_id, target = %target, "Received CONNECT");
handle_connect_stream(stream, &target, allow_egress).await;
}
_ => {
warn!(
stream_id,
cmd = %String::from_utf8_lossy(line),
"Unknown command",
);
let _ = stream.write_all(b"ERROR unknown command\n").await;
let _ = stream.finish().await;
}
}
}
async fn handle_forward_stream(
stream: wispers_connect::QuicStream,
port: u16,
allowed_ports: &Option<AllowedPorts>,
) {
let allowed = match allowed_ports {
None => false,
Some(ports) => ports.is_allowed(port),
};
if !allowed {
warn!(port, "FORWARD denied: port not in allowlist");
let _ = stream.write_all(b"ERROR port not allowed\n").await;
let _ = stream.finish().await;
return;
}
let stream = Arc::new(stream);
let tcp = match TcpStream::connect(format!("127.0.0.1:{}", port)).await {
Ok(tcp) => {
if let Err(e) = stream.write_all(b"OK\n").await {
warn!(error = %e, "Failed to send OK");
return;
}
tcp
}
Err(e) => {
let msg = format!("ERROR {}\n", e);
let _ = stream.write_all(msg.as_bytes()).await;
let _ = stream.finish().await;
return;
}
};
let (mut tcp_read, mut tcp_write) = tcp.into_split();
let stream_read = Arc::clone(&stream);
let stream_write = Arc::clone(&stream);
let quic_to_tcp = async move {
let mut buf = [0u8; 8192];
loop {
match stream_read.read(&mut buf).await {
Ok(0) => break,
Ok(n) => {
if let Err(e) = tcp_write.write_all(&buf[..n]).await {
debug!(error = %e, "TCP write error");
break;
}
}
Err(e) => {
debug!(error = %e, "QUIC read error");
break;
}
}
}
let _ = tcp_write.shutdown().await;
};
let tcp_to_quic = async move {
let mut buf = [0u8; 8192];
loop {
match tcp_read.read(&mut buf).await {
Ok(0) => break,
Ok(n) => {
if let Err(e) = stream_write.write_all(&buf[..n]).await {
debug!(error = %e, "QUIC write error");
break;
}
}
Err(e) => {
debug!(error = %e, "TCP read error");
break;
}
}
}
let _ = stream_write.finish().await;
};
tokio::join!(quic_to_tcp, tcp_to_quic);
}
async fn handle_connect_stream(
stream: wispers_connect::QuicStream,
target: &str,
allow_egress: bool,
) {
if !allow_egress {
warn!(target, "CONNECT denied: egress not enabled");
let _ = stream.write_all(b"ERROR egress not allowed\n").await;
let _ = stream.finish().await;
return;
}
let (host, port) = match parse_host_port(target) {
Some(hp) => hp,
None => {
warn!(target, "CONNECT invalid target");
let _ = stream.write_all(b"ERROR invalid target format\n").await;
let _ = stream.finish().await;
return;
}
};
let stream = Arc::new(stream);
let tcp = match TcpStream::connect((host.as_str(), port)).await {
Ok(tcp) => {
if let Err(e) = stream.write_all(b"OK\n").await {
warn!(error = %e, "Failed to send OK");
return;
}
tcp
}
Err(e) => {
let msg = format!("ERROR {}\n", e);
let _ = stream.write_all(msg.as_bytes()).await;
let _ = stream.finish().await;
return;
}
};
let (mut tcp_read, mut tcp_write) = tcp.into_split();
let stream_read = Arc::clone(&stream);
let stream_write = Arc::clone(&stream);
let quic_to_tcp = async move {
let mut buf = [0u8; 8192];
loop {
match stream_read.read(&mut buf).await {
Ok(0) => break,
Ok(n) => {
if let Err(e) = tcp_write.write_all(&buf[..n]).await {
debug!(error = %e, "TCP write error");
break;
}
}
Err(e) => {
debug!(error = %e, "QUIC read error");
break;
}
}
}
let _ = tcp_write.shutdown().await;
};
let tcp_to_quic = async move {
let mut buf = [0u8; 8192];
loop {
match tcp_read.read(&mut buf).await {
Ok(0) => break,
Ok(n) => {
if let Err(e) = stream_write.write_all(&buf[..n]).await {
debug!(error = %e, "QUIC write error");
break;
}
}
Err(e) => {
debug!(error = %e, "TCP read error");
break;
}
}
}
let _ = stream_write.finish().await;
};
tokio::join!(quic_to_tcp, tcp_to_quic);
}
fn parse_host_port(target: &str) -> Option<(String, u16)> {
let colon_pos = target.rfind(':')?;
let host = &target[..colon_pos];
let port_str = &target[colon_pos + 1..];
if host.is_empty() {
return None;
}
let port: u16 = port_str.parse().ok()?;
Some((host.to_string(), port))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_allowed_ports_parse_empty() {
let ports = AllowedPorts::parse("").unwrap();
assert!(matches!(ports, AllowedPorts::All));
assert!(ports.is_allowed(80));
assert!(ports.is_allowed(443));
assert!(ports.is_allowed(8080));
}
#[test]
fn test_allowed_ports_parse_single() {
let ports = AllowedPorts::parse("80").unwrap();
assert!(ports.is_allowed(80));
assert!(!ports.is_allowed(443));
}
#[test]
fn test_allowed_ports_parse_multiple() {
let ports = AllowedPorts::parse("80,443,8080").unwrap();
assert!(ports.is_allowed(80));
assert!(ports.is_allowed(443));
assert!(ports.is_allowed(8080));
assert!(!ports.is_allowed(22));
}
#[test]
fn test_allowed_ports_parse_with_spaces() {
let ports = AllowedPorts::parse("80, 443, 8080").unwrap();
assert!(ports.is_allowed(80));
assert!(ports.is_allowed(443));
assert!(ports.is_allowed(8080));
}
#[test]
fn test_allowed_ports_parse_invalid() {
assert!(AllowedPorts::parse("abc").is_err());
assert!(AllowedPorts::parse("80,abc").is_err());
assert!(AllowedPorts::parse("99999").is_err()); }
#[test]
fn test_parse_host_port_basic() {
let (host, port) = parse_host_port("example.com:443").unwrap();
assert_eq!(host, "example.com");
assert_eq!(port, 443);
}
#[test]
fn test_parse_host_port_localhost() {
let (host, port) = parse_host_port("localhost:8080").unwrap();
assert_eq!(host, "localhost");
assert_eq!(port, 8080);
}
#[test]
fn test_parse_host_port_ipv4() {
let (host, port) = parse_host_port("192.168.1.1:80").unwrap();
assert_eq!(host, "192.168.1.1");
assert_eq!(port, 80);
}
#[test]
fn test_parse_host_port_ipv6() {
let (host, port) = parse_host_port("[::1]:8080").unwrap();
assert_eq!(host, "[::1]");
assert_eq!(port, 8080);
}
#[test]
fn test_parse_host_port_invalid() {
assert!(parse_host_port("example.com").is_none()); assert!(parse_host_port(":8080").is_none()); assert!(parse_host_port("example.com:abc").is_none()); assert!(parse_host_port("example.com:99999").is_none()); }
}