From 56e599ef31ff1fd6bce9cc400b0abf5a0e037c30 Mon Sep 17 00:00:00 2001 From: Phil Denhoff Date: Mon, 21 Sep 2026 16:02:14 -0700 Subject: [PATCH] fix(opds): classify listener failures, keep serving through interface hiccups - Bind failures are now classified: address-in-use and permission errors are permanent (clear message, no retry loop); other failures retry with exponential backoff and only surface as errors after repeated attempts. - A failing listener no longer tears down the others; partial success serves the bound addresses. - Transient interface-enumeration failures no longer stop running listeners; advertisement refreshes on the next successful snapshot. - plan_bindings takes a BindPolicy; globally-routable addresses are excluded from binding by default (reserved for auth-enabled sharing). --- crates/citadel-opds/src/network.rs | 64 ++++- crates/citadel-opds/src/service.rs | 392 +++++++++++++++++++++++------ 2 files changed, 367 insertions(+), 89 deletions(-) diff --git a/crates/citadel-opds/src/network.rs b/crates/citadel-opds/src/network.rs index 6dc424a1..35e14f36 100644 --- a/crates/citadel-opds/src/network.rs +++ b/crates/citadel-opds/src/network.rs @@ -65,9 +65,11 @@ pub(crate) struct InterfaceSnapshot { } impl InterfaceSnapshot { - pub fn bindable_ips(&self) -> impl Iterator + '_ { + pub fn bindable_ips(&self, policy: BindPolicy) -> impl Iterator + '_ { + let allow_global = policy.allow_global; self.addresses .iter() + .filter(move |address| allow_global || address.scope != AddressScope::Global) .filter_map(|address| match (address.address, &address.scope) { ( IpAddr::V4(ip), @@ -82,8 +84,8 @@ impl InterfaceSnapshot { }) } - pub fn bindable_addresses(&self, port: u16) -> Vec { - self.bindable_ips() + pub fn bindable_addresses(&self, port: u16, policy: BindPolicy) -> Vec { + self.bindable_ips(policy) .map(|ip| match ip { IpAddr::V4(ip) => SocketAddr::new(IpAddr::V4(ip), port), IpAddr::V6(ip) => SocketAddr::V6(SocketAddrV6::new(ip, port, 0, 0)), @@ -93,7 +95,7 @@ impl InterfaceSnapshot { pub fn public(&self) -> OpdsNetworkInterface { let addresses = self - .bindable_ips() + .bindable_ips(BindPolicy::default()) .map(|address| address.to_string()) .collect::>(); OpdsNetworkInterface { @@ -334,6 +336,21 @@ pub(crate) enum BindPlan { Wait(WaitingReason), } +/// What address scopes sharing may reach. Globally-routable addresses stay +/// excluded until the caller can vouch for them (auth-enabled sharing). +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(crate) struct BindPolicy { + pub allow_global: bool, +} + +impl Default for BindPolicy { + fn default() -> Self { + Self { + allow_global: false, + } + } +} + /// Pure and cheap: the service layer re-snapshots interfaces and re-plans on /// every change (DHCP renew, roam, sleep/wake), so callers must not assume a /// plan outlives the interface state it was computed from. @@ -341,6 +358,7 @@ pub(crate) fn plan_bindings( interfaces: &[InterfaceSnapshot], target: &OpdsBindTarget, port: u16, + policy: BindPolicy, ) -> BindPlan { let selected = match target { OpdsBindTarget::AllLocalNetworks => interfaces @@ -366,7 +384,7 @@ pub(crate) fn plan_bindings( let addresses = selected .into_iter() - .flat_map(|interface| interface.bindable_addresses(port)) + .flat_map(|interface| interface.bindable_addresses(port, policy)) .collect::>(); if addresses.is_empty() { BindPlan::Wait(match target { @@ -463,6 +481,7 @@ mod tests { &interfaces, &OpdsBindTarget::AllLocalNetworks, 8080, + BindPolicy::default(), )); assert_eq!( @@ -490,7 +509,7 @@ mod tests { "en1", OpdsInterfaceKind::Lan, OpdsInterfaceState::Up, - vec![ipv4([192, 168, 2, 42]), ipv6("2001:db8::42")], + vec![ipv4([192, 168, 2, 42]), ipv6("fd12:3456::42")], ), ]; @@ -500,13 +519,14 @@ mod tests { id: "en1".to_string(), }, 8080, + BindPolicy::default(), )); assert_eq!( addresses, BTreeSet::from([ "192.168.2.42:8080".parse().unwrap(), - "[2001:db8::42]:8080".parse().unwrap(), + "[fd12:3456::42]:8080".parse().unwrap(), ]) ); } @@ -519,6 +539,7 @@ mod tests { id: "en0".to_string(), }, 8080, + BindPolicy::default(), ); assert_eq!( missing, @@ -536,6 +557,7 @@ mod tests { id: "en0".to_string(), }, 8080, + BindPolicy::default(), ); assert_eq!( down, @@ -553,6 +575,7 @@ mod tests { id: "en0".to_string(), }, 8080, + BindPolicy::default(), ); assert_eq!( unusable, @@ -578,6 +601,7 @@ mod tests { id: "lo0".to_string(), }, 8080, + BindPolicy::default(), ); assert_eq!( @@ -608,7 +632,7 @@ mod tests { ); assert_eq!( - snapshot.bindable_addresses(8080), + snapshot.bindable_addresses(8080, BindPolicy::default()), vec!["[fd12::1]:8080".parse().unwrap()] ); } @@ -690,11 +714,35 @@ mod tests { id: "br0".to_string(), }, 8080, + BindPolicy::default(), ), BindPlan::Listen(BTreeSet::from(["192.168.1.7:8080".parse().unwrap()])) ); } + #[test] + fn global_addresses_bind_only_when_the_policy_allows_them() { + let snapshot = interface( + "en0", + OpdsInterfaceKind::Lan, + OpdsInterfaceState::Up, + vec![ipv4([192, 168, 1, 42]), ipv6("2001:db8::42")], + ); + + assert_eq!( + snapshot + .bindable_addresses(8080, BindPolicy::default()) + .len(), + 1 + ); + assert_eq!( + snapshot + .bindable_addresses(8080, BindPolicy { allow_global: true }) + .len(), + 2 + ); + } + #[test] fn classify_routes_interfaces_by_name_label_and_description() { struct Row { diff --git a/crates/citadel-opds/src/service.rs b/crates/citadel-opds/src/service.rs index 1dc8bed1..cf9a1494 100644 --- a/crates/citadel-opds/src/service.rs +++ b/crates/citadel-opds/src/service.rs @@ -6,7 +6,7 @@ use std::{ atomic::{AtomicUsize, Ordering}, Arc, Mutex, }, - time::Duration, + time::{Duration, Instant}, }; use futures_util::future::join_all; @@ -18,7 +18,7 @@ use tokio::{ use super::{ network::{ - advertised_url, plan_bindings, InterfaceProvider, InterfaceSnapshot, + advertised_url, plan_bindings, BindPolicy, InterfaceProvider, InterfaceSnapshot, NetdevInterfaceProvider, OpdsNetworkInterface, WaitingReason, }, router, CatalogSource, @@ -27,6 +27,8 @@ use super::{ const MONITOR_INTERVAL: Duration = Duration::from_secs(2); const LISTENER_SHUTDOWN_TIMEOUT: Duration = Duration::from_secs(5); const WORKER_SHUTDOWN_TIMEOUT: Duration = Duration::from_secs(6); +const MAX_RETRY_BACKOFF: Duration = Duration::from_secs(60); +const TRANSIENT_ERROR_THRESHOLD: usize = 2; pub use crate::network::OpdsBindTarget; @@ -305,6 +307,9 @@ impl OpdsService { self.inner.source.clone(), &self.inner.status, Some(library_id.clone()), + &mut BTreeMap::new(), + self.inner.dependencies.monitor_interval, + BindPolicy::default(), ) .await; } @@ -445,6 +450,76 @@ async fn active_library_id(source: Arc) -> Option { .flatten() } +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +enum BindFailureClass { + Persistent, + Transient, +} + +struct ListenerFailure { + class: BindFailureClass, + consecutive: usize, + next_retry: Option, +} + +impl Default for ListenerFailure { + fn default() -> Self { + Self { + class: BindFailureClass::Transient, + consecutive: 0, + next_retry: None, + } + } +} + +/// Address-in-use and permission errors do not heal on their own; everything +/// else (an address vanishing mid-bind, transient socket exhaustion) is worth +/// retrying. +fn classify_bind_error(error: &io::Error) -> BindFailureClass { + match error.kind() { + io::ErrorKind::AddrInUse | io::ErrorKind::PermissionDenied => BindFailureClass::Persistent, + _ => BindFailureClass::Transient, + } +} + +fn bind_failure_error(address: SocketAddr, class: BindFailureClass) -> OpdsStatusError { + match class { + BindFailureClass::Persistent => OpdsStatusError { + code: OpdsErrorCode::PortUnavailable, + message: format!( + "The port Citadel uses for {address} is unavailable or restricted. Pick a different port." + ), + }, + BindFailureClass::Transient => OpdsStatusError { + code: OpdsErrorCode::ListenerFailed, + message: format!( + "Citadel could not open the listener on {address}. It will keep trying." + ), + }, + } +} + +fn retry_delay(attempts: usize, base: Duration) -> Duration { + let shift = attempts.saturating_sub(1).min(6) as u32; + base.saturating_mul(1u32 << shift).min(MAX_RETRY_BACKOFF) +} + +fn record_failure( + failures: &mut BTreeMap, + address: SocketAddr, + class: BindFailureClass, + now: Instant, + retry_base: Duration, +) { + let entry = failures.entry(address).or_default(); + entry.class = class; + entry.consecutive = entry.consecutive.saturating_add(1); + entry.next_retry = match class { + BindFailureClass::Persistent => None, + BindFailureClass::Transient => Some(now + retry_delay(entry.consecutive, retry_base)), + }; +} + async fn apply_plan( interfaces: &[InterfaceSnapshot], config: &OpdsStartConfig, @@ -455,10 +530,14 @@ async fn apply_plan( source: Arc, status: &Mutex, library_id: Option, + failures: &mut BTreeMap, + retry_base: Duration, + policy: BindPolicy, ) { - match plan_bindings(interfaces, &config.target, port) { + match plan_bindings(interfaces, &config.target, port, policy) { super::network::BindPlan::Wait(reason) => { shutdown_servers(servers).await; + failures.clear(); set_configured_status( status, config, @@ -479,56 +558,100 @@ async fn apply_plan( if let Some(server) = servers.remove(&address) { stale_servers.insert(address, server); } + failures.remove(&address); } shutdown_servers(&mut stale_servers).await; - if let Err(error) = start_missing_servers(servers, &desired, listeners, tracker, source) - { - shutdown_servers(servers).await; + let now = Instant::now(); + let mut persistent_failure = None; + let mut worst_transient = 0usize; + for address in &desired { + if servers.contains_key(address) { + continue; + } + if failures + .get(address) + .is_some_and(|failure| match failure.class { + BindFailureClass::Persistent => true, + BindFailureClass::Transient => { + failure.next_retry.is_some_and(|at| at > now) + } + }) + { + continue; + } + match listeners.start(*address, source.clone(), tracker.clone()) { + Ok(task) => { + servers.insert(*address, task); + failures.remove(address); + } + Err(error) => { + let class = classify_bind_error(&error); + let consecutive = failures + .get(address) + .map(|failure| failure.consecutive) + .unwrap_or_default() + + 1; + record_failure(failures, *address, class, now, retry_base); + match class { + BindFailureClass::Persistent => { + if persistent_failure.is_none() { + persistent_failure = Some((*address, class)); + } + } + BindFailureClass::Transient => { + worst_transient = worst_transient.max(consecutive); + } + } + } + } + } + + let bound: BTreeSet = servers.keys().copied().collect(); + if let Some((address, class)) = persistent_failure { set_configured_status( status, config, OpdsLifecycleState::Error, - Some(bind_error(error)), + Some(bind_failure_error(address, class)), + urls(&bound), + library_id, + ); + } else if worst_transient >= TRANSIENT_ERROR_THRESHOLD { + set_configured_status( + status, + config, + OpdsLifecycleState::Error, + Some(OpdsStatusError { + code: OpdsErrorCode::ListenerFailed, + message: + "Citadel is having trouble opening its sharing listeners. It will keep trying." + .to_string(), + }), + urls(&bound), + library_id, + ); + } else if !bound.is_empty() { + set_configured_status( + status, + config, + OpdsLifecycleState::Running, + None, + urls(&bound), + library_id, + ); + } else { + set_configured_status( + status, + config, + OpdsLifecycleState::Starting, + None, Vec::new(), library_id, ); - return; } - set_configured_status( - status, - config, - OpdsLifecycleState::Running, - None, - urls(&desired), - library_id, - ); - } - } -} - -fn start_missing_servers( - servers: &mut BTreeMap, - desired: &BTreeSet, - listeners: &dyn ListenerFactory, - tracker: Arc, - source: Arc, -) -> io::Result<()> { - for address in desired { - if address.ip().is_unspecified() { - return Err(io::Error::new( - io::ErrorKind::InvalidInput, - "wildcard listeners are forbidden", - )); - } - if !servers.contains_key(address) { - servers.insert( - *address, - listeners.start(*address, source.clone(), tracker.clone())?, - ); } } - Ok(()) } async fn monitor_service( @@ -546,44 +669,38 @@ async fn monitor_service( let mut ticker = tokio::time::interval(interval); ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); ticker.tick().await; + let mut failures: BTreeMap = BTreeMap::new(); loop { tokio::select! { _ = &mut stop => break, _ = ticker.tick() => {} } - if servers - .values() - .any(|server| server.task.as_ref().is_some_and(JoinHandle::is_finished)) - { - shutdown_servers(&mut servers).await; - set_configured_status( - &status, - &config, - OpdsLifecycleState::Error, - Some(OpdsStatusError { - code: OpdsErrorCode::ListenerFailed, - message: "An OPDS listener stopped unexpectedly; Citadel will retry." - .to_string(), - }), - Vec::new(), - None, + let dead = servers + .iter() + .filter(|(_, server)| server.task.as_ref().is_some_and(JoinHandle::is_finished)) + .map(|(address, _)| *address) + .collect::>(); + for address in dead { + if let Some(mut server) = servers.remove(&address) { + if let Some(task) = server.task.take() { + let _ = task.await; + } + } + record_failure( + &mut failures, + address, + BindFailureClass::Transient, + Instant::now(), + interval, ); - continue; } let snapshot = match snapshot_interfaces(interfaces.clone()).await { Ok(snapshot) => snapshot, Err(_) => { - shutdown_servers(&mut servers).await; - set_configured_status( - &status, - &config, - OpdsLifecycleState::Error, - Some(interface_enumeration_error(true)), - Vec::new(), - None, - ); + // Keep the existing listeners serving on the last known plan; + // advertisement refreshes on the next successful snapshot. continue; } }; @@ -598,6 +715,9 @@ async fn monitor_service( source.clone(), &status, library_id, + &mut failures, + interval, + BindPolicy::default(), ) .await; } @@ -731,20 +851,6 @@ fn interface_enumeration_error(retrying: bool) -> OpdsStatusError { } } -fn bind_error(error: io::Error) -> OpdsStatusError { - if error.kind() == io::ErrorKind::AddrInUse { - OpdsStatusError { - code: OpdsErrorCode::PortUnavailable, - message: "That port is already in use on the selected network interface.".to_string(), - } - } else { - OpdsStatusError { - code: OpdsErrorCode::Unexpected, - message: "Citadel could not open the requested OPDS listener.".to_string(), - } - } -} - #[cfg(test)] mod tests { use std::{ @@ -824,6 +930,8 @@ mod tests { fail_bind: AtomicBool, fail_task: AtomicBool, hold_shutdown: Arc, + fail_every: Mutex>, + fail_addrs: Mutex>, } impl ListenerFactory for FakeListeners { @@ -833,6 +941,12 @@ mod tests { _source: Arc, tracker: Arc, ) -> io::Result { + if let Some(kind) = *self.fail_every.lock().unwrap() { + return Err(io::Error::from(kind)); + } + if self.fail_addrs.lock().unwrap().contains(&address) { + return Err(io::Error::from(io::ErrorKind::AddrInUse)); + } if self.fail_bind.swap(false, Ordering::AcqRel) { return Err(io::Error::from(io::ErrorKind::AddrInUse)); } @@ -936,6 +1050,7 @@ mod tests { InterfaceSnapshot { id: "en0".to_string(), label: "Ethernet".to_string(), + description: None, state, kind: OpdsInterfaceKind::Lan, addresses: vec![InterfaceAddress { @@ -1112,13 +1227,128 @@ mod tests { Duration::from_millis(50), ); + service.start(config(8080)).await.unwrap(); + while !listeners.active.lock().unwrap().is_empty() { + tokio::time::sleep(Duration::from_millis(5)).await; + } + let mut restarted = None; + for _ in 0..200 { + let status = service.status().await; + if status.state == OpdsLifecycleState::Running + && listeners.active.lock().unwrap().len() == 1 + { + restarted = Some(status); + break; + } + tokio::time::sleep(Duration::from_millis(5)).await; + } + let restarted = restarted.expect("listener was not restarted"); + assert!(restarted.error.is_none()); + service.stop().await; + assert!(listeners.active.lock().unwrap().is_empty()); + } + + #[tokio::test(flavor = "multi_thread")] + async fn repeated_transient_bind_failures_surface_an_error_then_recover() { + let source = test_source(); + let interfaces = Arc::new(FakeInterfaces::new(vec![lan( + OpdsInterfaceState::Up, + [192, 168, 1, 5], + )])); + let listeners = Arc::new(FakeListeners::default()); + *listeners.fail_every.lock().unwrap() = Some(io::ErrorKind::AddrNotAvailable); + let service = service_with( + source, + interfaces, + listeners.clone(), + Duration::from_millis(10), + ); + service.start(config(8080)).await.unwrap(); let failed = wait_for_state(&service, OpdsLifecycleState::Error).await; assert_eq!(failed.error.unwrap().code, OpdsErrorCode::ListenerFailed); + + *listeners.fail_every.lock().unwrap() = None; wait_for_state(&service, OpdsLifecycleState::Running).await; + service.stop().await; + } + + #[tokio::test(flavor = "multi_thread")] + async fn persistent_bind_failure_keeps_the_other_listeners_serving() { + let source = test_source(); + let other = InterfaceSnapshot { + id: "en1".to_string(), + label: "en1 label".to_string(), + description: None, + state: OpdsInterfaceState::Up, + kind: OpdsInterfaceKind::Lan, + addresses: vec![InterfaceAddress { + address: IpAddr::V4(Ipv4Addr::new(192, 168, 2, 5)), + scope: AddressScope::Private, + deprecated: false, + tentative: false, + duplicated: false, + }], + }; + let interfaces = Arc::new(FakeInterfaces::new(vec![ + lan(OpdsInterfaceState::Up, [192, 168, 1, 5]), + other, + ])); + let listeners = Arc::new(FakeListeners::default()); + listeners + .fail_addrs + .lock() + .unwrap() + .insert("192.168.2.5:8080".parse().unwrap()); + let service = service_with( + source, + interfaces, + listeners.clone(), + Duration::from_millis(10), + ); + + let failed = service + .start(OpdsStartConfig { + target: OpdsBindTarget::AllLocalNetworks, + port: 8080, + }) + .await + .unwrap(); + assert_eq!(failed.state, OpdsLifecycleState::Error); + assert_eq!(failed.error.unwrap().code, OpdsErrorCode::PortUnavailable); + assert!(!failed.urls.is_empty()); assert_eq!(listeners.active.lock().unwrap().len(), 1); service.stop().await; - assert!(listeners.active.lock().unwrap().is_empty()); + } + + #[tokio::test(flavor = "multi_thread")] + async fn interface_enumeration_hiccup_keeps_listeners_serving() { + let source = test_source(); + let interfaces = Arc::new(FakeInterfaces::new(vec![lan( + OpdsInterfaceState::Up, + [192, 168, 1, 5], + )])); + let listeners = Arc::new(FakeListeners::default()); + let service = service_with( + source, + interfaces.clone(), + listeners.clone(), + Duration::from_millis(10), + ); + + service.start(config(8080)).await.unwrap(); + interfaces.set_error(io::ErrorKind::Other); + tokio::time::sleep(Duration::from_millis(60)).await; + assert_eq!( + service.status().await.state, + OpdsLifecycleState::Running, + "a transient snapshot failure must not stop serving" + ); + assert_eq!(listeners.active.lock().unwrap().len(), 1); + + interfaces.set(vec![lan(OpdsInterfaceState::Up, [192, 168, 1, 5])]); + wait_for_state(&service, OpdsLifecycleState::Running).await; + service.stop().await; } #[tokio::test(flavor = "multi_thread")]