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