Skip to main content

praxis_protocol/
pipelines.rs

1// SPDX-License-Identifier: Apache-2.0
2// Copyright (c) 2024 Praxis Contributors
3
4//! Hot-swappable pipeline storage for protocol adapters.
5//!
6//! [`ListenerPipelines`] lives in the protocol crate because it is
7//! the interface between protocol adapters (which invoke filter
8//! execution) and the filter engine (which owns [`FilterPipeline`]).
9//!
10//! Each pipeline is wrapped in [`ArcSwap`] for lock-free atomic
11//! replacement during hot reloads. In-flight requests hold an
12//! [`Arc`] guard to the old pipeline, so they drain safely while
13//! new requests pick up the replacement.
14//!
15//! [`FilterPipeline`]: praxis_filter::FilterPipeline
16//! [`ArcSwap`]: arc_swap::ArcSwap
17
18use std::{collections::HashMap, sync::Arc};
19
20use arc_swap::ArcSwap;
21use praxis_filter::FilterPipeline;
22
23// -----------------------------------------------------------------------------
24// ListenerPipelines
25// -----------------------------------------------------------------------------
26
27/// Maps listener names to their resolved [`FilterPipeline`]s.
28///
29/// Each pipeline is wrapped in [`ArcSwap`] so it can be atomically
30/// replaced at runtime without blocking in-flight requests.
31///
32/// ```
33/// use std::{collections::HashMap, sync::Arc};
34///
35/// use praxis_filter::{FilterPipeline, FilterRegistry};
36/// use praxis_protocol::ListenerPipelines;
37///
38/// let registry = FilterRegistry::with_builtins();
39/// let pipeline = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
40///
41/// let mut map = HashMap::new();
42/// map.insert("web".to_owned(), pipeline);
43/// let pipelines = ListenerPipelines::new(map);
44///
45/// assert!(pipelines.get("web").is_some());
46/// assert!(pipelines.get("missing").is_none());
47/// ```
48///
49/// [`ArcSwap`]: arc_swap::ArcSwap
50pub struct ListenerPipelines {
51    /// Maps listener names to their swappable filter pipelines.
52    pipelines: HashMap<String, Arc<ArcSwap<FilterPipeline>>>,
53}
54
55impl ListenerPipelines {
56    /// Create from a map of listener name to pipeline.
57    ///
58    /// Called by the server during startup after building filter pipelines
59    /// from the loaded configuration. Each pipeline is wrapped in [`ArcSwap`]
60    /// for atomic replacement during hot reloads.
61    pub fn new(pipelines: HashMap<String, Arc<FilterPipeline>>) -> Self {
62        let swappable = pipelines
63            .into_iter()
64            .map(|(name, p)| (name, Arc::new(ArcSwap::from(p))))
65            .collect();
66        Self { pipelines: swappable }
67    }
68
69    /// Get the swappable pipeline for a listener by name.
70    ///
71    /// Called by protocol adapters on every request to access the filter
72    /// pipeline. The returned [`ArcSwap`] reference allows protocol adapters
73    /// to load the current pipeline without blocking reload operations.
74    pub fn get(&self, listener_name: &str) -> Option<&Arc<ArcSwap<FilterPipeline>>> {
75        self.pipelines.get(listener_name)
76    }
77
78    /// Atomically replace the pipeline for a listener.
79    ///
80    /// No-op if the listener name is not present.
81    ///
82    /// ```
83    /// use std::{collections::HashMap, sync::Arc};
84    ///
85    /// use praxis_filter::{FilterPipeline, FilterRegistry};
86    /// use praxis_protocol::ListenerPipelines;
87    ///
88    /// let registry = FilterRegistry::with_builtins();
89    /// let old = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
90    /// let new = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
91    ///
92    /// let mut map = HashMap::new();
93    /// map.insert("web".to_owned(), old);
94    /// let pipelines = ListenerPipelines::new(map);
95    ///
96    /// pipelines.swap("web", new);
97    /// pipelines.swap(
98    ///     "nonexistent",
99    ///     Arc::new(FilterPipeline::build(&mut [], &registry).unwrap()),
100    /// );
101    /// ```
102    pub fn swap(&self, listener_name: &str, new_pipeline: Arc<FilterPipeline>) {
103        if let Some(slot) = self.pipelines.get(listener_name) {
104            slot.store(new_pipeline);
105        }
106    }
107
108    /// Every filesystem path any filter in any listener's pipeline reads
109    /// configuration from, de-duplicated.
110    ///
111    /// Two listeners can share a filter chain, so the same document would
112    /// otherwise appear more than once and be watched and hashed repeatedly.
113    pub fn referenced_files(&self) -> Vec<std::path::PathBuf> {
114        let mut seen = std::collections::BTreeSet::new();
115        for name in self.listener_names() {
116            if let Some(slot) = self.get(name) {
117                for path in slot.load().referenced_files() {
118                    seen.insert(path);
119                }
120            }
121        }
122        seen.into_iter().collect()
123    }
124
125    /// Returns an iterator over listener names.
126    ///
127    /// Used during config reload to iterate over all listeners when
128    /// swapping pipelines or collecting referenced files for watching.
129    pub fn listener_names(&self) -> impl Iterator<Item = &str> {
130        self.pipelines.keys().map(String::as_str)
131    }
132}
133
134// -----------------------------------------------------------------------------
135// Tests
136// -----------------------------------------------------------------------------
137
138#[cfg(test)]
139#[expect(clippy::allow_attributes, reason = "blanket test suppressions")]
140#[allow(
141    clippy::unwrap_used,
142    clippy::expect_used,
143    clippy::indexing_slicing,
144    clippy::too_many_lines,
145    reason = "tests"
146)]
147mod tests {
148    use praxis_filter::FilterRegistry;
149
150    use super::*;
151
152    #[test]
153    fn get_returns_pipeline() {
154        let pipelines = make_pipelines(&["web"]);
155        assert!(pipelines.get("web").is_some(), "should find 'web' pipeline");
156    }
157
158    #[test]
159    fn get_returns_none_for_missing() {
160        let pipelines = make_pipelines(&["web"]);
161        assert!(pipelines.get("missing").is_none(), "should return None for missing");
162    }
163
164    #[test]
165    fn swap_replaces_pipeline_pointer() {
166        let pipelines = make_pipelines(&["web"]);
167        let old_ptr = Arc::as_ptr(&pipelines.get("web").unwrap().load());
168
169        let registry = FilterRegistry::with_builtins();
170        let new_pipeline = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
171        pipelines.swap("web", Arc::clone(&new_pipeline));
172
173        let new_ptr = Arc::as_ptr(&pipelines.get("web").unwrap().load());
174        assert_ne!(old_ptr, new_ptr, "swap should replace the pipeline pointer");
175    }
176
177    #[test]
178    fn old_guard_remains_valid_after_swap() {
179        let pipelines = make_pipelines(&["web"]);
180        let old_guard = pipelines.get("web").unwrap().load();
181        let old_ptr = Arc::as_ptr(&old_guard);
182
183        let registry = FilterRegistry::with_builtins();
184        let new_pipeline = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
185        pipelines.swap("web", new_pipeline);
186
187        let still_old_ptr = Arc::as_ptr(&old_guard);
188        assert_eq!(
189            old_ptr, still_old_ptr,
190            "old guard should still point to the original pipeline"
191        );
192    }
193
194    #[test]
195    fn swap_nonexistent_is_noop() {
196        let pipelines = make_pipelines(&["web"]);
197        let registry = FilterRegistry::with_builtins();
198        let new_pipeline = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
199        pipelines.swap("nonexistent", new_pipeline);
200        assert!(pipelines.get("web").is_some(), "existing pipeline should be unaffected");
201    }
202
203    #[test]
204    fn get_returns_arcswap_reference() {
205        let pipelines = make_pipelines(&["web"]);
206        let slot: &Arc<ArcSwap<FilterPipeline>> = pipelines.get("web").unwrap();
207        let _loaded: arc_swap::Guard<Arc<FilterPipeline>> = slot.load();
208    }
209
210    #[test]
211    fn referenced_files_empty_without_listeners() {
212        let pipelines = make_pipelines(&[]);
213        assert!(
214            pipelines.referenced_files().is_empty(),
215            "no listeners means no referenced documents"
216        );
217    }
218
219    #[test]
220    fn referenced_files_empty_when_no_filter_declares_one() {
221        let pipelines = make_pipelines(&["web"]);
222        assert!(
223            pipelines.referenced_files().is_empty(),
224            "a pipeline of non-declaring filters contributes nothing"
225        );
226    }
227
228    #[test]
229    fn referenced_files_collects_across_listeners() {
230        let pipelines =
231            make_pipelines_with_documents(&[("web", "/etc/praxis/web.yaml"), ("api", "/etc/praxis/api.yaml")]);
232        assert_eq!(
233            pipelines.referenced_files(),
234            vec![
235                std::path::PathBuf::from("/etc/praxis/api.yaml"),
236                std::path::PathBuf::from("/etc/praxis/web.yaml"),
237            ],
238            "every listener's documents must be collected, sorted by the BTreeSet"
239        );
240    }
241
242    #[test]
243    fn referenced_files_dedupes_a_document_shared_by_two_listeners() {
244        let shared = "/etc/praxis/shared.yaml";
245        let pipelines = make_pipelines_with_documents(&[("web", shared), ("api", shared)]);
246        assert_eq!(
247            pipelines.referenced_files(),
248            vec![std::path::PathBuf::from(shared)],
249            "a shared document must appear once"
250        );
251    }
252
253    // -------------------------------------------------------------------------
254    // Test Utilities
255    // -------------------------------------------------------------------------
256
257    /// A filter that declares the document named by its `document:` config key.
258    struct DocumentReaderFilter {
259        document: std::path::PathBuf,
260    }
261
262    #[async_trait::async_trait]
263    impl praxis_filter::HttpFilter for DocumentReaderFilter {
264        fn name(&self) -> &'static str {
265            "document_reader"
266        }
267
268        fn referenced_files(&self) -> Vec<std::path::PathBuf> {
269            vec![self.document.clone()]
270        }
271
272        async fn on_request(
273            &self,
274            _ctx: &mut praxis_filter::HttpFilterContext<'_>,
275        ) -> Result<praxis_filter::FilterAction, praxis_filter::FilterError> {
276            Ok(praxis_filter::FilterAction::Continue)
277        }
278    }
279
280    /// Build [`ListenerPipelines`] where each named listener runs a single filter
281    /// declaring the given document.
282    fn make_pipelines_with_documents(listeners: &[(&str, &str)]) -> ListenerPipelines {
283        let mut registry = FilterRegistry::with_builtins();
284        registry
285            .register(
286                "document_reader",
287                praxis_filter::FilterFactory::Http(Arc::new(|cfg: &serde_yaml::Value| {
288                    let document = cfg
289                        .get("document")
290                        .and_then(serde_yaml::Value::as_str)
291                        .ok_or_else(|| praxis_filter::FilterError::from("document_reader: missing document"))?;
292                    let filter: Box<dyn praxis_filter::HttpFilter> = Box::new(DocumentReaderFilter {
293                        document: std::path::PathBuf::from(document),
294                    });
295                    Ok(filter)
296                })),
297            )
298            .unwrap();
299
300        let mut map = HashMap::new();
301        for (listener, document) in listeners {
302            let yaml = format!("- filter: document_reader\n  document: {document}\n");
303            let mut entries: Vec<praxis_core::config::FilterEntry> = serde_yaml::from_str(&yaml).unwrap();
304            let pipeline = Arc::new(FilterPipeline::build(&mut entries, &registry).unwrap());
305            map.insert((*listener).to_owned(), pipeline);
306        }
307        ListenerPipelines::new(map)
308    }
309
310    /// Build [`ListenerPipelines`] with empty pipelines for the given names.
311    fn make_pipelines(names: &[&str]) -> ListenerPipelines {
312        let registry = FilterRegistry::with_builtins();
313        let mut map = HashMap::new();
314        for name in names {
315            let pipeline = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
316            map.insert((*name).to_owned(), pipeline);
317        }
318        ListenerPipelines::new(map)
319    }
320}