Skip to main content

cloud_sdk_reqwest/asynchronous/
client.rs

1use 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/// Hardened provider-neutral reqwest asynchronous bearer transport.
27///
28/// The adapter uses reqwest's Tokio-based execution internally but does not
29/// install or own a runtime. Callers must poll it from a compatible executor.
30#[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    /// Captures the current generation without exposing token bytes.
56    pub fn credential_snapshot(&self) -> Result<BearerCredentialSnapshot, CredentialStateError> {
57        self.credentials.snapshot()
58    }
59
60    /// Atomically replaces the token while retaining immutable scope.
61    pub fn rotate_bearer_token(
62        &self,
63        replacement: BearerToken,
64    ) -> Result<CredentialGeneration, CredentialUpdateError> {
65        self.credentials.rotate(replacement)
66    }
67
68    /// Validates and rotates mutable bytes, clearing the complete source.
69    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    /// Validates and rotates guarded storage, which clears on return.
77    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    /// Installs a refresh only if its captured generation is still current.
85    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    /// Validates refreshed mutable bytes, clears them, and rejects stale work.
94    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    /// Consumes guarded refreshed storage and rejects stale work.
103    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}