1mod local;
4
5pub use local::LocalMockTransport;
6
7use core::fmt;
8use core::sync::atomic::{AtomicUsize, Ordering};
9
10use cloud_sdk::Method;
11use cloud_sdk::authentication::{
12 AsyncAuthenticatedTransport, AuthenticatedRequest, BlockingAuthenticatedTransport,
13};
14use cloud_sdk::transport::{
15 AsyncResponseStaging, AsyncTransport, BlockingTransport, BoundTransport, EndpointIdentity,
16 EndpointIdentityError, HeaderSensitivity, RequestHeaders, RequestTarget, ResponseAttempt,
17 ResponseCompletion, ResponseContentType, ResponseHeaders, ResponseMetadata, ResponseWriter,
18 TransportRequest,
19};
20
21use crate::{FixtureBodyError, ResponseFixture};
22
23#[derive(Clone, Copy)]
25pub struct ExpectedRequest<'a> {
26 method: Method,
27 target: RequestTarget<'a>,
28 body: &'a [u8],
29 headers: RequestHeaders<'a>,
30}
31
32impl<'a> ExpectedRequest<'a> {
33 #[must_use]
35 pub const fn new(method: Method, target: RequestTarget<'a>) -> Self {
36 Self {
37 method,
38 target,
39 body: &[],
40 headers: RequestHeaders::EMPTY,
41 }
42 }
43
44 #[must_use]
46 pub const fn with_body(mut self, body: &'a [u8]) -> Self {
47 self.body = body;
48 self
49 }
50
51 #[must_use]
53 pub const fn with_headers(mut self, headers: RequestHeaders<'a>) -> Self {
54 self.headers = headers;
55 self
56 }
57
58 const fn method(self) -> Method {
59 self.method
60 }
61
62 const fn target(self) -> RequestTarget<'a> {
63 self.target
64 }
65
66 const fn body(self) -> &'a [u8] {
67 self.body
68 }
69
70 const fn headers(self) -> RequestHeaders<'a> {
71 self.headers
72 }
73}
74
75impl fmt::Debug for ExpectedRequest<'_> {
76 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
77 formatter
78 .debug_struct("ExpectedRequest")
79 .field("method", &self.method)
80 .field("target", &"[redacted]")
81 .field("body", &"[redacted]")
82 .field("headers", &self.headers)
83 .finish()
84 }
85}
86
87#[derive(Debug)]
89pub struct MockExchange<'a> {
90 request: ExpectedRequest<'a>,
91 response: ResponseFixture<'a>,
92}
93
94impl<'a> MockExchange<'a> {
95 #[must_use]
97 pub const fn new(request: ExpectedRequest<'a>, response: ResponseFixture<'a>) -> Self {
98 Self { request, response }
99 }
100}
101
102#[derive(Clone, Copy, Debug, Eq, PartialEq)]
104pub enum MockError {
105 Exhausted,
107 MethodMismatch,
109 TargetMismatch,
111 BodyMismatch,
113 HeadersMismatch,
115 ResponseBufferTooSmall,
117 CursorOverflow,
119 ConcurrentRequest,
121 InvalidFixtureMetadata,
123 ResponseWriterRejected,
125}
126
127impl_static_error!(MockError,
128 Self::Exhausted => "mock transport has no expected exchange remaining",
129 Self::MethodMismatch => "mock request method differs from expectation",
130 Self::TargetMismatch => "mock request target differs from expectation",
131 Self::BodyMismatch => "mock request body differs from expectation",
132 Self::HeadersMismatch => "mock request headers differ from expectation",
133 Self::ResponseBufferTooSmall => "mock response buffer is too small",
134 Self::CursorOverflow => "mock transport cursor overflowed",
135 Self::ConcurrentRequest => "mock transport cursor changed concurrently",
136 Self::InvalidFixtureMetadata => "mock fixture metadata is invalid",
137 Self::ResponseWriterRejected => "mock response writer rejected output",
138);
139
140pub struct MockTransport<'a> {
142 exchanges: &'a [MockExchange<'a>],
143 cursor: AtomicUsize,
144 endpoint: Option<EndpointIdentity<'a>>,
145}
146
147impl<'a> MockTransport<'a> {
148 #[must_use]
150 pub const fn new(exchanges: &'a [MockExchange<'a>]) -> Self {
151 Self {
152 exchanges,
153 cursor: AtomicUsize::new(0),
154 endpoint: None,
155 }
156 }
157
158 #[must_use]
160 pub const fn with_endpoint(mut self, endpoint: EndpointIdentity<'a>) -> Self {
161 self.endpoint = Some(endpoint);
162 self
163 }
164
165 #[must_use]
167 pub fn remaining(&self) -> usize {
168 self.exchanges
169 .len()
170 .saturating_sub(self.cursor.load(Ordering::Acquire))
171 }
172
173 #[must_use]
175 pub fn is_complete(&self) -> bool {
176 self.remaining() == 0
177 }
178
179 fn send_inner(
180 &self,
181 request: TransportRequest<'_>,
182 response: &mut ResponseWriter<'_>,
183 ) -> Result<(), MockError> {
184 if response.is_committed() {
185 return Err(MockError::ResponseWriterRejected);
186 }
187 let mut response = response
188 .begin_attempt()
189 .map_err(|_| MockError::ResponseWriterRejected)?;
190 let completion = self.stage_inner(request, &mut response)?;
191 response
192 .commit_completion(completion)
193 .map_err(|_| MockError::ResponseWriterRejected)
194 }
195
196 fn stage_inner<'buffer>(
197 &self,
198 request: TransportRequest<'_>,
199 response: &mut impl MockResponseSink<'buffer>,
200 ) -> Result<ResponseCompletion, MockError> {
201 let cursor = self.cursor.load(Ordering::Acquire);
202 let exchange = self.exchanges.get(cursor).ok_or(MockError::Exhausted)?;
203 if request.method() != exchange.request.method() {
204 return Err(MockError::MethodMismatch);
205 }
206 if request.target() != exchange.request.target() {
207 return Err(MockError::TargetMismatch);
208 }
209 if request.body() != exchange.request.body() {
210 return Err(MockError::BodyMismatch);
211 }
212 if !request_headers_match(request.headers(), exchange.request.headers()) {
213 return Err(MockError::HeadersMismatch);
214 }
215 let next_cursor = cursor.checked_add(1).ok_or(MockError::CursorOverflow)?;
216 let _content_type = exchange
217 .response
218 .content_type()
219 .map(ResponseContentType::new)
220 .transpose()
221 .map_err(|_| MockError::InvalidFixtureMetadata)?;
222 let rate_limit = exchange
223 .response
224 .rate_limit()
225 .map(|value| value.into_rate_limit())
226 .transpose()
227 .map_err(|_| MockError::InvalidFixtureMetadata)?;
228 {
229 let response_headers = response
230 .headers_mut()
231 .map_err(|_| MockError::ResponseWriterRejected)?;
232 if let Some(source) = exchange.response.headers() {
233 for header in source.iter() {
234 response_headers
235 .try_push(header.name(), header.value(), header.sensitivity())
236 .map_err(|_| MockError::InvalidFixtureMetadata)?;
237 }
238 }
239 if let Some(value) = exchange.response.content_type() {
240 response_headers
241 .try_push("content-type", value.as_bytes(), HeaderSensitivity::Public)
242 .map_err(|_| MockError::InvalidFixtureMetadata)?;
243 }
244 if let Some(value) = rate_limit {
245 push_rate_limit_headers(response_headers, value)
246 .map_err(|_| MockError::InvalidFixtureMetadata)?;
247 }
248 }
249 let body_len = exchange
250 .response
251 .body()
252 .write_to(
253 response
254 .body_mut()
255 .map_err(|_| MockError::ResponseWriterRejected)?,
256 )
257 .map_err(|error| match error {
258 FixtureBodyError::OutputTooSmall | FixtureBodyError::TooLarge => {
259 MockError::ResponseBufferTooSmall
260 }
261 })?;
262 let mut metadata = ResponseMetadata::EMPTY;
263 if let Some(value) = rate_limit {
264 metadata = metadata.with_rate_limit(value);
265 }
266 self.cursor
267 .compare_exchange(cursor, next_cursor, Ordering::AcqRel, Ordering::Acquire)
268 .map_err(|_| MockError::ConcurrentRequest)?;
269 Ok(ResponseCompletion::new(
270 exchange.response.status(),
271 body_len,
272 metadata,
273 ))
274 }
275}
276
277trait MockResponseSink<'buffer> {
278 fn body_mut(&mut self) -> Result<&mut [u8], MockError>;
279 fn headers_mut(&mut self) -> Result<&mut ResponseHeaders<'buffer>, MockError>;
280}
281
282impl<'buffer> MockResponseSink<'buffer> for ResponseAttempt<'_, 'buffer> {
283 fn body_mut(&mut self) -> Result<&mut [u8], MockError> {
284 ResponseAttempt::body_mut(self).map_err(|_| MockError::ResponseWriterRejected)
285 }
286
287 fn headers_mut(&mut self) -> Result<&mut ResponseHeaders<'buffer>, MockError> {
288 ResponseAttempt::headers_mut(self).map_err(|_| MockError::ResponseWriterRejected)
289 }
290}
291
292impl<'buffer> MockResponseSink<'buffer> for AsyncResponseStaging<'_, 'buffer> {
293 fn body_mut(&mut self) -> Result<&mut [u8], MockError> {
294 AsyncResponseStaging::body_mut(self).map_err(|_| MockError::ResponseWriterRejected)
295 }
296
297 fn headers_mut(&mut self) -> Result<&mut ResponseHeaders<'buffer>, MockError> {
298 AsyncResponseStaging::headers_mut(self).map_err(|_| MockError::ResponseWriterRejected)
299 }
300}
301
302fn request_headers_match(actual: RequestHeaders<'_>, expected: RequestHeaders<'_>) -> bool {
303 let actual = actual.as_slice();
304 let expected = expected.as_slice();
305 actual.len() == expected.len()
306 && actual.iter().zip(expected).all(|(actual, expected)| {
307 actual.name() == expected.name()
308 && actual.value().as_str().as_bytes() == expected.value().as_str().as_bytes()
309 && actual.sensitivity() == expected.sensitivity()
310 })
311}
312
313fn push_rate_limit_headers(
314 headers: &mut ResponseHeaders,
315 rate_limit: cloud_sdk::rate_limit::RateLimit,
316) -> Result<(), ()> {
317 let mut storage = [0_u8; 20];
318 for (name, value) in [
319 ("ratelimit-limit", rate_limit.limit()),
320 ("ratelimit-remaining", rate_limit.remaining()),
321 ("ratelimit-reset", rate_limit.reset_epoch_seconds()),
322 ] {
323 let text = write_decimal(value, &mut storage).ok_or(())?;
324 headers
325 .try_push(name, text.as_bytes(), HeaderSensitivity::Public)
326 .map_err(|_| ())?;
327 }
328 Ok(())
329}
330
331fn write_decimal(value: u64, output: &mut [u8; 20]) -> Option<&str> {
332 let mut value = value;
333 let mut cursor = output.len();
334 loop {
335 cursor = cursor.checked_sub(1)?;
336 let digit = u8::try_from(value % 10).ok()?;
337 *output.get_mut(cursor)? = b'0'.checked_add(digit)?;
338 value /= 10;
339 if value == 0 {
340 break;
341 }
342 }
343 core::str::from_utf8(output.get(cursor..)?).ok()
344}
345
346impl BlockingTransport for MockTransport<'_> {
347 type Error = MockError;
348
349 fn send(
350 &self,
351 request: TransportRequest<'_>,
352 response: &mut ResponseWriter<'_>,
353 ) -> Result<(), Self::Error> {
354 self.send_inner(request, response)
355 }
356}
357
358impl BlockingAuthenticatedTransport for MockTransport<'_> {
359 type Error = MockError;
360
361 fn send_authenticated(
362 &self,
363 request: AuthenticatedRequest<'_, '_>,
364 response: &mut ResponseWriter<'_>,
365 ) -> Result<(), Self::Error> {
366 self.send_inner(request.transport_request(), response)
367 }
368}
369
370impl AsyncTransport for MockTransport<'_> {
371 type Error = MockError;
372
373 async fn send<'transport, 'request, 'writer, 'buffer>(
374 &'transport self,
375 request: TransportRequest<'request>,
376 mut response: AsyncResponseStaging<'writer, 'buffer>,
377 ) -> Result<ResponseCompletion, Self::Error>
378 where
379 'transport: 'writer,
380 'request: 'writer,
381 'buffer: 'writer,
382 {
383 self.stage_inner(request, &mut response)
384 }
385}
386
387impl AsyncAuthenticatedTransport for MockTransport<'_> {
388 type Error = MockError;
389
390 async fn send_authenticated<'transport, 'request, 'policy, 'writer, 'buffer>(
391 &'transport self,
392 request: AuthenticatedRequest<'request, 'policy>,
393 mut response: AsyncResponseStaging<'writer, 'buffer>,
394 ) -> Result<ResponseCompletion, Self::Error>
395 where
396 'transport: 'writer,
397 'request: 'writer,
398 'policy: 'writer,
399 'buffer: 'writer,
400 {
401 self.stage_inner(request.transport_request(), &mut response)
402 }
403}
404
405impl BoundTransport for MockTransport<'_> {
406 fn endpoint_identity(&self) -> Result<EndpointIdentity<'_>, EndpointIdentityError> {
407 self.endpoint.ok_or(EndpointIdentityError::UnboundTransport)
408 }
409}
410
411impl fmt::Debug for MockTransport<'_> {
412 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
413 formatter
414 .debug_struct("MockTransport")
415 .field("remaining", &self.remaining())
416 .finish_non_exhaustive()
417 }
418}