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,
7 ResponseStorageSanitizer, StatusCode, TransportRequest, TransportResponse,
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<'buffer>(
70 &self,
71 request: TransportRequest<'_>,
72 response_body: &'buffer mut [u8],
73 ) -> Result<TransportResponse<'buffer>, TransportError> {
74 sanitize_bytes(response_body);
75 let url = self
76 .endpoint
77 .compose(request.target())
78 .map_err(|_| TransportError::TargetRejected)?;
79 let token_snapshot = self
80 .credentials
81 .snapshot()
82 .map_err(|_| TransportError::CredentialStateUnavailable)?;
83 let authorization = token_snapshot
84 .header_value()
85 .map_err(|_| TransportError::HeaderRejected)?;
86 let mut outbound = self
87 .client
88 .request(map_method(request.method())?, url)
89 .header(AUTHORIZATION, authorization);
90
91 for header in request.headers().as_slice() {
92 let name = HeaderName::from_bytes(header.name().as_str().as_bytes())
93 .map_err(|_| TransportError::HeaderRejected)?;
94 let mut value = HeaderValue::from_str(header.value().as_str())
95 .map_err(|_| TransportError::HeaderRejected)?;
96 value.set_sensitive(matches!(
97 header.sensitivity(),
98 cloud_sdk::transport::HeaderSensitivity::Sensitive
99 ));
100 outbound = outbound.header(name, value);
101 }
102 if !request.body().is_empty() && request.headers().get("content-type").is_none() {
103 return Err(TransportError::MissingContentType);
104 }
105
106 if !request.body().is_empty() {
107 let body = SanitizedBuffer::copy_from(request.body())
108 .map_err(|_| TransportError::RequestBodyAllocationFailed)?;
109 let _ = u64::try_from(request.body().len())
110 .map_err(|_| TransportError::RequestBodyTooLarge)?;
111 outbound = outbound.body(Body::from(body.into_bytes()));
112 }
113
114 let mut response = outbound.send().await.map_err(classify_reqwest_error)?;
115 self.endpoint
116 .verify_origin(response.url())
117 .map_err(|_| TransportError::ResponseOriginChanged)?;
118 if response.content_length().is_some_and(|length| {
119 u64::try_from(response_body.len()).map_or(true, |cap| length > cap)
120 }) {
121 return Err(TransportError::ResponseTooLarge);
122 }
123 let status =
124 StatusCode::new(response.status().as_u16()).ok_or(TransportError::InvalidStatus)?;
125 let buffered = read_response(&mut response, response_body.len()).await?;
126 let headers = capture_response_headers(response.headers())?;
127 let rate_limit = parse_rate_limit(&headers)?;
128 let content_type = parse_response_content_type(&headers)?;
129 let body_len = buffered.len();
130 let initialized = response_body
131 .get_mut(..body_len)
132 .ok_or(TransportError::ResponseReadFailed)?;
133 initialized.copy_from_slice(buffered.as_ref());
134 let response = TransportResponse::new(status, initialized).with_headers(headers);
135 let response = content_type.map_or(response, |value| response.with_content_type(value));
136 drop(token_snapshot);
137 Ok(rate_limit.map_or(response, |value| response.with_rate_limit(value)))
138 }
139}
140
141impl AsyncTransport for AsyncClient {
142 type Error = TransportError;
143
144 async fn send<'transport, 'request, 'buffer>(
145 &'transport self,
146 request: TransportRequest<'request>,
147 response_body: &'buffer mut [u8],
148 ) -> Result<TransportResponse<'buffer>, Self::Error>
149 where
150 'request: 'transport,
151 'buffer: 'transport,
152 {
153 self.send_inner(request, response_body).await
154 }
155}
156
157impl ResponseStorageSanitizer for AsyncClient {
158 fn sanitize_response_storage(&self, response_storage: &mut [u8]) {
159 sanitize_bytes(response_storage);
160 }
161}
162
163impl BoundTransport for AsyncClient {
164 fn endpoint_identity(&self) -> Result<EndpointIdentity<'_>, EndpointIdentityError> {
165 self.endpoint.identity()
166 }
167}
168
169impl fmt::Debug for AsyncClient {
170 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
171 formatter
172 .debug_struct("AsyncClient")
173 .field("endpoint", &"[redacted]")
174 .field("credentials", &"[redacted]")
175 .finish_non_exhaustive()
176 }
177}
178
179async fn read_response(
180 response: &mut reqwest::Response,
181 limit: usize,
182) -> Result<SanitizedBuffer, TransportError> {
183 let mut buffered = SanitizedBuffer::with_capacity(limit)
184 .map_err(|_| TransportError::ResponseBodyAllocationFailed)?;
185 loop {
186 let chunk = response
187 .chunk()
188 .await
189 .map_err(|_| TransportError::ResponseReadFailed)?;
190 let Some(chunk) = chunk else { break };
191 buffered
192 .extend_bounded(&chunk, limit)
193 .map_err(|_| TransportError::ResponseTooLarge)?;
194 }
195 Ok(buffered)
196}
197
198fn map_method(method: Method) -> Result<reqwest::Method, TransportError> {
199 reqwest::Method::from_bytes(method.as_str().as_bytes())
200 .map_err(|_| TransportError::MethodRejected)
201}
202
203fn classify_reqwest_error(error: reqwest::Error) -> TransportError {
204 if error.is_timeout() {
205 TransportError::TimedOut
206 } else if error.is_connect() {
207 TransportError::ConnectFailed
208 } else {
209 TransportError::RequestFailed
210 }
211}