Skip to main content

confium_coordinator/coordinator/
middleware.rs

1//! Middleware pipeline — unified request processing chain.
2
3/// A request context passed through the pipeline.
4#[derive(Debug, Clone)]
5pub struct RequestContext {
6    pub request_type: String,
7    pub signer_id: Option<String>,
8    pub quorum_id: Option<String>,
9    pub session_id: Option<String>,
10    pub payload_size: usize,
11}
12
13/// Middleware result: continue or reject.
14#[derive(Debug, Clone)]
15pub enum MiddlewareResult {
16    Continue,
17    Reject(String),
18}
19
20/// Middleware trait — each stage processes the request.
21pub trait Middleware: Send + Sync {
22    fn name(&self) -> &str;
23    fn process(&self, ctx: &RequestContext) -> MiddlewareResult;
24}
25
26/// The pipeline: runs middlewares in order, stops on first rejection.
27pub struct Pipeline {
28    middlewares: Vec<Box<dyn Middleware>>,
29}
30
31impl Pipeline {
32    pub fn new() -> Self {
33        Self {
34            middlewares: Vec::new(),
35        }
36    }
37
38    pub fn add(&mut self, mw: Box<dyn Middleware>) -> &mut Self {
39        self.middlewares.push(mw);
40        self
41    }
42
43    pub fn execute(&self, ctx: &RequestContext) -> Result<(), String> {
44        for mw in &self.middlewares {
45            match mw.process(ctx) {
46                MiddlewareResult::Continue => {}
47                MiddlewareResult::Reject(reason) => {
48                    return Err(format!("{} rejected: {}", mw.name(), reason));
49                }
50            }
51        }
52        Ok(())
53    }
54
55    pub fn middleware_count(&self) -> usize {
56        self.middlewares.len()
57    }
58}
59
60impl Default for Pipeline {
61    fn default() -> Self {
62        Self::new()
63    }
64}
65
66// Built-in middlewares
67
68/// Rate limiting middleware (uses the rate limiter).
69pub struct RateLimitMiddleware {
70    pub max_per_second: u32,
71}
72
73impl Middleware for RateLimitMiddleware {
74    fn name(&self) -> &str {
75        "rate-limiter"
76    }
77    fn process(&self, ctx: &RequestContext) -> MiddlewareResult {
78        // Simplified: check payload size as a proxy for load
79        if ctx.payload_size > 1_000_000 {
80            MiddlewareResult::Reject("payload too large".into())
81        } else {
82            MiddlewareResult::Continue
83        }
84    }
85}
86
87/// Authentication middleware.
88pub struct AuthMiddleware {
89    pub required: bool,
90}
91
92impl Middleware for AuthMiddleware {
93    fn name(&self) -> &str {
94        "auth"
95    }
96    fn process(&self, ctx: &RequestContext) -> MiddlewareResult {
97        if self.required && ctx.signer_id.is_none() {
98            MiddlewareResult::Reject("authentication required".into())
99        } else {
100            MiddlewareResult::Continue
101        }
102    }
103}
104
105/// Policy enforcement middleware.
106pub struct PolicyMiddleware;
107
108impl Middleware for PolicyMiddleware {
109    fn name(&self) -> &str {
110        "policy"
111    }
112    fn process(&self, ctx: &RequestContext) -> MiddlewareResult {
113        if ctx.quorum_id.is_none() && ctx.request_type != "health_check" {
114            MiddlewareResult::Reject("quorum_id required".into())
115        } else {
116            MiddlewareResult::Continue
117        }
118    }
119}
120
121/// Backpressure middleware.
122pub struct BackpressureMiddleware {
123    pub max_payload: usize,
124}
125
126impl Middleware for BackpressureMiddleware {
127    fn name(&self) -> &str {
128        "backpressure"
129    }
130    fn process(&self, ctx: &RequestContext) -> MiddlewareResult {
131        if ctx.payload_size > self.max_payload {
132            MiddlewareResult::Reject(format!(
133                "payload {} exceeds max {}",
134                ctx.payload_size, self.max_payload
135            ))
136        } else {
137            MiddlewareResult::Continue
138        }
139    }
140}
141
142/// Logging middleware (always continues).
143pub struct LoggingMiddleware;
144
145impl Middleware for LoggingMiddleware {
146    fn name(&self) -> &str {
147        "logging"
148    }
149    fn process(&self, _ctx: &RequestContext) -> MiddlewareResult {
150        MiddlewareResult::Continue
151    }
152}
153
154#[cfg(test)]
155mod tests {
156    use super::*;
157
158    fn make_ctx(req_type: &str) -> RequestContext {
159        RequestContext {
160            request_type: req_type.into(),
161            signer_id: Some("alice".into()),
162            quorum_id: Some("q1".into()),
163            session_id: None,
164            payload_size: 100,
165        }
166    }
167
168    #[test]
169    fn empty_pipeline_always_passes() {
170        let pipeline = Pipeline::new();
171        assert!(pipeline.execute(&make_ctx("test")).is_ok());
172    }
173
174    #[test]
175    fn auth_rejects_unauthenticated() {
176        let mut pipeline = Pipeline::new();
177        pipeline.add(Box::new(AuthMiddleware { required: true }));
178        let mut ctx = make_ctx("test");
179        ctx.signer_id = None;
180        assert!(pipeline.execute(&ctx).is_err());
181    }
182
183    #[test]
184    fn auth_allows_authenticated() {
185        let mut pipeline = Pipeline::new();
186        pipeline.add(Box::new(AuthMiddleware { required: true }));
187        assert!(pipeline.execute(&make_ctx("test")).is_ok());
188    }
189
190    #[test]
191    fn full_pipeline_passes_valid_request() {
192        let mut pipeline = Pipeline::new();
193        pipeline
194            .add(Box::new(RateLimitMiddleware {
195                max_per_second: 100,
196            }))
197            .add(Box::new(AuthMiddleware { required: true }))
198            .add(Box::new(PolicyMiddleware))
199            .add(Box::new(BackpressureMiddleware {
200                max_payload: 10_000,
201            }))
202            .add(Box::new(LoggingMiddleware));
203        assert!(pipeline.execute(&make_ctx("sign")).is_ok());
204        assert_eq!(pipeline.middleware_count(), 5);
205    }
206
207    #[test]
208    fn backpressure_rejects_large_payload() {
209        let mut pipeline = Pipeline::new();
210        pipeline.add(Box::new(BackpressureMiddleware { max_payload: 100 }));
211        let mut ctx = make_ctx("test");
212        ctx.payload_size = 200;
213        assert!(pipeline.execute(&ctx).is_err());
214    }
215
216    #[test]
217    fn policy_allows_health_check_without_quorum() {
218        let mut pipeline = Pipeline::new();
219        pipeline.add(Box::new(PolicyMiddleware));
220        let mut ctx = make_ctx("health_check");
221        ctx.quorum_id = None;
222        assert!(pipeline.execute(&ctx).is_ok());
223    }
224
225    #[test]
226    fn policy_rejects_sign_without_quorum() {
227        let mut pipeline = Pipeline::new();
228        pipeline.add(Box::new(PolicyMiddleware));
229        let mut ctx = make_ctx("sign");
230        ctx.quorum_id = None;
231        assert!(pipeline.execute(&ctx).is_err());
232    }
233
234    #[test]
235    fn first_rejection_stops_pipeline() {
236        let mut pipeline = Pipeline::new();
237        pipeline.add(Box::new(AuthMiddleware { required: true }));
238        pipeline.add(Box::new(LoggingMiddleware));
239        let mut ctx = make_ctx("test");
240        ctx.signer_id = None;
241        let result = pipeline.execute(&ctx);
242        assert!(result.is_err());
243        assert!(result.unwrap_err().contains("auth"));
244    }
245
246    #[test]
247    fn rate_limit_rejects_large_payload() {
248        let mut pipeline = Pipeline::new();
249        pipeline.add(Box::new(RateLimitMiddleware { max_per_second: 10 }));
250        let mut ctx = make_ctx("test");
251        ctx.payload_size = 2_000_000;
252        assert!(pipeline.execute(&ctx).is_err());
253    }
254}