Skip to main content

volli_server/
lib.rs

1#![cfg_attr(test, allow(unused_crate_dependencies))]
2use crate::keys::{load_csk, save_csk};
3use base64::Engine as _;
4use base64::engine::general_purpose::STANDARD_NO_PAD;
5use ed25519_dalek::SigningKey;
6use ed25519_dalek::VerifyingKey;
7use eyre::Report;
8use hex::encode as hex_encode;
9use ipnet::IpNet;
10use quinn::{ClientConfig, Endpoint, ServerConfig};
11use rand::rngs::OsRng;
12use rcgen::generate_simple_self_signed;
13use rustls::{Certificate, PrivateKey, RootCertStore, ServerConfig as TlsServerConfig, ServerName};
14use sha2::{Digest, Sha256};
15use std::collections::HashMap;
16use std::{
17    fs,
18    net::ToSocketAddrs,
19    sync::{
20        Arc,
21        atomic::{AtomicUsize, Ordering},
22    },
23    time::Duration,
24};
25use tokio::{
26    io::{AsyncBufReadExt, AsyncWriteExt, BufReader},
27    net::{TcpListener, TcpStream, UnixListener, UnixStream},
28    sync::{Mutex, broadcast},
29    time,
30};
31use tokio_rustls::{TlsAcceptor, TlsConnector};
32use tracing::{debug, info};
33use volli_agent::{AgentConfig, Protocol, Role};
34use volli_core::token::{decode_token, verify_token};
35use volli_core::{AliveEntry, CoordAnnounce, DEFAULT_QUIC_PORT, DEFAULT_TCP_PORT, Message};
36use volli_transport::{QuicTransport, TcpTransport, Transport};
37
38type AliveTable = Arc<Mutex<HashMap<String, AliveEntry>>>;
39type AliveTx = broadcast::Sender<CoordAnnounce>;
40
41async fn update_alive(
42    table: &AliveTable,
43    tx: &AliveTx,
44    profile: &str,
45    meta: CoordAnnounce,
46) -> bool {
47    let mut map = table.lock().await;
48    let now = now_millis();
49    match map.get(&meta.coord_id) {
50        Some(entry) if entry.meta.ts >= meta.ts => {
51            map.get_mut(&meta.coord_id).unwrap().last_seen = now;
52            false
53        }
54        _ => {
55            tracing::info!(id=%meta.coord_id, host=%meta.host, "discovered peer");
56            let entry = JoinHostEntry {
57                coord_id: Some(meta.coord_id.clone()),
58                host: meta.host.clone(),
59                tcp_port: Some(meta.tcp),
60                quic_port: Some(meta.quic),
61                token: None,
62                cert: meta.tls_cert.clone(),
63                fingerprint: meta.tls_fp.clone(),
64                last_ok: Some(now_secs()),
65                last_fail: None,
66            };
67            let _ = add_join_host(profile, entry);
68            map.insert(
69                meta.coord_id.clone(),
70                AliveEntry {
71                    meta: meta.clone(),
72                    last_seen: now,
73                },
74            );
75            let _ = tx.send(meta);
76            true
77        }
78    }
79}
80
81fn now_millis() -> u64 {
82    use std::time::{SystemTime, UNIX_EPOCH};
83    SystemTime::now()
84        .duration_since(UNIX_EPOCH)
85        .unwrap()
86        .as_millis() as u64
87}
88
89fn now_secs() -> u64 {
90    use std::time::{SystemTime, UNIX_EPOCH};
91    SystemTime::now()
92        .duration_since(UNIX_EPOCH)
93        .unwrap()
94        .as_secs()
95}
96
97async fn sweep_dead(table: AliveTable, profile: String) {
98    loop {
99        time::sleep(Duration::from_secs(30)).await;
100        let now = now_millis();
101        let mut map = table.lock().await;
102        let before = map.len();
103        map.retain(|_id, v| {
104            let alive = now.saturating_sub(v.last_seen) <= 5 * 60_000;
105            if !alive {
106                tracing::info!(id=%v.meta.coord_id, "peer expired");
107                let _ = remove_join_host(&profile, &v.meta.host);
108            }
109            alive
110        });
111        if before != map.len() {
112            tracing::debug!(removed = before - map.len(), "swept dead peers");
113        }
114    }
115}
116
117pub mod keys;
118pub use keys::{
119    CoordProfileExport, JoinHostEntry, add_join_host, add_join_host_from_token, bootstrap_keypair,
120    default_secret_dir, delete_profile, export_profile, import_profile, list_profiles,
121    load_agent_whitelist, load_bind_host, load_bootstrap, load_coord_whitelist, load_join_hosts,
122    load_profile_host, load_quic_port, load_signing_key, load_tcp_port, load_verifying_key,
123    profile_exists, remove_join_host, remove_join_host_index, rename_profile, save_agent_whitelist,
124    save_bind_host, save_bootstrap, save_coord_whitelist, save_join_hosts, save_profile_host,
125    save_quic_port, save_tcp_port, secret_dir,
126};
127
128pub fn cmd_socket_path(profile: &str) -> std::path::PathBuf {
129    std::env::temp_dir().join(format!("volli-{profile}.sock"))
130}
131
132pub struct ServerConfigOpts {
133    pub advertise_host: String,
134    pub bind: String,
135    pub tcp_port: u16,
136    pub quic_port: u16,
137    pub cert: Option<String>,
138    pub key: Option<String>,
139    pub token: Option<String>,
140    pub secret_dir: Option<String>,
141    pub profile: String,
142    pub max_connections: usize,
143    pub agent_whitelist: Option<Vec<String>>,
144    pub coord_whitelist: Option<Vec<String>>,
145    pub join_hosts: Vec<JoinHostEntry>,
146}
147
148impl Default for ServerConfigOpts {
149    fn default() -> Self {
150        Self {
151            advertise_host: "127.0.0.1".into(),
152            bind: "127.0.0.1".into(),
153            tcp_port: DEFAULT_TCP_PORT,
154            quic_port: DEFAULT_QUIC_PORT,
155            cert: None,
156            key: None,
157            token: None,
158            secret_dir: None,
159            profile: "default".into(),
160            max_connections: 1000,
161            agent_whitelist: None,
162            coord_whitelist: None,
163            join_hosts: Vec::new(),
164        }
165    }
166}
167fn init_signing(
168    secret_dir: &std::path::Path,
169    show_bootstrap: &mut bool,
170) -> Result<
171    (
172        Arc<SigningKey>,
173        bool,
174        Option<std::path::PathBuf>,
175        Option<std::path::PathBuf>,
176        String,
177        [u8; 32],
178    ),
179    Report,
180> {
181    let sk_path = secret_dir.join("coord_sk");
182    let pk_path = secret_dir.join("coord_pk");
183    if sk_path.exists() && pk_path.exists() {
184        let key = load_signing_key(Some(secret_dir))?;
185        let verifying: VerifyingKey = key.verifying_key();
186        let id = hex_encode(verifying.to_bytes());
187        let fp: [u8; 32] = Sha256::digest(verifying.to_bytes()).into();
188        Ok((Arc::new(key), false, Some(sk_path), Some(pk_path), id, fp))
189    } else {
190        *show_bootstrap = true;
191        let key = SigningKey::generate(&mut OsRng);
192        let verifying: VerifyingKey = key.verifying_key();
193        let id = hex_encode(verifying.to_bytes());
194        let fp: [u8; 32] = Sha256::digest(verifying.to_bytes()).into();
195        Ok((Arc::new(key), true, Some(sk_path), Some(pk_path), id, fp))
196    }
197}
198
199fn init_csk(profile: &str) -> Result<([u8; 32], u32, bool), Report> {
200    match load_csk(profile)? {
201        Some(v) => Ok((v.0, v.1, false)),
202        None => {
203            let mut k = [0u8; 32];
204            getrandom::getrandom(&mut k)?;
205            Ok((k, 1, true))
206        }
207    }
208}
209
210fn init_cert(
211    cfg: &ServerConfigOpts,
212    secret_dir: &std::path::Path,
213) -> Result<
214    (
215        Vec<Certificate>,
216        PrivateKey,
217        String,
218        Vec<u8>,
219        Vec<u8>,
220        Option<std::path::PathBuf>,
221        Option<std::path::PathBuf>,
222        bool,
223    ),
224    Report,
225> {
226    let (cert_fs, key_fs) = if cfg.cert.is_none() && cfg.key.is_none() {
227        (
228            Some(secret_dir.join("tls_cert.der")),
229            Some(secret_dir.join("tls_key.der")),
230        )
231    } else {
232        (
233            cfg.cert.as_ref().map(std::path::PathBuf::from),
234            cfg.key.as_ref().map(std::path::PathBuf::from),
235        )
236    };
237    let (c_der, k_der, persist) = if let (Some(c), Some(k)) = (cert_fs.as_ref(), key_fs.as_ref()) {
238        if c.exists() && k.exists() {
239            (std::fs::read(c)?, std::fs::read(k)?, false)
240        } else {
241            let cert = generate_simple_self_signed(vec!["volli".into()])?;
242            (
243                cert.serialize_der()?,
244                cert.serialize_private_key_der(),
245                true,
246            )
247        }
248    } else {
249        let cert = generate_simple_self_signed(vec!["volli".into()])?;
250        (
251            cert.serialize_der()?,
252            cert.serialize_private_key_der(),
253            false,
254        )
255    };
256    let fp = hex_encode(Sha256::digest(&c_der));
257    Ok((
258        vec![Certificate(c_der.clone())],
259        PrivateKey(k_der.clone()),
260        fp,
261        c_der,
262        k_der,
263        cert_fs,
264        key_fs,
265        persist,
266    ))
267}
268
269fn parse_nets(list: &Option<Vec<String>>) -> Arc<Vec<IpNet>> {
270    Arc::new(
271        list.clone()
272            .unwrap_or_default()
273            .into_iter()
274            .filter_map(|s| s.parse().ok())
275            .collect(),
276    )
277}
278
279pub async fn run(
280    mut cfg: ServerConfigOpts,
281    mut on_ready: Option<Box<dyn FnOnce() + Send>>,
282    wait_for_join: bool,
283) -> Result<(), Report> {
284    let mut show_bootstrap = false;
285    let secret_base = cfg
286        .secret_dir
287        .as_deref()
288        .map(std::path::PathBuf::from)
289        .unwrap_or_else(default_secret_dir);
290    let secret_dir: &std::path::Path = &secret_base;
291    let (signing, persist_keys, sk_path, pk_path, coord_id, pub_fp) =
292        init_signing(secret_dir, &mut show_bootstrap)?;
293    let (csk, csk_ver, persist_csk) = init_csk(&cfg.profile)?;
294    let (cert_chain, key, fingerprint, cert_der, key_der, cert_path, key_path, persist_cert) =
295        init_cert(&cfg, secret_dir)?;
296    let quic_endpoint = setup_quic(&cfg.bind, cfg.quic_port, &cert_chain, &key)?;
297    let quic_port = quic_endpoint.local_addr()?.port();
298    let tls_acceptor = setup_tls_acceptor(&cert_chain, &key)?;
299    let tcp_listener = TcpListener::bind((cfg.bind.as_str(), cfg.tcp_port)).await?;
300    let tcp_port = tcp_listener.local_addr()?.port();
301    cfg.quic_port = quic_port;
302    cfg.tcp_port = tcp_port;
303
304    let self_meta = CoordAnnounce {
305        tenant: "self".into(),
306        cluster: "default".into(),
307        coord_id: coord_id.clone(),
308        host: cfg.advertise_host.clone(),
309        quic: cfg.quic_port,
310        tcp: cfg.tcp_port,
311        pub_fp,
312        csk_ver,
313        ts: now_millis(),
314        tls_cert: Some(STANDARD_NO_PAD.encode(&cert_der)),
315        tls_fp: Some(fingerprint.clone()),
316    };
317    let peers: AliveTable = Arc::new(Mutex::new(HashMap::new()));
318    let (alive_tx, _) = broadcast::channel(16);
319    {
320        let mut map = peers.lock().await;
321        map.insert(
322            coord_id.clone(),
323            AliveEntry {
324                meta: self_meta.clone(),
325                last_seen: now_millis(),
326            },
327        );
328    }
329    tokio::spawn(sweep_dead(peers.clone(), cfg.profile.clone()));
330    let (join_tx, join_rx) = if wait_for_join {
331        let (tx, rx) = tokio::sync::oneshot::channel();
332        (Some(tx), Some(rx))
333    } else {
334        (None, None)
335    };
336    {
337        let jc = AgentConfig {
338            role: Role::Coordinator,
339            ..Default::default()
340        };
341        let peers_clone = peers.clone();
342        let tx_clone = alive_tx.clone();
343        let self_meta_clone = self_meta.clone();
344        let profile_clone = cfg.profile.clone();
345        let hosts = cfg.join_hosts.clone();
346        let join_tx_opt = if wait_for_join { join_tx } else { None };
347        tokio::spawn(async move {
348            if let Err(e) = run_coord_mesh(
349                jc,
350                hosts,
351                self_meta_clone,
352                peers_clone,
353                tx_clone,
354                profile_clone,
355                join_tx_opt,
356            )
357            .await
358            {
359                tracing::error!("join error: {}", e);
360            }
361        });
362    }
363    let active_connections = Arc::new(AtomicUsize::new(0));
364    let agent_nets = parse_nets(&cfg.agent_whitelist);
365    let coord_nets = parse_nets(&cfg.coord_whitelist);
366    let token = cfg
367        .token
368        .take()
369        .map(|t| volli_core::token::decode_token(&t).unwrap())
370        .unwrap_or_else(|| {
371            volli_core::token::issue_token(&csk, "self", "default", "agent", 86_400).unwrap()
372        });
373
374    let secret = volli_core::BootstrapSecret {
375        host: cfg.advertise_host.clone(),
376        quic_port: cfg.quic_port,
377        tcp_port: cfg.tcp_port,
378        token: token.clone(),
379        cert: cert_der.clone(),
380    };
381    let secret_encoded = secret.encode()?;
382    info!(
383        "Server listening on {} TCP:{} QUIC:{}",
384        cfg.bind, cfg.tcp_port, cfg.quic_port
385    );
386    if show_bootstrap {
387        println!("Agent bootstrap command:\n  volli agent --join {secret_encoded}");
388        println!(
389            "Hint: run 'volli --profile {} admin coord-token' to get a coordinator join command",
390            cfg.profile
391        );
392    }
393    let socket_path = cmd_socket_path(&cfg.profile);
394    info!("command socket path={} ", socket_path.display());
395    tokio::spawn(command_socket(
396        socket_path.clone(),
397        csk,
398        cfg.advertise_host.clone(),
399        cfg.quic_port,
400        cfg.tcp_port,
401        cert_der.clone(),
402    ));
403
404    if wait_for_join {
405        if let Some(rx) = join_rx {
406            let prof = cfg.profile.clone();
407            let sk_p = sk_path.clone();
408            let pk_p = pk_path.clone();
409            let cert_p = cert_path.clone();
410            let key_p = key_path.clone();
411            let sd = secret_dir.to_path_buf();
412            let signing_cl = signing.clone();
413            let cert_der_cl = cert_der.clone();
414            let key_der_cl = key_der.clone();
415            tokio::spawn(async move {
416                if rx.await.is_ok() {
417                    if persist_csk {
418                        save_csk(&prof, &csk, csk_ver).ok();
419                    }
420                    if persist_keys {
421                        if let (Some(sk), Some(pk)) = (sk_p.as_ref(), pk_p.as_ref()) {
422                            fs::create_dir_all(&sd).ok();
423                            fs::write(sk, hex::encode(signing_cl.to_bytes())).ok();
424                            let ver_bytes = signing_cl.verifying_key().to_bytes();
425                            fs::write(pk, hex::encode(ver_bytes)).ok();
426                        }
427                    }
428                    if persist_cert {
429                        if let (Some(c), Some(k)) = (cert_p.as_ref(), key_p.as_ref()) {
430                            if let Some(parent) = c.parent() {
431                                fs::create_dir_all(parent).ok();
432                            }
433                            fs::write(c, &cert_der_cl).ok();
434                            fs::write(k, &key_der_cl).ok();
435                        }
436                    }
437                    if let Some(cb) = on_ready {
438                        cb();
439                    }
440                }
441            });
442        }
443    } else {
444        if persist_csk {
445            save_csk(&cfg.profile, &csk, csk_ver).ok();
446        }
447        if persist_keys {
448            if let (Some(sk), Some(pk)) = (sk_path.as_ref(), pk_path.as_ref()) {
449                fs::create_dir_all(secret_dir)?;
450                fs::write(sk, hex::encode(signing.to_bytes()))?;
451                let ver_bytes = signing.verifying_key().to_bytes();
452                fs::write(pk, hex::encode(ver_bytes))?;
453                info!("Generated coordinator keypair at {}", secret_dir.display());
454            }
455        }
456        if persist_cert {
457            if let (Some(c), Some(k)) = (cert_path.as_ref(), key_path.as_ref()) {
458                if let Some(parent) = c.parent() {
459                    fs::create_dir_all(parent)?;
460                }
461                fs::write(c, &cert_der)?;
462                fs::write(k, &key_der)?;
463            }
464        }
465        if let Some(cb) = on_ready.take() {
466            cb();
467        }
468    }
469    loop {
470        tokio::select! {
471            Ok((stream, addr)) = tcp_listener.accept() => {
472                if active_connections.load(Ordering::SeqCst) >= cfg.max_connections {
473                    info!(peer=%addr, "connection limit reached");
474                    drop(stream);
475                    continue;
476                }
477                active_connections.fetch_add(1, Ordering::SeqCst);
478                info!(peer=%addr, "Accepted TCP connection");
479                let fp = fingerprint.clone();
480                let id = coord_id.clone();
481                let signing_clone = signing.clone();
482                let acceptor = tls_acceptor.clone();
483                let counter = active_connections.clone();
484                let agent_nets = agent_nets.clone();
485                let coord_nets = coord_nets.clone();
486                let peers_clone = peers.clone();
487                let self_meta_clone = self_meta.clone();
488                let profile_clone = cfg.profile.clone();
489                let tx = alive_tx.clone();
490                tokio::spawn(async move {
491                    if let Ok(tls) = acceptor.accept(stream).await {
492                        let proto = tls.get_ref().1.alpn_protocol().map(|p| String::from_utf8_lossy(p).to_string());
493                        match proto.as_deref() {
494                            Some("volli/agent") => {
495                                handle_agent_client(
496                                    Box::new(TcpTransport::new(tls)),
497                                    signing_clone,
498                                    csk,
499                                    fp,
500                                    id,
501                                    addr,
502                                    agent_nets,
503                                )
504                                .await
505                                .ok();
506                            }
507                            Some("volli/coord") => {
508                                let peer_fp = tls
509                                    .get_ref()
510                                    .1
511                                    .peer_certificates()
512                                    .and_then(|c| c.first().cloned());
513                                let fp = peer_fp.as_ref().map(|c| hex::encode(Sha256::digest(&c.0)));
514                                let cert = peer_fp.map(|c| c.0);
515                                handle_coord_client(
516                                    Box::new(TcpTransport::new(tls)),
517                                    csk,
518                                    self_meta_clone,
519                                    peers_clone,
520                                    tx,
521                                    profile_clone,
522                                    addr,
523                                    coord_nets,
524                                    cert,
525                                    fp,
526                                )
527                                .await
528                                .ok();
529                            }
530                            _ => {}
531                        }
532                    }
533                    counter.fetch_sub(1, Ordering::SeqCst);
534                });
535            }
536            Some(connecting) = quic_endpoint.accept() => {
537                let addr = connecting.remote_address();
538                if active_connections.load(Ordering::SeqCst) >= cfg.max_connections {
539                    info!(peer=%addr, "connection limit reached");
540                    drop(connecting);
541                    continue;
542                }
543                active_connections.fetch_add(1, Ordering::SeqCst);
544                info!(peer=%addr, "Accepted QUIC connection");
545                let fp = fingerprint.clone();
546                let id = coord_id.clone();
547                let signing_clone = signing.clone();
548                let counter = active_connections.clone();
549                let agent_nets = agent_nets.clone();
550                let coord_nets = coord_nets.clone();
551                let peers_clone = peers.clone();
552                let self_meta_clone = self_meta.clone();
553                let profile_clone = cfg.profile.clone();
554                let tx = alive_tx.clone();
555                tokio::spawn(async move {
556                    if let Ok(conn) = connecting.await {
557                        debug!("Opening QUIC connection");
558                        let protocol = conn
559                            .handshake_data()
560                            .and_then(|d| d.downcast::<quinn::crypto::rustls::HandshakeData>().ok())
561                            .and_then(|hd| hd.protocol.clone());
562                        if let Ok((send, recv)) = conn.accept_bi().await {
563                            match protocol.as_deref() {
564                                Some(b"volli/agent") => {
565                                    info!("Accepted agent connection");
566                                    handle_agent_client(
567                                        Box::new(QuicTransport::new(send, recv)),
568                                        signing_clone,
569                                        csk,
570                                        fp,
571                                        id,
572                                        addr,
573                                        agent_nets,
574                                    )
575                                    .await
576                                    .ok();
577                                }
578                                Some(b"volli/coord") => {
579                                    info!("Accepted coordinator connection");
580                                    let cert_opt = conn
581                                        .peer_identity()
582                                        .and_then(|i| i.downcast::<Vec<Certificate>>().ok())
583                                        .and_then(|mut v| v.pop());
584                                    let fp = cert_opt.as_ref().map(|c| hex::encode(Sha256::digest(&c.0)));
585                                    let cert = cert_opt.map(|c| c.0);
586                                    handle_coord_client(
587                                        Box::new(QuicTransport::new(send, recv)),
588                                        csk,
589                                        self_meta_clone,
590                                        peers_clone,
591                                        tx,
592                                        profile_clone,
593                                        addr,
594                                        coord_nets,
595                                        cert,
596                                        fp,
597                                    )
598                                    .await
599                                    .ok();
600                                }
601                                _ => {}
602                            }
603                        }
604                    }
605                    counter.fetch_sub(1, Ordering::SeqCst);
606                });
607            }
608        }
609    }
610}
611
612async fn handle_agent_client(
613    mut transport: Box<dyn Transport>,
614    signing: Arc<SigningKey>,
615    csk: [u8; 32],
616    fingerprint: String,
617    coord_id: String,
618    peer: std::net::SocketAddr,
619    whitelist: Arc<Vec<IpNet>>,
620) -> Result<(), Report> {
621    if !whitelist.is_empty() && !whitelist.iter().any(|n| n.contains(&peer.ip())) {
622        info!(%peer, "agent connection rejected by whitelist");
623        return Ok(());
624    }
625    let peer_str = peer.to_string();
626    match transport.recv().await? {
627        Some(Message::Auth { token: recv_token }) => {
628            let token = decode_token(&recv_token)?;
629            if let Err(e) = verify_token(&token, &csk) {
630                transport.send(&Message::AuthErr).await.ok();
631                return Err(e);
632            }
633            if fingerprint.is_empty() {
634                // no fp check
635            }
636            if token.payload.agent_id.is_empty() {
637                transport.send(&Message::AuthErr).await.ok();
638                return Err(eyre::eyre!("invalid agent"));
639            }
640            transport.send(&Message::AuthOk).await?;
641            info!(target: "connection", %peer, "agent authenticated");
642        }
643        _ => {
644            transport.send(&Message::AuthErr).await.ok();
645            return Ok(());
646        }
647    }
648    // send hello
649    let mut nonce = [0u8; 32];
650    getrandom::getrandom(&mut nonce)?;
651    let sig = volli_core::handshake::sign_nonce(&signing, &nonce);
652    let hello = Message::Hello {
653        coord_id: coord_id.clone(),
654        nonce,
655        sig: sig.clone(),
656    };
657    transport.send(&hello).await?;
658    match transport.recv().await? {
659        Some(Message::Welcome {
660            coord_id: cid,
661            nonce: rnonce,
662            sig: rsig,
663        }) => {
664            if cid != coord_id || rnonce != nonce || rsig != sig {
665                return Err(eyre::eyre!("handshake mismatch"));
666            }
667        }
668        _ => return Err(eyre::eyre!("handshake failed")),
669    }
670    info!(target: "connection", %peer, "coordinator authenticated");
671    let mut interval = time::interval(Duration::from_secs(5));
672    tracing::debug!(target: "connection", %peer_str, "sending ping to agent");
673    transport.send(&Message::Ping).await?;
674    loop {
675        tokio::select! {
676            msg = transport.recv() => {
677                if let Some(Message::Pong { mac }) = msg? {
678                    info!(target: "connection", %peer_str, %mac, "received pong from agent");
679                }
680            }
681            _ = interval.tick() => {
682                tracing::debug!(target: "connection", %peer_str, "sending ping to agent");
683                transport.send(&Message::Ping).await?;
684            }
685        }
686    }
687}
688
689#[allow(clippy::too_many_arguments)]
690async fn handle_coord_client(
691    mut transport: Box<dyn Transport>,
692    csk: [u8; 32],
693    self_meta: CoordAnnounce,
694    peers: AliveTable,
695    alive_tx: AliveTx,
696    profile: String,
697    peer: std::net::SocketAddr,
698    whitelist: Arc<Vec<IpNet>>,
699    peer_cert: Option<Vec<u8>>,
700    peer_fp: Option<String>,
701) -> Result<(), Report> {
702    if !whitelist.is_empty() && !whitelist.iter().any(|n| n.contains(&peer.ip())) {
703        info!(%peer, "coord connection rejected by whitelist");
704        return Ok(());
705    }
706    let stored_token: Option<String>;
707    match transport.recv().await? {
708        Some(Message::Auth { token: recv_token }) => {
709            info!(%peer, "coord connection received auth token");
710            let token = decode_token(&recv_token)?;
711            if let Err(e) = verify_token(&token, &csk) {
712                info!(%peer, "coord connection rejected by auth");
713                transport.send(&Message::AuthErr).await.ok();
714                return Err(e);
715            }
716            stored_token = Some(recv_token);
717            transport.send(&Message::AuthOk).await?;
718        }
719        _ => {
720            info!(%peer, "coord connection rejected by auth");
721            transport.send(&Message::AuthErr).await.ok();
722            return Ok(());
723        }
724    }
725    let mut interval = time::interval(Duration::from_secs(5));
726    let mut rx = alive_tx.subscribe();
727    let mut pending: Option<CoordAnnounce> = None;
728    let mut saved = false;
729    tracing::debug!(target: "connection", %peer, "sending heartbeat to coordinator");
730    transport
731        .send(&Message::CoordAnnounce {
732            meta: self_meta.clone(),
733            diff: Box::new(None),
734        })
735        .await?;
736    loop {
737        tokio::select! {
738            msg = transport.recv() => {
739                match msg? {
740                    Some(Message::CoordAnnounce { meta, diff }) => {
741                        tracing::debug!(target: "connection", %peer, id=%meta.coord_id, "received heartbeat");
742                        if !saved {
743                            let entry = JoinHostEntry {
744                                coord_id: Some(meta.coord_id.clone()),
745                                host: meta.host.clone(),
746                                tcp_port: Some(meta.tcp),
747                                quic_port: Some(meta.quic),
748                                token: stored_token.clone(),
749                                cert: meta
750                                    .tls_cert
751                                    .clone()
752                                    .or_else(|| peer_cert.as_ref().map(|c| STANDARD_NO_PAD.encode(c))),
753                                fingerprint: meta.tls_fp.clone().or(peer_fp.clone()),
754                                last_ok: Some(now_secs()),
755                                last_fail: None,
756                            };
757                            add_join_host(&profile, entry).ok();
758                            saved = true;
759                        }
760                        update_alive(&peers, &alive_tx, &profile, meta).await;
761                        if let Some(d) = *diff { update_alive(&peers, &alive_tx, &profile, d).await; }
762                    }
763                    Some(Message::Ping) => {
764                        transport.send(&Message::Pong { mac: String::new() }).await.ok();
765                    }
766                    Some(Message::Pong { .. }) => {}
767                    _ => {}
768                }
769            }
770            _ = interval.tick() => {
771                if pending.is_none() {
772                    match rx.try_recv() {
773                        Ok(upd) => pending = Some(upd),
774                        Err(broadcast::error::TryRecvError::Closed) => return Ok(()),
775                        _ => {}
776                    }
777                }
778                tracing::debug!(target: "connection", %peer, "sending heartbeat to coordinator");
779                transport
780                    .send(&Message::CoordAnnounce { meta: self_meta.clone(), diff: Box::new(pending.take()) })
781                    .await?;
782            }
783        }
784    }
785}
786
787async fn handle_coord_peer(
788    mut transport: Box<dyn Transport>,
789    cfg: &AgentConfig,
790    peer: String,
791    self_meta: CoordAnnounce,
792    peers: AliveTable,
793    alive_tx: AliveTx,
794    profile: String,
795    join_notify: Option<tokio::sync::oneshot::Sender<()>>,
796) -> Result<(), Report> {
797    transport
798        .send(&Message::Auth {
799            token: cfg.token.clone(),
800        })
801        .await?;
802    match transport.recv().await? {
803        Some(Message::AuthOk) => {
804            tracing::info!(target: "connection", %peer, "coordinator authenticated");
805            if let Some(tx) = join_notify {
806                let _ = tx.send(());
807            }
808        }
809        _ => return Err(eyre::eyre!("authentication failed")),
810    }
811    let mut interval = time::interval(Duration::from_secs(5));
812    let mut rx = alive_tx.subscribe();
813    let mut pending: Option<CoordAnnounce> = None;
814    tracing::debug!(target: "connection", %peer, "sending heartbeat to coordinator");
815    transport
816        .send(&Message::CoordAnnounce {
817            meta: self_meta.clone(),
818            diff: Box::new(None),
819        })
820        .await?;
821
822    loop {
823        tokio::select! {
824            msg = transport.recv() => {
825                match msg? {
826                    Some(Message::CoordAnnounce { meta, diff }) => {
827                        tracing::debug!(target: "connection", %peer, id=%meta.coord_id, "received heartbeat");
828                        update_alive(&peers, &alive_tx, &profile, meta).await;
829                        if let Some(d) = *diff { update_alive(&peers, &alive_tx, &profile, d).await; }
830                    }
831                    Some(Message::Ping) => {
832                      tracing::debug!(target: "connection", %peer, "received ping");
833                      transport.send(&Message::Pong { mac: String::new() }).await.ok();
834                    }
835                    Some(Message::Pong { .. }) => {
836                      tracing::debug!(target: "connection", %peer, "received pong");
837                    }
838                    _ => {}
839                }
840            }
841            _ = interval.tick() => {
842                if pending.is_none() {
843                    match rx.try_recv() {
844                        Ok(upd) => pending = Some(upd),
845                        Err(broadcast::error::TryRecvError::Closed) => {
846                          tracing::debug!(target: "connection", %peer, "connection closed");
847                          return Ok(())
848                        },
849                        _ => {}
850                    }
851                }
852                tracing::debug!(target: "connection", %peer, "sending heartbeat to coordinator");
853                transport
854                    .send(&Message::CoordAnnounce { meta: self_meta.clone(), diff: Box::new(pending.take()) })
855                    .await?;
856            }
857
858        }
859    }
860}
861
862fn configure_client(cert: &[u8], alpn: &str) -> Result<ClientConfig, Report> {
863    let mut roots = rustls::RootCertStore::empty();
864    roots.add(&Certificate(cert.to_vec()))?;
865    let mut crypto = rustls::ClientConfig::builder()
866        .with_safe_defaults()
867        .with_root_certificates(roots)
868        .with_no_client_auth();
869    crypto.alpn_protocols = vec![alpn.as_bytes().to_vec()];
870    Ok(ClientConfig::new(Arc::new(crypto)))
871}
872
873async fn connect_coord_tcp(cfg: &AgentConfig) -> Result<(Box<dyn Transport>, String), Report> {
874    let addr = format!("{}:{}", cfg.host, cfg.tcp_port);
875    let mut addrs = addr.to_socket_addrs()?;
876    let addr = addrs
877        .find(|a| a.is_ipv4())
878        .or_else(|| addrs.next())
879        .ok_or_else(|| eyre::eyre!("invalid addr"))?;
880    let stream = TcpStream::connect(addr).await?;
881    let alpn = "volli/coord";
882    let mut roots = RootCertStore::empty();
883    roots.add(&Certificate(cfg.cert.clone()))?;
884    let mut root = rustls::ClientConfig::builder()
885        .with_safe_defaults()
886        .with_root_certificates(roots)
887        .with_no_client_auth();
888    root.alpn_protocols = vec![alpn.as_bytes().to_vec()];
889    let connector = TlsConnector::from(Arc::new(root));
890    let tls = connector
891        .connect(ServerName::try_from("volli")?, stream)
892        .await?;
893    if let Some(certs) = tls.get_ref().1.peer_certificates() {
894        if let Some(cert) = certs.first() {
895            let hash = Sha256::digest(&cert.0);
896            if hex::encode(hash) != cfg.fingerprint {
897                return Err(eyre::eyre!("server fingerprint mismatch"));
898            }
899        }
900    }
901    let peer = tls.get_ref().0.peer_addr()?.to_string();
902    Ok((Box::new(TcpTransport::new(tls)), peer))
903}
904
905async fn connect_coord_quic(cfg: &AgentConfig) -> Result<(Box<dyn Transport>, String), Report> {
906    let addr = format!("{}:{}", cfg.host, cfg.quic_port);
907    let mut addrs = addr.to_socket_addrs()?;
908    let addr = addrs
909        .find(|a| a.is_ipv4())
910        .or_else(|| addrs.next())
911        .ok_or_else(|| eyre::eyre!("invalid addr"))?;
912    let mut endpoint = Endpoint::client("0.0.0.0:0".parse()?)?;
913    let quinn_cfg = configure_client(&cfg.cert, "volli/coord")?;
914    endpoint.set_default_client_config(quinn_cfg);
915    let connection = endpoint.connect(addr, "volli")?.await?;
916    if let Some(identity) = connection.peer_identity() {
917        if let Ok(certs) = identity.downcast::<Vec<Certificate>>() {
918            if let Some(cert) = certs.first() {
919                let hash = Sha256::digest(&cert.0);
920                if hex::encode(hash) != cfg.fingerprint {
921                    return Err(eyre::eyre!("server fingerprint mismatch"));
922                }
923            }
924        }
925    }
926    let peer = connection.remote_address().to_string();
927    let (send, recv) = connection.open_bi().await?;
928    Ok((Box::new(QuicTransport::new(send, recv)), peer))
929}
930
931async fn run_coord_mesh(
932    mut cfg: AgentConfig,
933    mut hosts: Vec<JoinHostEntry>,
934    self_meta: CoordAnnounce,
935    peers: AliveTable,
936    alive_tx: AliveTx,
937    profile: String,
938    mut join_notify: Option<tokio::sync::oneshot::Sender<()>>,
939) -> Result<(), Report> {
940    let proto_pref = cfg.protocol.take();
941    if hosts.is_empty() {
942        return Ok(());
943    }
944    let mut idx = 0usize;
945    let mut backoff = 1u64;
946    loop {
947        let entry = hosts.get(idx).cloned().unwrap();
948        if let Some(ref id) = entry.coord_id {
949            if self_meta.coord_id >= *id {
950                idx = (idx + 1) % hosts.len();
951                tokio::time::sleep(Duration::from_secs(backoff)).await;
952                info!("Skipping host with lower ID");
953                continue;
954            }
955        }
956        if entry.token.is_none() || entry.cert.is_none() || entry.fingerprint.is_none() {
957            idx = (idx + 1) % hosts.len();
958            tokio::time::sleep(Duration::from_secs(backoff)).await;
959            info!("Skipping host with missing metadata");
960            continue;
961        }
962        cfg.host = entry.host.clone();
963        if let Some(p) = entry.tcp_port {
964            cfg.tcp_port = p;
965        }
966        if let Some(p) = entry.quic_port {
967            cfg.quic_port = p;
968        }
969        cfg.token = entry.token.clone().unwrap();
970        let cert_bytes = match STANDARD_NO_PAD.decode(entry.cert.as_ref().unwrap().as_bytes()) {
971            Ok(c) => c,
972            Err(e) => {
973                tracing::warn!(host=%entry.host, "invalid stored cert: {e}");
974                hosts.remove(idx);
975                save_join_hosts(&profile, &hosts).ok();
976                if hosts.is_empty() {
977                    info!("No hosts left to join");
978                    return Ok(());
979                }
980                idx %= hosts.len();
981                tokio::time::sleep(Duration::from_secs(backoff)).await;
982                info!("Skipping host with invalid cert");
983                continue;
984            }
985        };
986        cfg.cert = cert_bytes;
987        cfg.fingerprint = entry.fingerprint.clone().unwrap();
988        info!("Connecting to host {:?}", cfg);
989        let res = match proto_pref.as_ref().unwrap_or(&Protocol::Quic) {
990            Protocol::Quic => match connect_coord_quic(&cfg).await {
991                Ok((tr, peer)) => {
992                    info!("Connected to host over QUIC {:?}", cfg);
993                    hosts[idx].last_ok = Some(now_secs());
994                    save_join_hosts(&profile, &hosts).ok();
995                    handle_coord_peer(
996                        tr,
997                        &cfg,
998                        peer,
999                        self_meta.clone(),
1000                        peers.clone(),
1001                        alive_tx.clone(),
1002                        profile.clone(),
1003                        join_notify.take(),
1004                    )
1005                    .await
1006                }
1007                Err(e) => {
1008                    tracing::warn!("quic connect error: {}", e);
1009                    match connect_coord_tcp(&cfg).await {
1010                        Ok((tr, peer)) => {
1011                            hosts[idx].last_ok = Some(now_secs());
1012                            save_join_hosts(&profile, &hosts).ok();
1013                            handle_coord_peer(
1014                                tr,
1015                                &cfg,
1016                                peer,
1017                                self_meta.clone(),
1018                                peers.clone(),
1019                                alive_tx.clone(),
1020                                profile.clone(),
1021                                join_notify.take(),
1022                            )
1023                            .await
1024                        }
1025                        Err(e) => Err(e),
1026                    }
1027                }
1028            },
1029            Protocol::Tcp => match connect_coord_tcp(&cfg).await {
1030                Ok((tr, peer)) => {
1031                    info!("Connected to host over TCP {:?}", cfg);
1032                    hosts[idx].last_ok = Some(now_secs());
1033                    save_join_hosts(&profile, &hosts).ok();
1034                    handle_coord_peer(
1035                        tr,
1036                        &cfg,
1037                        peer,
1038                        self_meta.clone(),
1039                        peers.clone(),
1040                        alive_tx.clone(),
1041                        profile.clone(),
1042                        join_notify.take(),
1043                    )
1044                    .await
1045                }
1046                Err(e) => Err(e),
1047            },
1048        };
1049
1050        match res {
1051            Ok(_) => {
1052                backoff = 1;
1053                idx = 0;
1054            }
1055            Err(e) => {
1056                tracing::error!("coord connection error: {}", e);
1057                hosts[idx].last_fail = Some(now_secs());
1058                save_join_hosts(&profile, &hosts).ok();
1059                backoff = (backoff * 2).min(32);
1060                idx = (idx + 1) % hosts.len();
1061            }
1062        }
1063
1064        info!("Sleeping for {} seconds", backoff);
1065        tokio::time::sleep(Duration::from_secs(backoff)).await;
1066    }
1067}
1068
1069fn setup_quic(
1070    host: &str,
1071    port: u16,
1072    certs: &[Certificate],
1073    key: &PrivateKey,
1074) -> Result<Endpoint, Report> {
1075    let mut tls = TlsServerConfig::builder()
1076        .with_safe_defaults()
1077        .with_no_client_auth()
1078        .with_single_cert(certs.to_vec(), key.clone())?;
1079    tls.alpn_protocols = vec![b"volli/agent".to_vec(), b"volli/coord".to_vec()];
1080    let mut server_config = ServerConfig::with_crypto(Arc::new(tls));
1081    let mut transport = quinn::TransportConfig::default();
1082    transport.max_idle_timeout(Some(Duration::from_secs(5).try_into()?));
1083    server_config.transport = Arc::new(transport);
1084    server_config.retry_token_lifetime(Duration::from_millis(1000));
1085    let addr = format!("{host}:{port}").to_socket_addrs()?.next().unwrap();
1086    let endpoint = Endpoint::server(server_config, addr)?;
1087    Ok(endpoint)
1088}
1089
1090pub fn load_or_generate_cert(
1091    cert_path: Option<&str>,
1092    key_path: Option<&str>,
1093    secret_dir: Option<&std::path::Path>,
1094) -> Result<(Vec<Certificate>, PrivateKey, String), Report> {
1095    let (cert_path_fs, key_path_fs) = if cert_path.is_none() && key_path.is_none() {
1096        let dir = secret_dir
1097            .map(std::path::PathBuf::from)
1098            .unwrap_or_else(default_secret_dir);
1099        (
1100            Some(dir.join("tls_cert.der")),
1101            Some(dir.join("tls_key.der")),
1102        )
1103    } else {
1104        (
1105            cert_path.map(std::path::PathBuf::from),
1106            key_path.map(std::path::PathBuf::from),
1107        )
1108    };
1109
1110    let (cert_der, key_der) =
1111        if let (Some(c), Some(k)) = (cert_path_fs.as_ref(), key_path_fs.as_ref()) {
1112            if c.exists() && k.exists() {
1113                (fs::read(c)?, fs::read(k)?)
1114            } else {
1115                fs::create_dir_all(c.parent().unwrap())?;
1116                let cert = generate_simple_self_signed(vec!["volli".into()])?;
1117                let cert_der = cert.serialize_der()?;
1118                let key_der = cert.serialize_private_key_der();
1119                fs::write(c, &cert_der)?;
1120                fs::write(k, &key_der)?;
1121                (cert_der, key_der)
1122            }
1123        } else {
1124            let cert = generate_simple_self_signed(vec!["volli".into()])?;
1125            (cert.serialize_der()?, cert.serialize_private_key_der())
1126        };
1127
1128    let fingerprint = hex_encode(Sha256::digest(&cert_der));
1129    Ok((
1130        vec![Certificate(cert_der)],
1131        PrivateKey(key_der),
1132        fingerprint,
1133    ))
1134}
1135
1136fn setup_tls_acceptor(certs: &[Certificate], key: &PrivateKey) -> Result<TlsAcceptor, Report> {
1137    let mut config = TlsServerConfig::builder()
1138        .with_safe_defaults()
1139        .with_no_client_auth()
1140        .with_single_cert(certs.to_vec(), key.clone())?;
1141    config.alpn_protocols = vec![b"volli/agent".to_vec(), b"volli/coord".to_vec()];
1142    Ok(TlsAcceptor::from(Arc::new(config)))
1143}
1144
1145fn build_secret(
1146    csk: &[u8; 32],
1147    host: &str,
1148    quic_port: u16,
1149    tcp_port: u16,
1150    cert_der: &[u8],
1151) -> Result<String, Report> {
1152    let token = volli_core::token::issue_token(csk, "self", "default", "agent", 86_400)?;
1153    let secret = volli_core::BootstrapSecret {
1154        host: host.to_string(),
1155        quic_port,
1156        tcp_port,
1157        token,
1158        cert: cert_der.to_vec(),
1159    };
1160    secret.encode()
1161}
1162
1163pub async fn command_socket(
1164    socket_path: std::path::PathBuf,
1165    csk: [u8; 32],
1166    host: String,
1167    quic: u16,
1168    tcp: u16,
1169    cert_der: Vec<u8>,
1170) -> Result<(), Report> {
1171    let _ = std::fs::remove_file(&socket_path);
1172    let listener = UnixListener::bind(&socket_path)?;
1173    loop {
1174        let (stream, _) = listener.accept().await?;
1175        let host = host.clone();
1176        let cert = cert_der.clone();
1177        tokio::spawn(async move {
1178            handle_cmd(stream, csk, host, quic, tcp, cert).await.ok();
1179        });
1180    }
1181}
1182
1183#[allow(clippy::too_many_arguments)]
1184async fn handle_cmd(
1185    stream: UnixStream,
1186    csk: [u8; 32],
1187    host: String,
1188    quic: u16,
1189    tcp: u16,
1190    cert: Vec<u8>,
1191) -> Result<(), Report> {
1192    let mut reader = BufReader::new(stream);
1193    let mut line = String::new();
1194    reader.read_line(&mut line).await?;
1195    let cmd = line.trim();
1196    let secret = build_secret(&csk, &host, quic, tcp, &cert)?;
1197    let resp = match cmd {
1198        "agent_token" => format!("volli agent --join {secret}\n"),
1199        "coord_token" => format!("volli serve --join {secret}\n"),
1200        _ => "unknown\n".to_string(),
1201    };
1202    reader.get_mut().write_all(resp.as_bytes()).await?;
1203    Ok(())
1204}