use std::{
net::SocketAddr,
path::PathBuf,
sync::{
Arc,
atomic::{AtomicBool, Ordering},
},
time::{Duration, Instant},
};
use eyre::Result;
use future_form::Sendable;
use iroh::{EndpointAddr, endpoint::presets};
use sedimentree_core::depth::CountLeadingZeroBytes;
use subduction_core::{
authenticated::Authenticated,
handshake::{
self,
audience::{Audience, DiscoveryId},
},
nonce_cache::NonceCache,
peer::{
counter::{PeerCounter, wall_clock_seed},
id::PeerId,
},
storage::metrics::{MetricsStorage, RefreshMetrics},
subduction::{Subduction, builder::SubductionBuilder},
timeout::call::CallTimeout,
timestamp::TimestampSeconds,
transport::message::MessageTransport,
};
use subduction_crypto::{nonce::Nonce, signer::memory::MemorySigner};
use subduction_http_longpoll::server::LongPollHandler;
use subduction_redb_storage::RedbStorage;
use subduction_websocket::{
DEFAULT_MAX_MESSAGE_SIZE,
handshake::WebSocketHandshake,
sleep::TokioSleeper,
timeout::FuturesTimerTimeout,
tokio::{TokioSpawn, unified::UnifiedWebSocket},
websocket::{KeepAlive, WebSocket},
};
use tokio::{net::TcpListener, task::JoinSet, time};
use tokio_util::sync::CancellationToken;
use tungstenite::{http::Uri, protocol::WebSocketConfig};
use subduction_ephemeral::{
clock::std_clock::StdClock, config::EphemeralConfig, handler::EphemeralHandler,
policy::OpenEphemeralPolicy,
};
use subduction_keyhive::{connection::KeyhiveConnection, runtime::init_sendable_keyhive};
use crate::{
handler::{
CliConn, CliEphemeralHandler, CliHandler, CliHandlerOpenPolicy, CliKeyhiveHandler,
CliKeyhiveProtocol, CliSyncHandler, CliWireHandler,
},
key,
keyhive::{CliConnKeyhiveAdapter, FsKeyhiveStorage, KEYHIVE_DIR},
metrics,
policy::CliKeyhivePolicyHandle,
transport::UnifiedTransport,
};
type CliSubduction<H> = Arc<
Subduction<
'static,
future_form::Sendable,
MetricsStorage<RedbStorage>,
CliConn,
H,
CliKeyhivePolicyHandle,
MemorySigner,
FuturesTimerTimeout,
TokioSpawn,
CountLeadingZeroBytes,
>,
>;
#[derive(Debug, clap::Parser)]
#[allow(clippy::struct_excessive_bools)]
pub(crate) struct ServerArgs {
#[arg(short, long, default_value = "0.0.0.0:8080")]
pub(crate) socket: String,
#[arg(short, long)]
pub(crate) data_dir: Option<PathBuf>,
#[command(flatten)]
pub(crate) key: key::KeyArgs,
#[arg(long, default_value = "600")]
pub(crate) handshake_max_drift: u64,
#[arg(long)]
pub(crate) service_name: Option<String>,
#[arg(short, long, default_value = "5")]
pub(crate) timeout: u64,
#[arg(long, default_value_t = DEFAULT_MAX_MESSAGE_SIZE)]
pub(crate) max_message_size: usize,
#[arg(long = "max-frame-size", value_name = "MAX_FRAME_SIZE")]
pub(crate) max_frame_size_override: Option<usize>,
#[arg(long, default_value = "9090")]
pub(crate) metrics_port: u16,
#[arg(long, default_value_t = false)]
pub(crate) metrics: bool,
#[arg(long, value_name = "ADDR")]
pub(crate) admin_addr: Option<SocketAddr>,
#[arg(long, default_value_t = DEFAULT_METRICS_REFRESH_SECS)]
pub(crate) metrics_refresh_interval: u64,
#[arg(long, value_name = "MAX_RESIDENT_TREES")]
pub(crate) max_resident_trees: Option<usize>,
#[arg(
long,
default_value_t = true,
action = clap::ArgAction::Set,
num_args = 0..=1,
default_missing_value = "true"
)]
pub(crate) websocket: bool,
#[arg(
long,
default_value_t = true,
action = clap::ArgAction::Set,
num_args = 0..=1,
default_missing_value = "true"
)]
pub(crate) longpoll: bool,
#[arg(long = "ws-peer", value_name = "URL")]
pub(crate) ws_peers: Vec<String>,
#[arg(long, default_value_t = false)]
pub(crate) iroh: bool,
#[arg(long = "iroh-peer", value_name = "NODE_ID")]
pub(crate) iroh_peers: Vec<String>,
#[arg(long = "iroh-peer-addr", value_name = "IP:PORT")]
pub(crate) iroh_peer_addrs: Vec<SocketAddr>,
#[arg(long = "iroh-direct-only")]
pub(crate) iroh_direct_only: bool,
#[arg(long = "iroh-relay-url", value_name = "URL")]
pub(crate) iroh_relay_url: Option<String>,
#[arg(long = "ready-file", value_name = "PATH")]
pub(crate) ready_file: Option<PathBuf>,
#[arg(long, value_enum, default_value_t = AuthMode::Keyhive)]
pub(crate) auth: AuthMode,
#[arg(long, action = clap::ArgAction::Set, default_value_t = true)]
pub(crate) keyhive_cache_refresh: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, clap::ValueEnum)]
pub(crate) enum AuthMode {
Keyhive,
Open,
}
impl AuthMode {
pub(crate) const fn keyhive_enabled(self) -> bool {
matches!(self, AuthMode::Keyhive)
}
}
impl ServerArgs {
pub(crate) fn max_frame_size(&self) -> usize {
self.max_frame_size_override
.unwrap_or(self.max_message_size)
}
}
const DEFAULT_METRICS_REFRESH_SECS: u64 = 60;
pub(crate) async fn run(args: ServerArgs, token: CancellationToken) -> Result<()> {
if args.auth.keyhive_enabled() {
run_with_keyhive(args, token).await
} else {
run_open(args, token).await
}
}
async fn run_with_keyhive(args: ServerArgs, token: CancellationToken) -> Result<()> {
tracing::warn!(version = env!("CARGO_PKG_VERSION"), "Subduction server");
let common = SetupCommon::init(&args, &token).await?;
let keyhive_signer = key::keyhive_signer_from_seed(&common.seed);
let keyhive_root = common.data_dir.join(KEYHIVE_DIR);
tracing::info!(root = ?keyhive_root, "Initializing keyhive storage");
let fs_keyhive_storage = FsKeyhiveStorage::new(keyhive_root)?;
let (keyhive_instance, kh_peer_id, contact_card) = init_sendable_keyhive(keyhive_signer)
.await
.map_err(|e| eyre::eyre!(e))?;
let shared_keyhive = Arc::new(async_lock::Mutex::new(keyhive_instance));
let keyhive_protocol: CliKeyhiveProtocol = Arc::new(subduction_keyhive::KeyhiveProtocol::new(
Arc::clone(&shared_keyhive),
fs_keyhive_storage,
kh_peer_id,
contact_card,
));
if let Err(e) = keyhive_protocol.ingest_from_storage().await {
tracing::warn!(error = %e, "keyhive ingest_from_storage failed");
}
let storage_policy = Arc::new(CliKeyhivePolicyHandle::new(Arc::clone(&shared_keyhive)));
if args.keyhive_cache_refresh {
let refresh_proto = Arc::clone(&keyhive_protocol);
let refresh_cancel = token.clone();
tokio::spawn(async move {
let mut tick = time::interval(Duration::from_secs(2));
tick.tick().await;
loop {
tokio::select! {
() = refresh_cancel.cancelled() => break,
_ = tick.tick() => {
if let Err(e) = refresh_proto.refresh_cache().await {
tracing::warn!(error = %e, "refresh_cache failed");
}
}
}
}
tracing::debug!("keyhive cache refresh task shutting down");
});
}
let keyhive_for_handler = Arc::clone(&keyhive_protocol);
serve(
args,
token,
common,
storage_policy,
Some(keyhive_protocol),
move |sync, ephemeral| {
let keyhive = CliKeyhiveHandler::new(keyhive_for_handler, CliConnKeyhiveAdapter::new);
Arc::new(CliHandler::new(sync, ephemeral, keyhive))
},
)
.await
}
async fn run_open(args: ServerArgs, token: CancellationToken) -> Result<()> {
tracing::warn!(version = env!("CARGO_PKG_VERSION"), "Subduction server");
tracing::info!("Keyhive disabled (--auth=open); using open (allow-all) storage policy");
let common = SetupCommon::init(&args, &token).await?;
let storage_policy = Arc::new(CliKeyhivePolicyHandle::open());
serve(
args,
token,
common,
storage_policy,
None,
|sync, ephemeral| Arc::new(CliHandlerOpenPolicy::new(sync, ephemeral)),
)
.await
}
struct SetupCommon {
addr: SocketAddr,
data_dir: PathBuf,
storage: MetricsStorage<RedbStorage>,
seed: [u8; 32],
signer: MemorySigner,
peer_id: PeerId,
handshake_max_drift: Duration,
service_name: String,
discovery_id: Option<DiscoveryId>,
discovery_audience: Option<Audience>,
}
impl SetupCommon {
async fn init(args: &ServerArgs, token: &CancellationToken) -> Result<Self> {
let addr: SocketAddr = args.socket.parse()?;
let data_dir = args
.data_dir
.clone()
.unwrap_or_else(|| PathBuf::from("./data"));
if args.metrics {
let metrics_handle = metrics::init_metrics();
let metrics_addr: SocketAddr = ([0, 0, 0, 0], args.metrics_port).into();
metrics::start_metrics_server(metrics_addr, metrics_handle).await?;
subduction_core::metrics::set_build_info(
env!("CARGO_PKG_VERSION"),
env!("SUBDUCTION_GIT_SHA"),
);
}
tracing::info!(dir = ?data_dir, "Initializing redb storage");
let redb_storage = RedbStorage::new(data_dir.clone())?;
if let Some(admin_addr) = args.admin_addr {
crate::admin::start_admin_server(admin_addr, redb_storage.clone(), token.clone())
.await?;
}
let storage = MetricsStorage::new(redb_storage);
if args.metrics {
let startup_tree_count = storage.refresh_metrics().await?;
tracing::info!(
sedimentrees = startup_tree_count,
"Loaded sedimentrees from durable storage at startup"
);
let metrics_storage = storage.clone();
let metrics_token = token.clone();
#[cfg(feature = "os-metrics")]
let metrics_data_dir = data_dir.clone();
let refresh_interval = Duration::from_secs(args.metrics_refresh_interval);
tokio::spawn(async move {
#[cfg(feature = "os-metrics")]
let process_collector = {
let collector = metrics_process::Collector::default();
collector.describe();
collector
};
let mut interval = time::interval(refresh_interval);
interval.tick().await;
loop {
tokio::select! {
_ = interval.tick() => {
if let Err(e) = metrics_storage.refresh_metrics().await {
tracing::warn!(error = %e, "Failed to refresh storage metrics");
}
#[cfg(feature = "os-metrics")]
{
process_collector.collect();
publish_disk_usage(&metrics_data_dir);
}
#[cfg(tokio_unstable)]
{
let rt = tokio::runtime::Handle::current().metrics();
subduction_core::metrics::set_tokio_runtime(
subduction_core::metrics::TokioRuntimeSample {
workers: rt.num_workers(),
alive_tasks: rt.num_alive_tasks(),
blocking_threads: rt.num_blocking_threads(),
idle_blocking_threads: rt.num_idle_blocking_threads(),
blocking_queue_depth: rt.blocking_queue_depth(),
global_queue_depth: rt.global_queue_depth(),
},
);
}
}
() = metrics_token.cancelled() => {
tracing::debug!("Stopping metrics refresh task");
break;
}
}
}
});
}
let seed = key::resolve_key_seed(&args.key)?;
let signer = key::signer_from_seed(&seed);
let peer_id = PeerId::from(signer.verifying_key());
let handshake_max_drift = Duration::from_secs(args.handshake_max_drift);
let service_name = args
.service_name
.clone()
.unwrap_or_else(|| args.socket.clone());
let discovery_id = Some(DiscoveryId::new(service_name.as_bytes()));
let discovery_audience: Option<Audience> = discovery_id.map(Audience::discover_id);
Ok(Self {
addr,
data_dir,
storage,
seed,
signer,
peer_id,
handshake_max_drift,
service_name,
discovery_id,
discovery_audience,
})
}
}
#[allow(clippy::too_many_lines)]
async fn serve<H, F>(
args: ServerArgs,
token: CancellationToken,
common: SetupCommon,
storage_policy: Arc<CliKeyhivePolicyHandle>,
keyhive_protocol: Option<CliKeyhiveProtocol>,
make_handler: F,
) -> Result<()>
where
H: CliWireHandler,
F: FnOnce(CliSyncHandler, CliEphemeralHandler) -> Arc<H>,
{
let SetupCommon {
addr,
storage,
signer,
peer_id,
handshake_max_drift,
service_name,
discovery_id,
discovery_audience,
..
} = common;
let builder = SubductionBuilder::new()
.signer(signer.clone())
.storage(storage, storage_policy)
.spawner(TokioSpawn)
.timer(FuturesTimerTimeout)
.send_counter(PeerCounter::with_seed(wall_clock_seed))
.roundtrip_timeout(Duration::from_secs(args.timeout));
let builder = if let Some(max) = args.max_resident_trees {
tracing::info!(
max_resident_trees = max,
"bounding in-memory sedimentree cache"
);
builder.max_resident_trees(max)
} else {
builder
};
let builder = if let Some(id) = discovery_id {
builder.discovery_id(id)
} else {
builder
};
let mut requestor_tally = None;
let (subduction, listener_fut, manager_fut, ephemeral): (CliSubduction<H>, _, _, _) = builder
.build_composed(|sync_handler| {
requestor_tally = Some(sync_handler.requestor_tally());
let connections = sync_handler.connections();
let (ephemeral_handler, ephemeral_rx) = EphemeralHandler::new(
connections,
OpenEphemeralPolicy,
EphemeralConfig::default(),
StdClock,
TokioSpawn,
);
tokio::spawn(async move {
while let Ok(event) = ephemeral_rx.recv().await {
tracing::debug!(
sender = %event.sender,
topic = %event.id,
nonce = event.nonce,
payload_size = event.payload.len(),
"ephemeral event relayed"
);
}
});
let handler = make_handler(sync_handler, ephemeral_handler.clone());
(handler, ephemeral_handler)
});
let server_peer_id = subduction.peer_id();
{
let resident_subduction = subduction.clone();
let resident_token = token.clone();
let refresh_interval = Duration::from_secs(args.metrics_refresh_interval);
tokio::spawn(async move {
let mut interval = time::interval(refresh_interval);
interval.tick().await;
loop {
tokio::select! {
_ = interval.tick() => {
let resident = resident_subduction.resident_sedimentree_count().await;
subduction_core::metrics::set_sedimentree_cache_resident(resident);
subduction_core::metrics::set_connections_active(
resident_subduction.total_connection_count().await,
);
if let Some(tally) = &requestor_tally {
let ranked = tally.take_window().await;
let counts: Vec<u64> =
ranked.iter().map(|(_, count)| *count).collect();
let total: u64 = counts.iter().sum();
subduction_core::metrics::set_top_requestors(&counts, total);
if !ranked.is_empty() {
let top: Vec<String> = ranked
.iter()
.take(10)
.map(|(peer, count)| format!("{peer}={count}"))
.collect();
tracing::info!(
window_secs = refresh_interval.as_secs(),
total_requestors = ranked.len(),
top = ?top,
"top requestors by batch-sync requests"
);
}
}
}
() = resident_token.cancelled() => break,
}
}
});
}
let lp_handler = LongPollHandler::new(
signer.clone(),
Arc::new(NonceCache::default()),
server_peer_id,
discovery_audience,
handshake_max_drift,
FuturesTimerTimeout,
);
let tcp_listener = TcpListener::bind(addr).await?;
let assigned_address = tcp_listener.local_addr()?;
let ws_enabled = args.websocket;
let lp_enabled = args.longpoll;
let iroh_enabled = args.iroh;
if !ws_enabled && !lp_enabled && !iroh_enabled {
eyre::bail!("At least one transport must be enabled (--websocket, --longpoll, or --iroh)");
}
tracing::info!(
addr = %assigned_address,
transports = %[
ws_enabled.then_some("WebSocket"),
lp_enabled.then_some("HTTP long-poll"),
iroh_enabled.then_some("Iroh (QUIC)"),
]
.into_iter()
.flatten()
.collect::<Vec<&str>>()
.join(" + "),
"Server started"
);
tracing::info!(peer = %peer_id, "Peer ID");
let supervised_failure = Arc::new(AtomicBool::new(false));
let actor_cancel = token.clone();
let listener_cancel = token.clone();
let manager_supervisor = actor_cancel.clone();
let manager_failed = Arc::clone(&supervised_failure);
tokio::spawn(async move {
tokio::select! {
_ = manager_fut => {
if !manager_supervisor.is_cancelled() {
tracing::error!(
"connection manager exited unexpectedly; \
shutting down for supervised restart"
);
manager_failed.store(true, Ordering::Release);
manager_supervisor.cancel();
}
},
() = actor_cancel.cancelled() => {}
}
});
let listener_supervisor = listener_cancel.clone();
let listener_failed = Arc::clone(&supervised_failure);
tokio::spawn(async move {
tokio::select! {
_ = listener_fut => {
if !listener_supervisor.is_cancelled() {
tracing::error!(
"dispatch listener exited unexpectedly; \
shutting down for supervised restart"
);
listener_failed.store(true, Ordering::Release);
listener_supervisor.cancel();
}
},
() = listener_cancel.cancelled() => {}
}
});
let accept_cancel = token.child_token();
let accept_subduction = subduction.clone();
let accept_ephemeral = ephemeral.clone();
let accept_handler = lp_handler;
let max_message_size = args.max_message_size;
let max_frame_size = args.max_frame_size();
let ws_keepalive = KeepAlive::balanced();
let accept_keyhive = keyhive_protocol.clone();
let accept_task = tokio::spawn(async move {
accept_loop(
tcp_listener,
accept_subduction,
accept_ephemeral,
accept_handler,
accept_keyhive,
accept_cancel,
handshake_max_drift,
max_message_size,
max_frame_size,
ws_keepalive,
server_peer_id,
discovery_audience,
ws_enabled,
lp_enabled,
)
.await;
});
let mut iroh_node_id: Option<String> = None;
let mut iroh_addrs: Vec<SocketAddr> = Vec::new();
let iroh_accept_task = if iroh_enabled {
let relay_mode = if args.iroh_direct_only {
iroh::endpoint::RelayMode::Disabled
} else if let Some(url) = &args.iroh_relay_url {
let relay_map =
iroh::RelayMap::try_from_iter([url.as_str()]).map_err(|e| eyre::eyre!(e))?;
iroh::endpoint::RelayMode::Custom(relay_map)
} else {
iroh::endpoint::RelayMode::Default
};
let iroh_endpoint = iroh::Endpoint::builder(presets::N0)
.alpns(vec![subduction_iroh::ALPN.to_vec()])
.relay_mode(relay_mode)
.bind()
.await?;
let iroh_addr = iroh_endpoint.addr();
iroh_node_id = Some(iroh_addr.id.to_string());
iroh_addrs = iroh_addr.ip_addrs().copied().collect();
tracing::info!(node_id = %iroh_addr.id, "Iroh endpoint bound");
for addr in &iroh_addr.addrs {
tracing::info!(addr = ?addr, "transport address");
}
let iroh_subduction = subduction.clone();
let iroh_signer = signer.clone();
let iroh_nonce_cache = NonceCache::default();
let iroh_ep = iroh_endpoint.clone();
let iroh_cancel = token.child_token();
let iroh_ephemeral = ephemeral.clone();
let iroh_keyhive_proto = keyhive_protocol.clone();
let task = tokio::spawn({
let cancel = iroh_cancel.clone();
async move {
loop {
tokio::select! {
() = cancel.cancelled() => {
tracing::info!("iroh accept loop canceled");
break;
}
result = subduction_iroh::server::accept_one(
&iroh_ep,
&iroh_signer,
&iroh_nonce_cache,
server_peer_id,
discovery_audience,
handshake_max_drift,
) => {
match result {
Ok(accepted) => {
let remote = accepted.authenticated.peer_id();
tokio::spawn(accepted.listener_task);
tokio::spawn(accepted.sender_task);
let auth = accepted.authenticated.map(|c| MessageTransport::new(UnifiedTransport::Iroh(c)));
let auth_for_keyhive = auth.clone();
match iroh_subduction.add_connection(auth).await {
Ok(_) => {
iroh_ephemeral.subscribe_peer(remote).await;
notify_peer_connect(iroh_keyhive_proto.as_ref(), auth_for_keyhive).await;
iroh_subduction.full_sync_with_peer(&remote, true, CallTimeout::Default).await;
tracing::info!(peer = %remote, "iroh: added peer");
}
Err(e) => {
tracing::error!(error = %e, "failed to add iroh connection");
}
}
}
Err(e) => {
tracing::warn!(error = %e, "iroh accept error");
}
}
}
}
}
}
});
for iroh_peer_str in &args.iroh_peers {
let node_id: iroh::PublicKey = match iroh_peer_str.parse() {
Ok(id) => id,
Err(e) => {
tracing::error!(node_id = %iroh_peer_str, error = %e, "invalid iroh peer node ID");
continue;
}
};
let mut peer_addr = EndpointAddr::new(node_id);
for addr in &args.iroh_peer_addrs {
peer_addr = peer_addr.with_ip_addr(*addr);
}
let peer_ep = iroh_endpoint.clone();
let peer_subduction = subduction.clone();
let peer_ephemeral = ephemeral.clone();
let peer_signer = signer.clone();
let peer_cancel = token.clone();
let peer_service_name = service_name.clone();
let peer_keyhive = keyhive_protocol.clone();
tokio::spawn(async move {
match try_connect_iroh(
&peer_ep,
peer_addr,
&peer_subduction,
&peer_ephemeral,
peer_keyhive.as_ref(),
&peer_signer,
&peer_service_name,
peer_cancel,
)
.await
{
Ok(remote_id) => {
tracing::info!(
node_id = %node_id,
peer = %remote_id,
"iroh: connected to peer"
);
}
Err(e) => {
tracing::error!(node_id = %node_id, error = %e, "iroh: failed to connect to peer");
}
}
});
}
Some(task)
} else {
None
};
for peer_url in &args.ws_peers {
let uri: Uri = match peer_url.parse() {
Ok(uri) => uri,
Err(e) => {
tracing::error!(url = %peer_url, error = %e, "Invalid peer URL");
continue;
}
};
let peer_subduction = subduction.clone();
let peer_ephemeral = ephemeral.clone();
let peer_signer = signer.clone();
let peer_service_name = service_name.clone();
let peer_cancel = token.clone();
let peer_max_message_size = args.max_message_size;
let peer_max_frame_size = args.max_frame_size();
let peer_keyhive = keyhive_protocol.clone();
let peer_keepalive = KeepAlive::balanced();
tokio::spawn(async move {
match try_connect_ws(
uri.clone(),
&peer_subduction,
&peer_ephemeral,
peer_keyhive.as_ref(),
&peer_signer,
&peer_service_name,
peer_cancel,
peer_max_message_size,
peer_max_frame_size,
peer_keepalive,
)
.await
{
Ok(remote_id) => {
tracing::info!(uri = %uri, peer = %remote_id, "Connected to peer");
}
Err(e) => {
tracing::error!(uri = %uri, error = %e, "Failed to connect to peer");
}
}
});
}
if let Some(ref ready_path) = args.ready_file {
let iroh_line = iroh_node_id
.as_deref()
.map_or(String::new(), |id| format!("iroh_node_id={id}\n"));
let iroh_addrs_line = if iroh_addrs.is_empty() {
String::new()
} else {
let addrs: Vec<String> = iroh_addrs.iter().map(ToString::to_string).collect();
format!("iroh_addrs={}\n", addrs.join(","))
};
let content = format!(
"port={}\npeer_id={}\n{iroh_line}{iroh_addrs_line}",
assigned_address.port(),
peer_id,
);
std::fs::write(ready_path, content)
.map_err(|e| eyre::eyre!("failed to write ready file: {e}"))?;
tracing::info!(path = %ready_path.display(), "Ready file written");
}
token.cancelled().await;
tracing::info!("Shutting down server...");
accept_task.abort();
if let Some(iroh_task) = iroh_accept_task {
iroh_task.abort();
}
if supervised_failure.load(Ordering::Acquire) {
return Err(eyre::eyre!(
"critical background task (connection manager or listener) exited unexpectedly"
));
}
Ok(())
}
#[cfg(feature = "os-metrics")]
fn publish_disk_usage(data_dir: &std::path::Path) {
let redb_bytes = std::fs::metadata(data_dir.join(subduction_redb_storage::DB_FILE_NAME))
.map_or(0, |m| m.len());
#[cfg(unix)]
{
match nix::sys::statvfs::statvfs(data_dir) {
Ok(stat) => {
let frsize = stat.fragment_size();
#[allow(clippy::useless_conversion)]
let (free, total) = (
u64::from(stat.blocks_available()) * frsize,
u64::from(stat.blocks()) * frsize,
);
subduction_core::metrics::set_disk_usage(free, total, redb_bytes);
}
Err(e) => {
tracing::warn!(error = %e, "statvfs on data dir failed");
subduction_core::metrics::set_redb_file_bytes(redb_bytes);
}
}
}
#[cfg(not(unix))]
subduction_core::metrics::set_redb_file_bytes(redb_bytes);
}
#[allow(clippy::too_many_arguments)]
async fn accept_loop<H: CliWireHandler>(
tcp_listener: TcpListener,
subduction: CliSubduction<H>,
ephemeral: CliEphemeralHandler,
lp_handler: LongPollHandler<MemorySigner, FuturesTimerTimeout>,
keyhive_proto: Option<CliKeyhiveProtocol>,
cancel: CancellationToken,
handshake_max_drift: Duration,
max_message_size: usize,
max_frame_size: usize,
ws_keepalive: KeepAlive,
server_peer_id: PeerId,
discovery_audience: Option<Audience>,
ws_enabled: bool,
lp_enabled: bool,
) {
let mut conns = JoinSet::new();
loop {
tokio::select! {
() = cancel.cancelled() => {
tracing::info!("accept loop canceled");
break;
}
res = tcp_listener.accept() => {
match res {
Ok((tcp, addr)) => {
tracing::info!(addr = %addr, "new TCP connection");
let task_subduction = subduction.clone();
let task_ephemeral = ephemeral.clone();
let task_handler = lp_handler.clone();
let task_keyhive = keyhive_proto.clone();
let task_discovery = discovery_audience;
conns.spawn(async move {
let mut peek_buf = [0u8; 4];
match tcp.peek(&mut peek_buf).await {
Ok(n) if n >= 3 => {}
Ok(_) | Err(_) => {
tracing::warn!(addr = %addr, "failed to peek TCP stream");
return;
}
}
let is_http = peek_buf.starts_with(b"POST")
|| peek_buf.starts_with(b"OPTI");
if peek_buf.starts_with(b"GET") && ws_enabled {
handle_websocket(
tcp,
addr,
task_subduction,
task_ephemeral.clone(),
task_keyhive,
handshake_max_drift,
max_message_size,
max_frame_size,
ws_keepalive,
server_peer_id,
task_discovery,
)
.await;
} else if is_http && lp_enabled {
handle_http_longpoll(
tcp,
addr,
task_subduction,
task_ephemeral,
task_handler,
task_keyhive,
)
.await;
} else if peek_buf.starts_with(b"GET") {
tracing::warn!(addr = %addr, "WebSocket connection rejected (transport disabled)");
} else if is_http {
tracing::warn!(addr = %addr, "HTTP long-poll connection rejected (transport disabled)");
} else {
tracing::warn!(
addr = %addr,
peek = ?&peek_buf,
"unknown protocol"
);
}
});
}
Err(e) => tracing::error!(error = %e, "Accept error"),
}
}
}
}
while (conns.join_next().await).is_some() {}
}
#[allow(clippy::too_many_arguments, clippy::too_many_lines)]
async fn handle_websocket<H: CliWireHandler>(
tcp: tokio::net::TcpStream,
addr: SocketAddr,
subduction: CliSubduction<H>,
ephemeral: CliEphemeralHandler,
keyhive_proto: Option<CliKeyhiveProtocol>,
handshake_max_drift: Duration,
max_message_size: usize,
max_frame_size: usize,
keepalive: KeepAlive,
server_peer_id: PeerId,
discovery_audience: Option<Audience>,
) {
let mut ws_config = WebSocketConfig::default();
ws_config.max_message_size = Some(max_message_size);
ws_config.max_frame_size = Some(max_frame_size);
let mut forwarded_for: Option<String> = None;
let ws_stream = match async_tungstenite::tokio::accept_hdr_async_with_config(
tcp,
|req: &tungstenite::handshake::server::Request,
resp: tungstenite::handshake::server::Response| {
forwarded_for = req
.headers()
.get("x-forwarded-for")
.and_then(|v| v.to_str().ok())
.map(ToOwned::to_owned);
Ok(resp)
},
Some(ws_config),
)
.await
{
Ok(ws) => ws,
Err(e) => {
tracing::error!(addr = %addr, error = %e, "WebSocket upgrade error");
return;
}
};
tracing::debug!(addr = %addr, "WebSocket upgrade complete");
let now = TimestampSeconds::now();
let handshake_started = Instant::now();
let result = handshake::respond::<future_form::Sendable, _, _, _, _>(
WebSocketHandshake::new(ws_stream),
|ws_handshake, peer_id| {
let (ws, sender_fut, keepalive_task) = WebSocket::new_with_keepalive(
ws_handshake.into_inner(),
peer_id,
keepalive,
TokioSleeper,
);
let listen_ws = ws.clone();
tokio::spawn(async move {
if let Err(e) = listen_ws.listen().await {
tracing::info!(error = %e, "WebSocket listener disconnected");
}
});
tokio::spawn(async move {
if let Err(e) = sender_fut.await {
tracing::info!(error = %e, "WebSocket sender disconnected");
}
});
tokio::spawn(async move {
let outcome = keepalive_task.await;
tracing::debug!(?outcome, "WebSocket keepalive task exited");
});
let unified_ws = UnifiedWebSocket::Accepted(ws);
(
MessageTransport::new(UnifiedTransport::WebSocket(unified_ws)),
(),
)
},
subduction.signer(),
subduction.nonce_cache(),
server_peer_id,
discovery_audience,
now,
handshake_max_drift,
)
.await;
let authenticated = match result {
Ok((auth, ())) => {
subduction_core::metrics::handshake_outcome("ok");
subduction_core::metrics::handshake_duration(
"ok",
handshake_started.elapsed().as_secs_f64(),
);
tracing::info!(
peer = %auth.peer_id(),
addr = %addr,
forwarded_for = forwarded_for.as_deref().unwrap_or("-"),
"WebSocket handshake complete"
);
auth
}
Err(e) => {
use handshake::{
AuthenticateError as AE, HandshakeError, challenge::ChallengeValidationError,
};
#[allow(clippy::wildcard_enum_match_arm)]
let outcome = match &e {
AE::Decode(_) => "decode",
AE::Transport(_) => "io",
AE::ConnectionClosed => "closed",
AE::Handshake(HandshakeError::ChallengeValidation(
ChallengeValidationError::ClockDrift { .. },
)) => "drift",
_ => "rejected",
};
subduction_core::metrics::handshake_outcome(outcome);
subduction_core::metrics::handshake_duration(
"err",
handshake_started.elapsed().as_secs_f64(),
);
tracing::warn!(addr = %addr, error = %e, "WebSocket handshake failed");
return;
}
};
let peer_id = authenticated.peer_id();
let auth_for_keyhive = authenticated.clone();
if let Err(e) = subduction.add_connection(authenticated).await {
tracing::error!(error = %e, "Failed to add WebSocket connection");
} else {
ephemeral.subscribe_peer(peer_id).await;
notify_peer_connect(keyhive_proto.as_ref(), auth_for_keyhive).await;
}
}
async fn handle_http_longpoll<H: CliWireHandler>(
tcp: tokio::net::TcpStream,
addr: SocketAddr,
subduction: CliSubduction<H>,
ephemeral: CliEphemeralHandler,
handler: LongPollHandler<MemorySigner, FuturesTimerTimeout>,
keyhive_proto: Option<CliKeyhiveProtocol>,
) {
use http_body_util::Full;
use hyper::{
body::Bytes,
header::{
ACCESS_CONTROL_ALLOW_HEADERS, ACCESS_CONTROL_ALLOW_METHODS,
ACCESS_CONTROL_ALLOW_ORIGIN, ACCESS_CONTROL_MAX_AGE, HeaderValue,
},
};
use hyper_util::rt::TokioIo;
let io = TokioIo::new(tcp);
let service = hyper::service::service_fn(move |req| {
let handler = handler.clone();
let subduction = subduction.clone();
let ephemeral = ephemeral.clone();
let keyhive_proto = keyhive_proto.clone();
async move {
if req.method() == hyper::Method::OPTIONS {
let mut resp = hyper::Response::new(Full::new(Bytes::new()));
*resp.status_mut() = hyper::StatusCode::NO_CONTENT;
resp.headers_mut()
.insert(ACCESS_CONTROL_ALLOW_ORIGIN, HeaderValue::from_static("*"));
resp.headers_mut().insert(
ACCESS_CONTROL_ALLOW_METHODS,
HeaderValue::from_static("POST, OPTIONS"),
);
resp.headers_mut().insert(
ACCESS_CONTROL_ALLOW_HEADERS,
HeaderValue::from_static("Content-Type, X-Session-Id"),
);
resp.headers_mut().insert(
hyper::header::ACCESS_CONTROL_EXPOSE_HEADERS,
HeaderValue::from_static("X-Session-Id"),
);
resp.headers_mut()
.insert(ACCESS_CONTROL_MAX_AGE, HeaderValue::from_static("86400"));
return Ok::<_, hyper::Error>(resp);
}
let resp = match handler.handle(req).await {
Ok(resp) => resp,
Err(e) => {
tracing::error!(error = %e, "fatal handler error");
hyper::Response::new(Full::new(Bytes::from(e.to_string())))
}
};
if resp.status() == hyper::StatusCode::OK
&& let Some(session_hdr) = resp
.headers()
.get(subduction_http_longpoll::SESSION_ID_HEADER)
&& let Ok(sid_str) = session_hdr.to_str()
&& let Some(sid) = subduction_http_longpoll::session::SessionId::from_hex(sid_str)
&& let Some(auth) = handler.take_authenticated(&sid).await
{
let peer_id = auth.peer_id();
let unified_auth =
auth.map(|lp| MessageTransport::new(UnifiedTransport::HttpLongPoll(lp)));
let auth_for_keyhive = unified_auth.clone();
if let Err(e) = subduction.add_connection(unified_auth).await {
tracing::error!(error = %e, "Failed to add HTTP long-poll connection");
} else {
ephemeral.subscribe_peer(peer_id).await;
notify_peer_connect(keyhive_proto.as_ref(), auth_for_keyhive).await;
}
}
let (mut parts, body) = resp.into_parts();
parts
.headers
.insert(ACCESS_CONTROL_ALLOW_ORIGIN, HeaderValue::from_static("*"));
parts.headers.insert(
ACCESS_CONTROL_ALLOW_METHODS,
HeaderValue::from_static("POST, OPTIONS"),
);
parts.headers.insert(
ACCESS_CONTROL_ALLOW_HEADERS,
HeaderValue::from_static("Content-Type, X-Session-Id"),
);
parts.headers.insert(
hyper::header::ACCESS_CONTROL_EXPOSE_HEADERS,
HeaderValue::from_static("X-Session-Id"),
);
Ok::<_, hyper::Error>(hyper::Response::from_parts(parts, body))
}
});
let builder =
hyper_util::server::conn::auto::Builder::new(hyper_util::rt::TokioExecutor::new());
let conn = builder.serve_connection(io, service);
if let Err(e) = conn.await {
tracing::debug!(addr = %addr, error = %e, "HTTP connection ended");
}
}
async fn notify_peer_connect(
protocol: Option<&CliKeyhiveProtocol>,
conn: Authenticated<CliConn, Sendable>,
) {
let Some(protocol) = protocol else {
return;
};
let adapter = CliConnKeyhiveAdapter::new(conn);
let kh_peer_id = adapter.peer_id();
protocol.add_peer(kh_peer_id, adapter).await;
}
#[allow(clippy::too_many_arguments)]
async fn try_connect_ws<H: CliWireHandler>(
uri: Uri,
subduction: &CliSubduction<H>,
ephemeral: &CliEphemeralHandler,
keyhive_proto: Option<&CliKeyhiveProtocol>,
signer: &MemorySigner,
service_name: &str,
cancel: CancellationToken,
max_message_size: usize,
max_frame_size: usize,
keepalive: KeepAlive,
) -> Result<PeerId, eyre::Error> {
let uri_str = uri.to_string();
tracing::info!(uri = %uri_str, service_name = %service_name, "Connecting to peer via discovery");
let mut ws_config = WebSocketConfig::default();
ws_config.max_message_size = Some(max_message_size);
ws_config.max_frame_size = Some(max_frame_size);
let (ws_stream, _resp) =
async_tungstenite::tokio::connect_async_with_config(uri.clone(), Some(ws_config)).await?;
let audience = Audience::discover(service_name.as_bytes());
let now = TimestampSeconds::now();
let nonce = Nonce::random();
let listen_uri = uri_str.clone();
let sender_uri = uri_str.clone();
let keepalive_uri = uri_str.clone();
let listen_cancel = cancel.clone();
let keepalive_cancel = cancel.clone();
let (authenticated, ()) = handshake::initiate::<future_form::Sendable, _, _, _, _>(
WebSocketHandshake::new(ws_stream),
move |ws_handshake, peer_id| {
let (ws, sender_fut, keepalive_task) = WebSocket::new_with_keepalive(
ws_handshake.into_inner(),
peer_id,
keepalive,
TokioSleeper,
);
let ws_conn = UnifiedWebSocket::Dialed(ws.clone());
let listen_ws = ws.clone();
tokio::spawn(async move {
tokio::select! {
() = listen_cancel.cancelled() => {
tracing::debug!(uri = %listen_uri, "Shutting down listener for peer");
}
result = listen_ws.listen() => {
if let Err(e) = result {
tracing::info!(uri = %listen_uri, error = %e, "WebSocket listener disconnected");
}
}
}
});
let sender_cancel = cancel;
tokio::spawn(async move {
tokio::select! {
() = sender_cancel.cancelled() => {
tracing::debug!(uri = %sender_uri, "Shutting down sender for peer");
}
result = sender_fut => {
if let Err(e) = result {
tracing::info!(uri = %sender_uri, error = %e, "WebSocket sender disconnected");
}
}
}
});
let keepalive_fut = keepalive_task.into_future();
tokio::spawn(async move {
tokio::select! {
() = keepalive_cancel.cancelled() => {
tracing::debug!(uri = %keepalive_uri, "Shutting down keepalive for peer");
}
outcome = keepalive_fut => {
tracing::debug!(uri = %keepalive_uri, ?outcome, "keepalive task for peer exited");
}
}
});
(
MessageTransport::new(UnifiedTransport::WebSocket(ws_conn)),
(),
)
},
signer,
audience,
now,
nonce,
)
.await?;
let remote_id = authenticated.peer_id();
tracing::info!(peer = %remote_id, "Handshake complete: connected to peer");
let auth_for_keyhive = authenticated.clone();
subduction.add_connection(authenticated).await?;
ephemeral.subscribe_peer(remote_id).await;
notify_peer_connect(keyhive_proto, auth_for_keyhive).await;
tracing::info!(uri = %uri_str, "Connected to peer");
Ok(remote_id)
}
#[allow(clippy::too_many_arguments)]
async fn try_connect_iroh<H: CliWireHandler>(
endpoint: &iroh::Endpoint,
addr: EndpointAddr,
subduction: &CliSubduction<H>,
ephemeral: &CliEphemeralHandler,
keyhive_proto: Option<&CliKeyhiveProtocol>,
signer: &MemorySigner,
service_name: &str,
cancel: CancellationToken,
) -> Result<PeerId, eyre::Error> {
let node_id = addr.id;
tracing::info!(node_id = %node_id, service_name = %service_name, "iroh: connecting via discovery");
let audience = Audience::discover(service_name.as_bytes());
let connect_result = subduction_iroh::client::connect(endpoint, addr, signer, audience).await?;
let authenticated = connect_result.authenticated;
let listener_task = connect_result.listener_task;
let sender_task = connect_result.sender_task;
let listener_cancel = cancel.clone();
let sender_cancel = cancel;
tokio::spawn(async move {
tokio::select! {
() = listener_cancel.cancelled() => {
tracing::debug!(node_id = %node_id, "iroh: shutting down listener for peer");
}
result = listener_task => {
if let Err(e) = result {
tracing::info!(node_id = %node_id, error = %e, "iroh: listener disconnected");
}
}
}
});
tokio::spawn(async move {
tokio::select! {
() = sender_cancel.cancelled() => {
tracing::debug!(node_id = %node_id, "iroh: shutting down sender for peer");
}
result = sender_task => {
if let Err(e) = result {
tracing::info!(node_id = %node_id, error = %e, "iroh: sender disconnected");
}
}
}
});
let remote_id = authenticated.peer_id();
let auth = authenticated.map(|c| MessageTransport::new(UnifiedTransport::Iroh(c)));
let auth_for_keyhive = auth.clone();
subduction.add_connection(auth).await?;
ephemeral.subscribe_peer(remote_id).await;
notify_peer_connect(keyhive_proto, auth_for_keyhive).await;
subduction
.full_sync_with_peer(&remote_id, true, CallTimeout::Default)
.await;
tracing::info!(node_id = %node_id, peer = %remote_id, "iroh: added peer");
Ok(remote_id)
}
#[cfg(test)]
#[allow(clippy::expect_used, clippy::panic)]
mod tests {
use clap::Parser;
use super::{AuthMode, ServerArgs};
fn parse(extra: &[&str]) -> ServerArgs {
let mut argv = vec!["server"];
argv.extend_from_slice(extra);
ServerArgs::try_parse_from(argv).expect("args parse")
}
#[test]
fn auth_defaults_to_keyhive() {
let args = parse(&[]);
assert_eq!(args.auth, AuthMode::Keyhive);
assert!(args.auth.keyhive_enabled());
assert!(args.keyhive_cache_refresh);
}
#[test]
fn auth_open_disables_keyhive() {
let args = parse(&["--auth", "open"]);
assert_eq!(args.auth, AuthMode::Open);
assert!(!args.auth.keyhive_enabled());
}
#[test]
fn keyhive_cache_refresh_is_settable() {
let args = parse(&["--keyhive-cache-refresh", "false"]);
assert!(!args.keyhive_cache_refresh);
assert!(args.auth.keyhive_enabled());
}
#[test]
fn auth_rejects_unknown_mode() {
assert!(ServerArgs::try_parse_from(["server", "--auth", "nope"]).is_err());
}
}