Skip to main content

codex_http_client/
transport.rs

1use crate::client::HttpClient;
2use crate::client::RequestBuilder;
3use crate::error::TransportError;
4use crate::request::Request;
5use crate::request::RequestBody;
6use crate::request::Response;
7use bytes::Bytes;
8use futures::StreamExt;
9use futures::stream::BoxStream;
10use http::HeaderMap;
11use http::Method;
12use http::StatusCode;
13use tracing::Level;
14use tracing::enabled;
15use tracing::trace;
16
17pub type ByteStream = BoxStream<'static, Result<Bytes, TransportError>>;
18
19pub struct StreamResponse {
20    pub status: StatusCode,
21    pub headers: HeaderMap,
22    pub bytes: ByteStream,
23}
24
25pub trait HttpTransport: Send + Sync {
26    fn execute(
27        &self,
28        req: Request,
29    ) -> impl std::future::Future<Output = Result<Response, TransportError>> + Send;
30    fn stream(
31        &self,
32        req: Request,
33    ) -> impl std::future::Future<Output = Result<StreamResponse, TransportError>> + Send;
34}
35
36#[derive(Clone, Debug)]
37pub struct ReqwestTransport {
38    client: HttpClient,
39}
40
41impl ReqwestTransport {
42    pub fn new(client: reqwest::Client) -> Self {
43        Self {
44            client: HttpClient::new(client),
45        }
46    }
47
48    pub fn from_http_client(client: HttpClient) -> Self {
49        Self { client }
50    }
51
52    fn build(&self, req: Request) -> Result<RequestBuilder, TransportError> {
53        let prepared = req.prepare_body_for_send().map_err(TransportError::Build)?;
54
55        let Request {
56            method,
57            url,
58            headers: _,
59            body: _,
60            compression: _,
61            timeout,
62        } = req;
63
64        let mut builder = self.client.request(
65            Method::from_bytes(method.as_str().as_bytes()).unwrap_or(Method::GET),
66            &url,
67        );
68
69        if let Some(timeout) = timeout {
70            builder = builder.timeout(timeout);
71        }
72
73        builder = builder.headers(prepared.headers);
74        if let Some(body) = prepared.body {
75            builder = builder.body(body);
76        }
77        Ok(builder)
78    }
79
80    fn map_error(err: reqwest::Error) -> TransportError {
81        if err.is_timeout() {
82            TransportError::Timeout
83        } else {
84            TransportError::Network(err.to_string())
85        }
86    }
87
88    fn trace_request(&self, req: &Request) {
89        if self.client.request_logging_enabled() && enabled!(Level::TRACE) {
90            trace!(
91                "{} to {}: {}",
92                req.method,
93                req.url,
94                request_body_for_trace(req)
95            );
96        }
97    }
98}
99
100fn request_body_for_trace(req: &Request) -> String {
101    match req.body.as_ref() {
102        Some(RequestBody::Json(body)) => body.to_string(),
103        Some(RequestBody::EncodedJson(body)) => {
104            String::from_utf8_lossy(body.trace_bytes()).into_owned()
105        }
106        Some(RequestBody::Raw(body)) => format!("<raw body: {} bytes>", body.len()),
107        None => String::new(),
108    }
109}
110
111impl HttpTransport for ReqwestTransport {
112    async fn execute(&self, req: Request) -> Result<Response, TransportError> {
113        self.trace_request(&req);
114
115        let url = req.url.clone();
116        let builder = self.build(req)?;
117        let resp = builder.send().await.map_err(Self::map_error)?;
118        let status = resp.status();
119        let headers = resp.headers().clone();
120        let bytes = resp.bytes().await.map_err(Self::map_error)?;
121        if !status.is_success() {
122            let body = String::from_utf8(bytes.to_vec()).ok();
123            return Err(TransportError::Http {
124                status,
125                url: Some(url),
126                headers: Some(headers),
127                body,
128            });
129        }
130        Ok(Response {
131            status,
132            headers,
133            body: bytes,
134        })
135    }
136
137    async fn stream(&self, req: Request) -> Result<StreamResponse, TransportError> {
138        self.trace_request(&req);
139
140        let url = req.url.clone();
141        let builder = self.build(req)?;
142        let resp = builder.send().await.map_err(Self::map_error)?;
143        let status = resp.status();
144        let headers = resp.headers().clone();
145        if !status.is_success() {
146            let body = resp.text().await.ok();
147            return Err(TransportError::Http {
148                status,
149                url: Some(url),
150                headers: Some(headers),
151                body,
152            });
153        }
154        let stream = resp
155            .bytes_stream()
156            .map(|result| result.map_err(Self::map_error));
157        Ok(StreamResponse {
158            status,
159            headers,
160            bytes: Box::pin(stream),
161        })
162    }
163}
164
165#[cfg(test)]
166#[path = "transport_tests.rs"]
167mod tests;