diff --git a/Cargo.lock b/Cargo.lock index 213e7bba..b13c6d61 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -658,6 +658,7 @@ name = "citadel-opds" version = "0.1.0" dependencies = [ "axum", + "base64 0.22.1", "bytes", "chrono", "diesel", @@ -670,6 +671,7 @@ dependencies = [ "specta", "tempfile", "tokio", + "tower", "urlencoding", "uuid", ] @@ -687,6 +689,7 @@ dependencies = [ "libcalibre", "log", "mobi", + "netdev", "quick-xml 0.38.4", "regex", "reqwest 0.12.28", @@ -708,6 +711,7 @@ dependencies = [ "tauri-plugin-webdriver-automation", "tauri-specta", "tempfile", + "tokio", "urlencoding", "uuid", "zip 0.6.6", diff --git a/crates/citadel-opds/Cargo.toml b/crates/citadel-opds/Cargo.toml index bad92709..1571f13c 100644 --- a/crates/citadel-opds/Cargo.toml +++ b/crates/citadel-opds/Cargo.toml @@ -7,6 +7,7 @@ description = "OPDS catalog, authentication, and networking runtime for Citadel" [dependencies] axum = "0.8.9" +base64 = "0.22" bytes = "1" chrono = { version = "0.4.31", features = ["serde"] } futures-util = "0.3" @@ -21,5 +22,6 @@ uuid = { version = "1.6.1", features = ["v4", "fast-rng"] } [dev-dependencies] diesel = { version = "2.2.4", features = ["sqlite"] } -reqwest = "0.12" +reqwest = { version = "0.12", features = ["json"] } tempfile = "3.8" +tower = "0.5" diff --git a/crates/citadel-opds/src/catalog.rs b/crates/citadel-opds/src/catalog.rs index 0cb07e9d..9b84df90 100644 --- a/crates/citadel-opds/src/catalog.rs +++ b/crates/citadel-opds/src/catalog.rs @@ -536,9 +536,11 @@ mod tests { impl CatalogSource for MemorySource { fn active_library_id(&self) -> Result { - Ok("550e8400-e29b-41d4-a716-446655440000".to_string()) + self.books + .first() + .map(|_| "memory-library".to_string()) + .ok_or(CalibreError::LibraryNotInitialized) } - fn book_page( &self, limit: i64, @@ -586,7 +588,6 @@ mod tests { fn active_library_id(&self) -> Result { Err(CalibreError::LibraryNotInitialized) } - fn book_page( &self, _limit: i64, @@ -612,7 +613,6 @@ mod tests { fn active_library_id(&self) -> Result { Err((self.0)()) } - fn book_page( &self, _limit: i64, diff --git a/crates/citadel-opds/src/lib.rs b/crates/citadel-opds/src/lib.rs index 1ef9e953..7af0585f 100644 --- a/crates/citadel-opds/src/lib.rs +++ b/crates/citadel-opds/src/lib.rs @@ -4,6 +4,12 @@ pub mod assets; pub mod catalog; mod identity; pub mod network; +pub mod service; mod xml; pub use catalog::{router, CatalogSource}; +pub use network::{OpdsInterfaceKind, OpdsInterfaceState, OpdsNetworkInterface}; +pub use service::{ + OpdsBindTarget, OpdsErrorCode, OpdsLifecycleState, OpdsService, OpdsServiceStatus, + OpdsStartConfig, OpdsStatusError, +}; diff --git a/crates/citadel-opds/src/service.rs b/crates/citadel-opds/src/service.rs new file mode 100644 index 00000000..1dc8bed1 --- /dev/null +++ b/crates/citadel-opds/src/service.rs @@ -0,0 +1,1252 @@ +use std::{ + collections::{BTreeMap, BTreeSet}, + io, + net::{SocketAddr, TcpListener as StdTcpListener}, + sync::{ + atomic::{AtomicUsize, Ordering}, + Arc, Mutex, + }, + time::Duration, +}; + +use futures_util::future::join_all; +use serde::{Deserialize, Serialize}; +use tokio::{ + sync::{oneshot, Notify}, + task::JoinHandle, +}; + +use super::{ + network::{ + advertised_url, plan_bindings, InterfaceProvider, InterfaceSnapshot, + NetdevInterfaceProvider, OpdsNetworkInterface, WaitingReason, + }, + router, CatalogSource, +}; + +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); + +pub use crate::network::OpdsBindTarget; + +#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize, specta::Type)] +#[serde(rename_all = "camelCase")] +pub struct OpdsStartConfig { + pub target: OpdsBindTarget, + pub port: u32, +} + +#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize, specta::Type)] +#[serde(rename_all = "camelCase")] +pub enum OpdsLifecycleState { + Stopped, + Starting, + Running, + WaitingForInterface, + Error, + Stopping, +} + +#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize, specta::Type)] +#[serde(rename_all = "camelCase")] +pub enum OpdsErrorCode { + InvalidPort, + LibraryNotReady, + ConfigurationConflict, + InterfaceUnavailable, + InterfaceEnumerationFailed, + PortUnavailable, + ListenerFailed, + Unexpected, +} + +#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize, specta::Type)] +#[serde(rename_all = "camelCase")] +pub struct OpdsStatusError { + pub code: OpdsErrorCode, + pub message: String, +} + +#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize, specta::Type)] +#[serde(rename_all = "camelCase")] +pub struct OpdsServiceStatus { + pub state: OpdsLifecycleState, + pub target: Option, + pub port: Option, + pub active_library_id: Option, + pub urls: Vec, + pub error: Option, +} + +impl Default for OpdsServiceStatus { + fn default() -> Self { + Self { + state: OpdsLifecycleState::Stopped, + target: None, + port: None, + active_library_id: None, + urls: Vec::new(), + error: None, + } + } +} + +trait InterfaceSnapshots: Send + Sync { + fn snapshot(&self) -> io::Result>; +} + +struct NetdevInterfaceSnapshots; + +impl InterfaceSnapshots for NetdevInterfaceSnapshots { + fn snapshot(&self) -> io::Result> { + let mut provider = NetdevInterfaceProvider; + provider.snapshot() + } +} + +trait ListenerFactory: Send + Sync { + fn start( + &self, + address: SocketAddr, + source: Arc, + tracker: Arc, + ) -> io::Result; +} + +struct TcpListenerFactory; + +impl ListenerFactory for TcpListenerFactory { + fn start( + &self, + address: SocketAddr, + source: Arc, + tracker: Arc, + ) -> io::Result { + let listener = StdTcpListener::bind(address)?; + listener.set_nonblocking(true)?; + let listener = tokio::net::TcpListener::from_std(listener)?; + let (shutdown, shutdown_receiver) = oneshot::channel(); + let app = router(source); + let completion = tracker.register(); + let task = tokio::spawn(async move { + let _completion = completion; + axum::serve(listener, app) + .with_graceful_shutdown(async move { + let _ = shutdown_receiver.await; + }) + .await + }); + Ok(ServerTask { + shutdown: Some(shutdown), + task: Some(task), + }) + } +} + +struct ServiceDependencies { + interfaces: Arc, + listeners: Arc, + monitor_interval: Duration, + worker_shutdown_timeout: Duration, +} + +impl Default for ServiceDependencies { + fn default() -> Self { + Self { + interfaces: Arc::new(NetdevInterfaceSnapshots), + listeners: Arc::new(TcpListenerFactory), + monitor_interval: MONITOR_INTERVAL, + worker_shutdown_timeout: WORKER_SHUTDOWN_TIMEOUT, + } + } +} + +struct RunningController { + config: OpdsStartConfig, + stop: oneshot::Sender<()>, + worker: JoinHandle<()>, +} + +struct ServiceInner { + status: Arc>, + controller: tokio::sync::Mutex>, + source: Arc, + dependencies: ServiceDependencies, + listener_tracker: Arc, +} + +#[derive(Clone)] +pub struct OpdsService { + inner: Arc, +} + +impl OpdsService { + pub fn new(source: Arc) -> Self { + Self::with_dependencies(source, ServiceDependencies::default()) + } + + fn with_dependencies( + source: Arc, + dependencies: ServiceDependencies, + ) -> Self { + Self { + inner: Arc::new(ServiceInner { + status: Arc::new(Mutex::new(OpdsServiceStatus::default())), + controller: tokio::sync::Mutex::new(None), + source, + dependencies, + listener_tracker: Arc::new(ListenerTracker::default()), + }), + } + } + + pub async fn status(&self) -> OpdsServiceStatus { + if let Ok(mut controller) = self.inner.controller.try_lock() { + reap_finished_controller(&mut controller, &self.inner.status).await; + } + + let mut status = status_snapshot(&self.inner.status); + if status.state == OpdsLifecycleState::Stopped { + status.active_library_id = None; + return status; + } + status.active_library_id = active_library_id(self.inner.source.clone()).await; + status + } + + pub async fn list_interfaces(&self) -> Result, OpdsStatusError> { + snapshot_interfaces(self.inner.dependencies.interfaces.clone()) + .await + .map(|interfaces| { + interfaces + .into_iter() + .map(|interface| interface.public()) + .collect() + }) + .map_err(|_| OpdsStatusError { + code: OpdsErrorCode::InterfaceEnumerationFailed, + message: "Citadel could not list network interfaces.".to_string(), + }) + } + + pub async fn start( + &self, + config: OpdsStartConfig, + ) -> Result { + let mut controller = self.inner.controller.lock().await; + reap_finished_controller(&mut controller, &self.inner.status).await; + if let Some(running) = controller.as_ref() { + if running.config == config { + drop(controller); + return Ok(self.status().await); + } + return Err(OpdsStatusError { + code: OpdsErrorCode::ConfigurationConflict, + message: "Sharing is already active with different network settings. Stop it before changing the interface or port.".to_string(), + }); + } + + set_configured_status( + &self.inner.status, + &config, + OpdsLifecycleState::Starting, + None, + Vec::new(), + None, + ); + let Ok(port) = u16::try_from(config.port) + .ok() + .filter(|port| *port != 0) + .ok_or(()) + else { + set_configured_status( + &self.inner.status, + &config, + OpdsLifecycleState::Error, + Some(OpdsStatusError { + code: OpdsErrorCode::InvalidPort, + message: "Choose a port between 1 and 65535.".to_string(), + }), + Vec::new(), + None, + ); + drop(controller); + return Ok(self.status().await); + }; + + let library_id = active_library_id(self.inner.source.clone()).await; + let Some(library_id) = library_id else { + set_configured_status( + &self.inner.status, + &config, + OpdsLifecycleState::Error, + Some(OpdsStatusError { + code: OpdsErrorCode::LibraryNotReady, + message: "Open a library before starting sharing.".to_string(), + }), + Vec::new(), + None, + ); + drop(controller); + return Ok(self.status().await); + }; + + let mut servers = BTreeMap::new(); + match snapshot_interfaces(self.inner.dependencies.interfaces.clone()).await { + Ok(interfaces) => { + apply_plan( + &interfaces, + &config, + port, + &mut servers, + self.inner.dependencies.listeners.as_ref(), + self.inner.listener_tracker.clone(), + self.inner.source.clone(), + &self.inner.status, + Some(library_id.clone()), + ) + .await; + } + Err(_) => { + set_configured_status( + &self.inner.status, + &config, + OpdsLifecycleState::Error, + Some(interface_enumeration_error(false)), + Vec::new(), + Some(library_id), + ); + } + } + + let (stop, stop_receiver) = oneshot::channel(); + let worker = tokio::spawn(monitor_service( + self.inner.dependencies.interfaces.clone(), + self.inner.dependencies.listeners.clone(), + self.inner.listener_tracker.clone(), + self.inner.source.clone(), + self.inner.status.clone(), + config.clone(), + port, + servers, + stop_receiver, + self.inner.dependencies.monitor_interval, + )); + *controller = Some(RunningController { + config, + stop, + worker, + }); + drop(controller); + Ok(self.status().await) + } + + pub async fn stop(&self) -> OpdsServiceStatus { + let mut controller = self.inner.controller.lock().await; + let Some(running) = controller.take() else { + set_stopped_status(&self.inner.status); + drop(controller); + return self.status().await; + }; + set_simple_status( + &self.inner.status, + OpdsLifecycleState::Stopping, + None, + Vec::new(), + ); + let stopped = stop_controller( + running, + self.inner.listener_tracker.clone(), + self.inner.dependencies.worker_shutdown_timeout, + ) + .await; + if stopped { + set_stopped_status(&self.inner.status); + } else { + set_simple_status( + &self.inner.status, + OpdsLifecycleState::Error, + Some(OpdsStatusError { + code: OpdsErrorCode::ListenerFailed, + message: "Citadel could not confirm that every OPDS listener stopped." + .to_string(), + }), + Vec::new(), + ); + } + drop(controller); + self.status().await + } +} + +struct ServerTask { + shutdown: Option>, + task: Option>>, +} + +#[derive(Default)] +struct ListenerTracker { + active: AtomicUsize, + empty: Notify, +} + +impl ListenerTracker { + fn register(self: &Arc) -> ListenerCompletion { + self.active.fetch_add(1, Ordering::AcqRel); + ListenerCompletion { + tracker: self.clone(), + } + } + + async fn wait_until_empty(&self) { + loop { + let notified = self.empty.notified(); + if self.active.load(Ordering::Acquire) == 0 { + return; + } + notified.await; + } + } +} + +struct ListenerCompletion { + tracker: Arc, +} + +impl Drop for ListenerCompletion { + fn drop(&mut self) { + if self.tracker.active.fetch_sub(1, Ordering::AcqRel) == 1 { + self.tracker.empty.notify_waiters(); + } + } +} + +impl Drop for ServerTask { + fn drop(&mut self) { + if let Some(task) = self.task.as_ref() { + task.abort(); + } + } +} + +async fn snapshot_interfaces( + interfaces: Arc, +) -> io::Result> { + tokio::task::spawn_blocking(move || interfaces.snapshot()) + .await + .map_err(|_| io::Error::other("interface snapshot task failed"))? +} + +async fn active_library_id(source: Arc) -> Option { + tokio::task::spawn_blocking(move || source.active_library_id().ok()) + .await + .ok() + .flatten() +} + +async fn apply_plan( + interfaces: &[InterfaceSnapshot], + config: &OpdsStartConfig, + port: u16, + servers: &mut BTreeMap, + listeners: &dyn ListenerFactory, + tracker: Arc, + source: Arc, + status: &Mutex, + library_id: Option, +) { + match plan_bindings(interfaces, &config.target, port) { + super::network::BindPlan::Wait(reason) => { + shutdown_servers(servers).await; + set_configured_status( + status, + config, + OpdsLifecycleState::WaitingForInterface, + Some(waiting_error(&reason)), + Vec::new(), + library_id, + ); + } + super::network::BindPlan::Listen(desired) => { + let stale = servers + .keys() + .filter(|address| !desired.contains(address)) + .copied() + .collect::>(); + let mut stale_servers = BTreeMap::new(); + for address in stale { + if let Some(server) = servers.remove(&address) { + stale_servers.insert(address, server); + } + } + shutdown_servers(&mut stale_servers).await; + + if let Err(error) = start_missing_servers(servers, &desired, listeners, tracker, source) + { + shutdown_servers(servers).await; + set_configured_status( + status, + config, + OpdsLifecycleState::Error, + Some(bind_error(error)), + 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( + interfaces: Arc, + listeners: Arc, + tracker: Arc, + source: Arc, + status: Arc>, + config: OpdsStartConfig, + port: u16, + mut servers: BTreeMap, + mut stop: oneshot::Receiver<()>, + interval: Duration, +) { + let mut ticker = tokio::time::interval(interval); + ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + ticker.tick().await; + 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, + ); + 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, + ); + continue; + } + }; + let library_id = active_library_id(source.clone()).await; + apply_plan( + &snapshot, + &config, + port, + &mut servers, + listeners.as_ref(), + tracker.clone(), + source.clone(), + &status, + library_id, + ) + .await; + } + shutdown_servers(&mut servers).await; +} + +async fn reap_finished_controller( + controller: &mut Option, + status: &Mutex, +) { + if !controller + .as_ref() + .is_some_and(|running| running.worker.is_finished()) + { + return; + } + let running = controller.take().expect("finished controller disappeared"); + let _ = running.worker.await; + set_simple_status( + status, + OpdsLifecycleState::Error, + Some(OpdsStatusError { + code: OpdsErrorCode::ListenerFailed, + message: "The OPDS service stopped unexpectedly. Start sharing again to retry." + .to_string(), + }), + Vec::new(), + ); +} + +async fn stop_controller( + controller: RunningController, + tracker: Arc, + worker_shutdown_timeout: Duration, +) -> bool { + let _ = controller.stop.send(()); + let mut worker = controller.worker; + if tokio::time::timeout(worker_shutdown_timeout, &mut worker) + .await + .is_err() + { + worker.abort(); + let _ = worker.await; + } + tokio::time::timeout(LISTENER_SHUTDOWN_TIMEOUT, tracker.wait_until_empty()) + .await + .is_ok() +} + +async fn shutdown_servers(servers: &mut BTreeMap) { + let mut servers = std::mem::take(servers).into_values().collect::>(); + for server in &mut servers { + if let Some(shutdown) = server.shutdown.take() { + let _ = shutdown.send(()); + } + } + let mut tasks = servers + .iter_mut() + .filter_map(|server| server.task.take()) + .collect::>(); + if tokio::time::timeout(LISTENER_SHUTDOWN_TIMEOUT, join_all(tasks.iter_mut())) + .await + .is_err() + { + for task in &tasks { + task.abort(); + } + let _ = join_all(tasks).await; + } +} + +fn status_snapshot(status: &Mutex) -> OpdsServiceStatus { + status.lock().expect("OPDS status mutex poisoned").clone() +} + +fn set_stopped_status(status: &Mutex) { + *status.lock().expect("OPDS status mutex poisoned") = OpdsServiceStatus::default(); +} + +fn set_simple_status( + status: &Mutex, + state: OpdsLifecycleState, + error: Option, + urls: Vec, +) { + let mut status = status.lock().expect("OPDS status mutex poisoned"); + status.state = state; + status.error = error; + status.urls = urls; +} + +fn set_configured_status( + status: &Mutex, + config: &OpdsStartConfig, + state: OpdsLifecycleState, + error: Option, + urls: Vec, + library_id: Option, +) { + let mut status = status.lock().expect("OPDS status mutex poisoned"); + status.state = state; + status.target = Some(config.target.clone()); + status.port = Some(config.port); + if library_id.is_some() { + status.active_library_id = library_id; + } + status.error = error; + status.urls = urls; +} + +fn urls(addresses: &BTreeSet) -> Vec { + addresses.iter().copied().map(advertised_url).collect() +} + +fn waiting_error(reason: &WaitingReason) -> OpdsStatusError { + OpdsStatusError { + code: OpdsErrorCode::InterfaceUnavailable, + message: reason.message(), + } +} + +fn interface_enumeration_error(retrying: bool) -> OpdsStatusError { + OpdsStatusError { + code: OpdsErrorCode::InterfaceEnumerationFailed, + message: if retrying { + "Citadel could not inspect network interfaces; it will retry." + } else { + "Citadel could not inspect network interfaces." + } + .to_string(), + } +} + +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::{ + net::{IpAddr, Ipv4Addr}, + path::Path, + sync::atomic::{AtomicBool, Ordering}, + }; + + use super::*; + use crate::network::{AddressScope, InterfaceAddress, OpdsInterfaceKind, OpdsInterfaceState}; + + struct FakeInterfaces { + current: Mutex, io::ErrorKind>>, + } + + struct BlockingInterfaces { + snapshot: Vec, + entered: AtomicBool, + release: AtomicBool, + } + + struct HangingMonitorInterfaces { + snapshot: Vec, + calls: AtomicUsize, + release: AtomicBool, + } + + impl InterfaceSnapshots for HangingMonitorInterfaces { + fn snapshot(&self) -> io::Result> { + if self.calls.fetch_add(1, Ordering::AcqRel) != 0 { + while !self.release.load(Ordering::Acquire) { + std::thread::sleep(Duration::from_millis(1)); + } + } + Ok(self.snapshot.clone()) + } + } + + impl InterfaceSnapshots for BlockingInterfaces { + fn snapshot(&self) -> io::Result> { + self.entered.store(true, Ordering::Release); + while !self.release.load(Ordering::Acquire) { + std::thread::sleep(Duration::from_millis(1)); + } + Ok(self.snapshot.clone()) + } + } + + impl FakeInterfaces { + fn new(snapshot: Vec) -> Self { + Self { + current: Mutex::new(Ok(snapshot)), + } + } + + fn set(&self, snapshot: Vec) { + *self.current.lock().unwrap() = Ok(snapshot); + } + + fn set_error(&self, kind: io::ErrorKind) { + *self.current.lock().unwrap() = Err(kind); + } + } + + impl InterfaceSnapshots for FakeInterfaces { + fn snapshot(&self) -> io::Result> { + match &*self.current.lock().unwrap() { + Ok(snapshot) => Ok(snapshot.clone()), + Err(kind) => Err(io::Error::from(*kind)), + } + } + } + + #[derive(Default)] + struct FakeListeners { + active: Arc>>, + fail_bind: AtomicBool, + fail_task: AtomicBool, + hold_shutdown: Arc, + } + + impl ListenerFactory for FakeListeners { + fn start( + &self, + address: SocketAddr, + _source: Arc, + tracker: Arc, + ) -> io::Result { + if self.fail_bind.swap(false, Ordering::AcqRel) { + return Err(io::Error::from(io::ErrorKind::AddrInUse)); + } + self.active.lock().unwrap().insert(address); + let active = self.active.clone(); + let self_hold_shutdown = self.hold_shutdown.clone(); + let task_should_fail = self.fail_task.swap(false, Ordering::AcqRel); + let completion = tracker.register(); + let (shutdown, shutdown_receiver) = oneshot::channel(); + let task = tokio::spawn(async move { + let _completion = completion; + let _active = FakeActiveListener { active, address }; + let result = if task_should_fail { + Err(io::Error::from(io::ErrorKind::BrokenPipe)) + } else { + let _ = shutdown_receiver.await; + while self_hold_shutdown.load(Ordering::Acquire) { + tokio::time::sleep(Duration::from_millis(1)).await; + } + Ok(()) + }; + result + }); + Ok(ServerTask { + shutdown: Some(shutdown), + task: Some(task), + }) + } + } + + struct FakeActiveListener { + active: Arc>>, + address: SocketAddr, + } + + impl Drop for FakeActiveListener { + fn drop(&mut self) { + self.active.lock().unwrap().remove(&self.address); + } + } + + struct TestSource { + library_id: Option, + } + + use chrono::NaiveDateTime; + + impl CatalogSource for TestSource { + fn active_library_id(&self) -> Result { + self.library_id + .clone() + .ok_or(libcalibre::CalibreError::LibraryNotInitialized) + } + + fn book_page( + &self, + _limit: i64, + _offset: i64, + ) -> Result<(String, Option, libcalibre::BookPage), libcalibre::CalibreError> + { + Ok(( + self.active_library_id()?, + None, + libcalibre::BookPage { + items: Vec::new(), + total: 0, + }, + )) + } + + fn book_file( + &self, + book_id: libcalibre::BookId, + format: &str, + ) -> Result { + Err(libcalibre::CalibreError::BookFileNotFound( + book_id, + format.to_string(), + )) + } + + fn book_cover( + &self, + book_id: libcalibre::BookId, + ) -> Result { + Err(libcalibre::CalibreError::BookCoverNotFound(book_id)) + } + } + + fn test_source() -> Arc { + Arc::new(TestSource { + library_id: Some("test-library".to_string()), + }) + } + + fn missing_source() -> Arc { + Arc::new(TestSource { library_id: None }) + } + + fn lan(state: OpdsInterfaceState, address: [u8; 4]) -> InterfaceSnapshot { + InterfaceSnapshot { + id: "en0".to_string(), + label: "Ethernet".to_string(), + state, + kind: OpdsInterfaceKind::Lan, + addresses: vec![InterfaceAddress { + address: IpAddr::V4(Ipv4Addr::from(address)), + scope: AddressScope::Private, + deprecated: false, + tentative: false, + duplicated: false, + }], + } + } + + fn config(port: u32) -> OpdsStartConfig { + OpdsStartConfig { + target: OpdsBindTarget::Interface { + id: "en0".to_string(), + }, + port, + } + } + + fn service_with( + source: Arc, + interfaces: Arc, + listeners: Arc, + interval: Duration, + ) -> OpdsService { + OpdsService::with_dependencies( + source, + ServiceDependencies { + interfaces, + listeners, + monitor_interval: interval, + worker_shutdown_timeout: Duration::from_millis(100), + }, + ) + } + + async fn wait_for_state( + service: &OpdsService, + expected: OpdsLifecycleState, + ) -> OpdsServiceStatus { + for _ in 0..100 { + let status = service.status().await; + if status.state == expected { + return status; + } + tokio::time::sleep(Duration::from_millis(5)).await; + } + panic!( + "timed out waiting for {expected:?}; current status: {:?}", + service.status().await + ); + } + + #[tokio::test(flavor = "multi_thread")] + async fn selected_interface_loss_waits_without_broadening_and_recovers() { + 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), + ); + + let running = service.start(config(8080)).await.unwrap(); + assert_eq!(running.state, OpdsLifecycleState::Running); + assert_eq!( + *listeners.active.lock().unwrap(), + BTreeSet::from(["192.168.1.5:8080".parse().unwrap()]) + ); + + interfaces.set(vec![lan(OpdsInterfaceState::Down, [192, 168, 1, 5])]); + let waiting = wait_for_state(&service, OpdsLifecycleState::WaitingForInterface).await; + assert_eq!( + waiting.error.unwrap().code, + OpdsErrorCode::InterfaceUnavailable + ); + assert!(listeners.active.lock().unwrap().is_empty()); + + interfaces.set(vec![lan(OpdsInterfaceState::Up, [192, 168, 1, 9])]); + let recovered = wait_for_state(&service, OpdsLifecycleState::Running).await; + assert_eq!(recovered.urls, vec!["http://192.168.1.9:8080/opds"]); + assert_eq!( + *listeners.active.lock().unwrap(), + BTreeSet::from(["192.168.1.9:8080".parse().unwrap()]) + ); + + assert_eq!(service.stop().await.state, OpdsLifecycleState::Stopped); + assert!(listeners.active.lock().unwrap().is_empty()); + } + + #[tokio::test(flavor = "multi_thread")] + async fn concurrent_same_config_starts_are_idempotent_and_conflicts_are_typed() { + 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, + listeners.clone(), + Duration::from_millis(10), + ); + + let requested = config(8080); + let (first, second) = tokio::join!( + service.start(requested.clone()), + service.start(requested.clone()) + ); + assert_eq!(first.unwrap().state, OpdsLifecycleState::Running); + assert_eq!(second.unwrap().state, OpdsLifecycleState::Running); + assert_eq!(listeners.active.lock().unwrap().len(), 1); + + let conflict = service.start(config(8081)).await.unwrap_err(); + assert_eq!(conflict.code, OpdsErrorCode::ConfigurationConflict); + assert_eq!(service.status().await.port, Some(8080)); + service.stop().await; + assert_eq!(service.stop().await.state, OpdsLifecycleState::Stopped); + } + + #[tokio::test(flavor = "multi_thread")] + async fn invalid_port_missing_library_and_occupied_port_are_actionable() { + let interfaces = Arc::new(FakeInterfaces::new(vec![lan( + OpdsInterfaceState::Up, + [192, 168, 1, 5], + )])); + let listeners = Arc::new(FakeListeners::default()); + let missing = service_with( + missing_source(), + interfaces.clone(), + listeners.clone(), + Duration::from_secs(1), + ); + let status = missing.start(config(8080)).await.unwrap(); + assert_eq!(status.error.unwrap().code, OpdsErrorCode::LibraryNotReady); + + let source = test_source(); + let service = service_with( + source, + interfaces, + listeners.clone(), + Duration::from_secs(1), + ); + let invalid = service.start(config(70_000)).await.unwrap(); + assert_eq!(invalid.error.unwrap().code, OpdsErrorCode::InvalidPort); + + listeners.fail_bind.store(true, Ordering::Release); + let occupied = service.start(config(8080)).await.unwrap(); + assert_eq!(occupied.error.unwrap().code, OpdsErrorCode::PortUnavailable); + service.stop().await; + } + + #[tokio::test(flavor = "multi_thread")] + async fn listener_task_failure_is_observable_and_then_retried() { + 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_task.store(true, Ordering::Release); + let service = service_with( + source, + interfaces, + listeners.clone(), + Duration::from_millis(50), + ); + + service.start(config(8080)).await.unwrap(); + let failed = wait_for_state(&service, OpdsLifecycleState::Error).await; + assert_eq!(failed.error.unwrap().code, OpdsErrorCode::ListenerFailed); + wait_for_state(&service, OpdsLifecycleState::Running).await; + 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_failure_is_actionable_and_retried() { + let source = test_source(); + let interfaces = Arc::new(FakeInterfaces::new(Vec::new())); + interfaces.set_error(io::ErrorKind::Other); + let listeners = Arc::new(FakeListeners::default()); + let service = service_with( + source, + interfaces.clone(), + listeners, + Duration::from_millis(10), + ); + + let failed = service.start(config(8080)).await.unwrap(); + assert_eq!( + failed.error.unwrap().code, + OpdsErrorCode::InterfaceEnumerationFailed + ); + 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")] + async fn starting_and_stopping_are_observable_during_blocked_dependencies() { + let source = test_source(); + let interfaces = Arc::new(BlockingInterfaces { + snapshot: vec![lan(OpdsInterfaceState::Up, [192, 168, 1, 5])], + entered: AtomicBool::new(false), + release: AtomicBool::new(false), + }); + let listeners = Arc::new(FakeListeners::default()); + let service = service_with( + source, + interfaces.clone(), + listeners.clone(), + Duration::from_secs(1), + ); + + let starting_service = service.clone(); + let start = tokio::spawn(async move { starting_service.start(config(8080)).await }); + while !interfaces.entered.load(Ordering::Acquire) { + tokio::task::yield_now().await; + } + assert_eq!(service.status().await.state, OpdsLifecycleState::Starting); + interfaces.release.store(true, Ordering::Release); + assert_eq!( + start.await.unwrap().unwrap().state, + OpdsLifecycleState::Running + ); + + listeners.hold_shutdown.store(true, Ordering::Release); + let stopping_service = service.clone(); + let stop = tokio::spawn(async move { stopping_service.stop().await }); + wait_for_state(&service, OpdsLifecycleState::Stopping).await; + listeners.hold_shutdown.store(false, Ordering::Release); + assert_eq!(stop.await.unwrap().state, OpdsLifecycleState::Stopped); + } + + #[tokio::test(flavor = "multi_thread")] + async fn real_listener_rejects_occupied_port_and_releases_it_on_stop() { + let occupied = StdTcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let address = occupied.local_addr().unwrap(); + let factory = TcpListenerFactory; + let tracker = Arc::new(ListenerTracker::default()); + let error = match factory.start(address, missing_source(), tracker.clone()) { + Ok(_) => panic!("occupied listener unexpectedly bound"), + Err(error) => error, + }; + assert_eq!(error.kind(), io::ErrorKind::AddrInUse); + drop(occupied); + + let mut servers = BTreeMap::from([( + address, + factory + .start(address, missing_source(), tracker.clone()) + .unwrap(), + )]); + shutdown_servers(&mut servers).await; + tracker.wait_until_empty().await; + assert!(StdTcpListener::bind(address).is_ok()); + } + + #[tokio::test(flavor = "multi_thread")] + async fn forced_worker_abort_awaits_listener_cancellation_and_releases_port() { + let probe = StdTcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let port = probe.local_addr().unwrap().port(); + drop(probe); + let interfaces = Arc::new(HangingMonitorInterfaces { + snapshot: vec![lan(OpdsInterfaceState::Up, [127, 0, 0, 1])], + calls: AtomicUsize::new(0), + release: AtomicBool::new(false), + }); + let source = test_source(); + let service = service_with( + source, + interfaces.clone(), + Arc::new(TcpListenerFactory), + Duration::from_millis(5), + ); + assert_eq!( + service.start(config(u32::from(port))).await.unwrap().state, + OpdsLifecycleState::Running + ); + while interfaces.calls.load(Ordering::Acquire) < 2 { + tokio::task::yield_now().await; + } + + assert_eq!(service.stop().await.state, OpdsLifecycleState::Stopped); + assert!(StdTcpListener::bind((Ipv4Addr::LOCALHOST, port)).is_ok()); + interfaces.release.store(true, Ordering::Release); + } + + #[tokio::test(flavor = "multi_thread")] + async fn listener_completion_cannot_be_lost_between_count_check_and_wait() { + for _ in 0..1_000 { + let tracker = Arc::new(ListenerTracker::default()); + let completion = tracker.register(); + let waiter_tracker = tracker.clone(); + let waiter = tokio::spawn(async move { waiter_tracker.wait_until_empty().await }); + tokio::task::yield_now().await; + drop(completion); + tokio::time::timeout(Duration::from_millis(100), waiter) + .await + .expect("listener completion notification was lost") + .unwrap(); + } + } +} diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 1bc24d6b..48023ad5 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -19,6 +19,7 @@ chrono = { version = "0.4.31", features = ["serde"] } diesel = { version = "2.1.0", features = ["sqlite", "chrono", "returning_clauses_for_sqlite_3_35"] } epub = "2.1.1" mobi = "0.8.0" +netdev = { version = "=0.45.0", default-features = false } citadel-core = { path = "../crates/citadel-core" } citadel-opds = { path = "../crates/citadel-opds" } libcalibre = { path = "../crates/libcalibre" } @@ -46,6 +47,7 @@ tauri-plugin-opener = "2" tauri-plugin-updater = "2" tauri-plugin-webdriver-automation = "0.1.3" image = { version = "0.25.6", default-features = false, features = ["jpeg", "png"] } +tokio = { version = "1.52.3", features = ["macros", "net", "rt-multi-thread", "sync", "time"] } [dev-dependencies] tempfile = "3.8" diff --git a/src-tauri/src/app_updates.rs b/src-tauri/src/app_updates.rs index 745a83dc..997cea44 100644 --- a/src-tauri/src/app_updates.rs +++ b/src-tauri/src/app_updates.rs @@ -1,5 +1,7 @@ use tauri_plugin_updater::UpdaterExt; +use citadel_opds::OpdsService; + #[derive(serde::Serialize, specta::Type)] pub struct UpdateCheckResult { pub has_update: bool, @@ -28,7 +30,10 @@ pub async fn clb_cmd_check_for_updates(app: tauri::AppHandle) -> Result Result { +pub async fn clb_cmd_install_update_if_available( + app: tauri::AppHandle, + opds: tauri::State<'_, OpdsService>, +) -> Result { let updater = app .updater() .map_err(|err| format!("Updater initialization failed: {err}"))?; @@ -42,6 +47,7 @@ pub async fn clb_cmd_install_update_if_available(app: tauri::AppHandle) -> Resul .await .map_err(|err| format!("Failed to download/install update: {err}"))?; + opds.stop().await; app.restart(); } Ok(None) => Ok("no-update".to_string()), diff --git a/src-tauri/src/main.rs b/src-tauri/src/main.rs index d741b232..875113dd 100644 --- a/src-tauri/src/main.rs +++ b/src-tauri/src/main.rs @@ -1,6 +1,11 @@ // Prevents additional console window on Windows in release, DO NOT REMOVE!! #![cfg_attr(not(debug_assertions), windows_subsystem = "windows")] +use std::sync::{ + atomic::{AtomicU8, Ordering}, + Arc, +}; + use libs::calibre; #[cfg(debug_assertions)] use specta_typescript::Typescript; @@ -16,6 +21,7 @@ pub mod libs { mod book; mod menu; mod metadata; +pub mod opds; mod state; fn run_tauri_backend() -> std::io::Result<()> { @@ -64,6 +70,11 @@ fn run_tauri_backend() -> std::io::Result<()> { metadata::commands::clb_query_metadata_by_isbn, app_updates::clb_cmd_check_for_updates, app_updates::clb_cmd_install_update_if_available, + // OPDS sharing commands + opds::commands::clb_query_opds_interfaces, + opds::commands::clb_cmd_start_opds, + opds::commands::clb_cmd_stop_opds, + opds::commands::clb_query_opds_status, // Window commands menu::clb_cmd_open_settings, ]); @@ -86,10 +97,13 @@ fn run_tauri_backend() -> std::io::Result<()> { tauri_builder = tauri_builder.plugin(tauri_plugin_webdriver_automation::init()); } - tauri_builder + let state = state::CitadelState::new(); + let opds_service = citadel_opds::OpdsService::new(Arc::new(state.clone())); + let app = tauri_builder .plugin(tauri_plugin_opener::init()) .plugin(tauri_plugin_updater::Builder::new().build()) - .manage(state::CitadelState::new()) + .manage(state) + .manage(opds_service) .invoke_handler(builder.invoke_handler()) .plugin(tauri_plugin_store::Builder::new().build()) .setup(move |app| { @@ -140,8 +154,42 @@ fn run_tauri_backend() -> std::io::Result<()> { .plugin(tauri_plugin_drag::init()) .plugin(tauri_plugin_shell::init()) .plugin(tauri_plugin_clipboard_manager::init()) - .run(tauri::generate_context!()) - .expect("error while running tauri application"); + .build(tauri::generate_context!()) + .expect("error while building tauri application"); + + const RUNNING: u8 = 0; + const STOPPING: u8 = 1; + const EXITING: u8 = 2; + let exit_phase = Arc::new(AtomicU8::new(RUNNING)); + app.run(move |app_handle, event| { + if let tauri::RunEvent::ExitRequested { code, api, .. } = event { + match exit_phase.compare_exchange( + RUNNING, + STOPPING, + Ordering::AcqRel, + Ordering::Acquire, + ) { + Ok(_) => { + api.prevent_exit(); + let app_handle = app_handle.clone(); + let opds_service = app_handle + .state::() + .inner() + .clone(); + let exit_phase = exit_phase.clone(); + let exit_code = code.unwrap_or(0); + tauri::async_runtime::spawn(async move { + opds_service.stop().await; + exit_phase.store(EXITING, Ordering::Release); + app_handle.exit(exit_code); + }); + } + Err(STOPPING) => api.prevent_exit(), + Err(EXITING) => {} + Err(_) => unreachable!(), + } + } + }); Ok(()) } diff --git a/src-tauri/src/opds/commands.rs b/src-tauri/src/opds/commands.rs new file mode 100644 index 00000000..91cbd242 --- /dev/null +++ b/src-tauri/src/opds/commands.rs @@ -0,0 +1,38 @@ +use citadel_opds::{ + network::OpdsNetworkInterface, + service::{OpdsServiceStatus, OpdsStartConfig, OpdsStatusError}, + OpdsService, +}; + +#[tauri::command] +#[specta::specta] +pub async fn clb_query_opds_interfaces( + service: tauri::State<'_, OpdsService>, +) -> Result, OpdsStatusError> { + service.list_interfaces().await +} + +#[tauri::command] +#[specta::specta] +pub async fn clb_cmd_start_opds( + service: tauri::State<'_, OpdsService>, + config: OpdsStartConfig, +) -> Result { + service.start(config).await +} + +#[tauri::command] +#[specta::specta] +pub async fn clb_cmd_stop_opds( + service: tauri::State<'_, OpdsService>, +) -> Result { + Ok(service.stop().await) +} + +#[tauri::command] +#[specta::specta] +pub async fn clb_query_opds_status( + service: tauri::State<'_, OpdsService>, +) -> Result { + Ok(service.status().await) +} diff --git a/src-tauri/src/opds/mod.rs b/src-tauri/src/opds/mod.rs new file mode 100644 index 00000000..a0da9712 --- /dev/null +++ b/src-tauri/src/opds/mod.rs @@ -0,0 +1,3 @@ +pub(crate) mod commands; + +pub use citadel_opds::*; diff --git a/src-tauri/src/state.rs b/src-tauri/src/state.rs index 954aba8d..87a2301c 100644 --- a/src-tauri/src/state.rs +++ b/src-tauri/src/state.rs @@ -1,37 +1,71 @@ -use std::sync::Mutex; +use std::sync::{ + atomic::{AtomicBool, Ordering}, + Arc, Mutex, +}; use chrono::NaiveDateTime; use libcalibre::{BookId, BookPage, CalibreError, Library, ResolvedBookAsset}; use citadel_opds::CatalogSource; +#[derive(Clone)] pub struct CitadelState { + inner: Arc, +} + +struct CitadelStateInner { library: Mutex>, current_library_path: Mutex>, + library_switch: Mutex<()>, + library_transitioning: AtomicBool, } impl CitadelState { pub fn new() -> Self { Self { - library: Mutex::new(None), - current_library_path: Mutex::new(None), + inner: Arc::new(CitadelStateInner { + library: Mutex::new(None), + current_library_path: Mutex::new(None), + library_switch: Mutex::new(()), + library_transitioning: AtomicBool::new(false), + }), } } /// Initialize or switch to a library pub fn init_library(&self, library_path: String) -> Result<(), String> { - let db_path = libcalibre::util::get_db_path(&library_path) - .ok_or_else(|| format!("Invalid library path: {}", library_path))?; - - let lib = Library::new(db_path).map_err(|e| format!("Failed to open library: {}", e))?; - - *self.library.lock().expect("Library mutex poisoned") = Some(lib); - *self - .current_library_path + let _switch = self + .inner + .library_switch .lock() - .expect("Library path mutex poisoned") = Some(library_path); + .expect("Library switch mutex poisoned"); + { + let _library = self.inner.library.lock().expect("Library mutex poisoned"); + self.inner + .library_transitioning + .store(true, Ordering::Release); + } + let result = (|| { + let db_path = libcalibre::util::get_db_path(&library_path) + .ok_or_else(|| format!("Invalid library path: {}", library_path))?; + let library = Library::new(db_path) + .map_err(|error| format!("Failed to open library: {error}"))?; - Ok(()) + *self.inner.library.lock().expect("Library mutex poisoned") = Some(library); + *self + .inner + .current_library_path + .lock() + .expect("Library path mutex poisoned") = Some(library_path); + Ok(()) + })(); + { + let _library = self.inner.library.lock().expect("Library mutex poisoned"); + self.inner + .library_transitioning + .store(false, Ordering::Release); + } + result } /// Execute a function with mutable access to the library @@ -39,7 +73,7 @@ impl CitadelState { where F: FnOnce(&mut Library) -> R, { - let mut lib_guard = self.library.lock().expect("Library mutex poisoned"); + let mut lib_guard = self.inner.library.lock().expect("Library mutex poisoned"); match lib_guard.as_mut() { Some(lib) => Ok(f(lib)), None => Err("No library initialized. Please load a library first.".to_string()), @@ -48,7 +82,8 @@ impl CitadelState { /// Get the current library path pub fn get_library_path(&self) -> Option { - self.current_library_path + self.inner + .current_library_path .lock() .expect("Library path mutex poisoned") .clone() @@ -56,10 +91,23 @@ impl CitadelState { /// Check if a library is currently loaded pub fn is_initialized(&self) -> bool { - self.library - .lock() - .expect("Library mutex poisoned") - .is_some() + let library = self.inner.library.lock().expect("Library mutex poisoned"); + !self.is_library_transitioning() && library.is_some() + } + + pub fn is_library_transitioning(&self) -> bool { + self.inner.library_transitioning.load(Ordering::Acquire) + } + + pub fn active_library_id(&self) -> Result { + let mut library = self.inner.library.lock().expect("Library mutex poisoned"); + if self.is_library_transitioning() { + return Err(CalibreError::LibraryNotInitialized); + } + library + .as_mut() + .ok_or(CalibreError::LibraryNotInitialized)? + .library_uuid() } /// Resolve under the library mutex and return only owned data. Opening and @@ -69,7 +117,10 @@ impl CitadelState { book_id: BookId, format: &str, ) -> Result { - let mut library = self.library.lock().expect("Library mutex poisoned"); + let mut library = self.inner.library.lock().expect("Library mutex poisoned"); + if self.is_library_transitioning() { + return Err(CalibreError::LibraryNotInitialized); + } library .as_mut() .ok_or(CalibreError::LibraryNotInitialized)? @@ -77,7 +128,10 @@ impl CitadelState { } pub fn resolve_book_cover(&self, book_id: BookId) -> Result { - let mut library = self.library.lock().expect("Library mutex poisoned"); + let mut library = self.inner.library.lock().expect("Library mutex poisoned"); + if self.is_library_transitioning() { + return Err(CalibreError::LibraryNotInitialized); + } library .as_mut() .ok_or(CalibreError::LibraryNotInitialized)? @@ -89,7 +143,10 @@ impl CitadelState { limit: i64, offset: i64, ) -> Result<(String, Option, BookPage), CalibreError> { - let mut library = self.library.lock().expect("Library mutex poisoned"); + let mut library = self.inner.library.lock().expect("Library mutex poisoned"); + if self.is_library_transitioning() { + return Err(CalibreError::LibraryNotInitialized); + } let library = library .as_mut() .ok_or(CalibreError::LibraryNotInitialized)?; @@ -121,3 +178,135 @@ impl CatalogSource for CitadelState { self.resolve_book_cover(book_id) } } + +#[cfg(test)] +mod tests { + use std::{collections::HashMap, path::Path, time::Duration}; + + use axum::http::StatusCode; + + use super::*; + + fn extract_empty_library() -> tempfile::TempDir { + let directory = tempfile::tempdir().unwrap(); + let zip_path = + Path::new(env!("CARGO_MANIFEST_DIR")).join("resources/empty_7_2_calibre_lib.zip"); + let file = std::fs::File::open(zip_path).unwrap(); + zip::ZipArchive::new(file) + .unwrap() + .extract(directory.path()) + .unwrap(); + directory + } + + fn add_asset(directory: &Path, byte: u8, length: usize) -> BookId { + let root = directory.to_string_lossy().into_owned(); + let database = libcalibre::util::get_db_path(&root).unwrap(); + let mut library = libcalibre::Library::new(database).unwrap(); + let source = directory.join(format!("source-{byte}.epub")); + std::fs::write(&source, vec![byte; length]).unwrap(); + let book = library + .add_book(libcalibre::BookAdd { + title: format!("Library {byte}"), + author_names: vec!["Citadel Test".to_string()], + tags: None, + series: None, + series_index: None, + publisher: None, + publication_date: None, + rating: None, + comments: None, + identifiers: HashMap::new(), + language: Some("eng".to_string()), + file_paths: vec![source], + }) + .unwrap(); + book.id + } + + #[tokio::test(flavor = "multi_thread")] + async fn running_router_switches_library_at_same_url_and_is_unavailable_during_transition() { + let first = extract_empty_library(); + let second = extract_empty_library(); + let first_book = add_asset(first.path(), b'A', 2 * 1024 * 1024); + let second_book = add_asset(second.path(), b'B', 2 * 1024 * 1024); + assert_eq!(first_book, second_book); + let second_path = second.path().to_string_lossy().into_owned(); + let second_db = libcalibre::util::get_db_path(&second_path).unwrap(); + libcalibre::Library::new(second_db) + .unwrap() + .randomize_library_uuid() + .unwrap(); + + let state = CitadelState::new(); + state + .init_library(first.path().to_string_lossy().into_owned()) + .unwrap(); + let first_id = state.active_library_id().unwrap(); + let listener = tokio::net::TcpListener::bind((std::net::Ipv4Addr::LOCALHOST, 0)) + .await + .unwrap(); + let base = format!("http://{}", listener.local_addr().unwrap()); + let server_state = state.clone(); + let server = tokio::spawn(async move { + axum::serve(listener, citadel_opds::router(Arc::new(server_state))) + .await + .unwrap(); + }); + let client = reqwest::Client::new(); + + let first_response = client.get(format!("{base}/opds")).send().await.unwrap(); + assert_eq!(first_response.status(), StatusCode::OK); + assert!(first_response.text().await.unwrap().contains(&first_id)); + + let asset_url = format!( + "{base}/opds/books/{}/files/EPUB/book.epub", + first_book.as_i32() + ); + let old_download = client.get(&asset_url).send().await.unwrap(); + assert_eq!(old_download.status(), StatusCode::OK); + let responsive_feed = tokio::time::timeout( + Duration::from_secs(1), + client.get(format!("{base}/opds")).send(), + ) + .await + .expect("a stalled large download must not block feed requests") + .unwrap(); + assert_eq!(responsive_feed.status(), StatusCode::OK); + + state + .inner + .library_transitioning + .store(true, Ordering::Release); + let transition_response = reqwest::get(format!("{base}/opds")).await.unwrap(); + assert_eq!( + transition_response.status(), + StatusCode::SERVICE_UNAVAILABLE + ); + state + .inner + .library_transitioning + .store(false, Ordering::Release); + + state.init_library(second_path).unwrap(); + let second_id = state.active_library_id().unwrap(); + assert_ne!(first_id, second_id); + let second_response = client.get(format!("{base}/opds")).send().await.unwrap(); + assert_eq!(second_response.status(), StatusCode::OK); + assert!(second_response.text().await.unwrap().contains(&second_id)); + let new_bytes = client + .get(&asset_url) + .send() + .await + .unwrap() + .bytes() + .await + .unwrap(); + assert_eq!(new_bytes.len(), 2 * 1024 * 1024); + assert!(new_bytes.iter().all(|byte| *byte == b'B')); + let old_bytes = old_download.bytes().await.unwrap(); + assert_eq!(old_bytes.len(), 2 * 1024 * 1024); + assert!(old_bytes.iter().all(|byte| *byte == b'A')); + server.abort(); + } +} diff --git a/src/bindings.ts b/src/bindings.ts index 7d82df85..052387af 100644 --- a/src/bindings.ts +++ b/src/bindings.ts @@ -289,6 +289,38 @@ async clbCmdInstallUpdateIfAvailable() : Promise> { else return { status: "error", error: e as any }; } }, +async clbQueryOpdsInterfaces() : Promise> { + try { + return { status: "ok", data: await TAURI_INVOKE("clb_query_opds_interfaces") }; +} catch (e) { + if(e instanceof Error) throw e; + else return { status: "error", error: e as any }; +} +}, +async clbCmdStartOpds(config: OpdsStartConfig) : Promise> { + try { + return { status: "ok", data: await TAURI_INVOKE("clb_cmd_start_opds", { config }) }; +} catch (e) { + if(e instanceof Error) throw e; + else return { status: "error", error: e as any }; +} +}, +async clbCmdStopOpds() : Promise> { + try { + return { status: "ok", data: await TAURI_INVOKE("clb_cmd_stop_opds") }; +} catch (e) { + if(e instanceof Error) throw e; + else return { status: "error", error: e as any }; +} +}, +async clbQueryOpdsStatus() : Promise> { + try { + return { status: "ok", data: await TAURI_INVOKE("clb_query_opds_status") }; +} catch (e) { + if(e instanceof Error) throw e; + else return { status: "error", error: e as any }; +} +}, async clbCmdOpenSettings() : Promise { await TAURI_INVOKE("clb_cmd_open_settings"); } @@ -490,6 +522,15 @@ export type LocalOrRemoteUrl = { kind: LocalOrRemote; url: string; local_path: s */ export type MetadataProvider = "hardcover" | "loc" | "dnb" | "k10plus" | "openlibrary" export type NewAuthor = { name: string; sortable_name: string | null } +export type OpdsBindTarget = { type: "allLocalNetworks" } | { type: "interface"; id: string } +export type OpdsErrorCode = "invalidPort" | "libraryNotReady" | "configurationConflict" | "interfaceUnavailable" | "interfaceEnumerationFailed" | "portUnavailable" | "listenerFailed" | "unexpected" +export type OpdsInterfaceKind = "lan" | "vpn" | "loopback" | "other" +export type OpdsInterfaceState = "up" | "down" +export type OpdsLifecycleState = "stopped" | "starting" | "running" | "waitingForInterface" | "error" | "stopping" +export type OpdsNetworkInterface = { id: string; label: string; kind: OpdsInterfaceKind; state: OpdsInterfaceState; addresses: string[]; shareable: boolean } +export type OpdsServiceStatus = { state: OpdsLifecycleState; target: OpdsBindTarget | null; port: number | null; activeLibraryId: string | null; urls: string[]; error: OpdsStatusError | null } +export type OpdsStartConfig = { target: OpdsBindTarget; port: number } +export type OpdsStatusError = { code: OpdsErrorCode; message: string } export type ProviderStatus = { provider: MetadataProvider; is_valid: boolean; message: string } export type RemoteFile = { url: string } /**