Skip to main content

cloud_sdk_reqwest/blocking/
client.rs

1use core::fmt;
2use std::io::Read;
3use std::sync::Arc;
4
5use cloud_sdk::Method;
6use cloud_sdk::transport::{
7    BlockingTransport, BoundTransport, EndpointIdentity, EndpointIdentityError,
8    ResponseStorageSanitizer, StatusCode, TransportRequest, TransportResponse,
9};
10use cloud_sdk_sanitization::{SecretBuffer, sanitize_bytes};
11use reqwest::blocking::{Body, Client};
12use reqwest::header::{AUTHORIZATION, HeaderName, HeaderValue};
13
14use super::body::{ReadBodyError, SanitizedRequestBody, read_bounded};
15use crate::shared::{
16    BearerToken, CredentialStateError, CredentialStore, HttpsEndpoint, TokenRotationError,
17    TransportError, capture_response_headers, parse_rate_limit, parse_response_content_type,
18};
19
20/// Hardened provider-neutral reqwest blocking transport.
21#[derive(Clone)]
22pub struct BlockingClient {
23    client: Client,
24    endpoint: HttpsEndpoint,
25    credentials: Arc<CredentialStore>,
26}
27
28impl BlockingClient {
29    pub(super) fn new(client: Client, endpoint: HttpsEndpoint, token: BearerToken) -> Self {
30        Self {
31            client,
32            endpoint,
33            credentials: Arc::new(CredentialStore::new(token)),
34        }
35    }
36
37    /// Atomically replaces the bearer token used by newly started requests.
38    ///
39    /// In-flight requests retain their previous snapshot. The retired token is
40    /// sanitized after its last request snapshot is dropped.
41    pub fn rotate_bearer_token(
42        &self,
43        replacement: BearerToken,
44    ) -> Result<(), CredentialStateError> {
45        self.credentials.rotate(replacement)
46    }
47
48    /// Validates and rotates from mutable bytes, clearing the complete source
49    /// on success or failure. Rejected input leaves the active token unchanged.
50    pub fn rotate_bearer_token_from_mut_bytes(
51        &self,
52        source: &mut [u8],
53    ) -> Result<(), TokenRotationError> {
54        self.credentials.rotate_from_mut_bytes(source)
55    }
56
57    /// Validates and rotates from guarded storage. Dropping the consumed guard
58    /// clears the complete source on success or failure.
59    pub fn rotate_bearer_token_from_secret_buffer(
60        &self,
61        source: SecretBuffer<'_>,
62    ) -> Result<(), TokenRotationError> {
63        self.credentials.rotate_from_secret_buffer(source)
64    }
65
66    fn send_inner<'buffer>(
67        &self,
68        request: TransportRequest<'_>,
69        response_body: &'buffer mut [u8],
70    ) -> Result<TransportResponse<'buffer>, TransportError> {
71        sanitize_bytes(response_body);
72        let url = self
73            .endpoint
74            .compose(request.target())
75            .map_err(|_| TransportError::TargetRejected)?;
76        let method = map_method(request.method())?;
77        let token_snapshot = self
78            .credentials
79            .snapshot()
80            .map_err(|_| TransportError::CredentialStateUnavailable)?;
81        let authorization = token_snapshot
82            .header_value()
83            .map_err(|_| TransportError::HeaderRejected)?;
84        let mut outbound = self
85            .client
86            .request(method, url)
87            .header(AUTHORIZATION, authorization);
88
89        for header in request.headers().as_slice() {
90            let name = HeaderName::from_bytes(header.name().as_str().as_bytes())
91                .map_err(|_| TransportError::HeaderRejected)?;
92            let mut value = HeaderValue::from_str(header.value().as_str())
93                .map_err(|_| TransportError::HeaderRejected)?;
94            value.set_sensitive(matches!(
95                header.sensitivity(),
96                cloud_sdk::transport::HeaderSensitivity::Sensitive
97            ));
98            outbound = outbound.header(name, value);
99        }
100        if !request.body().is_empty() && request.headers().get("content-type").is_none() {
101            return Err(TransportError::MissingContentType);
102        }
103
104        if !request.body().is_empty() {
105            let body = SanitizedRequestBody::new(request.body())
106                .map_err(|_| TransportError::RequestBodyAllocationFailed)?;
107            let body_len = u64::try_from(request.body().len())
108                .map_err(|_| TransportError::RequestBodyTooLarge)?;
109            outbound = outbound.body(Body::sized(body, body_len));
110        }
111
112        let mut response = outbound.send().map_err(classify_reqwest_error)?;
113        self.endpoint
114            .verify_origin(response.url())
115            .map_err(|_| TransportError::ResponseOriginChanged)?;
116        let headers = capture_response_headers(response.headers())?;
117        if response.content_length().is_some_and(|length| {
118            u64::try_from(response_body.len()).map_or(true, |cap| length > cap)
119        }) {
120            return Err(TransportError::ResponseTooLarge);
121        }
122        let status =
123            StatusCode::new(response.status().as_u16()).ok_or(TransportError::InvalidStatus)?;
124        let rate_limit = parse_rate_limit(&headers)?;
125        let content_type = parse_response_content_type(&headers)?;
126        let body_len = read_response(&mut response, response_body)?;
127        let initialized = response_body
128            .get(..body_len)
129            .ok_or(TransportError::ResponseReadFailed)?;
130        let response = TransportResponse::new(status, initialized).with_headers(headers);
131        let response = content_type.map_or(response, |value| response.with_content_type(value));
132        drop(token_snapshot);
133        Ok(rate_limit.map_or(response, |value| response.with_rate_limit(value)))
134    }
135}
136
137impl BlockingTransport for BlockingClient {
138    type Error = TransportError;
139
140    fn send<'buffer>(
141        &self,
142        request: TransportRequest<'_>,
143        response_body: &'buffer mut [u8],
144    ) -> Result<TransportResponse<'buffer>, Self::Error> {
145        self.send_inner(request, response_body)
146    }
147}
148
149impl ResponseStorageSanitizer for BlockingClient {
150    fn sanitize_response_storage(&self, response_storage: &mut [u8]) {
151        sanitize_bytes(response_storage);
152    }
153}
154
155impl BoundTransport for BlockingClient {
156    fn endpoint_identity(&self) -> Result<EndpointIdentity<'_>, EndpointIdentityError> {
157        self.endpoint.identity()
158    }
159}
160
161impl fmt::Debug for BlockingClient {
162    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
163        formatter
164            .debug_struct("BlockingClient")
165            .field("endpoint", &"[redacted]")
166            .field("credentials", &"[redacted]")
167            .finish_non_exhaustive()
168    }
169}
170
171fn read_response(response: &mut impl Read, output: &mut [u8]) -> Result<usize, TransportError> {
172    match read_bounded(response, output) {
173        Ok(len) => Ok(len),
174        Err(error) => {
175            sanitize_bytes(output);
176            Err(match error {
177                ReadBodyError::TooLarge => TransportError::ResponseTooLarge,
178                ReadBodyError::ReadFailed => TransportError::ResponseReadFailed,
179            })
180        }
181    }
182}
183
184fn map_method(method: Method) -> Result<reqwest::Method, TransportError> {
185    reqwest::Method::from_bytes(method.as_str().as_bytes())
186        .map_err(|_| TransportError::MethodRejected)
187}
188
189fn classify_reqwest_error(error: reqwest::Error) -> TransportError {
190    if error.is_timeout() {
191        TransportError::TimedOut
192    } else if error.is_connect() {
193        TransportError::ConnectFailed
194    } else {
195        TransportError::RequestFailed
196    }
197}