codex_http_client/
transport.rs1use 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;