Skip to main content

codex_http_client/
client.rs

1//! Reusable HTTP client and request-builder wrappers.
2
3use 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/// Reusable HTTP client wrapper with shared tracing and request-diagnostic behavior.
21///
22/// Product callers should obtain this through [`crate::HttpClientFactory`] for a fixed
23/// destination or use [`crate::RouteAwareClientPool`] when request and redirect URLs can vary.
24#[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    /// Creates a client that suppresses request URL and response-header diagnostics.
36    ///
37    /// Use this for endpoints whose URLs or headers may contain credentials that are redacted by
38    /// the caller above the HTTP transport boundary.
39    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}