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