Skip to main content

confik/
reloading.rs

1//! Hot-reloadable configuration support.
2//!
3//! This module provides [`ReloadingConfig`], which wraps a configuration value
4//! and allows it to be atomically reloaded at runtime.
5//!
6//! # Examples
7//!
8//! ```rust
9//! # #[cfg(feature = "toml")]
10//! # {
11//! use confik::{Configuration, ReloadableConfig, TomlSource};
12//!
13//! #[derive(Debug, Configuration)]
14//! struct AppConfig {
15//!     port: u16,
16//!     host: String,
17//! }
18//!
19//! impl ReloadableConfig for AppConfig {
20//!     type Error = confik::Error;
21//!
22//!     fn build() -> Result<Self, Self::Error> {
23//!         Self::builder()
24//!             .override_with(TomlSource::new(r#"port = 8080
25//! host = "localhost""#))
26//!             .try_build()
27//!     }
28//! }
29//!
30//! // Create a reloading config (no turbofish needed!)
31//! let config = AppConfig::reloading().unwrap();
32//!
33//! // Access the current config
34//! let current = config.load();
35//! assert_eq!(current.port, 8080);
36//!
37//! // Reload from sources
38//! config.reload().unwrap();
39//! # }
40//! ```
41
42use std::sync::Arc;
43
44use arc_swap::ArcSwap;
45
46/// Trait for invoking reload callbacks.
47///
48/// This trait allows both `()` (no callback) and `Fn()` types to be used
49/// as the callback type in `ReloadingConfig`.
50pub trait ReloadCallback {
51    /// Invokes the callback, if any.
52    fn invoke(&self);
53}
54
55impl ReloadCallback for () {
56    fn invoke(&self) {
57        // No-op for unit type
58    }
59}
60
61impl<F: Fn()> ReloadCallback for F {
62    fn invoke(&self) {
63        self()
64    }
65}
66
67/// Defines how to create a new instance of [`ReloadingConfig`].
68///
69/// This trait is typically implemented for configuration types that need to support
70/// hot-reloading. It specifies how to build a fresh instance of the configuration
71/// from its sources.
72pub trait ReloadableConfig: Sized {
73    /// The error type returned when building fails.
74    type Error;
75
76    /// Defines the way to build the configuration item.
77    ///
78    /// This method should include all the logic needed to construct the configuration
79    /// from its sources, including any required validations.
80    ///
81    /// # Examples
82    ///
83    /// ```rust
84    /// # #[cfg(feature = "toml")]
85    /// # {
86    /// # use confik::{Configuration, ReloadableConfig, TomlSource};
87    /// # #[derive(Debug, Configuration)]
88    /// # struct MyConfig { value: String }
89    /// impl ReloadableConfig for MyConfig {
90    ///     type Error = confik::Error;
91    ///
92    ///     fn build() -> Result<Self, Self::Error> {
93    ///         Self::builder()
94    ///             .override_with(TomlSource::new(r#"value = "test""#))
95    ///             .try_build()
96    ///     }
97    /// }
98    /// # }
99    /// ```
100    fn build() -> Result<Self, Self::Error>;
101
102    /// Creates a new [`ReloadingConfig`] for this configuration type.
103    ///
104    /// This is a convenience method that avoids needing to specify type parameters.
105    ///
106    /// # Errors
107    ///
108    /// Returns an error if the initial configuration build fails.
109    ///
110    /// # Examples
111    ///
112    /// ```rust
113    /// # #[cfg(feature = "toml")]
114    /// # {
115    /// # use confik::{Configuration, ReloadableConfig, TomlSource};
116    /// # #[derive(Debug, Configuration)]
117    /// # struct MyConfig { value: String }
118    /// # impl ReloadableConfig for MyConfig {
119    /// #     type Error = confik::Error;
120    /// #     fn build() -> Result<Self, Self::Error> {
121    /// #         Self::builder().override_with(TomlSource::new(r#"value = "test""#)).try_build()
122    /// #     }
123    /// # }
124    /// // Much cleaner than ReloadingConfig::<MyConfig, _>::build()
125    /// let config = MyConfig::reloading().unwrap();
126    /// # }
127    /// ```
128    fn reloading() -> Result<ReloadingConfig<Self, ()>, Self::Error> {
129        ReloadingConfig::build()
130    }
131}
132
133/// An instance of config that may reload itself.
134///
135/// This struct wraps a configuration value and allows it to be atomically swapped
136/// with a newly-loaded version. Cloning this object is cheap as it only clones
137/// the underlying `Arc` pointers.
138///
139/// # Type Parameters
140///
141/// * `T` - The configuration type that implements [`ReloadableConfig`]
142/// * `F` - The type of the callback invoked after successful reloads (defaults to `()`), see [`ReloadCallback`]
143#[derive(Debug)]
144pub struct ReloadingConfig<T, F> {
145    config: Arc<ArcSwap<T>>,
146    on_update: F,
147}
148
149impl<T, F> Clone for ReloadingConfig<T, F>
150where
151    F: Clone,
152{
153    fn clone(&self) -> Self {
154        ReloadingConfig {
155            config: Arc::clone(&self.config),
156            on_update: self.on_update.clone(),
157        }
158    }
159}
160
161impl<T> ReloadingConfig<T, ()>
162where
163    T: ReloadableConfig,
164{
165    /// Creates a new [`ReloadingConfig`] by building the initial configuration.
166    ///
167    /// # Errors
168    ///
169    /// Returns an error if the initial configuration build fails.
170    pub fn build() -> Result<Self, <T as ReloadableConfig>::Error> {
171        Ok(ReloadingConfig {
172            config: Arc::new(ArcSwap::new(Arc::new(T::build()?))),
173            on_update: (),
174        })
175    }
176}
177
178impl<T, F> ReloadingConfig<T, F> {
179    /// Replaces the update callback with a new one.
180    ///
181    /// See [`ReloadCallback`].
182    ///
183    /// # Examples
184    ///
185    /// ```rust
186    /// # #[cfg(feature = "toml")]
187    /// # {
188    /// # use confik::{Configuration, ReloadableConfig, ReloadingConfig, TomlSource};
189    /// # #[derive(Debug, Configuration)]
190    /// # struct MyConfig { value: String }
191    /// # impl ReloadableConfig for MyConfig {
192    /// #     type Error = confik::Error;
193    /// #     fn build() -> Result<Self, Self::Error> {
194    /// #         Self::builder().override_with(TomlSource::new(r#"value = "test""#)).try_build()
195    /// #     }
196    /// # }
197    /// let config = ReloadingConfig::<MyConfig, _>::build().unwrap()
198    ///     .with_on_update(|| println!("Config reloaded!"));
199    /// # }
200    /// ```
201    #[must_use]
202    pub fn with_on_update<U>(self, new: U) -> ReloadingConfig<T, U> {
203        ReloadingConfig {
204            config: self.config,
205            on_update: new,
206        }
207    }
208
209    /// Loads the current configuration.
210    ///
211    /// Returns an `Arc` to the current configuration value. This is a cheap operation
212    /// that doesn't block writers.
213    #[must_use]
214    pub fn load(&self) -> Arc<T> {
215        self.config.load_full()
216    }
217}
218
219impl<T, F> ReloadingConfig<T, F>
220where
221    T: ReloadableConfig,
222    F: ReloadCallback,
223{
224    /// Attempts to reload the configuration.
225    ///
226    /// On success, calls the stored update function (if any).
227    /// On error, leaves the configuration unchanged.
228    ///
229    /// # Errors
230    ///
231    /// Returns an error if building the new configuration fails. In this case,
232    /// the current configuration remains unchanged and the update callback is
233    /// not invoked.
234    pub fn reload(&self) -> Result<(), <T as ReloadableConfig>::Error> {
235        let config = T::build()?;
236        self.config.store(Arc::new(config));
237        self.on_update.invoke();
238        Ok(())
239    }
240}
241
242#[cfg(feature = "signal")]
243impl<T, F> ReloadingConfig<T, F>
244where
245    T: ReloadableConfig + Send + Sync + 'static,
246    F: ReloadCallback + Clone + Send + Sync + 'static,
247{
248    /// Sets a listener for SIGHUP.
249    ///
250    /// This spawns a thread and listens for a signal using the [`signal_hook`] crate,
251    /// with all of that crate's caveats. If you're setting signals already, you may wish to
252    /// configure this yourself.
253    ///
254    /// When a SIGHUP signal is received, the configuration will be reloaded. If the reload
255    /// fails and the `tracing` feature is enabled, an error will be logged.
256    ///
257    /// # Errors
258    ///
259    /// Returns an error if signal registration fails.
260    ///
261    /// # Examples
262    ///
263    /// ```rust,no_run
264    /// # #[cfg(all(feature = "signal", feature = "toml"))]
265    /// # {
266    /// # use confik::{Configuration, ReloadableConfig, ReloadingConfig, TomlSource};
267    /// # #[derive(Debug, Configuration)]
268    /// # struct MyConfig { value: String }
269    /// # impl ReloadableConfig for MyConfig {
270    /// #     type Error = confik::Error;
271    /// #     fn build() -> Result<Self, Self::Error> {
272    /// #         Self::builder().override_with(TomlSource::new(r#"value = "test""#)).try_build()
273    /// #     }
274    /// # }
275    /// let config = ReloadingConfig::<MyConfig, _>::build().unwrap();
276    ///
277    /// // Set up signal handler
278    /// let handle = config.spawn_signal_handler().unwrap();
279    ///
280    /// // The config will now reload when receiving SIGHUP
281    /// // handle.join().unwrap(); // Wait for the signal handler thread
282    /// # }
283    /// ```
284    pub fn spawn_signal_handler(&self) -> Result<std::thread::JoinHandle<()>, std::io::Error>
285    where
286        <T as ReloadableConfig>::Error: std::fmt::Display,
287    {
288        use signal_hook::{consts::SIGHUP, iterator::Signals};
289
290        let mut signals = Signals::new([SIGHUP])?;
291        let config = self.clone();
292        Ok(std::thread::spawn(move || {
293            for signal in &mut signals {
294                if signal == SIGHUP {
295                    if let Err(err) = config.reload() {
296                        #[cfg(feature = "tracing")]
297                        tracing::error!(%err, "Failed to reload configuration");
298
299                        #[cfg(not(feature = "tracing"))]
300                        {
301                            // Avoid unused variable warning
302                            let _ = err;
303                        }
304                    }
305                }
306            }
307        }))
308    }
309}
310
311#[cfg(test)]
312mod tests {
313    use super::*;
314    use crate::Configuration;
315
316    #[derive(Debug, Clone, PartialEq, Configuration)]
317    struct TestConfig {
318        value: u32,
319    }
320
321    impl ReloadableConfig for TestConfig {
322        type Error = &'static str;
323
324        fn build() -> Result<Self, Self::Error> {
325            Ok(TestConfig { value: 42 })
326        }
327    }
328
329    #[test]
330    fn test_build_and_load() {
331        let config = ReloadingConfig::<TestConfig, _>::build().unwrap();
332        let current = config.load();
333        assert_eq!(current.value, 42);
334    }
335
336    #[test]
337    fn test_reload_without_callback() {
338        let config = TestConfig::reloading().unwrap();
339        config.reload().unwrap();
340        let current = config.load();
341        assert_eq!(current.value, 42);
342    }
343
344    #[test]
345    fn test_reload_with_callback() {
346        use std::sync::atomic::{AtomicBool, Ordering};
347
348        let called = Arc::new(AtomicBool::new(false));
349        let called_clone = Arc::clone(&called);
350
351        let config = TestConfig::reloading().unwrap().with_on_update(move || {
352            called_clone.store(true, Ordering::SeqCst);
353        });
354
355        assert!(!called.load(Ordering::SeqCst));
356        config.reload().unwrap();
357        assert!(called.load(Ordering::SeqCst));
358    }
359
360    #[test]
361    fn test_reload_updates_all_clones() {
362        use std::sync::atomic::{AtomicU32, Ordering};
363
364        static COUNTER: AtomicU32 = AtomicU32::new(0);
365
366        #[derive(Debug, serde::Deserialize, Configuration)]
367        struct CountingConfig {
368            id: u32,
369        }
370
371        impl ReloadableConfig for CountingConfig {
372            type Error = std::convert::Infallible;
373
374            fn build() -> Result<Self, Self::Error> {
375                Ok(CountingConfig {
376                    id: COUNTER.fetch_add(1, Ordering::SeqCst),
377                })
378            }
379        }
380
381        let config1 = CountingConfig::reloading().unwrap();
382        let config2 = config1.clone();
383
384        assert_eq!(config1.load().id, 0);
385        assert_eq!(config2.load().id, 0);
386
387        config1.reload().unwrap();
388
389        assert_eq!(config1.load().id, 1);
390        assert_eq!(config2.load().id, 1);
391    }
392
393    #[test]
394    fn test_reload_error_leaves_config_unchanged() {
395        use std::sync::atomic::{AtomicBool, Ordering};
396
397        static SHOULD_FAIL: AtomicBool = AtomicBool::new(false);
398
399        #[derive(Debug, serde::Deserialize, Configuration)]
400        struct FallibleConfig {
401            value: u32,
402        }
403
404        impl ReloadableConfig for FallibleConfig {
405            type Error = &'static str;
406
407            fn build() -> Result<Self, Self::Error> {
408                if SHOULD_FAIL.load(Ordering::SeqCst) {
409                    Err("Build failed")
410                } else {
411                    Ok(FallibleConfig { value: 42 })
412                }
413            }
414        }
415
416        let config = FallibleConfig::reloading().unwrap();
417        assert_eq!(config.load().value, 42);
418
419        // Make the next build fail
420        SHOULD_FAIL.store(true, Ordering::SeqCst);
421
422        // Reload should fail and leave config unchanged
423        let result = config.reload();
424        assert!(result.is_err());
425        assert_eq!(config.load().value, 42); // Still the old value
426
427        // Make build succeed again
428        SHOULD_FAIL.store(false, Ordering::SeqCst);
429        config.reload().unwrap();
430        assert_eq!(config.load().value, 42);
431    }
432
433    #[test]
434    fn test_callback_not_invoked_on_reload_error() {
435        use std::sync::atomic::{AtomicBool, Ordering};
436
437        static SHOULD_FAIL: AtomicBool = AtomicBool::new(false);
438
439        #[derive(Debug, serde::Deserialize, Configuration)]
440        struct FallibleConfig {
441            value: u32,
442        }
443
444        impl ReloadableConfig for FallibleConfig {
445            type Error = &'static str;
446
447            fn build() -> Result<Self, Self::Error> {
448                if SHOULD_FAIL.load(Ordering::SeqCst) {
449                    Err("Build failed")
450                } else {
451                    Ok(FallibleConfig { value: 100 })
452                }
453            }
454        }
455
456        let callback_called = Arc::new(AtomicBool::new(false));
457        let callback_called_clone = Arc::clone(&callback_called);
458
459        let config = FallibleConfig::reloading()
460            .unwrap()
461            .with_on_update(move || {
462                callback_called_clone.store(true, Ordering::SeqCst);
463            });
464
465        // Initial value should be 100
466        assert_eq!(config.load().value, 100);
467
468        // Successful reload should call callback
469        config.reload().unwrap();
470        assert!(callback_called.load(Ordering::SeqCst));
471
472        // Reset flag
473        callback_called.store(false, Ordering::SeqCst);
474
475        // Make next build fail
476        SHOULD_FAIL.store(true, Ordering::SeqCst);
477
478        // Failed reload should NOT call callback
479        let result = config.reload();
480        assert!(result.is_err());
481        assert!(!callback_called.load(Ordering::SeqCst));
482    }
483}