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 completion = stage_response(&exchange.response, response)?;
217 self.cursor
218 .compare_exchange(cursor, next_cursor, Ordering::AcqRel, Ordering::Acquire)
219 .map_err(|_| MockError::ConcurrentRequest)?;
220 Ok(completion)
221 }
222}
223
224pub(crate) fn stage_response<'buffer>(
225 fixture: &ResponseFixture<'_>,
226 response: &mut impl MockResponseSink<'buffer>,
227) -> Result<ResponseCompletion, MockError> {
228 let _content_type = fixture
229 .content_type()
230 .map(ResponseContentType::new)
231 .transpose()
232 .map_err(|_| MockError::InvalidFixtureMetadata)?;
233 let rate_limit = fixture
234 .rate_limit()
235 .map(|value| value.into_rate_limit())
236 .transpose()
237 .map_err(|_| MockError::InvalidFixtureMetadata)?;
238 {
239 let response_headers = response
240 .headers_mut()
241 .map_err(|_| MockError::ResponseWriterRejected)?;
242 if let Some(source) = fixture.headers() {
243 for header in source.iter() {
244 response_headers
245 .try_push(header.name(), header.value(), header.sensitivity())
246 .map_err(|_| MockError::InvalidFixtureMetadata)?;
247 }
248 }
249 if let Some(value) = fixture.content_type() {
250 response_headers
251 .try_push("content-type", value.as_bytes(), HeaderSensitivity::Public)
252 .map_err(|_| MockError::InvalidFixtureMetadata)?;
253 }
254 if let Some(value) = rate_limit {
255 push_rate_limit_headers(response_headers, value)
256 .map_err(|_| MockError::InvalidFixtureMetadata)?;
257 }
258 }
259 let body_len = fixture
260 .body()
261 .write_to(
262 response
263 .body_mut()
264 .map_err(|_| MockError::ResponseWriterRejected)?,
265 )
266 .map_err(|error| match error {
267 FixtureBodyError::OutputTooSmall | FixtureBodyError::TooLarge => {
268 MockError::ResponseBufferTooSmall
269 }
270 })?;
271 let mut metadata = ResponseMetadata::EMPTY;
272 if let Some(value) = rate_limit {
273 metadata = metadata.with_rate_limit(value);
274 }
275 Ok(ResponseCompletion::new(
276 fixture.status(),
277 body_len,
278 metadata,
279 ))
280}
281
282pub(crate) trait MockResponseSink<'buffer> {
283 fn body_mut(&mut self) -> Result<&mut [u8], MockError>;
284 fn headers_mut(&mut self) -> Result<&mut ResponseHeaders<'buffer>, MockError>;
285}
286
287impl<'buffer> MockResponseSink<'buffer> for ResponseAttempt<'_, 'buffer> {
288 fn body_mut(&mut self) -> Result<&mut [u8], MockError> {
289 ResponseAttempt::body_mut(self).map_err(|_| MockError::ResponseWriterRejected)
290 }
291
292 fn headers_mut(&mut self) -> Result<&mut ResponseHeaders<'buffer>, MockError> {
293 ResponseAttempt::headers_mut(self).map_err(|_| MockError::ResponseWriterRejected)
294 }
295}
296
297impl<'buffer> MockResponseSink<'buffer> for AsyncResponseStaging<'_, 'buffer> {
298 fn body_mut(&mut self) -> Result<&mut [u8], MockError> {
299 AsyncResponseStaging::body_mut(self).map_err(|_| MockError::ResponseWriterRejected)
300 }
301
302 fn headers_mut(&mut self) -> Result<&mut ResponseHeaders<'buffer>, MockError> {
303 AsyncResponseStaging::headers_mut(self).map_err(|_| MockError::ResponseWriterRejected)
304 }
305}
306
307fn request_headers_match(actual: RequestHeaders<'_>, expected: RequestHeaders<'_>) -> bool {
308 let actual = actual.as_slice();
309 let expected = expected.as_slice();
310 actual.len() == expected.len()
311 && actual.iter().zip(expected).all(|(actual, expected)| {
312 actual.name() == expected.name()
313 && actual.value().as_str().as_bytes() == expected.value().as_str().as_bytes()
314 && actual.sensitivity() == expected.sensitivity()
315 })
316}
317
318fn push_rate_limit_headers(
319 headers: &mut ResponseHeaders,
320 rate_limit: cloud_sdk::rate_limit::RateLimit,
321) -> Result<(), ()> {
322 let mut storage = [0_u8; 20];
323 for (name, value) in [
324 ("ratelimit-limit", rate_limit.limit()),
325 ("ratelimit-remaining", rate_limit.remaining()),
326 ("ratelimit-reset", rate_limit.reset_epoch_seconds()),
327 ] {
328 let text = write_decimal(value, &mut storage).ok_or(())?;
329 headers
330 .try_push(name, text.as_bytes(), HeaderSensitivity::Public)
331 .map_err(|_| ())?;
332 }
333 Ok(())
334}
335
336fn write_decimal(value: u64, output: &mut [u8; 20]) -> Option<&str> {
337 let mut value = value;
338 let mut cursor = output.len();
339 loop {
340 cursor = cursor.checked_sub(1)?;
341 let digit = u8::try_from(value % 10).ok()?;
342 *output.get_mut(cursor)? = b'0'.checked_add(digit)?;
343 value /= 10;
344 if value == 0 {
345 break;
346 }
347 }
348 core::str::from_utf8(output.get(cursor..)?).ok()
349}
350
351impl BlockingTransport for MockTransport<'_> {
352 type Error = MockError;
353
354 fn send(
355 &self,
356 request: TransportRequest<'_>,
357 response: &mut ResponseWriter<'_>,
358 ) -> Result<(), Self::Error> {
359 self.send_inner(request, response)
360 }
361}
362
363impl BlockingAuthenticatedTransport for MockTransport<'_> {
364 type Error = MockError;
365
366 fn send_authenticated(
367 &self,
368 request: AuthenticatedRequest<'_, '_>,
369 response: &mut ResponseWriter<'_>,
370 ) -> Result<(), Self::Error> {
371 self.send_inner(request.transport_request(), response)
372 }
373}
374
375impl AsyncTransport for MockTransport<'_> {
376 type Error = MockError;
377
378 async fn send<'transport, 'request, 'writer, 'buffer>(
379 &'transport self,
380 request: TransportRequest<'request>,
381 mut response: AsyncResponseStaging<'writer, 'buffer>,
382 ) -> Result<ResponseCompletion, Self::Error>
383 where
384 'transport: 'writer,
385 'request: 'writer,
386 'buffer: 'writer,
387 {
388 self.stage_inner(request, &mut response)
389 }
390}
391
392impl AsyncAuthenticatedTransport for MockTransport<'_> {
393 type Error = MockError;
394
395 async fn send_authenticated<'transport, 'request, 'policy, 'writer, 'buffer>(
396 &'transport self,
397 request: AuthenticatedRequest<'request, 'policy>,
398 mut response: AsyncResponseStaging<'writer, 'buffer>,
399 ) -> Result<ResponseCompletion, Self::Error>
400 where
401 'transport: 'writer,
402 'request: 'writer,
403 'policy: 'writer,
404 'buffer: 'writer,
405 {
406 self.stage_inner(request.transport_request(), &mut response)
407 }
408}
409
410impl BoundTransport for MockTransport<'_> {
411 fn endpoint_identity(&self) -> Result<EndpointIdentity<'_>, EndpointIdentityError> {
412 self.endpoint.ok_or(EndpointIdentityError::UnboundTransport)
413 }
414}
415
416impl fmt::Debug for MockTransport<'_> {
417 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
418 formatter
419 .debug_struct("MockTransport")
420 .field("remaining", &self.remaining())
421 .finish_non_exhaustive()
422 }
423}