use std::collections::VecDeque;
use std::net::SocketAddr;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, RwLock};
use std::time::{SystemTime, UNIX_EPOCH};
use dashmap::DashMap;
use quinn::Connection;
use serde::Serialize;
use crate::common::counted::TunnelCounters;
use crate::common::remote::{Direction, RemoteKind, RemoteRequest};
pub const HISTORY_CAPACITY: usize = 256;
fn unix_ms(t: SystemTime) -> u64 {
t.duration_since(UNIX_EPOCH)
.map(|d| d.as_millis() as u64)
.unwrap_or(0)
}
#[derive(Debug)]
pub struct ClientEntry {
pub id: u64,
pub remote: SocketAddr,
pub connected_at: SystemTime,
pub tunnels: DashMap<u64, Arc<TunnelEntry>>,
#[allow(dead_code)]
pub conn: Connection,
}
impl ClientEntry {
pub fn totals(&self) -> ClientTotals {
let mut t = ClientTotals::default();
for entry in self.tunnels.iter() {
let tot = entry.value().totals();
t.active_in += tot.active_in;
t.active_out += tot.active_out;
t.cumulative_in += tot.cumulative_in;
t.cumulative_out += tot.cumulative_out;
t.active_conns += tot.active_conns;
t.total_conns += tot.total_conns;
}
t
}
}
#[derive(Debug)]
pub struct TunnelEntry {
pub id: u64,
pub client_id: u64,
pub direction: Direction,
pub kind: RemoteKind,
pub spec: String,
pub opened_at: SystemTime,
pub conns: DashMap<u64, Arc<ConnEntry>>,
cumulative_in: AtomicU64,
cumulative_out: AtomicU64,
total_conns: AtomicU64,
}
#[derive(Debug, Default, Clone, Copy)]
pub struct TunnelTotals {
pub active_in: u64,
pub active_out: u64,
pub cumulative_in: u64,
pub cumulative_out: u64,
pub active_conns: u64,
pub total_conns: u64,
}
#[derive(Debug, Default, Clone, Copy)]
pub struct ClientTotals {
pub active_in: u64,
pub active_out: u64,
pub cumulative_in: u64,
pub cumulative_out: u64,
pub active_conns: u64,
pub total_conns: u64,
}
impl TunnelEntry {
pub fn totals(&self) -> TunnelTotals {
let mut active_in = 0u64;
let mut active_out = 0u64;
for c in self.conns.iter() {
let (i, o) = c.value().counters.snapshot();
active_in += i;
active_out += o;
}
TunnelTotals {
active_in,
active_out,
cumulative_in: self.cumulative_in.load(Ordering::Relaxed),
cumulative_out: self.cumulative_out.load(Ordering::Relaxed),
active_conns: self.conns.len() as u64,
total_conns: self.total_conns.load(Ordering::Relaxed),
}
}
}
#[derive(Debug)]
pub struct ConnEntry {
pub id: u64,
pub tunnel_id: u64,
pub client_id: u64,
pub opened_at: SystemTime,
pub peer: Option<String>,
pub counters: Arc<TunnelCounters>,
}
#[derive(Debug, Clone, Serialize)]
pub struct HistoryEntry {
pub client_id: u64,
pub remote: String,
pub connected_at_ms: u64,
pub disconnected_at_ms: u64,
pub reason: String,
pub bytes_in: u64,
pub bytes_out: u64,
pub total_conns: u64,
}
#[derive(Debug, Clone)]
pub struct ServerState {
inner: Arc<Inner>,
}
#[derive(Debug)]
struct Inner {
started_at: SystemTime,
listen_addr: SocketAddr,
next_tunnel_id: AtomicU64,
next_conn_id: AtomicU64,
clients: DashMap<u64, Arc<ClientEntry>>,
tunnels: DashMap<u64, Arc<TunnelEntry>>,
conns: DashMap<u64, Arc<ConnEntry>>,
history: RwLock<VecDeque<HistoryEntry>>,
}
impl ServerState {
pub fn new(listen_addr: SocketAddr) -> Self {
Self {
inner: Arc::new(Inner {
started_at: SystemTime::now(),
listen_addr,
next_tunnel_id: AtomicU64::new(0),
next_conn_id: AtomicU64::new(0),
clients: DashMap::new(),
tunnels: DashMap::new(),
conns: DashMap::new(),
history: RwLock::new(VecDeque::with_capacity(HISTORY_CAPACITY)),
}),
}
}
pub fn started_at(&self) -> SystemTime {
self.inner.started_at
}
pub fn listen_addr(&self) -> SocketAddr {
self.inner.listen_addr
}
pub fn client_count(&self) -> usize {
self.inner.clients.len()
}
pub fn tunnel_count(&self) -> usize {
self.inner.tunnels.len()
}
pub fn conn_count(&self) -> usize {
self.inner.conns.len()
}
pub fn register_client(
&self,
id: u64,
remote: SocketAddr,
conn: Connection,
) -> Arc<ClientEntry> {
let entry = Arc::new(ClientEntry {
id,
remote,
connected_at: SystemTime::now(),
tunnels: DashMap::new(),
conn,
});
self.inner.clients.insert(id, entry.clone());
entry
}
pub fn deregister_client(&self, id: u64, reason: impl Into<String>) {
let Some((_, entry)) = self.inner.clients.remove(&id) else {
return;
};
let totals = entry.totals();
let tunnel_ids: Vec<u64> = entry.tunnels.iter().map(|t| t.value().id).collect();
for tunnel_id in &tunnel_ids {
if let Some((_, tunnel)) = self.inner.tunnels.remove(tunnel_id) {
for c in tunnel.conns.iter() {
self.inner.conns.remove(&c.value().id);
}
}
}
let h = HistoryEntry {
client_id: entry.id,
remote: entry.remote.to_string(),
connected_at_ms: unix_ms(entry.connected_at),
disconnected_at_ms: unix_ms(SystemTime::now()),
reason: reason.into(),
bytes_in: totals.active_in + totals.cumulative_in,
bytes_out: totals.active_out + totals.cumulative_out,
total_conns: totals.total_conns,
};
if let Ok(mut hist) = self.inner.history.write() {
if hist.len() == HISTORY_CAPACITY {
hist.pop_front();
}
hist.push_back(h);
}
}
pub fn register_tunnels(
&self,
client: &ClientEntry,
requests: &[RemoteRequest],
) -> Vec<Arc<TunnelEntry>> {
requests
.iter()
.map(|req| {
let id = self.inner.next_tunnel_id.fetch_add(1, Ordering::Relaxed) + 1;
let entry = Arc::new(TunnelEntry {
id,
client_id: client.id,
direction: req.direction,
kind: req.kind.clone(),
spec: req.to_string(),
opened_at: SystemTime::now(),
conns: DashMap::new(),
cumulative_in: AtomicU64::new(0),
cumulative_out: AtomicU64::new(0),
total_conns: AtomicU64::new(0),
});
client.tunnels.insert(id, entry.clone());
self.inner.tunnels.insert(id, entry.clone());
entry
})
.collect()
}
pub fn tunnel(&self, id: u64) -> Option<Arc<TunnelEntry>> {
self.inner.tunnels.get(&id).map(|e| e.value().clone())
}
pub fn tunnels_snapshot(&self) -> Vec<Arc<TunnelEntry>> {
self.inner
.tunnels
.iter()
.map(|e| e.value().clone())
.collect()
}
pub fn register_conn(&self, tunnel: &Arc<TunnelEntry>, peer: Option<String>) -> ConnGuard {
let id = self.inner.next_conn_id.fetch_add(1, Ordering::Relaxed) + 1;
let counters = TunnelCounters::new();
let entry = Arc::new(ConnEntry {
id,
tunnel_id: tunnel.id,
client_id: tunnel.client_id,
opened_at: SystemTime::now(),
peer,
counters: counters.clone(),
});
tunnel.conns.insert(id, entry.clone());
self.inner.conns.insert(id, entry.clone());
tunnel.total_conns.fetch_add(1, Ordering::Relaxed);
ConnGuard {
state: self.clone(),
tunnel: tunnel.clone(),
conn: entry,
}
}
pub fn conn(&self, id: u64) -> Option<Arc<ConnEntry>> {
self.inner.conns.get(&id).map(|e| e.value().clone())
}
pub fn conns_snapshot(&self) -> Vec<Arc<ConnEntry>> {
self.inner.conns.iter().map(|e| e.value().clone()).collect()
}
fn close_conn(&self, tunnel: &Arc<TunnelEntry>, conn: &Arc<ConnEntry>) {
let (i, o) = conn.counters.snapshot();
tunnel.cumulative_in.fetch_add(i, Ordering::Relaxed);
tunnel.cumulative_out.fetch_add(o, Ordering::Relaxed);
tunnel.conns.remove(&conn.id);
self.inner.conns.remove(&conn.id);
}
pub fn clients_snapshot(&self) -> Vec<Arc<ClientEntry>> {
self.inner
.clients
.iter()
.map(|e| e.value().clone())
.collect()
}
pub fn client(&self, id: u64) -> Option<Arc<ClientEntry>> {
self.inner.clients.get(&id).map(|e| e.value().clone())
}
pub fn history_snapshot(&self, limit: usize) -> Vec<HistoryEntry> {
let Ok(hist) = self.inner.history.read() else {
return Vec::new();
};
let take = limit.min(hist.len());
hist.iter().rev().take(take).cloned().collect()
}
}
pub struct ConnGuard {
state: ServerState,
tunnel: Arc<TunnelEntry>,
conn: Arc<ConnEntry>,
}
impl ConnGuard {
pub fn counters(&self) -> Arc<TunnelCounters> {
self.conn.counters.clone()
}
pub fn id(&self) -> u64 {
self.conn.id
}
}
impl Drop for ConnGuard {
fn drop(&mut self) {
self.state.close_conn(&self.tunnel, &self.conn);
}
}
#[derive(Clone)]
pub struct TunnelHandle {
state: ServerState,
tunnel: Arc<TunnelEntry>,
}
impl TunnelHandle {
pub fn new(state: ServerState, tunnel: Arc<TunnelEntry>) -> Self {
Self { state, tunnel }
}
pub fn open_conn(&self, peer: Option<String>) -> ConnGuard {
self.state.register_conn(&self.tunnel, peer)
}
}
#[derive(Debug, Serialize)]
pub struct ServerInfoDto {
pub version: &'static str,
pub listen_addr: String,
pub started_at_ms: u64,
pub uptime_ms: u64,
pub client_count: usize,
pub tunnel_count: usize,
pub active_conn_count: usize,
}
#[derive(Debug, Serialize)]
pub struct ClientSummaryDto {
pub id: u64,
pub remote: String,
pub connected_at_ms: u64,
pub tunnel_count: usize,
pub active_conn_count: u64,
pub total_conns: u64,
pub bytes_in: u64,
pub bytes_out: u64,
}
#[derive(Debug, Serialize)]
pub struct ClientDetailDto {
#[serde(flatten)]
pub summary: ClientSummaryDto,
pub tunnels: Vec<TunnelDto>,
}
#[derive(Debug, Serialize)]
pub struct TunnelDto {
pub id: u64,
pub client_id: u64,
pub direction: &'static str,
pub kind: &'static str,
pub spec: String,
pub opened_at_ms: u64,
pub active_conn_count: u64,
pub total_conns: u64,
pub active_bytes_in: u64,
pub active_bytes_out: u64,
pub bytes_in: u64,
pub bytes_out: u64,
}
#[derive(Debug, Serialize)]
pub struct TunnelDetailDto {
#[serde(flatten)]
pub summary: TunnelDto,
pub conns: Vec<ConnDto>,
}
#[derive(Debug, Serialize)]
pub struct ConnDto {
pub id: u64,
pub tunnel_id: u64,
pub client_id: u64,
pub opened_at_ms: u64,
pub peer: Option<String>,
pub bytes_in: u64,
pub bytes_out: u64,
}
impl ClientSummaryDto {
pub fn from_entry(entry: &ClientEntry) -> Self {
let t = entry.totals();
Self {
id: entry.id,
remote: entry.remote.to_string(),
connected_at_ms: unix_ms(entry.connected_at),
tunnel_count: entry.tunnels.len(),
active_conn_count: t.active_conns,
total_conns: t.total_conns,
bytes_in: t.active_in + t.cumulative_in,
bytes_out: t.active_out + t.cumulative_out,
}
}
}
impl TunnelDto {
pub fn from_entry(entry: &TunnelEntry) -> Self {
let t = entry.totals();
Self {
id: entry.id,
client_id: entry.client_id,
direction: match entry.direction {
Direction::Forward => "forward",
Direction::Reverse => "reverse",
},
kind: match entry.kind {
RemoteKind::Tcp { .. } => "tcp",
RemoteKind::Udp { .. } => "udp",
RemoteKind::Socks5 { .. } => "socks5",
},
spec: entry.spec.clone(),
opened_at_ms: unix_ms(entry.opened_at),
active_conn_count: t.active_conns,
total_conns: t.total_conns,
active_bytes_in: t.active_in,
active_bytes_out: t.active_out,
bytes_in: t.active_in + t.cumulative_in,
bytes_out: t.active_out + t.cumulative_out,
}
}
}
impl ConnDto {
pub fn from_entry(entry: &ConnEntry) -> Self {
let (i, o) = entry.counters.snapshot();
Self {
id: entry.id,
tunnel_id: entry.tunnel_id,
client_id: entry.client_id,
opened_at_ms: unix_ms(entry.opened_at),
peer: entry.peer.clone(),
bytes_in: i,
bytes_out: o,
}
}
}
pub fn server_info(state: &ServerState) -> ServerInfoDto {
let now = SystemTime::now();
let uptime_ms = now
.duration_since(state.started_at())
.map(|d| d.as_millis() as u64)
.unwrap_or(0);
ServerInfoDto {
version: env!("CARGO_PKG_VERSION"),
listen_addr: state.listen_addr().to_string(),
started_at_ms: unix_ms(state.started_at()),
uptime_ms,
client_count: state.client_count(),
tunnel_count: state.tunnel_count(),
active_conn_count: state.conn_count(),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn history_is_bounded_and_recent_first() {
let state = ServerState::new("127.0.0.1:0".parse().unwrap());
for i in 0..(HISTORY_CAPACITY + 5) {
let mut hist = state.inner.history.write().unwrap();
if hist.len() == HISTORY_CAPACITY {
hist.pop_front();
}
hist.push_back(HistoryEntry {
client_id: i as u64,
remote: format!("127.0.0.1:{i}"),
connected_at_ms: 0,
disconnected_at_ms: 0,
reason: "ok".into(),
bytes_in: 0,
bytes_out: 0,
total_conns: 0,
});
}
let snap = state.history_snapshot(10);
assert_eq!(snap.len(), 10);
assert!(snap[0].client_id > snap.last().unwrap().client_id);
}
}