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, 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 AuthenticationValidationError, BearerCredential, BearerCredentialScope,
18 BearerCredentialSnapshot, BearerRefreshHandoff, BearerToken, CredentialStateError,
19 CredentialStore, CredentialUpdateError, HttpsEndpoint, TokenRefreshError, TokenRotationError,
20 TransportError, capture_response_headers, 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 if response_writer.is_committed() {
117 return Err(TransportError::ResponseCommitFailed);
118 }
119 let mut response_writer = response_writer
120 .begin_attempt()
121 .map_err(|_| TransportError::ResponseCommitFailed)?;
122 let endpoint_identity = self
123 .endpoint
124 .identity()
125 .map_err(|_| TransportError::AuthenticationEndpointMismatch)?;
126 validate_bearer_authentication(
127 endpoint_identity,
128 &self.scope,
129 authenticated.policy(),
130 self.allow_insecure_loopback,
131 )
132 .map_err(map_authentication_error)?;
133 let request = authenticated.transport_request();
134 let url = self
135 .endpoint
136 .compose(request.target())
137 .map_err(|_| TransportError::TargetRejected)?;
138 let token_snapshot = self
139 .credentials
140 .snapshot()
141 .map_err(|_| TransportError::CredentialStateUnavailable)?;
142 let authorization = token_snapshot
143 .header_value()
144 .map_err(|_| TransportError::HeaderRejected)?;
145 let mut outbound = self
146 .client
147 .request(map_method(request.method())?, url)
148 .header(AUTHORIZATION, authorization);
149
150 for header in request.headers().as_slice() {
151 let name = HeaderName::from_bytes(header.name().as_str().as_bytes())
152 .map_err(|_| TransportError::HeaderRejected)?;
153 let mut value = HeaderValue::from_str(header.value().as_str())
154 .map_err(|_| TransportError::HeaderRejected)?;
155 value.set_sensitive(matches!(
156 header.sensitivity(),
157 cloud_sdk::transport::HeaderSensitivity::Sensitive
158 ));
159 outbound = outbound.header(name, value);
160 }
161 if !request.body().is_empty() && request.headers().get("content-type").is_none() {
162 return Err(TransportError::MissingContentType);
163 }
164 if !request.body().is_empty() {
165 let body = SanitizedBuffer::copy_from(request.body())
166 .map_err(|_| TransportError::RequestBodyAllocationFailed)?;
167 let _ = u64::try_from(request.body().len())
168 .map_err(|_| TransportError::RequestBodyTooLarge)?;
169 outbound = outbound.body(Body::from(body.into_bytes()));
170 }
171
172 let mut response = outbound.send().await.map_err(classify_reqwest_error)?;
173 self.endpoint
174 .verify_origin(response.url())
175 .map_err(|_| TransportError::ResponseOriginChanged)?;
176 if response.content_length().is_some_and(|length| {
177 u64::try_from(response_writer.body_capacity()).map_or(true, |cap| length > cap)
178 }) {
179 return Err(TransportError::ResponseTooLarge);
180 }
181 let status =
182 StatusCode::new(response.status().as_u16()).ok_or(TransportError::InvalidStatus)?;
183 let buffered = read_response(&mut response, response_writer.body_capacity()).await?;
184 capture_response_headers(
185 response.headers(),
186 response_writer
187 .headers_mut()
188 .map_err(|_| TransportError::ResponseCommitFailed)?,
189 )?;
190 let rate_limit = parse_rate_limit(response_writer.headers())?;
191 parse_response_content_type(response_writer.headers())?;
192 let body_len = buffered.len();
193 let initialized = response_writer
194 .body_mut()
195 .map_err(|_| TransportError::ResponseCommitFailed)?
196 .get_mut(..body_len)
197 .ok_or(TransportError::ResponseReadFailed)?;
198 initialized.copy_from_slice(buffered.as_ref());
199 let mut metadata = ResponseMetadata::EMPTY;
200 if let Some(value) = rate_limit {
201 metadata = metadata.with_rate_limit(value);
202 }
203 drop(token_snapshot);
204 response_writer
205 .commit(status, body_len, metadata)
206 .map_err(|_| TransportError::ResponseCommitFailed)
207 }
208}
209
210impl AsyncAuthenticatedTransport for AsyncClient {
211 type Error = TransportError;
212
213 async fn send_authenticated<'transport, 'request, 'policy, 'writer>(
214 &'transport self,
215 request: AuthenticatedRequest<'request, 'policy>,
216 response: &'writer mut ResponseWriter<'_>,
217 ) -> Result<(), Self::Error>
218 where
219 'transport: 'writer,
220 'request: 'writer,
221 'policy: 'writer,
222 {
223 self.send_inner(request, response).await
224 }
225}
226
227impl ResponseStorageSanitizer for AsyncClient {
228 fn sanitize_response_storage(&self, response_storage: &mut [u8]) {
229 sanitize_bytes(response_storage);
230 }
231}
232
233impl BoundTransport for AsyncClient {
234 fn endpoint_identity(&self) -> Result<EndpointIdentity<'_>, EndpointIdentityError> {
235 self.endpoint.identity()
236 }
237}
238
239impl fmt::Debug for AsyncClient {
240 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
241 formatter
242 .debug_struct("AsyncClient")
243 .field("endpoint", &"[redacted]")
244 .field("scope", &"[redacted]")
245 .field("credentials", &"[redacted]")
246 .finish_non_exhaustive()
247 }
248}
249
250fn map_authentication_error(error: AuthenticationValidationError) -> TransportError {
251 match error {
252 AuthenticationValidationError::InsecureEndpoint => {
253 TransportError::InsecureAuthenticationEndpoint
254 }
255 AuthenticationValidationError::EndpointMismatch => {
256 TransportError::AuthenticationEndpointMismatch
257 }
258 AuthenticationValidationError::IncompletePolicy => {
259 TransportError::AuthenticationScopeRejected
260 }
261 AuthenticationValidationError::ScopeRejected => TransportError::AuthenticationScopeRejected,
262 }
263}
264
265async fn read_response(
266 response: &mut reqwest::Response,
267 limit: usize,
268) -> Result<SanitizedBuffer, TransportError> {
269 let mut buffered = SanitizedBuffer::with_capacity(limit)
270 .map_err(|_| TransportError::ResponseBodyAllocationFailed)?;
271 loop {
272 let chunk = response
273 .chunk()
274 .await
275 .map_err(|_| TransportError::ResponseReadFailed)?;
276 let Some(chunk) = chunk else { break };
277 buffered
278 .extend_bounded(&chunk, limit)
279 .map_err(|_| TransportError::ResponseTooLarge)?;
280 }
281 Ok(buffered)
282}
283
284fn map_method(method: Method) -> Result<reqwest::Method, TransportError> {
285 reqwest::Method::from_bytes(method.as_str().as_bytes())
286 .map_err(|_| TransportError::MethodRejected)
287}
288
289fn classify_reqwest_error(error: reqwest::Error) -> TransportError {
290 if error.is_timeout() {
291 TransportError::TimedOut
292 } else if error.is_connect() {
293 TransportError::ConnectFailed
294 } else {
295 TransportError::RequestFailed
296 }
297}