Skip to main content

rest/
client.rs

1use crate::HttpClientMetrics;
2use reqwest::{Client, Method, Request, RequestBuilder, Response, StatusCode};
3use rust_zero_core::{
4    BreakerState, CircuitBreaker, CircuitBreakerConfig, CircuitBreakerError,
5    CircuitBreakerSnapshot, TraceContext,
6};
7use serde::{de::DeserializeOwned, Serialize};
8use std::{
9    fmt,
10    sync::Arc,
11    time::{Duration, Instant},
12};
13
14/// Production defaults for calls to a named HTTP dependency.
15#[derive(Debug, Clone)]
16pub struct HttpClientConfig {
17    pub service: String,
18    pub timeout: Duration,
19    pub max_response_bytes: usize,
20    pub breaker: CircuitBreakerConfig,
21}
22
23impl HttpClientConfig {
24    pub fn new(service: impl Into<String>) -> Self {
25        Self {
26            service: service.into(),
27            timeout: Duration::from_secs(10),
28            max_response_bytes: 10 * 1024 * 1024,
29            breaker: CircuitBreakerConfig::new(5, Duration::from_secs(30)),
30        }
31    }
32
33    pub fn with_timeout(mut self, timeout: Duration) -> Self {
34        assert!(!timeout.is_zero(), "HTTP timeout must be greater than zero");
35        self.timeout = timeout;
36        self
37    }
38
39    pub fn with_max_response_bytes(mut self, bytes: usize) -> Self {
40        assert!(bytes > 0, "HTTP response limit must be greater than zero");
41        self.max_response_bytes = bytes;
42        self
43    }
44
45    pub fn with_breaker(mut self, breaker: CircuitBreakerConfig) -> Self {
46        self.breaker = breaker;
47        self
48    }
49}
50
51/// Named HTTP service client with deadlines, W3C propagation, response limits,
52/// and circuit breaking that treats 5xx responses as dependency failures.
53#[derive(Clone)]
54pub struct HttpClient {
55    service: Arc<str>,
56    client: Client,
57    breaker: Arc<CircuitBreaker>,
58    max_response_bytes: usize,
59    metrics: Option<HttpClientMetrics>,
60}
61
62impl HttpClient {
63    pub fn new(config: HttpClientConfig) -> Result<Self, HttpClientError> {
64        if config.service.trim().is_empty() {
65            return Err(HttpClientError::InvalidServiceName);
66        }
67        let client = Client::builder()
68            .timeout(config.timeout)
69            .build()
70            .map_err(HttpClientError::Build)?;
71
72        Ok(Self {
73            service: Arc::from(config.service),
74            client,
75            breaker: Arc::new(CircuitBreaker::new(config.breaker)),
76            max_response_bytes: config.max_response_bytes,
77            metrics: None,
78        })
79    }
80
81    /// Records transport outcomes for this client in a shared metrics registry.
82    pub fn with_metrics(mut self, metrics: HttpClientMetrics) -> Self {
83        self.metrics = Some(metrics);
84        self
85    }
86
87    pub fn service(&self) -> &str {
88        &self.service
89    }
90
91    pub fn breaker_state(&self) -> BreakerState {
92        self.breaker.state()
93    }
94
95    pub fn breaker_snapshot(&self) -> CircuitBreakerSnapshot {
96        self.breaker.snapshot()
97    }
98
99    pub fn request(&self, method: Method, url: impl reqwest::IntoUrl) -> RequestBuilder {
100        self.client.request(method, url)
101    }
102
103    /// Executes a pre-built request through the service circuit breaker.
104    pub async fn execute(&self, request: Request) -> Result<Response, HttpClientError> {
105        let method = request.method().as_str().to_owned();
106        let started_at = Instant::now();
107        let _in_flight = self
108            .metrics
109            .as_ref()
110            .map(|metrics| metrics.track_in_flight(self.service.to_string(), method.clone()));
111        let result = self
112            .breaker
113            .execute_async_with_accept(
114                || self.client.execute(request),
115                |result| match result {
116                    Ok(response) => !response.status().is_server_error(),
117                    Err(_) => false,
118                },
119            )
120            .await
121            .map_err(|error| match error {
122                CircuitBreakerError::Open => HttpClientError::CircuitOpen {
123                    service: self.service.to_string(),
124                },
125                CircuitBreakerError::Operation(error) => HttpClientError::Transport(error),
126            });
127
128        if let Some(metrics) = &self.metrics {
129            let result_label = match &result {
130                Ok(response) => response.status().as_str().to_owned(),
131                Err(HttpClientError::CircuitOpen { .. }) => "circuit_open".to_owned(),
132                Err(HttpClientError::Transport(_)) => "transport_error".to_owned(),
133                Err(_) => "client_error".to_owned(),
134            };
135            metrics.record(
136                &self.service,
137                &method,
138                &result_label,
139                started_at.elapsed().as_secs_f64(),
140            );
141        }
142
143        result
144    }
145
146    /// Adds a child `traceparent` header and executes the request.
147    pub async fn execute_traced(
148        &self,
149        mut request: Request,
150        parent: &TraceContext,
151    ) -> Result<Response, HttpClientError> {
152        let child = parent.child();
153        request.headers_mut().insert(
154            "traceparent",
155            child
156                .traceparent()
157                .parse()
158                .expect("generated traceparent must be a valid header"),
159        );
160        self.execute(request).await
161    }
162
163    pub async fn get_json<T>(&self, url: impl reqwest::IntoUrl) -> Result<T, HttpClientError>
164    where
165        T: DeserializeOwned,
166    {
167        let request = self
168            .request(Method::GET, url)
169            .build()
170            .map_err(HttpClientError::Build)?;
171        let response = self.execute(request).await?;
172        self.decode_json(response).await
173    }
174
175    pub async fn post_json<B, T>(
176        &self,
177        url: impl reqwest::IntoUrl,
178        body: &B,
179    ) -> Result<T, HttpClientError>
180    where
181        B: Serialize + ?Sized,
182        T: DeserializeOwned,
183    {
184        let request = self
185            .request(Method::POST, url)
186            .json(body)
187            .build()
188            .map_err(HttpClientError::Build)?;
189        let response = self.execute(request).await?;
190        self.decode_json(response).await
191    }
192
193    pub async fn decode_json<T>(&self, response: Response) -> Result<T, HttpClientError>
194    where
195        T: DeserializeOwned,
196    {
197        let status = response.status();
198        if !status.is_success() {
199            return Err(HttpClientError::Status(status));
200        }
201        if response
202            .content_length()
203            .is_some_and(|length| length > self.max_response_bytes as u64)
204        {
205            return Err(HttpClientError::BodyTooLarge {
206                limit: self.max_response_bytes,
207            });
208        }
209
210        let bytes = response.bytes().await.map_err(HttpClientError::Transport)?;
211        if bytes.len() > self.max_response_bytes {
212            return Err(HttpClientError::BodyTooLarge {
213                limit: self.max_response_bytes,
214            });
215        }
216        serde_json::from_slice(&bytes).map_err(HttpClientError::Decode)
217    }
218}
219
220#[derive(Debug)]
221pub enum HttpClientError {
222    InvalidServiceName,
223    Build(reqwest::Error),
224    CircuitOpen { service: String },
225    Transport(reqwest::Error),
226    Status(StatusCode),
227    BodyTooLarge { limit: usize },
228    Decode(serde_json::Error),
229}
230
231impl fmt::Display for HttpClientError {
232    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
233        match self {
234            Self::InvalidServiceName => formatter.write_str("HTTP service name cannot be empty"),
235            Self::Build(error) => write!(formatter, "failed to build HTTP request: {error}"),
236            Self::CircuitOpen { service } => {
237                write!(formatter, "HTTP circuit for service {service} is open")
238            }
239            Self::Transport(error) => write!(formatter, "HTTP transport failed: {error}"),
240            Self::Status(status) => write!(formatter, "HTTP service returned {status}"),
241            Self::BodyTooLarge { limit } => {
242                write!(formatter, "HTTP response exceeds the {limit}-byte limit")
243            }
244            Self::Decode(error) => {
245                write!(formatter, "failed to decode HTTP JSON response: {error}")
246            }
247        }
248    }
249}
250
251impl std::error::Error for HttpClientError {
252    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
253        match self {
254            Self::Build(error) | Self::Transport(error) => Some(error),
255            Self::Decode(error) => Some(error),
256            _ => None,
257        }
258    }
259}
260
261#[cfg(test)]
262mod tests {
263    use super::*;
264    use actix_web::{web, App, HttpRequest, HttpResponse, HttpServer};
265    use futures::stream;
266    use rust_zero_core::{Metrics, TraceFlags};
267    use serde_json::{json, Value};
268
269    async fn spawn_server() -> (String, actix_web::dev::ServerHandle) {
270        let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
271        let address = listener.local_addr().unwrap();
272        let server = HttpServer::new(|| {
273            App::new()
274                .route(
275                    "/get",
276                    web::get().to(|| async { HttpResponse::Ok().json(json!({"method": "get"})) }),
277                )
278                .route(
279                    "/post",
280                    web::post().to(|body: web::Json<Value>| async move {
281                        HttpResponse::Ok().json(body.into_inner())
282                    }),
283                )
284                .route(
285                    "/trace",
286                    web::get().to(|request: HttpRequest| async move {
287                        HttpResponse::Ok().json(json!({
288                            "traceparent": request
289                                .headers()
290                                .get("traceparent")
291                                .unwrap()
292                                .to_str()
293                                .unwrap()
294                        }))
295                    }),
296                )
297                .route(
298                    "/failure",
299                    web::get().to(|| async { HttpResponse::ServiceUnavailable().finish() }),
300                )
301                .route(
302                    "/invalid",
303                    web::get().to(|| async { HttpResponse::Ok().body("not json") }),
304                )
305                .route(
306                    "/chunked",
307                    web::get().to(|| async {
308                        HttpResponse::Ok().streaming(stream::once(async {
309                            Ok::<_, actix_web::Error>(web::Bytes::from_static(b"123456"))
310                        }))
311                    }),
312                )
313        })
314        .listen(listener)
315        .unwrap()
316        .run();
317        let handle = server.handle();
318        actix_web::rt::spawn(server);
319        (format!("http://{address}"), handle)
320    }
321
322    #[test]
323    fn rejects_empty_service_names() {
324        assert!(matches!(
325            HttpClient::new(HttpClientConfig::new(" ")),
326            Err(HttpClientError::InvalidServiceName)
327        ));
328    }
329
330    #[test]
331    fn builds_requests_with_json_and_trace_headers() {
332        let client = HttpClient::new(HttpClientConfig::new("users")).unwrap();
333        let parent = TraceContext::root(TraceFlags::SAMPLED);
334        let mut request = client
335            .request(Method::GET, "http://localhost/users")
336            .build()
337            .unwrap();
338        let child = parent.child();
339        request
340            .headers_mut()
341            .insert("traceparent", child.traceparent().parse().unwrap());
342
343        assert!(request.headers().contains_key("traceparent"));
344    }
345
346    #[actix_web::test]
347    async fn gets_posts_and_propagates_trace_context() {
348        let (base_url, server) = spawn_server().await;
349        let client = HttpClient::new(
350            HttpClientConfig::new("users")
351                .with_timeout(Duration::from_secs(1))
352                .with_max_response_bytes(1024),
353        )
354        .unwrap();
355
356        assert_eq!(client.service(), "users");
357        assert_eq!(
358            client
359                .get_json::<Value>(format!("{base_url}/get"))
360                .await
361                .unwrap(),
362            json!({"method": "get"})
363        );
364        assert_eq!(
365            client
366                .post_json::<_, Value>(format!("{base_url}/post"), &json!({"id": 42}))
367                .await
368                .unwrap(),
369            json!({"id": 42})
370        );
371
372        let parent = TraceContext::root(TraceFlags::SAMPLED);
373        let request = client
374            .request(Method::GET, format!("{base_url}/trace"))
375            .build()
376            .unwrap();
377        let response: Value = client
378            .decode_json(client.execute_traced(request, &parent).await.unwrap())
379            .await
380            .unwrap();
381        let propagated = response["traceparent"].as_str().unwrap();
382        assert!(propagated.starts_with(&format!("00-{}-", parent.trace_id())));
383
384        server.stop(true).await;
385    }
386
387    #[actix_web::test]
388    async fn reports_status_decode_and_response_limit_errors() {
389        let (base_url, server) = spawn_server().await;
390        let client =
391            HttpClient::new(HttpClientConfig::new("users").with_max_response_bytes(4)).unwrap();
392
393        let status = client
394            .get_json::<Value>(format!("{base_url}/failure"))
395            .await
396            .unwrap_err();
397        assert!(matches!(
398            status,
399            HttpClientError::Status(StatusCode::SERVICE_UNAVAILABLE)
400        ));
401
402        let invalid = client
403            .get_json::<Value>(format!("{base_url}/invalid"))
404            .await
405            .unwrap_err();
406        assert!(matches!(
407            invalid,
408            HttpClientError::BodyTooLarge { limit: 4 }
409        ));
410
411        let chunked = client
412            .get_json::<Value>(format!("{base_url}/chunked"))
413            .await
414            .unwrap_err();
415        assert!(matches!(
416            chunked,
417            HttpClientError::BodyTooLarge { limit: 4 }
418        ));
419
420        let decode_client = HttpClient::new(HttpClientConfig::new("users")).unwrap();
421        let decode = decode_client
422            .get_json::<Value>(format!("{base_url}/invalid"))
423            .await
424            .unwrap_err();
425        assert!(matches!(decode, HttpClientError::Decode(_)));
426
427        server.stop(true).await;
428    }
429
430    #[actix_web::test]
431    async fn opens_the_circuit_after_a_server_failure() {
432        let (base_url, server) = spawn_server().await;
433        let metrics = Metrics::new();
434        let client = HttpClient::new(
435            HttpClientConfig::new("inventory")
436                .with_breaker(CircuitBreakerConfig::new(1, Duration::from_secs(60))),
437        )
438        .unwrap()
439        .with_metrics(HttpClientMetrics::new(&metrics, "test").unwrap());
440
441        let first = client
442            .execute(
443                client
444                    .request(Method::GET, format!("{base_url}/failure"))
445                    .build()
446                    .unwrap(),
447            )
448            .await
449            .unwrap();
450        assert_eq!(first.status(), StatusCode::SERVICE_UNAVAILABLE);
451
452        let second = client
453            .execute(
454                client
455                    .request(Method::GET, format!("{base_url}/get"))
456                    .build()
457                    .unwrap(),
458            )
459            .await
460            .unwrap_err();
461        assert!(matches!(
462            second,
463            HttpClientError::CircuitOpen { service } if service == "inventory"
464        ));
465
466        let rendered = metrics.render();
467        assert!(rendered.contains(
468            "test_http_client_requests_total{service=\"inventory\",method=\"GET\",result=\"503\"} 1"
469        ));
470        assert!(rendered.contains(
471            "test_http_client_requests_total{service=\"inventory\",method=\"GET\",result=\"circuit_open\"} 1"
472        ));
473        assert!(rendered.contains(
474            "test_http_client_requests_in_flight{service=\"inventory\",method=\"GET\"} 0"
475        ));
476
477        server.stop(true).await;
478    }
479
480    #[actix_web::test]
481    async fn reports_request_build_and_transport_errors() {
482        let client = HttpClient::new(HttpClientConfig::new("users")).unwrap();
483        let build = client.get_json::<Value>("not a URL").await.unwrap_err();
484        assert!(matches!(build, HttpClientError::Build(_)));
485        assert!(std::error::Error::source(&build).is_some());
486
487        let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
488        let address = listener.local_addr().unwrap();
489        drop(listener);
490        let request = client
491            .request(Method::GET, format!("http://{address}"))
492            .build()
493            .unwrap();
494        let transport = client.execute(request).await.unwrap_err();
495        assert!(matches!(transport, HttpClientError::Transport(_)));
496        assert!(std::error::Error::source(&transport).is_some());
497    }
498
499    #[test]
500    fn formats_public_errors() {
501        let invalid = HttpClientError::InvalidServiceName;
502        assert_eq!(invalid.to_string(), "HTTP service name cannot be empty");
503        assert!(std::error::Error::source(&invalid).is_none());
504
505        assert_eq!(
506            HttpClientError::CircuitOpen {
507                service: "users".to_owned()
508            }
509            .to_string(),
510            "HTTP circuit for service users is open"
511        );
512        assert_eq!(
513            HttpClientError::Status(StatusCode::BAD_GATEWAY).to_string(),
514            "HTTP service returned 502 Bad Gateway"
515        );
516        assert_eq!(
517            HttpClientError::BodyTooLarge { limit: 16 }.to_string(),
518            "HTTP response exceeds the 16-byte limit"
519        );
520    }
521
522    #[test]
523    #[should_panic(expected = "HTTP timeout must be greater than zero")]
524    fn rejects_zero_timeouts() {
525        let _ = HttpClientConfig::new("users").with_timeout(Duration::ZERO);
526    }
527
528    #[test]
529    #[should_panic(expected = "HTTP response limit must be greater than zero")]
530    fn rejects_zero_response_limits() {
531        let _ = HttpClientConfig::new("users").with_max_response_bytes(0);
532    }
533}