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_core::config::ProtocolKind;
22use praxis_filter::FilterPipeline;
23
24// -----------------------------------------------------------------------------
25// ListenerPipelines
26// -----------------------------------------------------------------------------
27
28/// Maps listener names to their resolved [`FilterPipeline`]s.
29///
30/// Each pipeline is wrapped in [`ArcSwap`] so it can be atomically
31/// replaced at runtime without blocking in-flight requests.
32///
33/// ```
34/// use std::{collections::HashMap, sync::Arc};
35///
36/// use praxis_filter::{FilterPipeline, FilterRegistry};
37/// use praxis_protocol::ListenerPipelines;
38///
39/// let registry = FilterRegistry::with_builtins();
40/// let pipeline = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
41///
42/// let mut map = HashMap::new();
43/// map.insert("web".to_owned(), pipeline);
44/// let pipelines = ListenerPipelines::new(map);
45///
46/// assert!(pipelines.get("web").is_some());
47/// assert!(pipelines.get("missing").is_none());
48/// ```
49///
50/// [`ArcSwap`]: arc_swap::ArcSwap
51pub struct ListenerPipelines {
52    /// Maps listener names to their swappable filter pipelines.
53    pipelines: HashMap<String, Arc<ArcSwap<FilterPipeline>>>,
54    /// The protocol each listener's pipelines are resolved for.
55    ///
56    /// A protocol handler is bound once, for the process lifetime, and
57    /// executes only filters of its own protocol. Swaps replace the
58    /// pipeline in a slot but never this record, so a reload can tell
59    /// whether a rebuilt pipeline still matches the handler it targets.
60    protocols: HashMap<String, ProtocolKind>,
61}
62
63impl ListenerPipelines {
64    /// Create from a map of listener name to pipeline.
65    ///
66    /// Called by the server during startup after building filter pipelines
67    /// from the loaded configuration. Each pipeline is wrapped in [`ArcSwap`]
68    /// for atomic replacement during hot reloads.
69    pub fn new(pipelines: HashMap<String, Arc<FilterPipeline>>) -> Self {
70        Self::with_protocols(pipelines, HashMap::new())
71    }
72
73    /// Create from a map of listener name to pipeline, recording the
74    /// protocol each listener's pipeline was resolved for.
75    ///
76    /// ```
77    /// use std::{collections::HashMap, sync::Arc};
78    ///
79    /// use praxis_core::config::ProtocolKind;
80    /// use praxis_filter::{FilterPipeline, FilterRegistry};
81    /// use praxis_protocol::ListenerPipelines;
82    ///
83    /// let registry = FilterRegistry::with_builtins();
84    /// let pipeline = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
85    ///
86    /// let mut map = HashMap::new();
87    /// map.insert("db".to_owned(), pipeline);
88    /// let mut protocols = HashMap::new();
89    /// protocols.insert("db".to_owned(), ProtocolKind::Tcp);
90    /// let pipelines = ListenerPipelines::with_protocols(map, protocols);
91    ///
92    /// assert_eq!(pipelines.protocol("db"), Some(ProtocolKind::Tcp));
93    /// ```
94    pub fn with_protocols(
95        pipelines: HashMap<String, Arc<FilterPipeline>>,
96        protocols: HashMap<String, ProtocolKind>,
97    ) -> Self {
98        let swappable = pipelines
99            .into_iter()
100            .map(|(name, p)| (name, Arc::new(ArcSwap::from(p))))
101            .collect();
102        Self {
103            pipelines: swappable,
104            protocols,
105        }
106    }
107
108    /// The protocol a listener's pipelines are resolved for, if recorded.
109    ///
110    /// Fixed at construction: a swap never changes it.
111    pub fn protocol(&self, listener_name: &str) -> Option<ProtocolKind> {
112        self.protocols.get(listener_name).copied()
113    }
114
115    /// Get the swappable pipeline for a listener by name.
116    ///
117    /// Called by protocol adapters on every request to access the filter
118    /// pipeline. The returned [`ArcSwap`] reference allows protocol adapters
119    /// to load the current pipeline without blocking reload operations.
120    pub fn get(&self, listener_name: &str) -> Option<&Arc<ArcSwap<FilterPipeline>>> {
121        self.pipelines.get(listener_name)
122    }
123
124    /// Atomically replace the pipeline for a listener.
125    ///
126    /// No-op if the listener name is not present.
127    ///
128    /// ```
129    /// use std::{collections::HashMap, sync::Arc};
130    ///
131    /// use praxis_filter::{FilterPipeline, FilterRegistry};
132    /// use praxis_protocol::ListenerPipelines;
133    ///
134    /// let registry = FilterRegistry::with_builtins();
135    /// let old = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
136    /// let new = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
137    ///
138    /// let mut map = HashMap::new();
139    /// map.insert("web".to_owned(), old);
140    /// let pipelines = ListenerPipelines::new(map);
141    ///
142    /// pipelines.swap("web", new);
143    /// pipelines.swap(
144    ///     "nonexistent",
145    ///     Arc::new(FilterPipeline::build(&mut [], &registry).unwrap()),
146    /// );
147    /// ```
148    pub fn swap(&self, listener_name: &str, new_pipeline: Arc<FilterPipeline>) {
149        if let Some(slot) = self.pipelines.get(listener_name) {
150            slot.store(new_pipeline);
151        }
152    }
153
154    /// Every filesystem path any filter in any listener's pipeline reads
155    /// configuration from, de-duplicated.
156    ///
157    /// Two listeners can share a filter chain, so the same document would
158    /// otherwise appear more than once and be watched and hashed repeatedly.
159    pub fn referenced_files(&self) -> Vec<std::path::PathBuf> {
160        let mut seen = std::collections::BTreeSet::new();
161        for name in self.listener_names() {
162            if let Some(slot) = self.get(name) {
163                for path in slot.load().referenced_files() {
164                    seen.insert(path);
165                }
166            }
167        }
168        seen.into_iter().collect()
169    }
170
171    /// Returns an iterator over listener names.
172    ///
173    /// Used during config reload to iterate over all listeners when
174    /// swapping pipelines or collecting referenced files for watching.
175    pub fn listener_names(&self) -> impl Iterator<Item = &str> {
176        self.pipelines.keys().map(String::as_str)
177    }
178}
179
180// -----------------------------------------------------------------------------
181// Tests
182// -----------------------------------------------------------------------------
183
184#[cfg(test)]
185#[expect(clippy::allow_attributes, reason = "blanket test suppressions")]
186#[allow(
187    clippy::unwrap_used,
188    clippy::expect_used,
189    clippy::indexing_slicing,
190    clippy::too_many_lines,
191    reason = "tests"
192)]
193mod tests {
194    use praxis_filter::FilterRegistry;
195
196    use super::*;
197
198    #[test]
199    fn get_returns_pipeline() {
200        let pipelines = make_pipelines(&["web"]);
201        assert!(pipelines.get("web").is_some(), "should find 'web' pipeline");
202    }
203
204    #[test]
205    fn get_returns_none_for_missing() {
206        let pipelines = make_pipelines(&["web"]);
207        assert!(pipelines.get("missing").is_none(), "should return None for missing");
208    }
209
210    #[test]
211    fn swap_replaces_pipeline_pointer() {
212        let pipelines = make_pipelines(&["web"]);
213        let old_ptr = Arc::as_ptr(&pipelines.get("web").unwrap().load());
214
215        let registry = FilterRegistry::with_builtins();
216        let new_pipeline = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
217        pipelines.swap("web", Arc::clone(&new_pipeline));
218
219        let new_ptr = Arc::as_ptr(&pipelines.get("web").unwrap().load());
220        assert_ne!(old_ptr, new_ptr, "swap should replace the pipeline pointer");
221    }
222
223    #[test]
224    fn old_guard_remains_valid_after_swap() {
225        let pipelines = make_pipelines(&["web"]);
226        let old_guard = pipelines.get("web").unwrap().load();
227        let old_ptr = Arc::as_ptr(&old_guard);
228
229        let registry = FilterRegistry::with_builtins();
230        let new_pipeline = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
231        pipelines.swap("web", new_pipeline);
232
233        let still_old_ptr = Arc::as_ptr(&old_guard);
234        assert_eq!(
235            old_ptr, still_old_ptr,
236            "old guard should still point to the original pipeline"
237        );
238    }
239
240    #[test]
241    fn swap_nonexistent_is_noop() {
242        let pipelines = make_pipelines(&["web"]);
243        let registry = FilterRegistry::with_builtins();
244        let new_pipeline = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
245        pipelines.swap("nonexistent", new_pipeline);
246        assert!(pipelines.get("web").is_some(), "existing pipeline should be unaffected");
247    }
248
249    #[test]
250    fn get_returns_arcswap_reference() {
251        let pipelines = make_pipelines(&["web"]);
252        let slot: &Arc<ArcSwap<FilterPipeline>> = pipelines.get("web").unwrap();
253        let _loaded: arc_swap::Guard<Arc<FilterPipeline>> = slot.load();
254    }
255
256    #[test]
257    fn protocol_is_recorded_at_construction() {
258        let registry = FilterRegistry::with_builtins();
259        let mut map = HashMap::new();
260        map.insert(
261            "web".to_owned(),
262            Arc::new(FilterPipeline::build(&mut [], &registry).unwrap()),
263        );
264        let mut protocols = HashMap::new();
265        protocols.insert("web".to_owned(), ProtocolKind::Tcp);
266
267        let pipelines = ListenerPipelines::with_protocols(map, protocols);
268
269        assert_eq!(pipelines.protocol("web"), Some(ProtocolKind::Tcp));
270        assert_eq!(
271            pipelines.protocol("missing"),
272            None,
273            "unknown listeners have no protocol"
274        );
275    }
276
277    #[test]
278    fn protocol_survives_swap() {
279        let registry = FilterRegistry::with_builtins();
280        let mut map = HashMap::new();
281        map.insert(
282            "web".to_owned(),
283            Arc::new(FilterPipeline::build(&mut [], &registry).unwrap()),
284        );
285        let mut protocols = HashMap::new();
286        protocols.insert("web".to_owned(), ProtocolKind::Http);
287        let pipelines = ListenerPipelines::with_protocols(map, protocols);
288
289        pipelines.swap("web", Arc::new(FilterPipeline::build(&mut [], &registry).unwrap()));
290
291        assert_eq!(
292            pipelines.protocol("web"),
293            Some(ProtocolKind::Http),
294            "a swap replaces the pipeline, never the protocol the handler was bound for"
295        );
296    }
297
298    #[test]
299    fn new_records_no_protocol() {
300        let pipelines = make_pipelines(&["web"]);
301        assert_eq!(pipelines.protocol("web"), None);
302    }
303
304    #[test]
305    fn referenced_files_empty_without_listeners() {
306        let pipelines = make_pipelines(&[]);
307        assert!(
308            pipelines.referenced_files().is_empty(),
309            "no listeners means no referenced documents"
310        );
311    }
312
313    #[test]
314    fn referenced_files_empty_when_no_filter_declares_one() {
315        let pipelines = make_pipelines(&["web"]);
316        assert!(
317            pipelines.referenced_files().is_empty(),
318            "a pipeline of non-declaring filters contributes nothing"
319        );
320    }
321
322    #[test]
323    fn referenced_files_collects_across_listeners() {
324        let pipelines =
325            make_pipelines_with_documents(&[("web", "/etc/praxis/web.yaml"), ("api", "/etc/praxis/api.yaml")]);
326        assert_eq!(
327            pipelines.referenced_files(),
328            vec![
329                std::path::PathBuf::from("/etc/praxis/api.yaml"),
330                std::path::PathBuf::from("/etc/praxis/web.yaml"),
331            ],
332            "every listener's documents must be collected, sorted by the BTreeSet"
333        );
334    }
335
336    #[test]
337    fn referenced_files_dedupes_a_document_shared_by_two_listeners() {
338        let shared = "/etc/praxis/shared.yaml";
339        let pipelines = make_pipelines_with_documents(&[("web", shared), ("api", shared)]);
340        assert_eq!(
341            pipelines.referenced_files(),
342            vec![std::path::PathBuf::from(shared)],
343            "a shared document must appear once"
344        );
345    }
346
347    // -------------------------------------------------------------------------
348    // Test Utilities
349    // -------------------------------------------------------------------------
350
351    /// A filter that declares the document named by its `document:` config key.
352    struct DocumentReaderFilter {
353        document: std::path::PathBuf,
354    }
355
356    #[async_trait::async_trait]
357    impl praxis_filter::HttpFilter for DocumentReaderFilter {
358        fn name(&self) -> &'static str {
359            "document_reader"
360        }
361
362        fn referenced_files(&self) -> Vec<std::path::PathBuf> {
363            vec![self.document.clone()]
364        }
365
366        async fn on_request(
367            &self,
368            _ctx: &mut praxis_filter::HttpFilterContext<'_>,
369        ) -> Result<praxis_filter::FilterAction, praxis_filter::FilterError> {
370            Ok(praxis_filter::FilterAction::Continue)
371        }
372    }
373
374    /// Build [`ListenerPipelines`] where each named listener runs a single filter
375    /// declaring the given document.
376    fn make_pipelines_with_documents(listeners: &[(&str, &str)]) -> ListenerPipelines {
377        let mut registry = FilterRegistry::with_builtins();
378        registry
379            .register(
380                "document_reader",
381                praxis_filter::FilterFactory::Http(Arc::new(|cfg: &serde_yaml::Value| {
382                    let document = cfg
383                        .get("document")
384                        .and_then(serde_yaml::Value::as_str)
385                        .ok_or_else(|| praxis_filter::FilterError::from("document_reader: missing document"))?;
386                    let filter: Box<dyn praxis_filter::HttpFilter> = Box::new(DocumentReaderFilter {
387                        document: std::path::PathBuf::from(document),
388                    });
389                    Ok(filter)
390                })),
391            )
392            .unwrap();
393
394        let mut map = HashMap::new();
395        for (listener, document) in listeners {
396            let yaml = format!("- filter: document_reader\n  document: {document}\n");
397            let mut entries: Vec<praxis_core::config::FilterEntry> = serde_yaml::from_str(&yaml).unwrap();
398            let pipeline = Arc::new(FilterPipeline::build(&mut entries, &registry).unwrap());
399            map.insert((*listener).to_owned(), pipeline);
400        }
401        ListenerPipelines::new(map)
402    }
403
404    /// Build [`ListenerPipelines`] with empty pipelines for the given names.
405    fn make_pipelines(names: &[&str]) -> ListenerPipelines {
406        let registry = FilterRegistry::with_builtins();
407        let mut map = HashMap::new();
408        for name in names {
409            let pipeline = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
410            map.insert((*name).to_owned(), pipeline);
411        }
412        ListenerPipelines::new(map)
413    }
414}