use std::sync::atomic::{AtomicBool, AtomicU32, Ordering};
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::{watch, Mutex, RwLock};
use tokio::task::JoinSet;
use crate::connection::{self, CallError, Session};
use crate::dht;
use crate::frame::{CallResponse, PublishSpec};
use crate::identity::KeyPair;
use crate::transport::Trust;
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct Seed {
pub host: String,
pub port: u16,
}
impl Seed {
pub fn new(host: impl Into<String>, port: u16) -> Self {
Self {
host: host.into(),
port,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum LinkSelection {
#[default]
Auto,
FirstSuccess,
Random,
}
fn resolve_link_selection(
configured: LinkSelection,
station_discovery_enabled: bool,
) -> LinkSelection {
match configured {
LinkSelection::Auto if station_discovery_enabled => LinkSelection::Random,
LinkSelection::Auto => LinkSelection::FirstSuccess,
other => other,
}
}
fn select_links(
mut connected: Vec<Arc<PooledLink>>,
resolved: LinkSelection,
) -> Vec<Arc<PooledLink>> {
if resolved != LinkSelection::Random || connected.len() <= 1 {
return connected;
}
use rand::seq::SliceRandom;
connected.shuffle(&mut rand::rng());
connected
}
#[derive(Debug, Clone)]
pub struct StationDiscoveryOptions {
pub enabled: bool,
pub refresh_interval: Duration,
pub max_links: usize,
}
impl Default for StationDiscoveryOptions {
fn default() -> Self {
Self {
enabled: false,
refresh_interval: Duration::from_secs(30 * 60),
max_links: 5,
}
}
}
#[derive(Debug, Clone)]
pub struct PoolOptions {
pub link_selection: LinkSelection,
pub station_discovery: StationDiscoveryOptions,
pub respawn_delay: Duration,
pub call_timeout: Duration,
pub replication_factor: usize,
}
impl Default for PoolOptions {
fn default() -> Self {
Self {
link_selection: LinkSelection::default(),
station_discovery: StationDiscoveryOptions::default(),
respawn_delay: DEFAULT_RESPAWN_DELAY,
call_timeout: connection::DEFAULT_CALL_TIMEOUT,
replication_factor: 1,
}
}
}
pub const DEFAULT_RESPAWN_DELAY: Duration = Duration::from_secs(1);
const DISCOVERY_LINK_MAX_RESPAWN_ATTEMPTS: u32 = 5;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum LinkOrigin {
Bootstrap,
Discovered,
}
fn should_give_up(origin: LinkOrigin, consecutive_failures: u32) -> bool {
origin == LinkOrigin::Discovered && consecutive_failures >= DISCOVERY_LINK_MAX_RESPAWN_ATTEMPTS
}
struct LinkState {
session: Option<Session>,
peer_node_id: Option<[u8; 32]>,
}
pub struct PooledLink {
pub seed: Seed,
origin: LinkOrigin,
trust: Trust,
state: Mutex<LinkState>,
connected: AtomicBool,
consecutive_failures: AtomicU32,
gave_up: AtomicBool,
redialing: AtomicBool,
}
impl PooledLink {
fn new(seed: Seed, origin: LinkOrigin, trust: Trust) -> Arc<Self> {
Arc::new(Self {
seed,
origin,
trust,
state: Mutex::new(LinkState {
session: None,
peer_node_id: None,
}),
connected: AtomicBool::new(false),
consecutive_failures: AtomicU32::new(0),
gave_up: AtomicBool::new(false),
redialing: AtomicBool::new(false),
})
}
pub fn is_connected(&self) -> bool {
self.connected.load(Ordering::Acquire)
}
fn has_given_up(&self) -> bool {
self.gave_up.load(Ordering::Acquire)
}
async fn peer_node_id(&self) -> Option<[u8; 32]> {
self.state.lock().await.peer_node_id
}
}
#[derive(Debug, Clone)]
pub struct LinkInfo {
pub seed: Seed,
pub connected: bool,
pub node_id: Option<[u8; 32]>,
}
#[derive(Debug, Clone, Copy)]
pub struct PoolStatus {
pub healthy_links: usize,
pub total_links: usize,
}
impl PoolStatus {
pub fn is_healthy(&self) -> bool {
self.healthy_links > 0
}
}
#[derive(Debug)]
pub enum PoolCallError {
NoHealthyStation,
AllFailed(CallError),
}
impl std::fmt::Display for PoolCallError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
PoolCallError::NoHealthyStation => {
write!(f, "pool: no link has completed its CONNECT/HELLO handshake")
}
PoolCallError::AllFailed(e) => {
write!(f, "pool: every connected link's call failed: {e}")
}
}
}
}
impl std::error::Error for PoolCallError {}
#[derive(Debug)]
pub enum PoolPublishError {
NoHealthyStation,
AllFailed(usize),
}
impl std::fmt::Display for PoolPublishError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
PoolPublishError::NoHealthyStation => {
write!(f, "pool: no link has completed its CONNECT/HELLO handshake")
}
PoolPublishError::AllFailed(n) => {
write!(f, "pool: publish failed on all {n} targeted link(s)")
}
}
}
}
impl std::error::Error for PoolPublishError {}
pub struct Pool {
identity: Arc<KeyPair>,
trust: Trust,
options: PoolOptions,
links: RwLock<Vec<Arc<PooledLink>>>,
stop_tx: watch::Sender<bool>,
tasks: Mutex<JoinSet<()>>,
}
impl Pool {
pub fn connect(
seeds: Vec<Seed>,
trust: Trust,
identity: KeyPair,
options: PoolOptions,
) -> Arc<Pool> {
assert!(!seeds.is_empty(), "at least one seed is required");
let (stop_tx, _stop_rx) = watch::channel(false);
let identity = Arc::new(identity);
let pool = Arc::new(Pool {
identity: identity.clone(),
trust,
options,
links: RwLock::new(Vec::new()),
stop_tx,
tasks: Mutex::new(JoinSet::new()),
});
let bootstrap_links: Vec<Arc<PooledLink>> = seeds
.into_iter()
.map(|seed| PooledLink::new(seed, LinkOrigin::Bootstrap, trust))
.collect();
{
let mut links = pool
.links
.try_write()
.expect("no other task can hold this lock before Pool::connect returns");
links.extend(bootstrap_links.iter().cloned());
}
{
let mut tasks = pool
.tasks
.try_lock()
.expect("no other task can hold this lock before Pool::connect returns");
for link in bootstrap_links {
if try_claim_redial(&link) {
tasks.spawn(link_lifecycle_future(pool.clone(), link));
}
}
if pool.options.station_discovery.enabled {
tasks.spawn(discover_stations_loop(pool.clone()));
}
}
pool
}
pub async fn call(
self: &Arc<Self>,
procedure: &str,
realm: [u8; 32],
payload: crate::cbor::Value,
deadline_ms: i128,
) -> Result<CallResponse, PoolCallError> {
let candidates = self.select_connected_links().await;
if candidates.is_empty() {
return Err(PoolCallError::NoHealthyStation);
}
let mut last_err = None;
for link in candidates {
let mut state = link.state.lock().await;
let Some(session) = state.session.as_mut() else {
continue; };
match session
.call(
procedure,
realm,
payload.clone(),
deadline_ms,
&self.identity,
self.options.call_timeout,
)
.await
{
Ok(resp) => return Ok(resp),
Err(e) => {
drop(state);
self.mark_disconnected(&link).await;
last_err = Some(e);
}
}
}
Err(last_err
.map(PoolCallError::AllFailed)
.unwrap_or(PoolCallError::NoHealthyStation))
}
pub async fn publish(self: &Arc<Self>, spec: &PublishSpec) -> Result<(), PoolPublishError> {
let candidates = self.select_connected_links().await;
if candidates.is_empty() {
return Err(PoolPublishError::NoHealthyStation);
}
let targets: Vec<_> = candidates
.into_iter()
.take(self.options.replication_factor.max(1))
.collect();
let attempted = targets.len();
let mut successes = 0usize;
for link in targets {
let mut state = link.state.lock().await;
let Some(session) = state.session.as_mut() else {
continue;
};
match session.publish(spec, &self.identity).await {
Ok(()) => successes += 1,
Err(_) => {
drop(state);
self.mark_disconnected(&link).await;
}
}
}
if successes > 0 {
Ok(())
} else {
Err(PoolPublishError::AllFailed(attempted))
}
}
pub async fn status(&self) -> PoolStatus {
let links = self.links.read().await;
let healthy_links = links.iter().filter(|l| l.is_connected()).count();
PoolStatus {
healthy_links,
total_links: links.len(),
}
}
pub async fn links(&self) -> Vec<LinkInfo> {
let links = self.links.read().await;
let mut out = Vec::with_capacity(links.len());
for link in links.iter() {
out.push(LinkInfo {
seed: link.seed.clone(),
connected: link.is_connected(),
node_id: if link.is_connected() {
link.peer_node_id().await
} else {
None
},
});
}
out
}
pub async fn close(&self, reason: &str, detail: Option<&str>) {
let _ = self.stop_tx.send(true);
self.tasks.lock().await.shutdown().await;
let mut links = self.links.write().await;
for link in links.drain(..) {
let mut state = link.state.lock().await;
if let Some(session) = state.session.take() {
session.close(reason, detail, &self.identity).await;
}
link.connected.store(false, Ordering::Release);
}
}
async fn select_connected_links(&self) -> Vec<Arc<PooledLink>> {
let links = self.links.read().await;
let connected: Vec<Arc<PooledLink>> =
links.iter().filter(|l| l.is_connected()).cloned().collect();
drop(links);
let resolved = resolve_link_selection(
self.options.link_selection,
self.options.station_discovery.enabled,
);
select_links(connected, resolved)
}
async fn mark_disconnected(self: &Arc<Self>, link: &Arc<PooledLink>) {
let mut state = link.state.lock().await;
state.session = None;
link.connected.store(false, Ordering::Release);
drop(state);
if try_claim_redial(link) {
self.tasks
.lock()
.await
.spawn(link_lifecycle_future(self.clone(), link.clone()));
}
}
}
fn try_claim_redial(link: &PooledLink) -> bool {
link.redialing
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
}
async fn link_lifecycle_future(pool: Arc<Pool>, link: Arc<PooledLink>) {
run_link_lifecycle(&pool, &link).await;
link.redialing.store(false, Ordering::Release);
}
async fn run_link_lifecycle(pool: &Arc<Pool>, link: &Arc<PooledLink>) {
let mut stop_rx = pool.stop_tx.subscribe();
loop {
if *stop_rx.borrow() {
return;
}
match connection::connect(&link.seed.host, link.seed.port, link.trust, &pool.identity).await
{
Ok(session) => {
let node_id = session.station.node_id;
let mut state = link.state.lock().await;
state.session = Some(session);
state.peer_node_id = Some(node_id);
drop(state);
link.connected.store(true, Ordering::Release);
link.consecutive_failures.store(0, Ordering::Release);
return;
}
Err(_e) => {
let failures = link.consecutive_failures.fetch_add(1, Ordering::AcqRel) + 1;
if should_give_up(link.origin, failures) {
link.gave_up.store(true, Ordering::Release);
return;
}
tokio::select! {
_ = tokio::time::sleep(pool.options.respawn_delay) => {}
_ = stop_rx.changed() => {
if *stop_rx.borrow() {
return;
}
}
}
}
}
}
}
const LIST_STATIONS_PROCEDURE: &str = "hecate_stations.list_stations";
const DISCOVERY_CALL_DEADLINE: Duration = Duration::from_secs(5);
async fn discover_stations_loop(pool: Arc<Pool>) {
if matches!(pool.trust, Trust::Pinned(_)) {
return;
}
let mut stop_rx = pool.stop_tx.subscribe();
if !wait_for_any_healthy_link(&pool, &mut stop_rx).await {
return;
}
loop {
if *stop_rx.borrow() {
return;
}
discover_once(&pool).await;
tokio::select! {
_ = tokio::time::sleep(pool.options.station_discovery.refresh_interval) => {}
_ = stop_rx.changed() => {
if *stop_rx.borrow() {
return;
}
}
}
}
}
async fn wait_for_any_healthy_link(pool: &Arc<Pool>, stop_rx: &mut watch::Receiver<bool>) -> bool {
loop {
if *stop_rx.borrow() {
return false;
}
if pool.status().await.healthy_links > 0 {
return true;
}
tokio::select! {
_ = tokio::time::sleep(Duration::from_millis(200)) => {}
_ = stop_rx.changed() => {
if *stop_rx.borrow() {
return false;
}
}
}
}
}
async fn discover_once(pool: &Arc<Pool>) {
let Some(realm) = resolve_list_stations_realm(pool).await else {
return;
};
let deadline = now_ms() + DISCOVERY_CALL_DEADLINE.as_millis() as i128;
let Ok(CallResponse::Result { payload, .. }) = pool
.call(
LIST_STATIONS_PROCEDURE,
realm,
crate::cbor::Value::Map(vec![]),
deadline,
)
.await
else {
return;
};
let crate::cbor::Value::Map(fields) = &payload else {
return;
};
let Some(crate::cbor::Value::List(stations)) = fields
.iter()
.find(|(k, _)| matches!(k, crate::cbor::Value::Text(t) if t == "stations"))
.map(|(_, v)| v.clone())
else {
return;
};
add_discovered_links(pool, &stations).await;
}
async fn resolve_list_stations_realm(pool: &Arc<Pool>) -> Option<[u8; 32]> {
let links = pool.links.read().await;
let link = links.iter().find(|l| l.is_connected())?.clone();
drop(links);
let mut state = link.state.lock().await;
let session = state.session.as_mut()?;
let records =
dht::find_records_by_type(session, &pool.identity, dht::TYPE_PROCEDURE_ADVERTISEMENT)
.await
.ok()?;
drop(state);
for record in records {
if dht::verify(&record).is_err() {
continue;
}
let Ok(advertisement) = dht::read_procedure_advertisement(&record) else {
continue;
};
if let Some(realm) = try_match_list_stations_realm(&advertisement.procedure_uri) {
return Some(realm);
}
}
None
}
fn try_match_list_stations_realm(procedure_uri: &str) -> Option<[u8; 32]> {
let suffix = format!("/{LIST_STATIONS_PROCEDURE}");
let hex_realm = procedure_uri.strip_suffix(&suffix)?;
if hex_realm.len() != 64 {
return None;
}
let bytes = hex_decode(hex_realm)?;
bytes.try_into().ok()
}
fn hex_decode(s: &str) -> Option<Vec<u8>> {
if s.len() % 2 != 0 {
return None;
}
(0..s.len())
.step_by(2)
.map(|i| u8::from_str_radix(&s[i..i + 2], 16).ok())
.collect()
}
fn count_occupied_discovery_slots(links: &[Arc<PooledLink>]) -> usize {
links
.iter()
.filter(|l| l.origin == LinkOrigin::Discovered && !l.has_given_up())
.count()
}
async fn add_discovered_links(pool: &Arc<Pool>, stations: &[crate::cbor::Value]) {
let is_web_pki = matches!(pool.trust, Trust::WebPki);
for station in stations {
let links = pool.links.read().await;
let occupied_slots = count_occupied_discovery_slots(&links);
drop(links);
if occupied_slots >= pool.options.station_discovery.max_links {
break;
}
let Some((host, port, node_id)) = dial_target_from_station_row(station) else {
continue;
};
let is_bare_ip = host.parse::<std::net::IpAddr>().is_ok();
let link_trust = if is_bare_ip && is_web_pki {
match node_id {
Some(id) => Trust::Pinned(id),
None => continue, }
} else {
pool.trust
};
if let Some(id) = node_id {
if has_link_for_node_id(pool, id).await {
continue;
}
}
let seed = Seed::new(host, port);
spawn_seed_link_if_absent(pool, seed, link_trust).await;
}
}
pub(crate) fn dial_target_from_station_row(
row: &crate::cbor::Value,
) -> Option<(String, u16, Option<[u8; 32]>)> {
let crate::cbor::Value::Map(fields) = row else {
return None;
};
let get = |name: &str| {
fields
.iter()
.find(|(k, _)| matches!(k, crate::cbor::Value::Text(t) if t == name))
.map(|(_, v)| v)
};
let port = match get("quic_port") {
Some(crate::cbor::Value::Int(n)) if (1..=65535).contains(n) => *n as u16,
_ => return None,
};
let hostname = match get("hostname") {
Some(crate::cbor::Value::Text(t)) if !t.is_empty() => Some(t.clone()),
Some(crate::cbor::Value::Bytes(b)) if !b.is_empty() => String::from_utf8(b.clone()).ok(),
_ => None,
};
let host_advertised = match get("host_advertised") {
Some(crate::cbor::Value::List(items)) => items.iter().find_map(|item| match item {
crate::cbor::Value::Bytes(b) => String::from_utf8(b.clone()).ok(),
crate::cbor::Value::Text(t) => Some(t.clone()),
_ => None,
}),
_ => None,
};
let host = hostname.or(host_advertised)?;
let node_id = match get("node_id") {
Some(crate::cbor::Value::Bytes(b)) => b.as_slice().try_into().ok(),
_ => None,
};
Some((host, port, node_id))
}
async fn has_link_for_node_id(pool: &Arc<Pool>, node_id: [u8; 32]) -> bool {
let links = pool.links.read().await;
for link in links.iter() {
if link.is_connected() {
if let Some(known) = link.peer_node_id().await {
if known == node_id {
return true;
}
}
}
}
false
}
async fn spawn_seed_link_if_absent(pool: &Arc<Pool>, seed: Seed, trust: Trust) {
let mut links = pool.links.write().await;
if links.iter().any(|l| l.seed == seed) {
return;
}
let link = PooledLink::new(seed, LinkOrigin::Discovered, trust);
links.push(link.clone());
drop(links);
if try_claim_redial(&link) {
pool.tasks
.lock()
.await
.spawn(link_lifecycle_future(pool.clone(), link));
}
}
fn now_ms() -> i128 {
use std::time::{SystemTime, UNIX_EPOCH};
SystemTime::now()
.duration_since(UNIX_EPOCH)
.expect("system clock before 1970")
.as_millis() as i128
}
#[cfg(test)]
mod tests {
use super::*;
fn row(fields: Vec<(&str, crate::cbor::Value)>) -> crate::cbor::Value {
crate::cbor::Value::Map(
fields
.into_iter()
.map(|(k, v)| (crate::cbor::Value::Text(k.to_string()), v))
.collect(),
)
}
#[test]
fn resolve_link_selection_auto_pairs_with_discovery() {
assert_eq!(
resolve_link_selection(LinkSelection::Auto, false),
LinkSelection::FirstSuccess
);
assert_eq!(
resolve_link_selection(LinkSelection::Auto, true),
LinkSelection::Random
);
}
#[test]
fn resolve_link_selection_explicit_survives_either_way() {
assert_eq!(
resolve_link_selection(LinkSelection::FirstSuccess, true),
LinkSelection::FirstSuccess
);
assert_eq!(
resolve_link_selection(LinkSelection::Random, false),
LinkSelection::Random
);
}
#[test]
fn dial_target_prefers_hostname_over_bare_ip() {
let r = row(vec![
(
"hostname",
crate::cbor::Value::Text("station-de-frankfurt.macula.io".into()),
),
(
"host_advertised",
crate::cbor::Value::List(vec![crate::cbor::Value::Bytes(
b"2a01:7e01::f03c:94ff:fe22:719e".to_vec(),
)]),
),
("quic_port", crate::cbor::Value::Int(4433)),
]);
let (host, port, _) = dial_target_from_station_row(&r).expect("should parse");
assert_eq!(host, "station-de-frankfurt.macula.io");
assert_eq!(port, 4433);
}
#[test]
fn dial_target_falls_back_to_host_advertised_when_hostname_absent() {
let r = row(vec![
(
"host_advertised",
crate::cbor::Value::List(vec![crate::cbor::Value::Bytes(
b"2600:3c0b::2000:1fff:fe35:416b".to_vec(),
)]),
),
("quic_port", crate::cbor::Value::Int(4433)),
]);
let (host, _, _) = dial_target_from_station_row(&r).expect("should parse");
assert_eq!(host, "2600:3c0b::2000:1fff:fe35:416b");
}
#[test]
fn dial_target_rejects_missing_port() {
let r = row(vec![("hostname", crate::cbor::Value::Text("x".into()))]);
assert!(dial_target_from_station_row(&r).is_none());
}
#[test]
fn dial_target_extracts_node_id() {
let node_id = [0xABu8; 32];
let r = row(vec![
("hostname", crate::cbor::Value::Text("x".into())),
("quic_port", crate::cbor::Value::Int(4433)),
("node_id", crate::cbor::Value::Bytes(node_id.to_vec())),
]);
let (_, _, id) = dial_target_from_station_row(&r).expect("should parse");
assert_eq!(id, Some(node_id));
}
#[test]
fn should_give_up_never_applies_to_a_bootstrap_link() {
assert!(!should_give_up(
LinkOrigin::Bootstrap,
DISCOVERY_LINK_MAX_RESPAWN_ATTEMPTS
));
assert!(!should_give_up(LinkOrigin::Bootstrap, 1_000_000));
}
#[test]
fn should_give_up_applies_to_a_discovered_link_at_the_threshold() {
assert!(!should_give_up(
LinkOrigin::Discovered,
DISCOVERY_LINK_MAX_RESPAWN_ATTEMPTS - 1
));
assert!(should_give_up(
LinkOrigin::Discovered,
DISCOVERY_LINK_MAX_RESPAWN_ATTEMPTS
));
}
fn synthetic_link(origin: LinkOrigin, gave_up: bool) -> Arc<PooledLink> {
let link = PooledLink::new(Seed::new("x.example", 4433), origin, Trust::WebPki);
link.gave_up.store(gave_up, Ordering::Relaxed);
link
}
#[test]
fn discovery_slot_count_excludes_bootstrap_links_entirely() {
let links = vec![
synthetic_link(LinkOrigin::Bootstrap, false),
synthetic_link(LinkOrigin::Bootstrap, false),
synthetic_link(LinkOrigin::Bootstrap, false),
];
assert_eq!(count_occupied_discovery_slots(&links), 0);
}
#[test]
fn discovery_slot_count_excludes_a_given_up_discovered_link() {
let links = vec![
synthetic_link(LinkOrigin::Discovered, false),
synthetic_link(LinkOrigin::Discovered, true), ];
assert_eq!(count_occupied_discovery_slots(&links), 1);
}
#[test]
fn discovery_slot_count_mixed_origins() {
let links = vec![
synthetic_link(LinkOrigin::Bootstrap, false),
synthetic_link(LinkOrigin::Bootstrap, false),
synthetic_link(LinkOrigin::Discovered, false),
synthetic_link(LinkOrigin::Discovered, true),
];
assert_eq!(count_occupied_discovery_slots(&links), 1);
}
#[test]
fn try_claim_redial_only_lets_one_caller_win() {
let link = synthetic_link(LinkOrigin::Discovered, false);
assert!(try_claim_redial(&link));
assert!(
!try_claim_redial(&link),
"a second claim must fail while the first is outstanding"
);
link.redialing.store(false, Ordering::Release);
assert!(
try_claim_redial(&link),
"releasing the claim must allow a fresh one"
);
}
#[test]
fn try_match_list_stations_realm_matches_expected_format() {
let hex_realm = "0".repeat(64);
let uri = format!("{hex_realm}/hecate_stations.list_stations");
assert_eq!(try_match_list_stations_realm(&uri), Some([0u8; 32]));
}
#[test]
fn try_match_list_stations_realm_rejects_a_different_procedure() {
let hex_realm = "0".repeat(64);
let uri = format!("{hex_realm}/some.other_procedure");
assert_eq!(try_match_list_stations_realm(&uri), None);
}
#[test]
fn select_links_first_success_returns_input_unshuffled() {
let empty: Vec<Arc<PooledLink>> = Vec::new();
assert_eq!(select_links(empty, LinkSelection::FirstSuccess).len(), 0);
}
}