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