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::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/// Hardened provider-neutral reqwest asynchronous transport.
21///
22/// The adapter uses reqwest's Tokio-based execution internally but does not
23/// install or own a runtime. Callers must poll it from a compatible executor.
24#[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    /// Atomically replaces the bearer token used by newly started requests.
41    ///
42    /// In-flight requests retain their previous snapshot. No credential lock
43    /// is held across network I/O or `.await`.
44    pub fn rotate_bearer_token(
45        &self,
46        replacement: BearerToken,
47    ) -> Result<(), CredentialStateError> {
48        self.credentials.rotate(replacement)
49    }
50
51    /// Validates and rotates from mutable bytes, clearing the complete source
52    /// on success or failure. Rejected input leaves the active token unchanged.
53    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    /// Validates and rotates from guarded storage. Dropping the consumed guard
61    /// clears the complete source on success or failure.
62    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}