cloud_sdk_reqwest/asynchronous/
client.rs1use core::fmt;
2use std::sync::Arc;
3
4use cloud_sdk::Method;
5use cloud_sdk::transport::{
6 AsyncTransport, BoundTransport, EndpointIdentity, EndpointIdentityError, ResponseMetadata,
7 ResponseStorageSanitizer, ResponseWriter, StatusCode, TransportRequest,
8};
9use cloud_sdk_sanitization::{SecretBuffer, sanitize_bytes};
10use reqwest::header::{AUTHORIZATION, HeaderName, HeaderValue};
11use reqwest::{Body, Client};
12
13use crate::shared::{
14 BearerToken, CredentialStateError, CredentialStore, HttpsEndpoint, TokenRotationError,
15 TransportError, capture_response_headers, parse_rate_limit, parse_response_content_type,
16};
17
18use super::body::SanitizedBuffer;
19
20#[derive(Clone)]
25pub struct AsyncClient {
26 client: Client,
27 endpoint: HttpsEndpoint,
28 credentials: Arc<CredentialStore>,
29}
30
31impl AsyncClient {
32 pub(super) fn new(client: Client, endpoint: HttpsEndpoint, token: BearerToken) -> Self {
33 Self {
34 client,
35 endpoint,
36 credentials: Arc::new(CredentialStore::new(token)),
37 }
38 }
39
40 pub fn rotate_bearer_token(
45 &self,
46 replacement: BearerToken,
47 ) -> Result<(), CredentialStateError> {
48 self.credentials.rotate(replacement)
49 }
50
51 pub fn rotate_bearer_token_from_mut_bytes(
54 &self,
55 source: &mut [u8],
56 ) -> Result<(), TokenRotationError> {
57 self.credentials.rotate_from_mut_bytes(source)
58 }
59
60 pub fn rotate_bearer_token_from_secret_buffer(
63 &self,
64 source: SecretBuffer<'_>,
65 ) -> Result<(), TokenRotationError> {
66 self.credentials.rotate_from_secret_buffer(source)
67 }
68
69 async fn send_inner(
70 &self,
71 request: TransportRequest<'_>,
72 response_writer: &mut ResponseWriter<'_>,
73 ) -> Result<(), TransportError> {
74 if response_writer.is_committed() {
75 return Err(TransportError::ResponseCommitFailed);
76 }
77 let mut response_writer = response_writer
78 .begin_attempt()
79 .map_err(|_| TransportError::ResponseCommitFailed)?;
80 let url = self
81 .endpoint
82 .compose(request.target())
83 .map_err(|_| TransportError::TargetRejected)?;
84 let token_snapshot = self
85 .credentials
86 .snapshot()
87 .map_err(|_| TransportError::CredentialStateUnavailable)?;
88 let authorization = token_snapshot
89 .header_value()
90 .map_err(|_| TransportError::HeaderRejected)?;
91 let mut outbound = self
92 .client
93 .request(map_method(request.method())?, url)
94 .header(AUTHORIZATION, authorization);
95
96 for header in request.headers().as_slice() {
97 let name = HeaderName::from_bytes(header.name().as_str().as_bytes())
98 .map_err(|_| TransportError::HeaderRejected)?;
99 let mut value = HeaderValue::from_str(header.value().as_str())
100 .map_err(|_| TransportError::HeaderRejected)?;
101 value.set_sensitive(matches!(
102 header.sensitivity(),
103 cloud_sdk::transport::HeaderSensitivity::Sensitive
104 ));
105 outbound = outbound.header(name, value);
106 }
107 if !request.body().is_empty() && request.headers().get("content-type").is_none() {
108 return Err(TransportError::MissingContentType);
109 }
110
111 if !request.body().is_empty() {
112 let body = SanitizedBuffer::copy_from(request.body())
113 .map_err(|_| TransportError::RequestBodyAllocationFailed)?;
114 let _ = u64::try_from(request.body().len())
115 .map_err(|_| TransportError::RequestBodyTooLarge)?;
116 outbound = outbound.body(Body::from(body.into_bytes()));
117 }
118
119 let mut response = outbound.send().await.map_err(classify_reqwest_error)?;
120 self.endpoint
121 .verify_origin(response.url())
122 .map_err(|_| TransportError::ResponseOriginChanged)?;
123 if response.content_length().is_some_and(|length| {
124 u64::try_from(response_writer.body_capacity()).map_or(true, |cap| length > cap)
125 }) {
126 return Err(TransportError::ResponseTooLarge);
127 }
128 let status =
129 StatusCode::new(response.status().as_u16()).ok_or(TransportError::InvalidStatus)?;
130 let buffered = read_response(&mut response, response_writer.body_capacity()).await?;
131 capture_response_headers(
132 response.headers(),
133 response_writer
134 .headers_mut()
135 .map_err(|_| TransportError::ResponseCommitFailed)?,
136 )?;
137 let rate_limit = parse_rate_limit(response_writer.headers())?;
138 parse_response_content_type(response_writer.headers())?;
139 let body_len = buffered.len();
140 let initialized = response_writer
141 .body_mut()
142 .map_err(|_| TransportError::ResponseCommitFailed)?
143 .get_mut(..body_len)
144 .ok_or(TransportError::ResponseReadFailed)?;
145 initialized.copy_from_slice(buffered.as_ref());
146 let mut metadata = ResponseMetadata::EMPTY;
147 if let Some(value) = rate_limit {
148 metadata = metadata.with_rate_limit(value);
149 }
150 drop(token_snapshot);
151 response_writer
152 .commit(status, body_len, metadata)
153 .map_err(|_| TransportError::ResponseCommitFailed)
154 }
155}
156
157impl AsyncTransport for AsyncClient {
158 type Error = TransportError;
159
160 async fn send<'transport, 'request, 'writer>(
161 &'transport self,
162 request: TransportRequest<'request>,
163 response: &'writer mut ResponseWriter<'_>,
164 ) -> Result<(), Self::Error>
165 where
166 'transport: 'writer,
167 'request: 'writer,
168 {
169 self.send_inner(request, response).await
170 }
171}
172
173impl ResponseStorageSanitizer for AsyncClient {
174 fn sanitize_response_storage(&self, response_storage: &mut [u8]) {
175 sanitize_bytes(response_storage);
176 }
177}
178
179impl BoundTransport for AsyncClient {
180 fn endpoint_identity(&self) -> Result<EndpointIdentity<'_>, EndpointIdentityError> {
181 self.endpoint.identity()
182 }
183}
184
185impl fmt::Debug for AsyncClient {
186 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
187 formatter
188 .debug_struct("AsyncClient")
189 .field("endpoint", &"[redacted]")
190 .field("credentials", &"[redacted]")
191 .finish_non_exhaustive()
192 }
193}
194
195async fn read_response(
196 response: &mut reqwest::Response,
197 limit: usize,
198) -> Result<SanitizedBuffer, TransportError> {
199 let mut buffered = SanitizedBuffer::with_capacity(limit)
200 .map_err(|_| TransportError::ResponseBodyAllocationFailed)?;
201 loop {
202 let chunk = response
203 .chunk()
204 .await
205 .map_err(|_| TransportError::ResponseReadFailed)?;
206 let Some(chunk) = chunk else { break };
207 buffered
208 .extend_bounded(&chunk, limit)
209 .map_err(|_| TransportError::ResponseTooLarge)?;
210 }
211 Ok(buffered)
212}
213
214fn map_method(method: Method) -> Result<reqwest::Method, TransportError> {
215 reqwest::Method::from_bytes(method.as_str().as_bytes())
216 .map_err(|_| TransportError::MethodRejected)
217}
218
219fn classify_reqwest_error(error: reqwest::Error) -> TransportError {
220 if error.is_timeout() {
221 TransportError::TimedOut
222 } else if error.is_connect() {
223 TransportError::ConnectFailed
224 } else {
225 TransportError::RequestFailed
226 }
227}