cloud_sdk_reqwest/asynchronous/
client.rs1use core::fmt;
2use std::sync::Arc;
3
4use cloud_sdk::Method;
5use cloud_sdk::authentication::{
6 AsyncAuthenticatedTransport, AuthenticatedRequest, CredentialGeneration,
7};
8use cloud_sdk::transport::{
9 BoundTransport, EndpointIdentity, EndpointIdentityError, ResponseAttempt, ResponseMetadata,
10 ResponseStorageSanitizer, ResponseWriter, StatusCode,
11};
12use cloud_sdk_sanitization::{SecretBuffer, sanitize_bytes};
13use reqwest::header::{AUTHORIZATION, HeaderName, HeaderValue};
14use reqwest::{Body, Client};
15
16use crate::shared::{
17 BearerCredential, BearerCredentialScope, BearerCredentialSnapshot, BearerRefreshHandoff,
18 BearerToken, CredentialStateError, CredentialStore, CredentialUpdateError, HttpsEndpoint,
19 TokenRefreshError, TokenRotationError, TransportError, capture_response_headers,
20 map_authentication_error, parse_rate_limit, parse_response_content_type,
21 validate_bearer_authentication,
22};
23
24use super::body::SanitizedBuffer;
25
26#[derive(Clone)]
31pub struct AsyncClient {
32 client: Client,
33 endpoint: HttpsEndpoint,
34 scope: Arc<BearerCredentialScope>,
35 credentials: Arc<CredentialStore>,
36 allow_insecure_loopback: bool,
37}
38
39impl AsyncClient {
40 pub(super) fn new(
41 client: Client,
42 endpoint: HttpsEndpoint,
43 credential: BearerCredential,
44 allow_insecure_loopback: bool,
45 ) -> Self {
46 Self {
47 client,
48 endpoint,
49 scope: Arc::new(credential.scope),
50 credentials: Arc::new(CredentialStore::new(credential.token)),
51 allow_insecure_loopback,
52 }
53 }
54
55 pub fn credential_snapshot(&self) -> Result<BearerCredentialSnapshot, CredentialStateError> {
57 self.credentials.snapshot()
58 }
59
60 pub fn rotate_bearer_token(
62 &self,
63 replacement: BearerToken,
64 ) -> Result<CredentialGeneration, CredentialUpdateError> {
65 self.credentials.rotate(replacement)
66 }
67
68 pub fn rotate_bearer_token_from_mut_bytes(
70 &self,
71 source: &mut [u8],
72 ) -> Result<CredentialGeneration, TokenRotationError> {
73 self.credentials.rotate_from_mut_bytes(source)
74 }
75
76 pub fn rotate_bearer_token_from_secret_buffer(
78 &self,
79 source: SecretBuffer<'_>,
80 ) -> Result<CredentialGeneration, TokenRotationError> {
81 self.credentials.rotate_from_secret_buffer(source)
82 }
83
84 pub fn refresh_bearer_token(
86 &self,
87 handoff: BearerRefreshHandoff,
88 replacement: BearerToken,
89 ) -> Result<CredentialGeneration, TokenRefreshError> {
90 self.credentials.refresh(handoff, replacement)
91 }
92
93 pub fn refresh_bearer_token_from_mut_bytes(
95 &self,
96 handoff: BearerRefreshHandoff,
97 source: &mut [u8],
98 ) -> Result<CredentialGeneration, TokenRefreshError> {
99 self.credentials.refresh_from_mut_bytes(handoff, source)
100 }
101
102 pub fn refresh_bearer_token_from_secret_buffer(
104 &self,
105 handoff: BearerRefreshHandoff,
106 source: SecretBuffer<'_>,
107 ) -> Result<CredentialGeneration, TokenRefreshError> {
108 self.credentials.refresh_from_secret_buffer(handoff, source)
109 }
110
111 async fn send_inner(
112 &self,
113 authenticated: AuthenticatedRequest<'_, '_>,
114 response_writer: &mut ResponseWriter<'_>,
115 ) -> Result<(), TransportError> {
116 let mut response_attempt = response_writer
117 .begin_attempt()
118 .map_err(|_| TransportError::ResponseCommitFailed)?;
119 let endpoint_identity = self
120 .endpoint
121 .identity()
122 .map_err(|_| TransportError::AuthenticationEndpointMismatch)?;
123 validate_bearer_authentication(
124 endpoint_identity,
125 &self.scope,
126 authenticated.policy(),
127 self.allow_insecure_loopback,
128 )
129 .map_err(map_authentication_error)?;
130 let token_snapshot = self
131 .credentials
132 .snapshot()
133 .map_err(|_| TransportError::CredentialStateUnavailable)?;
134 let authorization = token_snapshot
135 .header_value()
136 .map_err(|_| TransportError::HeaderRejected)?;
137 drop(token_snapshot);
138 execute(
139 &self.client,
140 &self.endpoint,
141 authorization,
142 authenticated,
143 &mut response_attempt,
144 )
145 .await
146 }
147}
148
149pub(super) async fn execute(
150 client: &Client,
151 endpoint: &HttpsEndpoint,
152 authorization: HeaderValue,
153 authenticated: AuthenticatedRequest<'_, '_>,
154 response_writer: &mut ResponseAttempt<'_, '_>,
155) -> Result<(), TransportError> {
156 let request = authenticated.transport_request();
157 let url = endpoint
158 .compose(request.target())
159 .map_err(|_| TransportError::TargetRejected)?;
160 let mut outbound = client
161 .request(map_method(request.method())?, url)
162 .header(AUTHORIZATION, authorization);
163
164 for header in request.headers().as_slice() {
165 let name = HeaderName::from_bytes(header.name().as_str().as_bytes())
166 .map_err(|_| TransportError::HeaderRejected)?;
167 let mut value = HeaderValue::from_str(header.value().as_str())
168 .map_err(|_| TransportError::HeaderRejected)?;
169 value.set_sensitive(matches!(
170 header.sensitivity(),
171 cloud_sdk::transport::HeaderSensitivity::Sensitive
172 ));
173 outbound = outbound.header(name, value);
174 }
175 if !request.body().is_empty() && request.headers().get("content-type").is_none() {
176 return Err(TransportError::MissingContentType);
177 }
178 if !request.body().is_empty() {
179 let body = SanitizedBuffer::copy_from(request.body())
180 .map_err(|_| TransportError::RequestBodyAllocationFailed)?;
181 let _ =
182 u64::try_from(request.body().len()).map_err(|_| TransportError::RequestBodyTooLarge)?;
183 outbound = outbound.body(Body::from(body.into_bytes()));
184 }
185
186 let mut response = outbound.send().await.map_err(classify_reqwest_error)?;
187 endpoint
188 .verify_origin(response.url())
189 .map_err(|_| TransportError::ResponseOriginChanged)?;
190 if response.content_length().is_some_and(|length| {
191 u64::try_from(response_writer.body_capacity()).map_or(true, |cap| length > cap)
192 }) {
193 return Err(TransportError::ResponseTooLarge);
194 }
195 let status =
196 StatusCode::new(response.status().as_u16()).ok_or(TransportError::InvalidStatus)?;
197 let buffered = read_response(&mut response, response_writer.body_capacity()).await?;
198 capture_response_headers(
199 response.headers(),
200 response_writer
201 .headers_mut()
202 .map_err(|_| TransportError::ResponseCommitFailed)?,
203 )?;
204 let rate_limit = parse_rate_limit(response_writer.headers())?;
205 parse_response_content_type(response_writer.headers())?;
206 let body_len = buffered.len();
207 let initialized = response_writer
208 .body_mut()
209 .map_err(|_| TransportError::ResponseCommitFailed)?
210 .get_mut(..body_len)
211 .ok_or(TransportError::ResponseReadFailed)?;
212 initialized.copy_from_slice(buffered.as_ref());
213 let mut metadata = ResponseMetadata::EMPTY;
214 if let Some(value) = rate_limit {
215 metadata = metadata.with_rate_limit(value);
216 }
217 response_writer
218 .commit(status, body_len, metadata)
219 .map_err(|_| TransportError::ResponseCommitFailed)
220}
221
222impl AsyncAuthenticatedTransport for AsyncClient {
223 type Error = TransportError;
224
225 async fn send_authenticated<'transport, 'request, 'policy, 'writer>(
226 &'transport self,
227 request: AuthenticatedRequest<'request, 'policy>,
228 response: &'writer mut ResponseWriter<'_>,
229 ) -> Result<(), Self::Error>
230 where
231 'transport: 'writer,
232 'request: 'writer,
233 'policy: 'writer,
234 {
235 self.send_inner(request, response).await
236 }
237}
238
239impl ResponseStorageSanitizer for AsyncClient {
240 fn sanitize_response_storage(&self, response_storage: &mut [u8]) {
241 sanitize_bytes(response_storage);
242 }
243}
244
245impl BoundTransport for AsyncClient {
246 fn endpoint_identity(&self) -> Result<EndpointIdentity<'_>, EndpointIdentityError> {
247 self.endpoint.identity()
248 }
249}
250
251impl fmt::Debug for AsyncClient {
252 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
253 formatter
254 .debug_struct("AsyncClient")
255 .field("endpoint", &"[redacted]")
256 .field("scope", &"[redacted]")
257 .field("credentials", &"[redacted]")
258 .finish_non_exhaustive()
259 }
260}
261
262async fn read_response(
263 response: &mut reqwest::Response,
264 limit: usize,
265) -> Result<SanitizedBuffer, TransportError> {
266 let mut buffered = SanitizedBuffer::with_capacity(limit)
267 .map_err(|_| TransportError::ResponseBodyAllocationFailed)?;
268 loop {
269 let chunk = response
270 .chunk()
271 .await
272 .map_err(|_| TransportError::ResponseReadFailed)?;
273 let Some(chunk) = chunk else { break };
274 buffered
275 .extend_bounded(&chunk, limit)
276 .map_err(|_| TransportError::ResponseTooLarge)?;
277 }
278 Ok(buffered)
279}
280
281fn map_method(method: Method) -> Result<reqwest::Method, TransportError> {
282 reqwest::Method::from_bytes(method.as_str().as_bytes())
283 .map_err(|_| TransportError::MethodRejected)
284}
285
286fn classify_reqwest_error(error: reqwest::Error) -> TransportError {
287 if error.is_timeout() {
288 TransportError::TimedOut
289 } else if error.is_connect() {
290 TransportError::ConnectFailed
291 } else {
292 TransportError::RequestFailed
293 }
294}