Skip to main content

praxis_protocol/
pipelines.rs

1// SPDX-License-Identifier: MIT
2// Copyright (c) 2024 Praxis Contributors
3
4//! Maps listener names to their resolved [`FilterPipeline`].
5//!
6//! [`FilterPipeline`]: praxis_filter::FilterPipeline
7
8use std::{collections::HashMap, sync::Arc};
9
10use arc_swap::ArcSwap;
11use praxis_filter::FilterPipeline;
12
13// -----------------------------------------------------------------------------
14// ListenerPipelines
15// -----------------------------------------------------------------------------
16
17/// Maps listener names to their resolved [`FilterPipeline`]s.
18///
19/// Each pipeline is wrapped in [`ArcSwap`] so it can be atomically
20/// replaced at runtime without blocking in-flight requests.
21///
22/// ```
23/// use std::{collections::HashMap, sync::Arc};
24///
25/// use praxis_filter::{FilterPipeline, FilterRegistry};
26/// use praxis_protocol::ListenerPipelines;
27///
28/// let registry = FilterRegistry::with_builtins();
29/// let pipeline = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
30///
31/// let mut map = HashMap::new();
32/// map.insert("web".to_owned(), pipeline);
33/// let pipelines = ListenerPipelines::new(map);
34///
35/// assert!(pipelines.get("web").is_some());
36/// assert!(pipelines.get("missing").is_none());
37/// ```
38///
39/// [`ArcSwap`]: arc_swap::ArcSwap
40pub struct ListenerPipelines {
41    /// Maps listener names to their swappable filter pipelines.
42    pipelines: HashMap<String, Arc<ArcSwap<FilterPipeline>>>,
43}
44
45impl ListenerPipelines {
46    /// Create from a map of listener name to pipeline.
47    pub fn new(pipelines: HashMap<String, Arc<FilterPipeline>>) -> Self {
48        let swappable = pipelines
49            .into_iter()
50            .map(|(name, p)| (name, Arc::new(ArcSwap::from(p))))
51            .collect();
52        Self { pipelines: swappable }
53    }
54
55    /// Get the swappable pipeline for a listener by name.
56    pub fn get(&self, listener_name: &str) -> Option<&Arc<ArcSwap<FilterPipeline>>> {
57        self.pipelines.get(listener_name)
58    }
59
60    /// Atomically replace the pipeline for a listener.
61    ///
62    /// No-op if the listener name is not present.
63    ///
64    /// ```
65    /// use std::{collections::HashMap, sync::Arc};
66    ///
67    /// use praxis_filter::{FilterPipeline, FilterRegistry};
68    /// use praxis_protocol::ListenerPipelines;
69    ///
70    /// let registry = FilterRegistry::with_builtins();
71    /// let old = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
72    /// let new = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
73    ///
74    /// let mut map = HashMap::new();
75    /// map.insert("web".to_owned(), old);
76    /// let pipelines = ListenerPipelines::new(map);
77    ///
78    /// pipelines.swap("web", new);
79    /// pipelines.swap(
80    ///     "nonexistent",
81    ///     Arc::new(FilterPipeline::build(&mut [], &registry).unwrap()),
82    /// );
83    /// ```
84    pub fn swap(&self, listener_name: &str, new_pipeline: Arc<FilterPipeline>) {
85        if let Some(slot) = self.pipelines.get(listener_name) {
86            slot.store(new_pipeline);
87        }
88    }
89
90    /// Returns an iterator over listener names.
91    pub fn listener_names(&self) -> impl Iterator<Item = &str> {
92        self.pipelines.keys().map(String::as_str)
93    }
94}
95
96// -----------------------------------------------------------------------------
97// Tests
98// -----------------------------------------------------------------------------
99
100#[cfg(test)]
101#[expect(clippy::allow_attributes, reason = "blanket test suppressions")]
102#[allow(
103    clippy::unwrap_used,
104    clippy::expect_used,
105    clippy::indexing_slicing,
106    clippy::too_many_lines,
107    reason = "tests"
108)]
109mod tests {
110    use std::{collections::HashMap, sync::Arc};
111
112    use arc_swap::ArcSwap;
113    use praxis_filter::{FilterPipeline, FilterRegistry};
114
115    use super::*;
116
117    #[test]
118    fn get_returns_pipeline() {
119        let pipelines = make_pipelines(&["web"]);
120        assert!(pipelines.get("web").is_some(), "should find 'web' pipeline");
121    }
122
123    #[test]
124    fn get_returns_none_for_missing() {
125        let pipelines = make_pipelines(&["web"]);
126        assert!(pipelines.get("missing").is_none(), "should return None for missing");
127    }
128
129    #[test]
130    fn swap_replaces_pipeline_pointer() {
131        let pipelines = make_pipelines(&["web"]);
132        let old_ptr = Arc::as_ptr(&pipelines.get("web").unwrap().load());
133
134        let registry = FilterRegistry::with_builtins();
135        let new_pipeline = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
136        pipelines.swap("web", Arc::clone(&new_pipeline));
137
138        let new_ptr = Arc::as_ptr(&pipelines.get("web").unwrap().load());
139        assert_ne!(old_ptr, new_ptr, "swap should replace the pipeline pointer");
140    }
141
142    #[test]
143    fn old_guard_remains_valid_after_swap() {
144        let pipelines = make_pipelines(&["web"]);
145        let old_guard = pipelines.get("web").unwrap().load();
146        let old_ptr = Arc::as_ptr(&old_guard);
147
148        let registry = FilterRegistry::with_builtins();
149        let new_pipeline = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
150        pipelines.swap("web", new_pipeline);
151
152        let still_old_ptr = Arc::as_ptr(&old_guard);
153        assert_eq!(
154            old_ptr, still_old_ptr,
155            "old guard should still point to the original pipeline"
156        );
157    }
158
159    #[test]
160    fn swap_nonexistent_is_noop() {
161        let pipelines = make_pipelines(&["web"]);
162        let registry = FilterRegistry::with_builtins();
163        let new_pipeline = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
164        pipelines.swap("nonexistent", new_pipeline);
165        assert!(pipelines.get("web").is_some(), "existing pipeline should be unaffected");
166    }
167
168    #[test]
169    fn get_returns_arcswap_reference() {
170        let pipelines = make_pipelines(&["web"]);
171        let slot: &Arc<ArcSwap<FilterPipeline>> = pipelines.get("web").unwrap();
172        let _loaded: arc_swap::Guard<Arc<FilterPipeline>> = slot.load();
173    }
174
175    // -------------------------------------------------------------------------
176    // Test Utilities
177    // -------------------------------------------------------------------------
178
179    /// Build [`ListenerPipelines`] with empty pipelines for the given names.
180    fn make_pipelines(names: &[&str]) -> ListenerPipelines {
181        let registry = FilterRegistry::with_builtins();
182        let mut map = HashMap::new();
183        for name in names {
184            let pipeline = Arc::new(FilterPipeline::build(&mut [], &registry).unwrap());
185            map.insert((*name).to_owned(), pipeline);
186        }
187        ListenerPipelines::new(map)
188    }
189}