cloud_sdk_reqwest/blocking/
client.rs1use core::fmt;
2use std::io::Read;
3use std::sync::Arc;
4
5use cloud_sdk::Method;
6use cloud_sdk::transport::{
7 BlockingTransport, BoundTransport, EndpointIdentity, EndpointIdentityError, ResponseMetadata,
8 ResponseStorageSanitizer, ResponseWriter, StatusCode, TransportRequest,
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#[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 pub fn rotate_bearer_token(
42 &self,
43 replacement: BearerToken,
44 ) -> Result<(), CredentialStateError> {
45 self.credentials.rotate(replacement)
46 }
47
48 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 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(
67 &self,
68 request: TransportRequest<'_>,
69 response_writer: &mut ResponseWriter<'_>,
70 ) -> Result<(), TransportError> {
71 if response_writer.is_committed() {
72 return Err(TransportError::ResponseCommitFailed);
73 }
74 let mut response_writer = response_writer
75 .begin_attempt()
76 .map_err(|_| TransportError::ResponseCommitFailed)?;
77 let url = self
78 .endpoint
79 .compose(request.target())
80 .map_err(|_| TransportError::TargetRejected)?;
81 let method = map_method(request.method())?;
82 let token_snapshot = self
83 .credentials
84 .snapshot()
85 .map_err(|_| TransportError::CredentialStateUnavailable)?;
86 let authorization = token_snapshot
87 .header_value()
88 .map_err(|_| TransportError::HeaderRejected)?;
89 let mut outbound = self
90 .client
91 .request(method, url)
92 .header(AUTHORIZATION, authorization);
93
94 for header in request.headers().as_slice() {
95 let name = HeaderName::from_bytes(header.name().as_str().as_bytes())
96 .map_err(|_| TransportError::HeaderRejected)?;
97 let mut value = HeaderValue::from_str(header.value().as_str())
98 .map_err(|_| TransportError::HeaderRejected)?;
99 value.set_sensitive(matches!(
100 header.sensitivity(),
101 cloud_sdk::transport::HeaderSensitivity::Sensitive
102 ));
103 outbound = outbound.header(name, value);
104 }
105 if !request.body().is_empty() && request.headers().get("content-type").is_none() {
106 return Err(TransportError::MissingContentType);
107 }
108
109 if !request.body().is_empty() {
110 let body = SanitizedRequestBody::new(request.body())
111 .map_err(|_| TransportError::RequestBodyAllocationFailed)?;
112 let body_len = u64::try_from(request.body().len())
113 .map_err(|_| TransportError::RequestBodyTooLarge)?;
114 outbound = outbound.body(Body::sized(body, body_len));
115 }
116
117 let mut response = outbound.send().map_err(classify_reqwest_error)?;
118 self.endpoint
119 .verify_origin(response.url())
120 .map_err(|_| TransportError::ResponseOriginChanged)?;
121 capture_response_headers(
122 response.headers(),
123 response_writer
124 .headers_mut()
125 .map_err(|_| TransportError::ResponseCommitFailed)?,
126 )?;
127 if response.content_length().is_some_and(|length| {
128 u64::try_from(response_writer.body_capacity()).map_or(true, |cap| length > cap)
129 }) {
130 return Err(TransportError::ResponseTooLarge);
131 }
132 let status =
133 StatusCode::new(response.status().as_u16()).ok_or(TransportError::InvalidStatus)?;
134 let rate_limit = parse_rate_limit(response_writer.headers())?;
135 parse_response_content_type(response_writer.headers())?;
136 let body_len = read_response(
137 &mut response,
138 response_writer
139 .body_mut()
140 .map_err(|_| TransportError::ResponseCommitFailed)?,
141 )?;
142 let mut metadata = ResponseMetadata::EMPTY;
143 if let Some(value) = rate_limit {
144 metadata = metadata.with_rate_limit(value);
145 }
146 drop(token_snapshot);
147 response_writer
148 .commit(status, body_len, metadata)
149 .map_err(|_| TransportError::ResponseCommitFailed)
150 }
151}
152
153impl BlockingTransport for BlockingClient {
154 type Error = TransportError;
155
156 fn send(
157 &self,
158 request: TransportRequest<'_>,
159 response: &mut ResponseWriter<'_>,
160 ) -> Result<(), Self::Error> {
161 self.send_inner(request, response)
162 }
163}
164
165impl ResponseStorageSanitizer for BlockingClient {
166 fn sanitize_response_storage(&self, response_storage: &mut [u8]) {
167 sanitize_bytes(response_storage);
168 }
169}
170
171impl BoundTransport for BlockingClient {
172 fn endpoint_identity(&self) -> Result<EndpointIdentity<'_>, EndpointIdentityError> {
173 self.endpoint.identity()
174 }
175}
176
177impl fmt::Debug for BlockingClient {
178 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
179 formatter
180 .debug_struct("BlockingClient")
181 .field("endpoint", &"[redacted]")
182 .field("credentials", &"[redacted]")
183 .finish_non_exhaustive()
184 }
185}
186
187fn read_response(response: &mut impl Read, output: &mut [u8]) -> Result<usize, TransportError> {
188 match read_bounded(response, output) {
189 Ok(len) => Ok(len),
190 Err(error) => {
191 sanitize_bytes(output);
192 Err(match error {
193 ReadBodyError::TooLarge => TransportError::ResponseTooLarge,
194 ReadBodyError::ReadFailed => TransportError::ResponseReadFailed,
195 })
196 }
197 }
198}
199
200fn map_method(method: Method) -> Result<reqwest::Method, TransportError> {
201 reqwest::Method::from_bytes(method.as_str().as_bytes())
202 .map_err(|_| TransportError::MethodRejected)
203}
204
205fn classify_reqwest_error(error: reqwest::Error) -> TransportError {
206 if error.is_timeout() {
207 TransportError::TimedOut
208 } else if error.is_connect() {
209 TransportError::ConnectFailed
210 } else {
211 TransportError::RequestFailed
212 }
213}