1use http::Error as HttpRequestBuildError;
4use http::HeaderMap;
5use http::HeaderName;
6use http::HeaderValue;
7use opentelemetry::global;
8use opentelemetry::propagation::Injector;
9use reqwest::IntoUrl;
10use reqwest::Method;
11use serde::Serialize;
12use std::fmt::Display;
13use std::time::Duration;
14use tracing::Span;
15use tracing_opentelemetry::OpenTelemetrySpanExt;
16
17pub type HttpError = reqwest::Error;
18pub type HttpResponse = reqwest::Response;
19
20#[derive(Clone, Debug)]
25pub struct HttpClient {
26 inner: reqwest::Client,
27 request_logging: RequestLogging,
28}
29
30impl HttpClient {
31 pub fn new(inner: reqwest::Client) -> Self {
32 Self::from_parts(inner, RequestLogging::Enabled)
33 }
34
35 pub fn new_without_request_logging(inner: reqwest::Client) -> Self {
40 Self::from_parts(inner, RequestLogging::Disabled)
41 }
42
43 pub(crate) fn from_parts(inner: reqwest::Client, request_logging: RequestLogging) -> Self {
44 Self {
45 inner,
46 request_logging,
47 }
48 }
49
50 pub fn get<U>(&self, url: U) -> RequestBuilder
51 where
52 U: IntoUrl,
53 {
54 self.request(Method::GET, url)
55 }
56
57 pub fn head<U>(&self, url: U) -> RequestBuilder
58 where
59 U: IntoUrl,
60 {
61 self.request(Method::HEAD, url)
62 }
63
64 pub fn post<U>(&self, url: U) -> RequestBuilder
65 where
66 U: IntoUrl,
67 {
68 self.request(Method::POST, url)
69 }
70
71 pub fn delete<U>(&self, url: U) -> RequestBuilder
72 where
73 U: IntoUrl,
74 {
75 self.request(Method::DELETE, url)
76 }
77
78 pub fn request<U>(&self, method: Method, url: U) -> RequestBuilder
79 where
80 U: IntoUrl,
81 {
82 let url_str = url.as_str().to_string();
83 RequestBuilder::new(
84 self.inner.request(method.clone(), url),
85 method,
86 url_str,
87 self.request_logging,
88 )
89 }
90
91 pub(crate) async fn execute(
92 &self,
93 request: reqwest::Request,
94 ) -> Result<reqwest::Response, reqwest::Error> {
95 let method = request.method().clone();
96 let url = request.url().to_string();
97
98 match self.execute_without_request_logging(request).await {
99 Ok(response) => {
100 self.log_response(&method, &url, &response);
101 Ok(response)
102 }
103 Err(error) => {
104 self.log_error(&method, &url, &error);
105 Err(error)
106 }
107 }
108 }
109
110 pub(crate) async fn execute_without_request_logging(
111 &self,
112 mut request: reqwest::Request,
113 ) -> Result<reqwest::Response, reqwest::Error> {
114 request.headers_mut().extend(trace_headers());
115 self.inner.execute(request).await
116 }
117
118 pub(crate) fn log_response(&self, method: &Method, url: &str, response: &reqwest::Response) {
119 if self.request_logging == RequestLogging::Enabled {
120 tracing::debug!(
121 method = %method,
122 url = %url,
123 status = %response.status(),
124 headers = ?response.headers(),
125 version = ?response.version(),
126 "Request completed"
127 );
128 }
129 }
130
131 pub(crate) fn log_error(&self, method: &Method, url: &str, error: &reqwest::Error) {
132 if self.request_logging == RequestLogging::Enabled {
133 tracing::debug!(
134 method = %method,
135 url = %url,
136 status = error.status().map(|status| status.as_u16()),
137 error = %error,
138 "Request failed"
139 );
140 }
141 }
142 pub(crate) fn log_error_summary(&self, method: &Method, url: &str, error: &reqwest::Error) {
143 if self.request_logging == RequestLogging::Enabled {
144 tracing::debug!(
145 method = %method,
146 url = %url,
147 status = error.status().map(|status| status.as_u16()),
148 is_timeout = error.is_timeout(),
149 is_connect = error.is_connect(),
150 "Request failed"
151 );
152 }
153 }
154
155 pub(crate) const fn request_logging_enabled(&self) -> bool {
156 matches!(self.request_logging, RequestLogging::Enabled)
157 }
158}
159
160#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
161pub(crate) enum RequestLogging {
162 #[default]
163 Enabled,
164 Disabled,
165}
166
167#[must_use = "requests are not sent unless `send` is awaited"]
168#[derive(Debug)]
169pub struct RequestBuilder {
170 builder: reqwest::RequestBuilder,
171 method: Method,
172 url: String,
173 request_logging: RequestLogging,
174}
175
176impl RequestBuilder {
177 fn new(
178 builder: reqwest::RequestBuilder,
179 method: Method,
180 url: String,
181 request_logging: RequestLogging,
182 ) -> Self {
183 Self {
184 builder,
185 method,
186 url,
187 request_logging,
188 }
189 }
190
191 fn map(self, f: impl FnOnce(reqwest::RequestBuilder) -> reqwest::RequestBuilder) -> Self {
192 Self {
193 builder: f(self.builder),
194 method: self.method,
195 url: self.url,
196 request_logging: self.request_logging,
197 }
198 }
199
200 pub fn headers(self, headers: HeaderMap) -> Self {
201 self.map(|builder| builder.headers(headers))
202 }
203
204 pub fn header<K, V>(self, key: K, value: V) -> Self
205 where
206 HeaderName: TryFrom<K>,
207 <HeaderName as TryFrom<K>>::Error: Into<HttpRequestBuildError>,
208 HeaderValue: TryFrom<V>,
209 <HeaderValue as TryFrom<V>>::Error: Into<HttpRequestBuildError>,
210 {
211 self.map(|builder| builder.header(key, value))
212 }
213
214 pub fn bearer_auth<T>(self, token: T) -> Self
215 where
216 T: Display,
217 {
218 self.map(|builder| builder.bearer_auth(token))
219 }
220
221 pub fn timeout(self, timeout: Duration) -> Self {
222 self.map(|builder| builder.timeout(timeout))
223 }
224
225 pub fn json<T>(self, value: &T) -> Self
226 where
227 T: ?Sized + Serialize,
228 {
229 self.map(|builder| builder.json(value))
230 }
231
232 pub fn query<T>(self, query: &T) -> Self
233 where
234 T: ?Sized + Serialize,
235 {
236 self.map(|builder| builder.query(query))
237 }
238
239 pub fn body<B>(self, body: B) -> Self
240 where
241 B: Into<reqwest::Body>,
242 {
243 self.map(|builder| builder.body(body))
244 }
245
246 pub async fn send(self) -> Result<HttpResponse, HttpError> {
247 let headers = trace_headers();
248
249 match self.builder.headers(headers).send().await {
250 Ok(response) => {
251 if self.request_logging == RequestLogging::Enabled {
252 tracing::debug!(
253 method = %self.method,
254 url = %self.url,
255 status = %response.status(),
256 headers = ?response.headers(),
257 version = ?response.version(),
258 "Request completed"
259 );
260 }
261
262 Ok(response)
263 }
264 Err(error) => {
265 if self.request_logging == RequestLogging::Enabled {
266 let status = error.status();
267 tracing::debug!(
268 method = %self.method,
269 url = %self.url,
270 status = status.map(|s| s.as_u16()),
271 error = %error,
272 "Request failed"
273 );
274 }
275 Err(error)
276 }
277 }
278 }
279}
280
281struct HeaderMapInjector<'a>(&'a mut HeaderMap);
282
283impl<'a> Injector for HeaderMapInjector<'a> {
284 fn set(&mut self, key: &str, value: String) {
285 if let (Ok(name), Ok(val)) = (
286 HeaderName::from_bytes(key.as_bytes()),
287 HeaderValue::from_str(&value),
288 ) {
289 self.0.insert(name, val);
290 }
291 }
292}
293
294pub(crate) fn trace_headers() -> HeaderMap {
295 let mut headers = HeaderMap::new();
296 global::get_text_map_propagator(|prop| {
297 prop.inject_context(
298 &Span::current().context(),
299 &mut HeaderMapInjector(&mut headers),
300 );
301 });
302 headers
303}
304
305#[cfg(test)]
306mod tests {
307 use super::*;
308 use opentelemetry::propagation::Extractor;
309 use opentelemetry::propagation::TextMapPropagator;
310 use opentelemetry::trace::TraceContextExt;
311 use opentelemetry::trace::TracerProvider;
312 use opentelemetry_sdk::propagation::TraceContextPropagator;
313 use opentelemetry_sdk::trace::SdkTracerProvider;
314 use pretty_assertions::assert_eq;
315 use tracing::trace_span;
316 use tracing_subscriber::layer::SubscriberExt;
317 use tracing_subscriber::util::SubscriberInitExt;
318
319 #[test]
320 fn inject_trace_headers_uses_current_span_context() {
321 global::set_text_map_propagator(TraceContextPropagator::new());
322
323 let provider = SdkTracerProvider::builder().build();
324 let tracer = provider.tracer("test-tracer");
325 let subscriber =
326 tracing_subscriber::registry().with(tracing_opentelemetry::layer().with_tracer(tracer));
327 let _guard = subscriber.set_default();
328
329 let span = trace_span!("client_request");
330 let _entered = span.enter();
331 let span_context = span.context().span().span_context().clone();
332
333 let headers = trace_headers();
334
335 let extractor = HeaderMapExtractor(&headers);
336 let extracted = TraceContextPropagator::new().extract(&extractor);
337 let extracted_span = extracted.span();
338 let extracted_context = extracted_span.span_context();
339
340 assert!(extracted_context.is_valid());
341 assert_eq!(extracted_context.trace_id(), span_context.trace_id());
342 assert_eq!(extracted_context.span_id(), span_context.span_id());
343 }
344
345 struct HeaderMapExtractor<'a>(&'a HeaderMap);
346
347 impl<'a> Extractor for HeaderMapExtractor<'a> {
348 fn get(&self, key: &str) -> Option<&str> {
349 self.0.get(key).and_then(|value| value.to_str().ok())
350 }
351
352 fn keys(&self) -> Vec<&str> {
353 self.0.keys().map(HeaderName::as_str).collect()
354 }
355 }
356}