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#[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#[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 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 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 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}