Skip to main content

cloud_sdk_testkit/
mock.rs

1//! Deterministic no-allocation mock transport.
2
3mod 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/// Expected request fields for one mock exchange.
24#[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    /// Creates a bodyless expected request.
34    #[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    /// Adds the exact expected request body.
45    #[must_use]
46    pub const fn with_body(mut self, body: &'a [u8]) -> Self {
47        self.body = body;
48        self
49    }
50
51    /// Adds the exact expected ordered request headers.
52    #[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/// One expected request and deterministic response.
88#[derive(Debug)]
89pub struct MockExchange<'a> {
90    request: ExpectedRequest<'a>,
91    response: ResponseFixture<'a>,
92}
93
94impl<'a> MockExchange<'a> {
95    /// Creates one mock exchange.
96    #[must_use]
97    pub const fn new(request: ExpectedRequest<'a>, response: ResponseFixture<'a>) -> Self {
98        Self { request, response }
99    }
100}
101
102/// Deterministic mock transport failure.
103#[derive(Clone, Copy, Debug, Eq, PartialEq)]
104pub enum MockError {
105    /// No expected exchange remains.
106    Exhausted,
107    /// HTTP method differs from the next expectation.
108    MethodMismatch,
109    /// Request target differs from the next expectation.
110    TargetMismatch,
111    /// Request body differs from the next expectation.
112    BodyMismatch,
113    /// Request headers differ from the next expectation.
114    HeadersMismatch,
115    /// Caller response buffer cannot hold the complete fixture body.
116    ResponseBufferTooSmall,
117    /// Internal cursor arithmetic failed closed.
118    CursorOverflow,
119    /// Another request changed the ordered cursor during this exchange.
120    ConcurrentRequest,
121    /// Fixture metadata could not be represented by the core transport.
122    InvalidFixtureMetadata,
123    /// The core response writer rejected fixture output.
124    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
140/// Ordered no-allocation mock implementation of [`BlockingTransport`].
141pub struct MockTransport<'a> {
142    exchanges: &'a [MockExchange<'a>],
143    cursor: AtomicUsize,
144    endpoint: Option<EndpointIdentity<'a>>,
145}
146
147impl<'a> MockTransport<'a> {
148    /// Creates a mock over an ordered exchange slice.
149    #[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    /// Binds the mock permanently to one normalized endpoint identity.
159    #[must_use]
160    pub const fn with_endpoint(mut self, endpoint: EndpointIdentity<'a>) -> Self {
161        self.endpoint = Some(endpoint);
162        self
163    }
164
165    /// Returns the number of exchanges not yet consumed.
166    #[must_use]
167    pub fn remaining(&self) -> usize {
168        self.exchanges
169            .len()
170            .saturating_sub(self.cursor.load(Ordering::Acquire))
171    }
172
173    /// Reports whether every expected exchange was consumed.
174    #[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}