#![cfg_attr(test, allow(unused_crate_dependencies))]
use crate::keys::{load_csk, save_csk};
use base64::Engine as _;
use base64::engine::general_purpose::STANDARD_NO_PAD;
use ed25519_dalek::SigningKey;
use ed25519_dalek::VerifyingKey;
use eyre::Report;
use hex::encode as hex_encode;
use ipnet::IpNet;
use quinn::{ClientConfig, Endpoint, ServerConfig};
use rand::rngs::OsRng;
use rcgen::generate_simple_self_signed;
use rustls::{Certificate, PrivateKey, RootCertStore, ServerConfig as TlsServerConfig, ServerName};
use sha2::{Digest, Sha256};
use std::collections::HashMap;
use std::{
fs,
net::ToSocketAddrs,
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
time::Duration,
};
use tokio::{
io::{AsyncBufReadExt, AsyncWriteExt, BufReader},
net::{TcpListener, TcpStream, UnixListener, UnixStream},
sync::{Mutex, broadcast},
time,
};
use tokio_rustls::{TlsAcceptor, TlsConnector};
use tracing::{debug, info};
use volli_agent::{AgentConfig, Protocol, Role};
use volli_core::token::{decode_token, verify_token};
use volli_core::{AliveEntry, CoordAnnounce, DEFAULT_QUIC_PORT, DEFAULT_TCP_PORT, Message};
use volli_transport::{QuicTransport, TcpTransport, Transport};
type AliveTable = Arc<Mutex<HashMap<String, AliveEntry>>>;
type AliveTx = broadcast::Sender<CoordAnnounce>;
async fn update_alive(
table: &AliveTable,
tx: &AliveTx,
profile: &str,
meta: CoordAnnounce,
) -> bool {
let mut map = table.lock().await;
let now = now_millis();
match map.get(&meta.coord_id) {
Some(entry) if entry.meta.ts >= meta.ts => {
map.get_mut(&meta.coord_id).unwrap().last_seen = now;
false
}
_ => {
tracing::info!(id=%meta.coord_id, host=%meta.host, "discovered peer");
let entry = JoinHostEntry {
coord_id: Some(meta.coord_id.clone()),
host: meta.host.clone(),
tcp_port: Some(meta.tcp),
quic_port: Some(meta.quic),
token: None,
cert: meta.tls_cert.clone(),
fingerprint: meta.tls_fp.clone(),
last_ok: Some(now_secs()),
last_fail: None,
};
let _ = add_join_host(profile, entry);
map.insert(
meta.coord_id.clone(),
AliveEntry {
meta: meta.clone(),
last_seen: now,
},
);
let _ = tx.send(meta);
true
}
}
}
fn now_millis() -> u64 {
use std::time::{SystemTime, UNIX_EPOCH};
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_millis() as u64
}
fn now_secs() -> u64 {
use std::time::{SystemTime, UNIX_EPOCH};
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs()
}
async fn sweep_dead(table: AliveTable, profile: String) {
loop {
time::sleep(Duration::from_secs(30)).await;
let now = now_millis();
let mut map = table.lock().await;
let before = map.len();
map.retain(|_id, v| {
let alive = now.saturating_sub(v.last_seen) <= 5 * 60_000;
if !alive {
tracing::info!(id=%v.meta.coord_id, "peer expired");
let _ = remove_join_host(&profile, &v.meta.host);
}
alive
});
if before != map.len() {
tracing::debug!(removed = before - map.len(), "swept dead peers");
}
}
}
pub mod keys;
pub use keys::{
CoordProfileExport, JoinHostEntry, add_join_host, add_join_host_from_token, bootstrap_keypair,
default_secret_dir, delete_profile, export_profile, import_profile, list_profiles,
load_agent_whitelist, load_bind_host, load_bootstrap, load_coord_whitelist, load_join_hosts,
load_profile_host, load_quic_port, load_signing_key, load_tcp_port, load_verifying_key,
profile_exists, remove_join_host, remove_join_host_index, rename_profile, save_agent_whitelist,
save_bind_host, save_bootstrap, save_coord_whitelist, save_join_hosts, save_profile_host,
save_quic_port, save_tcp_port, secret_dir,
};
pub fn cmd_socket_path(profile: &str) -> std::path::PathBuf {
std::env::temp_dir().join(format!("volli-{profile}.sock"))
}
pub struct ServerConfigOpts {
pub advertise_host: String,
pub bind: String,
pub tcp_port: u16,
pub quic_port: u16,
pub cert: Option<String>,
pub key: Option<String>,
pub token: Option<String>,
pub secret_dir: Option<String>,
pub profile: String,
pub max_connections: usize,
pub agent_whitelist: Option<Vec<String>>,
pub coord_whitelist: Option<Vec<String>>,
pub join_hosts: Vec<JoinHostEntry>,
}
impl Default for ServerConfigOpts {
fn default() -> Self {
Self {
advertise_host: "127.0.0.1".into(),
bind: "127.0.0.1".into(),
tcp_port: DEFAULT_TCP_PORT,
quic_port: DEFAULT_QUIC_PORT,
cert: None,
key: None,
token: None,
secret_dir: None,
profile: "default".into(),
max_connections: 1000,
agent_whitelist: None,
coord_whitelist: None,
join_hosts: Vec::new(),
}
}
}
fn init_signing(
secret_dir: &std::path::Path,
show_bootstrap: &mut bool,
) -> Result<
(
Arc<SigningKey>,
bool,
Option<std::path::PathBuf>,
Option<std::path::PathBuf>,
String,
[u8; 32],
),
Report,
> {
let sk_path = secret_dir.join("coord_sk");
let pk_path = secret_dir.join("coord_pk");
if sk_path.exists() && pk_path.exists() {
let key = load_signing_key(Some(secret_dir))?;
let verifying: VerifyingKey = key.verifying_key();
let id = hex_encode(verifying.to_bytes());
let fp: [u8; 32] = Sha256::digest(verifying.to_bytes()).into();
Ok((Arc::new(key), false, Some(sk_path), Some(pk_path), id, fp))
} else {
*show_bootstrap = true;
let key = SigningKey::generate(&mut OsRng);
let verifying: VerifyingKey = key.verifying_key();
let id = hex_encode(verifying.to_bytes());
let fp: [u8; 32] = Sha256::digest(verifying.to_bytes()).into();
Ok((Arc::new(key), true, Some(sk_path), Some(pk_path), id, fp))
}
}
fn init_csk(profile: &str) -> Result<([u8; 32], u32, bool), Report> {
match load_csk(profile)? {
Some(v) => Ok((v.0, v.1, false)),
None => {
let mut k = [0u8; 32];
getrandom::getrandom(&mut k)?;
Ok((k, 1, true))
}
}
}
fn init_cert(
cfg: &ServerConfigOpts,
secret_dir: &std::path::Path,
) -> Result<
(
Vec<Certificate>,
PrivateKey,
String,
Vec<u8>,
Vec<u8>,
Option<std::path::PathBuf>,
Option<std::path::PathBuf>,
bool,
),
Report,
> {
let (cert_fs, key_fs) = if cfg.cert.is_none() && cfg.key.is_none() {
(
Some(secret_dir.join("tls_cert.der")),
Some(secret_dir.join("tls_key.der")),
)
} else {
(
cfg.cert.as_ref().map(std::path::PathBuf::from),
cfg.key.as_ref().map(std::path::PathBuf::from),
)
};
let (c_der, k_der, persist) = if let (Some(c), Some(k)) = (cert_fs.as_ref(), key_fs.as_ref()) {
if c.exists() && k.exists() {
(std::fs::read(c)?, std::fs::read(k)?, false)
} else {
let cert = generate_simple_self_signed(vec!["volli".into()])?;
(
cert.serialize_der()?,
cert.serialize_private_key_der(),
true,
)
}
} else {
let cert = generate_simple_self_signed(vec!["volli".into()])?;
(
cert.serialize_der()?,
cert.serialize_private_key_der(),
false,
)
};
let fp = hex_encode(Sha256::digest(&c_der));
Ok((
vec![Certificate(c_der.clone())],
PrivateKey(k_der.clone()),
fp,
c_der,
k_der,
cert_fs,
key_fs,
persist,
))
}
fn parse_nets(list: &Option<Vec<String>>) -> Arc<Vec<IpNet>> {
Arc::new(
list.clone()
.unwrap_or_default()
.into_iter()
.filter_map(|s| s.parse().ok())
.collect(),
)
}
pub async fn run(
mut cfg: ServerConfigOpts,
mut on_ready: Option<Box<dyn FnOnce() + Send>>,
wait_for_join: bool,
) -> Result<(), Report> {
let mut show_bootstrap = false;
let secret_base = cfg
.secret_dir
.as_deref()
.map(std::path::PathBuf::from)
.unwrap_or_else(default_secret_dir);
let secret_dir: &std::path::Path = &secret_base;
let (signing, persist_keys, sk_path, pk_path, coord_id, pub_fp) =
init_signing(secret_dir, &mut show_bootstrap)?;
let (csk, csk_ver, persist_csk) = init_csk(&cfg.profile)?;
let (cert_chain, key, fingerprint, cert_der, key_der, cert_path, key_path, persist_cert) =
init_cert(&cfg, secret_dir)?;
let quic_endpoint = setup_quic(&cfg.bind, cfg.quic_port, &cert_chain, &key)?;
let quic_port = quic_endpoint.local_addr()?.port();
let tls_acceptor = setup_tls_acceptor(&cert_chain, &key)?;
let tcp_listener = TcpListener::bind((cfg.bind.as_str(), cfg.tcp_port)).await?;
let tcp_port = tcp_listener.local_addr()?.port();
cfg.quic_port = quic_port;
cfg.tcp_port = tcp_port;
let self_meta = CoordAnnounce {
tenant: "self".into(),
cluster: "default".into(),
coord_id: coord_id.clone(),
host: cfg.advertise_host.clone(),
quic: cfg.quic_port,
tcp: cfg.tcp_port,
pub_fp,
csk_ver,
ts: now_millis(),
tls_cert: Some(STANDARD_NO_PAD.encode(&cert_der)),
tls_fp: Some(fingerprint.clone()),
};
let peers: AliveTable = Arc::new(Mutex::new(HashMap::new()));
let (alive_tx, _) = broadcast::channel(16);
{
let mut map = peers.lock().await;
map.insert(
coord_id.clone(),
AliveEntry {
meta: self_meta.clone(),
last_seen: now_millis(),
},
);
}
tokio::spawn(sweep_dead(peers.clone(), cfg.profile.clone()));
let (join_tx, join_rx) = if wait_for_join {
let (tx, rx) = tokio::sync::oneshot::channel();
(Some(tx), Some(rx))
} else {
(None, None)
};
{
let jc = AgentConfig {
role: Role::Coordinator,
..Default::default()
};
let peers_clone = peers.clone();
let tx_clone = alive_tx.clone();
let self_meta_clone = self_meta.clone();
let profile_clone = cfg.profile.clone();
let hosts = cfg.join_hosts.clone();
let join_tx_opt = if wait_for_join { join_tx } else { None };
tokio::spawn(async move {
if let Err(e) = run_coord_mesh(
jc,
hosts,
self_meta_clone,
peers_clone,
tx_clone,
profile_clone,
join_tx_opt,
)
.await
{
tracing::error!("join error: {}", e);
}
});
}
let active_connections = Arc::new(AtomicUsize::new(0));
let agent_nets = parse_nets(&cfg.agent_whitelist);
let coord_nets = parse_nets(&cfg.coord_whitelist);
let token = cfg
.token
.take()
.map(|t| volli_core::token::decode_token(&t).unwrap())
.unwrap_or_else(|| {
volli_core::token::issue_token(&csk, "self", "default", "agent", 86_400).unwrap()
});
let secret = volli_core::BootstrapSecret {
host: cfg.advertise_host.clone(),
quic_port: cfg.quic_port,
tcp_port: cfg.tcp_port,
token: token.clone(),
cert: cert_der.clone(),
};
let secret_encoded = secret.encode()?;
info!(
"Server listening on {} TCP:{} QUIC:{}",
cfg.bind, cfg.tcp_port, cfg.quic_port
);
if show_bootstrap {
println!("Agent bootstrap command:\n volli agent --join {secret_encoded}");
println!(
"Hint: run 'volli --profile {} admin coord-token' to get a coordinator join command",
cfg.profile
);
}
let socket_path = cmd_socket_path(&cfg.profile);
info!("command socket path={} ", socket_path.display());
tokio::spawn(command_socket(
socket_path.clone(),
csk,
cfg.advertise_host.clone(),
cfg.quic_port,
cfg.tcp_port,
cert_der.clone(),
));
if wait_for_join {
if let Some(rx) = join_rx {
let prof = cfg.profile.clone();
let sk_p = sk_path.clone();
let pk_p = pk_path.clone();
let cert_p = cert_path.clone();
let key_p = key_path.clone();
let sd = secret_dir.to_path_buf();
let signing_cl = signing.clone();
let cert_der_cl = cert_der.clone();
let key_der_cl = key_der.clone();
tokio::spawn(async move {
if rx.await.is_ok() {
if persist_csk {
save_csk(&prof, &csk, csk_ver).ok();
}
if persist_keys {
if let (Some(sk), Some(pk)) = (sk_p.as_ref(), pk_p.as_ref()) {
fs::create_dir_all(&sd).ok();
fs::write(sk, hex::encode(signing_cl.to_bytes())).ok();
let ver_bytes = signing_cl.verifying_key().to_bytes();
fs::write(pk, hex::encode(ver_bytes)).ok();
}
}
if persist_cert {
if let (Some(c), Some(k)) = (cert_p.as_ref(), key_p.as_ref()) {
if let Some(parent) = c.parent() {
fs::create_dir_all(parent).ok();
}
fs::write(c, &cert_der_cl).ok();
fs::write(k, &key_der_cl).ok();
}
}
if let Some(cb) = on_ready {
cb();
}
}
});
}
} else {
if persist_csk {
save_csk(&cfg.profile, &csk, csk_ver).ok();
}
if persist_keys {
if let (Some(sk), Some(pk)) = (sk_path.as_ref(), pk_path.as_ref()) {
fs::create_dir_all(secret_dir)?;
fs::write(sk, hex::encode(signing.to_bytes()))?;
let ver_bytes = signing.verifying_key().to_bytes();
fs::write(pk, hex::encode(ver_bytes))?;
info!("Generated coordinator keypair at {}", secret_dir.display());
}
}
if persist_cert {
if let (Some(c), Some(k)) = (cert_path.as_ref(), key_path.as_ref()) {
if let Some(parent) = c.parent() {
fs::create_dir_all(parent)?;
}
fs::write(c, &cert_der)?;
fs::write(k, &key_der)?;
}
}
if let Some(cb) = on_ready.take() {
cb();
}
}
loop {
tokio::select! {
Ok((stream, addr)) = tcp_listener.accept() => {
if active_connections.load(Ordering::SeqCst) >= cfg.max_connections {
info!(peer=%addr, "connection limit reached");
drop(stream);
continue;
}
active_connections.fetch_add(1, Ordering::SeqCst);
info!(peer=%addr, "Accepted TCP connection");
let fp = fingerprint.clone();
let id = coord_id.clone();
let signing_clone = signing.clone();
let acceptor = tls_acceptor.clone();
let counter = active_connections.clone();
let agent_nets = agent_nets.clone();
let coord_nets = coord_nets.clone();
let peers_clone = peers.clone();
let self_meta_clone = self_meta.clone();
let profile_clone = cfg.profile.clone();
let tx = alive_tx.clone();
tokio::spawn(async move {
if let Ok(tls) = acceptor.accept(stream).await {
let proto = tls.get_ref().1.alpn_protocol().map(|p| String::from_utf8_lossy(p).to_string());
match proto.as_deref() {
Some("volli/agent") => {
handle_agent_client(
Box::new(TcpTransport::new(tls)),
signing_clone,
csk,
fp,
id,
addr,
agent_nets,
)
.await
.ok();
}
Some("volli/coord") => {
let peer_fp = tls
.get_ref()
.1
.peer_certificates()
.and_then(|c| c.first().cloned());
let fp = peer_fp.as_ref().map(|c| hex::encode(Sha256::digest(&c.0)));
let cert = peer_fp.map(|c| c.0);
handle_coord_client(
Box::new(TcpTransport::new(tls)),
csk,
self_meta_clone,
peers_clone,
tx,
profile_clone,
addr,
coord_nets,
cert,
fp,
)
.await
.ok();
}
_ => {}
}
}
counter.fetch_sub(1, Ordering::SeqCst);
});
}
Some(connecting) = quic_endpoint.accept() => {
let addr = connecting.remote_address();
if active_connections.load(Ordering::SeqCst) >= cfg.max_connections {
info!(peer=%addr, "connection limit reached");
drop(connecting);
continue;
}
active_connections.fetch_add(1, Ordering::SeqCst);
info!(peer=%addr, "Accepted QUIC connection");
let fp = fingerprint.clone();
let id = coord_id.clone();
let signing_clone = signing.clone();
let counter = active_connections.clone();
let agent_nets = agent_nets.clone();
let coord_nets = coord_nets.clone();
let peers_clone = peers.clone();
let self_meta_clone = self_meta.clone();
let profile_clone = cfg.profile.clone();
let tx = alive_tx.clone();
tokio::spawn(async move {
if let Ok(conn) = connecting.await {
debug!("Opening QUIC connection");
let protocol = conn
.handshake_data()
.and_then(|d| d.downcast::<quinn::crypto::rustls::HandshakeData>().ok())
.and_then(|hd| hd.protocol.clone());
if let Ok((send, recv)) = conn.accept_bi().await {
match protocol.as_deref() {
Some(b"volli/agent") => {
info!("Accepted agent connection");
handle_agent_client(
Box::new(QuicTransport::new(send, recv)),
signing_clone,
csk,
fp,
id,
addr,
agent_nets,
)
.await
.ok();
}
Some(b"volli/coord") => {
info!("Accepted coordinator connection");
let cert_opt = conn
.peer_identity()
.and_then(|i| i.downcast::<Vec<Certificate>>().ok())
.and_then(|mut v| v.pop());
let fp = cert_opt.as_ref().map(|c| hex::encode(Sha256::digest(&c.0)));
let cert = cert_opt.map(|c| c.0);
handle_coord_client(
Box::new(QuicTransport::new(send, recv)),
csk,
self_meta_clone,
peers_clone,
tx,
profile_clone,
addr,
coord_nets,
cert,
fp,
)
.await
.ok();
}
_ => {}
}
}
}
counter.fetch_sub(1, Ordering::SeqCst);
});
}
}
}
}
async fn handle_agent_client(
mut transport: Box<dyn Transport>,
signing: Arc<SigningKey>,
csk: [u8; 32],
fingerprint: String,
coord_id: String,
peer: std::net::SocketAddr,
whitelist: Arc<Vec<IpNet>>,
) -> Result<(), Report> {
if !whitelist.is_empty() && !whitelist.iter().any(|n| n.contains(&peer.ip())) {
info!(%peer, "agent connection rejected by whitelist");
return Ok(());
}
let peer_str = peer.to_string();
match transport.recv().await? {
Some(Message::Auth { token: recv_token }) => {
let token = decode_token(&recv_token)?;
if let Err(e) = verify_token(&token, &csk) {
transport.send(&Message::AuthErr).await.ok();
return Err(e);
}
if fingerprint.is_empty() {
}
if token.payload.agent_id.is_empty() {
transport.send(&Message::AuthErr).await.ok();
return Err(eyre::eyre!("invalid agent"));
}
transport.send(&Message::AuthOk).await?;
info!(target: "connection", %peer, "agent authenticated");
}
_ => {
transport.send(&Message::AuthErr).await.ok();
return Ok(());
}
}
let mut nonce = [0u8; 32];
getrandom::getrandom(&mut nonce)?;
let sig = volli_core::handshake::sign_nonce(&signing, &nonce);
let hello = Message::Hello {
coord_id: coord_id.clone(),
nonce,
sig: sig.clone(),
};
transport.send(&hello).await?;
match transport.recv().await? {
Some(Message::Welcome {
coord_id: cid,
nonce: rnonce,
sig: rsig,
}) => {
if cid != coord_id || rnonce != nonce || rsig != sig {
return Err(eyre::eyre!("handshake mismatch"));
}
}
_ => return Err(eyre::eyre!("handshake failed")),
}
info!(target: "connection", %peer, "coordinator authenticated");
let mut interval = time::interval(Duration::from_secs(5));
tracing::debug!(target: "connection", %peer_str, "sending ping to agent");
transport.send(&Message::Ping).await?;
loop {
tokio::select! {
msg = transport.recv() => {
if let Some(Message::Pong { mac }) = msg? {
info!(target: "connection", %peer_str, %mac, "received pong from agent");
}
}
_ = interval.tick() => {
tracing::debug!(target: "connection", %peer_str, "sending ping to agent");
transport.send(&Message::Ping).await?;
}
}
}
}
#[allow(clippy::too_many_arguments)]
async fn handle_coord_client(
mut transport: Box<dyn Transport>,
csk: [u8; 32],
self_meta: CoordAnnounce,
peers: AliveTable,
alive_tx: AliveTx,
profile: String,
peer: std::net::SocketAddr,
whitelist: Arc<Vec<IpNet>>,
peer_cert: Option<Vec<u8>>,
peer_fp: Option<String>,
) -> Result<(), Report> {
if !whitelist.is_empty() && !whitelist.iter().any(|n| n.contains(&peer.ip())) {
info!(%peer, "coord connection rejected by whitelist");
return Ok(());
}
let stored_token: Option<String>;
match transport.recv().await? {
Some(Message::Auth { token: recv_token }) => {
info!(%peer, "coord connection received auth token");
let token = decode_token(&recv_token)?;
if let Err(e) = verify_token(&token, &csk) {
info!(%peer, "coord connection rejected by auth");
transport.send(&Message::AuthErr).await.ok();
return Err(e);
}
stored_token = Some(recv_token);
transport.send(&Message::AuthOk).await?;
}
_ => {
info!(%peer, "coord connection rejected by auth");
transport.send(&Message::AuthErr).await.ok();
return Ok(());
}
}
let mut interval = time::interval(Duration::from_secs(5));
let mut rx = alive_tx.subscribe();
let mut pending: Option<CoordAnnounce> = None;
let mut saved = false;
tracing::debug!(target: "connection", %peer, "sending heartbeat to coordinator");
transport
.send(&Message::CoordAnnounce {
meta: self_meta.clone(),
diff: Box::new(None),
})
.await?;
loop {
tokio::select! {
msg = transport.recv() => {
match msg? {
Some(Message::CoordAnnounce { meta, diff }) => {
tracing::debug!(target: "connection", %peer, id=%meta.coord_id, "received heartbeat");
if !saved {
let entry = JoinHostEntry {
coord_id: Some(meta.coord_id.clone()),
host: meta.host.clone(),
tcp_port: Some(meta.tcp),
quic_port: Some(meta.quic),
token: stored_token.clone(),
cert: meta
.tls_cert
.clone()
.or_else(|| peer_cert.as_ref().map(|c| STANDARD_NO_PAD.encode(c))),
fingerprint: meta.tls_fp.clone().or(peer_fp.clone()),
last_ok: Some(now_secs()),
last_fail: None,
};
add_join_host(&profile, entry).ok();
saved = true;
}
update_alive(&peers, &alive_tx, &profile, meta).await;
if let Some(d) = *diff { update_alive(&peers, &alive_tx, &profile, d).await; }
}
Some(Message::Ping) => {
transport.send(&Message::Pong { mac: String::new() }).await.ok();
}
Some(Message::Pong { .. }) => {}
_ => {}
}
}
_ = interval.tick() => {
if pending.is_none() {
match rx.try_recv() {
Ok(upd) => pending = Some(upd),
Err(broadcast::error::TryRecvError::Closed) => return Ok(()),
_ => {}
}
}
tracing::debug!(target: "connection", %peer, "sending heartbeat to coordinator");
transport
.send(&Message::CoordAnnounce { meta: self_meta.clone(), diff: Box::new(pending.take()) })
.await?;
}
}
}
}
async fn handle_coord_peer(
mut transport: Box<dyn Transport>,
cfg: &AgentConfig,
peer: String,
self_meta: CoordAnnounce,
peers: AliveTable,
alive_tx: AliveTx,
profile: String,
join_notify: Option<tokio::sync::oneshot::Sender<()>>,
) -> Result<(), Report> {
transport
.send(&Message::Auth {
token: cfg.token.clone(),
})
.await?;
match transport.recv().await? {
Some(Message::AuthOk) => {
tracing::info!(target: "connection", %peer, "coordinator authenticated");
if let Some(tx) = join_notify {
let _ = tx.send(());
}
}
_ => return Err(eyre::eyre!("authentication failed")),
}
let mut interval = time::interval(Duration::from_secs(5));
let mut rx = alive_tx.subscribe();
let mut pending: Option<CoordAnnounce> = None;
tracing::debug!(target: "connection", %peer, "sending heartbeat to coordinator");
transport
.send(&Message::CoordAnnounce {
meta: self_meta.clone(),
diff: Box::new(None),
})
.await?;
loop {
tokio::select! {
msg = transport.recv() => {
match msg? {
Some(Message::CoordAnnounce { meta, diff }) => {
tracing::debug!(target: "connection", %peer, id=%meta.coord_id, "received heartbeat");
update_alive(&peers, &alive_tx, &profile, meta).await;
if let Some(d) = *diff { update_alive(&peers, &alive_tx, &profile, d).await; }
}
Some(Message::Ping) => {
tracing::debug!(target: "connection", %peer, "received ping");
transport.send(&Message::Pong { mac: String::new() }).await.ok();
}
Some(Message::Pong { .. }) => {
tracing::debug!(target: "connection", %peer, "received pong");
}
_ => {}
}
}
_ = interval.tick() => {
if pending.is_none() {
match rx.try_recv() {
Ok(upd) => pending = Some(upd),
Err(broadcast::error::TryRecvError::Closed) => {
tracing::debug!(target: "connection", %peer, "connection closed");
return Ok(())
},
_ => {}
}
}
tracing::debug!(target: "connection", %peer, "sending heartbeat to coordinator");
transport
.send(&Message::CoordAnnounce { meta: self_meta.clone(), diff: Box::new(pending.take()) })
.await?;
}
}
}
}
fn configure_client(cert: &[u8], alpn: &str) -> Result<ClientConfig, Report> {
let mut roots = rustls::RootCertStore::empty();
roots.add(&Certificate(cert.to_vec()))?;
let mut crypto = rustls::ClientConfig::builder()
.with_safe_defaults()
.with_root_certificates(roots)
.with_no_client_auth();
crypto.alpn_protocols = vec![alpn.as_bytes().to_vec()];
Ok(ClientConfig::new(Arc::new(crypto)))
}
async fn connect_coord_tcp(cfg: &AgentConfig) -> Result<(Box<dyn Transport>, String), Report> {
let addr = format!("{}:{}", cfg.host, cfg.tcp_port);
let mut addrs = addr.to_socket_addrs()?;
let addr = addrs
.find(|a| a.is_ipv4())
.or_else(|| addrs.next())
.ok_or_else(|| eyre::eyre!("invalid addr"))?;
let stream = TcpStream::connect(addr).await?;
let alpn = "volli/coord";
let mut roots = RootCertStore::empty();
roots.add(&Certificate(cfg.cert.clone()))?;
let mut root = rustls::ClientConfig::builder()
.with_safe_defaults()
.with_root_certificates(roots)
.with_no_client_auth();
root.alpn_protocols = vec![alpn.as_bytes().to_vec()];
let connector = TlsConnector::from(Arc::new(root));
let tls = connector
.connect(ServerName::try_from("volli")?, stream)
.await?;
if let Some(certs) = tls.get_ref().1.peer_certificates() {
if let Some(cert) = certs.first() {
let hash = Sha256::digest(&cert.0);
if hex::encode(hash) != cfg.fingerprint {
return Err(eyre::eyre!("server fingerprint mismatch"));
}
}
}
let peer = tls.get_ref().0.peer_addr()?.to_string();
Ok((Box::new(TcpTransport::new(tls)), peer))
}
async fn connect_coord_quic(cfg: &AgentConfig) -> Result<(Box<dyn Transport>, String), Report> {
let addr = format!("{}:{}", cfg.host, cfg.quic_port);
let mut addrs = addr.to_socket_addrs()?;
let addr = addrs
.find(|a| a.is_ipv4())
.or_else(|| addrs.next())
.ok_or_else(|| eyre::eyre!("invalid addr"))?;
let mut endpoint = Endpoint::client("0.0.0.0:0".parse()?)?;
let quinn_cfg = configure_client(&cfg.cert, "volli/coord")?;
endpoint.set_default_client_config(quinn_cfg);
let connection = endpoint.connect(addr, "volli")?.await?;
if let Some(identity) = connection.peer_identity() {
if let Ok(certs) = identity.downcast::<Vec<Certificate>>() {
if let Some(cert) = certs.first() {
let hash = Sha256::digest(&cert.0);
if hex::encode(hash) != cfg.fingerprint {
return Err(eyre::eyre!("server fingerprint mismatch"));
}
}
}
}
let peer = connection.remote_address().to_string();
let (send, recv) = connection.open_bi().await?;
Ok((Box::new(QuicTransport::new(send, recv)), peer))
}
async fn run_coord_mesh(
mut cfg: AgentConfig,
mut hosts: Vec<JoinHostEntry>,
self_meta: CoordAnnounce,
peers: AliveTable,
alive_tx: AliveTx,
profile: String,
mut join_notify: Option<tokio::sync::oneshot::Sender<()>>,
) -> Result<(), Report> {
let proto_pref = cfg.protocol.take();
if hosts.is_empty() {
return Ok(());
}
let mut idx = 0usize;
let mut backoff = 1u64;
loop {
let entry = hosts.get(idx).cloned().unwrap();
if let Some(ref id) = entry.coord_id {
if self_meta.coord_id >= *id {
idx = (idx + 1) % hosts.len();
tokio::time::sleep(Duration::from_secs(backoff)).await;
info!("Skipping host with lower ID");
continue;
}
}
if entry.token.is_none() || entry.cert.is_none() || entry.fingerprint.is_none() {
idx = (idx + 1) % hosts.len();
tokio::time::sleep(Duration::from_secs(backoff)).await;
info!("Skipping host with missing metadata");
continue;
}
cfg.host = entry.host.clone();
if let Some(p) = entry.tcp_port {
cfg.tcp_port = p;
}
if let Some(p) = entry.quic_port {
cfg.quic_port = p;
}
cfg.token = entry.token.clone().unwrap();
let cert_bytes = match STANDARD_NO_PAD.decode(entry.cert.as_ref().unwrap().as_bytes()) {
Ok(c) => c,
Err(e) => {
tracing::warn!(host=%entry.host, "invalid stored cert: {e}");
hosts.remove(idx);
save_join_hosts(&profile, &hosts).ok();
if hosts.is_empty() {
info!("No hosts left to join");
return Ok(());
}
idx %= hosts.len();
tokio::time::sleep(Duration::from_secs(backoff)).await;
info!("Skipping host with invalid cert");
continue;
}
};
cfg.cert = cert_bytes;
cfg.fingerprint = entry.fingerprint.clone().unwrap();
info!("Connecting to host {:?}", cfg);
let res = match proto_pref.as_ref().unwrap_or(&Protocol::Quic) {
Protocol::Quic => match connect_coord_quic(&cfg).await {
Ok((tr, peer)) => {
info!("Connected to host over QUIC {:?}", cfg);
hosts[idx].last_ok = Some(now_secs());
save_join_hosts(&profile, &hosts).ok();
handle_coord_peer(
tr,
&cfg,
peer,
self_meta.clone(),
peers.clone(),
alive_tx.clone(),
profile.clone(),
join_notify.take(),
)
.await
}
Err(e) => {
tracing::warn!("quic connect error: {}", e);
match connect_coord_tcp(&cfg).await {
Ok((tr, peer)) => {
hosts[idx].last_ok = Some(now_secs());
save_join_hosts(&profile, &hosts).ok();
handle_coord_peer(
tr,
&cfg,
peer,
self_meta.clone(),
peers.clone(),
alive_tx.clone(),
profile.clone(),
join_notify.take(),
)
.await
}
Err(e) => Err(e),
}
}
},
Protocol::Tcp => match connect_coord_tcp(&cfg).await {
Ok((tr, peer)) => {
info!("Connected to host over TCP {:?}", cfg);
hosts[idx].last_ok = Some(now_secs());
save_join_hosts(&profile, &hosts).ok();
handle_coord_peer(
tr,
&cfg,
peer,
self_meta.clone(),
peers.clone(),
alive_tx.clone(),
profile.clone(),
join_notify.take(),
)
.await
}
Err(e) => Err(e),
},
};
match res {
Ok(_) => {
backoff = 1;
idx = 0;
}
Err(e) => {
tracing::error!("coord connection error: {}", e);
hosts[idx].last_fail = Some(now_secs());
save_join_hosts(&profile, &hosts).ok();
backoff = (backoff * 2).min(32);
idx = (idx + 1) % hosts.len();
}
}
info!("Sleeping for {} seconds", backoff);
tokio::time::sleep(Duration::from_secs(backoff)).await;
}
}
fn setup_quic(
host: &str,
port: u16,
certs: &[Certificate],
key: &PrivateKey,
) -> Result<Endpoint, Report> {
let mut tls = TlsServerConfig::builder()
.with_safe_defaults()
.with_no_client_auth()
.with_single_cert(certs.to_vec(), key.clone())?;
tls.alpn_protocols = vec![b"volli/agent".to_vec(), b"volli/coord".to_vec()];
let mut server_config = ServerConfig::with_crypto(Arc::new(tls));
let mut transport = quinn::TransportConfig::default();
transport.max_idle_timeout(Some(Duration::from_secs(5).try_into()?));
server_config.transport = Arc::new(transport);
server_config.retry_token_lifetime(Duration::from_millis(1000));
let addr = format!("{host}:{port}").to_socket_addrs()?.next().unwrap();
let endpoint = Endpoint::server(server_config, addr)?;
Ok(endpoint)
}
pub fn load_or_generate_cert(
cert_path: Option<&str>,
key_path: Option<&str>,
secret_dir: Option<&std::path::Path>,
) -> Result<(Vec<Certificate>, PrivateKey, String), Report> {
let (cert_path_fs, key_path_fs) = if cert_path.is_none() && key_path.is_none() {
let dir = secret_dir
.map(std::path::PathBuf::from)
.unwrap_or_else(default_secret_dir);
(
Some(dir.join("tls_cert.der")),
Some(dir.join("tls_key.der")),
)
} else {
(
cert_path.map(std::path::PathBuf::from),
key_path.map(std::path::PathBuf::from),
)
};
let (cert_der, key_der) =
if let (Some(c), Some(k)) = (cert_path_fs.as_ref(), key_path_fs.as_ref()) {
if c.exists() && k.exists() {
(fs::read(c)?, fs::read(k)?)
} else {
fs::create_dir_all(c.parent().unwrap())?;
let cert = generate_simple_self_signed(vec!["volli".into()])?;
let cert_der = cert.serialize_der()?;
let key_der = cert.serialize_private_key_der();
fs::write(c, &cert_der)?;
fs::write(k, &key_der)?;
(cert_der, key_der)
}
} else {
let cert = generate_simple_self_signed(vec!["volli".into()])?;
(cert.serialize_der()?, cert.serialize_private_key_der())
};
let fingerprint = hex_encode(Sha256::digest(&cert_der));
Ok((
vec![Certificate(cert_der)],
PrivateKey(key_der),
fingerprint,
))
}
fn setup_tls_acceptor(certs: &[Certificate], key: &PrivateKey) -> Result<TlsAcceptor, Report> {
let mut config = TlsServerConfig::builder()
.with_safe_defaults()
.with_no_client_auth()
.with_single_cert(certs.to_vec(), key.clone())?;
config.alpn_protocols = vec![b"volli/agent".to_vec(), b"volli/coord".to_vec()];
Ok(TlsAcceptor::from(Arc::new(config)))
}
fn build_secret(
csk: &[u8; 32],
host: &str,
quic_port: u16,
tcp_port: u16,
cert_der: &[u8],
) -> Result<String, Report> {
let token = volli_core::token::issue_token(csk, "self", "default", "agent", 86_400)?;
let secret = volli_core::BootstrapSecret {
host: host.to_string(),
quic_port,
tcp_port,
token,
cert: cert_der.to_vec(),
};
secret.encode()
}
pub async fn command_socket(
socket_path: std::path::PathBuf,
csk: [u8; 32],
host: String,
quic: u16,
tcp: u16,
cert_der: Vec<u8>,
) -> Result<(), Report> {
let _ = std::fs::remove_file(&socket_path);
let listener = UnixListener::bind(&socket_path)?;
loop {
let (stream, _) = listener.accept().await?;
let host = host.clone();
let cert = cert_der.clone();
tokio::spawn(async move {
handle_cmd(stream, csk, host, quic, tcp, cert).await.ok();
});
}
}
#[allow(clippy::too_many_arguments)]
async fn handle_cmd(
stream: UnixStream,
csk: [u8; 32],
host: String,
quic: u16,
tcp: u16,
cert: Vec<u8>,
) -> Result<(), Report> {
let mut reader = BufReader::new(stream);
let mut line = String::new();
reader.read_line(&mut line).await?;
let cmd = line.trim();
let secret = build_secret(&csk, &host, quic, tcp, &cert)?;
let resp = match cmd {
"agent_token" => format!("volli agent --join {secret}\n"),
"coord_token" => format!("volli serve --join {secret}\n"),
_ => "unknown\n".to_string(),
};
reader.get_mut().write_all(resp.as_bytes()).await?;
Ok(())
}