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, DeliveryClassified,
16 DeliveryPhase, EndpointIdentity, EndpointIdentityError, HeaderSensitivity, RequestHeaders,
17 RequestTarget, ResponseAttempt, ResponseCompletion, ResponseContentType, ResponseHeaders,
18 ResponseMetadata, ResponseWriter, 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
140impl DeliveryClassified for MockError {
141 fn delivery_phase(&self) -> DeliveryPhase {
142 DeliveryPhase::NotSent
144 }
145}
146
147pub struct MockTransport<'a> {
149 exchanges: &'a [MockExchange<'a>],
150 cursor: AtomicUsize,
151 endpoint: Option<EndpointIdentity<'a>>,
152}
153
154impl<'a> MockTransport<'a> {
155 #[must_use]
157 pub const fn new(exchanges: &'a [MockExchange<'a>]) -> Self {
158 Self {
159 exchanges,
160 cursor: AtomicUsize::new(0),
161 endpoint: None,
162 }
163 }
164
165 #[must_use]
167 pub const fn with_endpoint(mut self, endpoint: EndpointIdentity<'a>) -> Self {
168 self.endpoint = Some(endpoint);
169 self
170 }
171
172 #[must_use]
174 pub fn remaining(&self) -> usize {
175 self.exchanges
176 .len()
177 .saturating_sub(self.cursor.load(Ordering::Acquire))
178 }
179
180 #[must_use]
182 pub fn is_complete(&self) -> bool {
183 self.remaining() == 0
184 }
185
186 fn send_inner(
187 &self,
188 request: TransportRequest<'_>,
189 response: &mut ResponseWriter<'_>,
190 ) -> Result<(), MockError> {
191 if response.is_committed() {
192 return Err(MockError::ResponseWriterRejected);
193 }
194 let mut response = response
195 .begin_attempt()
196 .map_err(|_| MockError::ResponseWriterRejected)?;
197 let completion = self.stage_inner(request, &mut response)?;
198 response
199 .commit_completion(completion)
200 .map_err(|_| MockError::ResponseWriterRejected)
201 }
202
203 fn stage_inner<'buffer>(
204 &self,
205 request: TransportRequest<'_>,
206 response: &mut impl MockResponseSink<'buffer>,
207 ) -> Result<ResponseCompletion, MockError> {
208 let cursor = self.cursor.load(Ordering::Acquire);
209 let exchange = self.exchanges.get(cursor).ok_or(MockError::Exhausted)?;
210 if request.method() != exchange.request.method() {
211 return Err(MockError::MethodMismatch);
212 }
213 if request.target() != exchange.request.target() {
214 return Err(MockError::TargetMismatch);
215 }
216 if request.body() != exchange.request.body() {
217 return Err(MockError::BodyMismatch);
218 }
219 if !request_headers_match(request.headers(), exchange.request.headers()) {
220 return Err(MockError::HeadersMismatch);
221 }
222 let next_cursor = cursor.checked_add(1).ok_or(MockError::CursorOverflow)?;
223 let completion = stage_response(&exchange.response, response)?;
224 self.cursor
225 .compare_exchange(cursor, next_cursor, Ordering::AcqRel, Ordering::Acquire)
226 .map_err(|_| MockError::ConcurrentRequest)?;
227 Ok(completion)
228 }
229}
230
231pub(crate) fn stage_response<'buffer>(
232 fixture: &ResponseFixture<'_>,
233 response: &mut impl MockResponseSink<'buffer>,
234) -> Result<ResponseCompletion, MockError> {
235 let _content_type = fixture
236 .content_type()
237 .map(ResponseContentType::new)
238 .transpose()
239 .map_err(|_| MockError::InvalidFixtureMetadata)?;
240 let rate_limit = fixture
241 .rate_limit()
242 .map(|value| value.into_rate_limit())
243 .transpose()
244 .map_err(|_| MockError::InvalidFixtureMetadata)?;
245 {
246 let response_headers = response
247 .headers_mut()
248 .map_err(|_| MockError::ResponseWriterRejected)?;
249 if let Some(source) = fixture.headers() {
250 for header in source.iter() {
251 response_headers
252 .try_push(header.name(), header.value(), header.sensitivity())
253 .map_err(|_| MockError::InvalidFixtureMetadata)?;
254 }
255 }
256 if let Some(value) = fixture.content_type() {
257 response_headers
258 .try_push("content-type", value.as_bytes(), HeaderSensitivity::Public)
259 .map_err(|_| MockError::InvalidFixtureMetadata)?;
260 }
261 if let Some(value) = rate_limit {
262 push_rate_limit_headers(response_headers, value)
263 .map_err(|_| MockError::InvalidFixtureMetadata)?;
264 }
265 }
266 let body_len = fixture
267 .body()
268 .write_to(
269 response
270 .body_mut()
271 .map_err(|_| MockError::ResponseWriterRejected)?,
272 )
273 .map_err(|error| match error {
274 FixtureBodyError::OutputTooSmall | FixtureBodyError::TooLarge => {
275 MockError::ResponseBufferTooSmall
276 }
277 })?;
278 let mut metadata = ResponseMetadata::EMPTY;
279 if let Some(value) = rate_limit {
280 metadata = metadata.with_rate_limit(value);
281 }
282 Ok(ResponseCompletion::new(
283 fixture.status(),
284 body_len,
285 metadata,
286 ))
287}
288
289pub(crate) trait MockResponseSink<'buffer> {
290 fn body_mut(&mut self) -> Result<&mut [u8], MockError>;
291 fn headers_mut(&mut self) -> Result<&mut ResponseHeaders<'buffer>, MockError>;
292}
293
294impl<'buffer> MockResponseSink<'buffer> for ResponseAttempt<'_, 'buffer> {
295 fn body_mut(&mut self) -> Result<&mut [u8], MockError> {
296 ResponseAttempt::body_mut(self).map_err(|_| MockError::ResponseWriterRejected)
297 }
298
299 fn headers_mut(&mut self) -> Result<&mut ResponseHeaders<'buffer>, MockError> {
300 ResponseAttempt::headers_mut(self).map_err(|_| MockError::ResponseWriterRejected)
301 }
302}
303
304impl<'buffer> MockResponseSink<'buffer> for AsyncResponseStaging<'_, 'buffer> {
305 fn body_mut(&mut self) -> Result<&mut [u8], MockError> {
306 AsyncResponseStaging::body_mut(self).map_err(|_| MockError::ResponseWriterRejected)
307 }
308
309 fn headers_mut(&mut self) -> Result<&mut ResponseHeaders<'buffer>, MockError> {
310 AsyncResponseStaging::headers_mut(self).map_err(|_| MockError::ResponseWriterRejected)
311 }
312}
313
314fn request_headers_match(actual: RequestHeaders<'_>, expected: RequestHeaders<'_>) -> bool {
315 let actual = actual.as_slice();
316 let expected = expected.as_slice();
317 actual.len() == expected.len()
318 && actual.iter().zip(expected).all(|(actual, expected)| {
319 actual.name() == expected.name()
320 && actual.value().as_str().as_bytes() == expected.value().as_str().as_bytes()
321 && actual.sensitivity() == expected.sensitivity()
322 })
323}
324
325fn push_rate_limit_headers(
326 headers: &mut ResponseHeaders,
327 rate_limit: cloud_sdk::rate_limit::RateLimit,
328) -> Result<(), ()> {
329 let mut storage = [0_u8; 20];
330 for (name, value) in [
331 ("ratelimit-limit", rate_limit.limit()),
332 ("ratelimit-remaining", rate_limit.remaining()),
333 ("ratelimit-reset", rate_limit.reset_epoch_seconds()),
334 ] {
335 let text = write_decimal(value, &mut storage).ok_or(())?;
336 headers
337 .try_push(name, text.as_bytes(), HeaderSensitivity::Public)
338 .map_err(|_| ())?;
339 }
340 Ok(())
341}
342
343fn write_decimal(value: u64, output: &mut [u8; 20]) -> Option<&str> {
344 let mut value = value;
345 let mut cursor = output.len();
346 loop {
347 cursor = cursor.checked_sub(1)?;
348 let digit = u8::try_from(value % 10).ok()?;
349 *output.get_mut(cursor)? = b'0'.checked_add(digit)?;
350 value /= 10;
351 if value == 0 {
352 break;
353 }
354 }
355 core::str::from_utf8(output.get(cursor..)?).ok()
356}
357
358impl BlockingTransport for MockTransport<'_> {
359 type Error = MockError;
360
361 fn send(
362 &self,
363 request: TransportRequest<'_>,
364 response: &mut ResponseWriter<'_>,
365 ) -> Result<(), Self::Error> {
366 self.send_inner(request, response)
367 }
368}
369
370impl BlockingAuthenticatedTransport for MockTransport<'_> {
371 type Error = MockError;
372
373 fn send_authenticated(
374 &self,
375 request: AuthenticatedRequest<'_, '_>,
376 response: &mut ResponseWriter<'_>,
377 ) -> Result<(), Self::Error> {
378 self.send_inner(request.transport_request(), response)
379 }
380}
381
382impl AsyncTransport for MockTransport<'_> {
383 type Error = MockError;
384
385 async fn send<'transport, 'request, 'writer, 'buffer>(
386 &'transport self,
387 request: TransportRequest<'request>,
388 mut response: AsyncResponseStaging<'writer, 'buffer>,
389 ) -> Result<ResponseCompletion, Self::Error>
390 where
391 'transport: 'writer,
392 'request: 'writer,
393 'buffer: 'writer,
394 {
395 self.stage_inner(request, &mut response)
396 }
397}
398
399impl AsyncAuthenticatedTransport for MockTransport<'_> {
400 type Error = MockError;
401
402 async fn send_authenticated<'transport, 'request, 'policy, 'writer, 'buffer>(
403 &'transport self,
404 request: AuthenticatedRequest<'request, 'policy>,
405 mut response: AsyncResponseStaging<'writer, 'buffer>,
406 ) -> Result<ResponseCompletion, Self::Error>
407 where
408 'transport: 'writer,
409 'request: 'writer,
410 'policy: 'writer,
411 'buffer: 'writer,
412 {
413 self.stage_inner(request.transport_request(), &mut response)
414 }
415}
416
417impl BoundTransport for MockTransport<'_> {
418 fn endpoint_identity(&self) -> Result<EndpointIdentity<'_>, EndpointIdentityError> {
419 self.endpoint.ok_or(EndpointIdentityError::UnboundTransport)
420 }
421}
422
423impl fmt::Debug for MockTransport<'_> {
424 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
425 formatter
426 .debug_struct("MockTransport")
427 .field("remaining", &self.remaining())
428 .finish_non_exhaustive()
429 }
430}