Skip to main content

lean_ctx/core/providers/
hardened_http.rs

1use std::time::{Duration, Instant};
2
3const DEFAULT_RESOLVE_TIMEOUT: Duration = Duration::from_secs(5);
4const DEFAULT_CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
5const DEFAULT_RESPONSE_TIMEOUT: Duration = Duration::from_secs(30);
6const DEGRADED_MULTIPLIER: u32 = 2;
7const MAX_RESPONSE_BODY_BYTES: u64 = 10 * 1024 * 1024;
8
9#[derive(Debug, Clone)]
10pub enum HttpOutcome {
11    Success {
12        status: u16,
13        body: String,
14        elapsed_ms: u64,
15    },
16    Timeout {
17        elapsed_ms: u64,
18        phase: String,
19    },
20    NetworkError {
21        message: String,
22        elapsed_ms: u64,
23    },
24    HttpError {
25        status: u16,
26        body: String,
27        elapsed_ms: u64,
28    },
29}
30
31pub struct HardenedClient {
32    agent: ureq::Agent,
33    provider_id: String,
34}
35
36pub fn hardened_agent() -> ureq::Agent {
37    let multiplier = if crate::core::io_health::environment()
38        == crate::core::io_health::IoEnvironment::Degraded
39    {
40        DEGRADED_MULTIPLIER
41    } else {
42        1
43    };
44
45    crate::core::http_client::ureq_agent_with_timeouts(
46        Some(DEFAULT_RESOLVE_TIMEOUT * multiplier),
47        Some(DEFAULT_CONNECT_TIMEOUT * multiplier),
48        Some(DEFAULT_RESPONSE_TIMEOUT * multiplier),
49    )
50}
51
52impl HardenedClient {
53    pub fn new(provider_id: &str) -> Self {
54        Self {
55            agent: hardened_agent(),
56            provider_id: provider_id.to_owned(),
57        }
58    }
59
60    pub fn get(&self, url: &str) -> HttpOutcome {
61        let request = self.agent.get(url);
62        self.execute_request(|| request.call())
63    }
64
65    pub fn get_with_headers(&self, url: &str, headers: &[(&str, &str)]) -> HttpOutcome {
66        let mut request = self.agent.get(url);
67        for &(key, value) in headers {
68            request = request.header(key, value);
69        }
70        self.execute_request(|| request.call())
71    }
72
73    pub fn post(&self, url: &str, body: &str) -> HttpOutcome {
74        let request = self.agent.post(url);
75        self.execute_request(|| request.send(body))
76    }
77
78    pub fn post_with_headers(
79        &self,
80        url: &str,
81        body: &str,
82        headers: &[(&str, &str)],
83    ) -> HttpOutcome {
84        let mut request = self.agent.post(url);
85        for &(key, value) in headers {
86            request = request.header(key, value);
87        }
88        self.execute_request(|| request.send(body))
89    }
90
91    fn execute_request<F>(&self, request: F) -> HttpOutcome
92    where
93        F: FnOnce() -> Result<ureq::http::Response<ureq::Body>, ureq::Error>,
94    {
95        let started = Instant::now();
96        let response = match request() {
97            Ok(response) => response,
98            Err(error) => return self.error_outcome(error, elapsed_ms(started)),
99        };
100        let status = response.status().as_u16();
101        let body = match response
102            .into_body()
103            .into_with_config()
104            .limit(MAX_RESPONSE_BODY_BYTES)
105            .read_to_string()
106        {
107            Ok(body) => body,
108            Err(error) => return self.error_outcome(error, elapsed_ms(started)),
109        };
110        let elapsed_ms = elapsed_ms(started);
111
112        if status >= 400 {
113            HttpOutcome::HttpError {
114                status,
115                body,
116                elapsed_ms,
117            }
118        } else {
119            HttpOutcome::Success {
120                status,
121                body,
122                elapsed_ms,
123            }
124        }
125    }
126
127    fn error_outcome(&self, error: ureq::Error, elapsed_ms: u64) -> HttpOutcome {
128        match error {
129            ureq::Error::Timeout(timeout) => HttpOutcome::Timeout {
130                elapsed_ms,
131                phase: timeout.to_string(),
132            },
133            ureq::Error::StatusCode(status) => HttpOutcome::HttpError {
134                status,
135                body: String::new(),
136                elapsed_ms,
137            },
138            error => HttpOutcome::NetworkError {
139                message: format!("{}: {error}", self.provider_id),
140                elapsed_ms,
141            },
142        }
143    }
144}
145
146impl HttpOutcome {
147    pub fn into_body(self) -> Result<String, String> {
148        match self {
149            Self::Success { body, .. } => Ok(body),
150            Self::Timeout { phase, .. } => Err(format!("timeout during {phase}")),
151            Self::NetworkError { message, .. } => Err(message),
152            Self::HttpError { status, body, .. } => Err(format!("HTTP {status}: {body}")),
153        }
154    }
155
156    pub fn is_success(&self) -> bool {
157        matches!(self, Self::Success { .. })
158    }
159
160    pub fn body(&self) -> Option<&str> {
161        match self {
162            Self::Success { body, .. } => Some(body),
163            _ => None,
164        }
165    }
166
167    pub fn elapsed_ms(&self) -> u64 {
168        match self {
169            Self::Success { elapsed_ms, .. }
170            | Self::Timeout { elapsed_ms, .. }
171            | Self::NetworkError { elapsed_ms, .. }
172            | Self::HttpError { elapsed_ms, .. } => *elapsed_ms,
173        }
174    }
175
176    pub fn into_result(self) -> Result<String, String> {
177        match self {
178            Self::Success { body, .. } => Ok(body),
179            Self::Timeout { phase, .. } => Err(format!("HTTP request timed out during {phase}")),
180            Self::NetworkError { message, .. } => Err(message),
181            Self::HttpError { status, body, .. } => {
182                if body.is_empty() {
183                    Err(format!("HTTP request failed with status {status}"))
184                } else {
185                    Err(format!("HTTP request failed with status {status}: {body}"))
186                }
187            }
188        }
189    }
190}
191
192pub fn provider_get(provider_id: &str, url: &str) -> HttpOutcome {
193    HardenedClient::new(provider_id).get(url)
194}
195
196pub fn provider_get_with_headers(
197    provider_id: &str,
198    url: &str,
199    headers: &[(&str, &str)],
200) -> HttpOutcome {
201    HardenedClient::new(provider_id).get_with_headers(url, headers)
202}
203
204pub fn provider_post(provider_id: &str, url: &str, body: &str) -> HttpOutcome {
205    HardenedClient::new(provider_id).post(url, body)
206}
207
208fn elapsed_ms(started: Instant) -> u64 {
209    u64::try_from(started.elapsed().as_millis()).unwrap_or(u64::MAX)
210}
211
212#[cfg(test)]
213mod tests {
214    use super::{
215        DEFAULT_CONNECT_TIMEOUT, DEFAULT_RESOLVE_TIMEOUT, DEFAULT_RESPONSE_TIMEOUT,
216        DEGRADED_MULTIPLIER, HardenedClient, HttpOutcome, hardened_agent,
217    };
218    use std::io::{Read, Write};
219    use std::net::TcpListener;
220    use std::sync::mpsc::{self, Receiver};
221    use std::time::Duration;
222
223    #[test]
224    fn test_hardened_agent_creates_agent() {
225        let _agent = hardened_agent();
226    }
227
228    #[test]
229    fn test_http_outcome_success_is_success() {
230        let outcome = success_outcome(12);
231        assert!(outcome.is_success());
232    }
233
234    #[test]
235    fn test_http_outcome_timeout_not_success() {
236        let outcome = timeout_outcome(12);
237        assert!(!outcome.is_success());
238    }
239
240    #[test]
241    fn test_http_outcome_network_error_not_success() {
242        let outcome = HttpOutcome::NetworkError {
243            message: "offline".to_owned(),
244            elapsed_ms: 12,
245        };
246        assert!(!outcome.is_success());
247    }
248
249    #[test]
250    fn test_http_outcome_into_result_success() {
251        assert_eq!(success_outcome(12).into_result(), Ok("payload".to_owned()));
252    }
253
254    #[test]
255    fn test_http_outcome_into_result_timeout() {
256        let error = timeout_outcome(12)
257            .into_result()
258            .expect_err("timeout expected");
259        assert!(error.contains("timed out"));
260        assert!(error.contains("connect"));
261    }
262
263    #[test]
264    fn test_http_outcome_into_body_variants() {
265        assert_eq!(success_outcome(12).into_body(), Ok("payload".to_owned()));
266        assert_eq!(
267            timeout_outcome(12).into_body(),
268            Err("timeout during connect".to_owned())
269        );
270        assert_eq!(
271            HttpOutcome::NetworkError {
272                message: "offline".to_owned(),
273                elapsed_ms: 12,
274            }
275            .into_body(),
276            Err("offline".to_owned())
277        );
278        assert_eq!(
279            HttpOutcome::HttpError {
280                status: 503,
281                body: "unavailable".to_owned(),
282                elapsed_ms: 12,
283            }
284            .into_body(),
285            Err("HTTP 503: unavailable".to_owned())
286        );
287    }
288
289    #[test]
290    fn test_get_with_headers() {
291        let (get_url, get_request) = serve_once();
292        let client = HardenedClient::new("test");
293        assert!(
294            client
295                .get_with_headers(&get_url, &[("X-Test-Header", "get-value")])
296                .is_success()
297        );
298        assert!(
299            get_request
300                .recv()
301                .expect("GET request expected")
302                .to_ascii_lowercase()
303                .contains("x-test-header: get-value")
304        );
305    }
306
307    #[test]
308    fn test_http_outcome_body_extraction() {
309        let success = success_outcome(12);
310        let error = timeout_outcome(12);
311        assert_eq!(success.body(), Some("payload"));
312        assert_eq!(error.body(), None);
313    }
314
315    #[test]
316    fn test_http_outcome_elapsed_tracking() {
317        let outcomes = [
318            success_outcome(1),
319            timeout_outcome(2),
320            HttpOutcome::NetworkError {
321                message: "offline".to_owned(),
322                elapsed_ms: 3,
323            },
324            HttpOutcome::HttpError {
325                status: 503,
326                body: "unavailable".to_owned(),
327                elapsed_ms: 4,
328            },
329        ];
330        let elapsed: Vec<u64> = outcomes.iter().map(HttpOutcome::elapsed_ms).collect();
331        assert_eq!(elapsed, [1, 2, 3, 4]);
332    }
333
334    #[test]
335    fn test_degraded_multiplier_value() {
336        assert_eq!(DEGRADED_MULTIPLIER, 2);
337    }
338
339    #[test]
340    fn test_default_timeouts_reasonable() {
341        assert_eq!(DEFAULT_RESOLVE_TIMEOUT, Duration::from_secs(5));
342        assert_eq!(DEFAULT_CONNECT_TIMEOUT, Duration::from_secs(10));
343        assert_eq!(DEFAULT_RESPONSE_TIMEOUT, Duration::from_secs(30));
344    }
345
346    fn success_outcome(elapsed_ms: u64) -> HttpOutcome {
347        HttpOutcome::Success {
348            status: 200,
349            body: "payload".to_owned(),
350            elapsed_ms,
351        }
352    }
353
354    fn timeout_outcome(elapsed_ms: u64) -> HttpOutcome {
355        HttpOutcome::Timeout {
356            elapsed_ms,
357            phase: "connect".to_owned(),
358        }
359    }
360
361    fn serve_once() -> (String, Receiver<String>) {
362        let listener = TcpListener::bind("127.0.0.1:0").expect("test listener should bind");
363        let address = listener
364            .local_addr()
365            .expect("test listener address should exist");
366        let (sender, receiver) = mpsc::channel();
367        std::thread::spawn(move || {
368            let (mut stream, _) = listener.accept().expect("request should connect");
369            let mut buffer = [0_u8; 4096];
370            let length = stream
371                .read(&mut buffer)
372                .expect("request should be readable");
373            sender
374                .send(String::from_utf8_lossy(&buffer[..length]).into_owned())
375                .expect("request should be received");
376            stream
377                .write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok")
378                .expect("response should be writable");
379        });
380        (format!("http://{address}/test"), receiver)
381    }
382}