Skip to main content

rig_http/
test_utils.rs

1//! HTTP client doubles for tests of code that sends through an
2//! [`HttpClientExt`].
3
4use std::{
5    collections::VecDeque,
6    future::{self, Future},
7    sync::{Arc, Mutex, MutexGuard},
8};
9
10use bytes::Bytes;
11
12use crate::{
13    http_client::{
14        self, HttpClientExt, LazyBody, MultipartForm, Request, Response, StreamingResponse,
15    },
16    wasm_compat::WasmCompatSend,
17};
18
19/// The reply a double gives on a surface it does not script: the request
20/// reached a transport that has no answer for it.
21fn not_implemented() -> http_client::Error {
22    http_client::Error::non_success_with_details(
23        http::StatusCode::NOT_IMPLEMENTED,
24        http::HeaderMap::new(),
25        String::new(),
26    )
27}
28
29/// The unary half of a streaming-only double.
30///
31/// Four doubles below script `send_streaming` and nothing else; that they
32/// have no unary surface is one fact about streaming-only doubles, so it is
33/// written here once and each impl states it by name.
34macro_rules! no_unary_surface {
35    () => {
36        fn send<T, U>(
37            &self,
38            _req: Request<T>,
39        ) -> impl Future<Output = http_client::Result<Response<LazyBody<U>>>>
40        + WasmCompatSend
41        + 'static
42        where
43            T: Into<Bytes> + WasmCompatSend,
44            U: From<Bytes> + WasmCompatSend + 'static,
45        {
46            future::ready(Err(not_implemented()))
47        }
48
49        fn send_multipart<U>(
50            &self,
51            _req: Request<MultipartForm>,
52        ) -> impl Future<Output = http_client::Result<Response<LazyBody<U>>>>
53        + WasmCompatSend
54        + 'static
55        where
56            U: From<Bytes> + WasmCompatSend + 'static,
57        {
58            future::ready(Err(not_implemented()))
59        }
60    };
61}
62
63/// Request data captured by [`RecordingHttpClient`].
64#[derive(Debug, Clone, PartialEq, Eq)]
65pub struct CapturedHttpRequest {
66    /// Request URI.
67    pub uri: String,
68    /// Request headers.
69    pub headers: http::HeaderMap,
70    /// Request body bytes.
71    pub body: Bytes,
72}
73
74/// Response scripted for [`RecordingHttpClient`].
75#[derive(Clone, Debug)]
76pub enum MockHttpResponse {
77    /// Return this body with a successful HTTP status and no headers.
78    Success(Bytes),
79    /// Return this body with a successful HTTP status and these headers —
80    /// the reply whose `Content-Type` decides how it may be framed.
81    SuccessWithHeaders(Bytes, http::HeaderMap),
82    /// Return an HTTP response with the given (typically non-success) status
83    /// and body, instead of a transport-level error.
84    ErrorResponse(http::StatusCode, Bytes),
85    /// Return a status-code error that preserved the failed response's
86    /// headers, exactly as the bundled reqwest transport does (rig#2210).
87    ErrorWithHeaders(http::StatusCode, String, http::HeaderMap),
88    /// Return an HTTP response with the given (typically non-success) status,
89    /// body, and response headers, instead of a transport-level error.
90    ErrorResponseWithHeaders(http::StatusCode, Bytes, http::HeaderMap),
91}
92
93impl MockHttpResponse {
94    /// Create a successful response from bytes.
95    pub fn success(body: impl Into<Bytes>) -> Self {
96        Self::Success(body.into())
97    }
98
99    /// Create a successful response carrying `content_type`.
100    pub fn success_typed(body: impl Into<Bytes>, content_type: &'static str) -> Self {
101        let mut headers = http::HeaderMap::new();
102        headers.insert(
103            http::header::CONTENT_TYPE,
104            http::HeaderValue::from_static(content_type),
105        );
106        Self::SuccessWithHeaders(body.into(), headers)
107    }
108
109    /// Create a transport-level status error whose reply carried no headers
110    /// of interest: the same shape as [`Self::error_with_headers`] with an
111    /// empty map, since a rejection always comes with its headers.
112    pub fn error(status: http::StatusCode, message: impl Into<String>) -> Self {
113        Self::error_with_headers(status, message, http::HeaderMap::new())
114    }
115
116    /// Create a transport-level status error that preserved the failed
117    /// response's headers, as the bundled reqwest transport does.
118    pub fn error_with_headers(
119        status: http::StatusCode,
120        message: impl Into<String>,
121        headers: http::HeaderMap,
122    ) -> Self {
123        Self::ErrorWithHeaders(status, message.into(), headers)
124    }
125}
126
127impl Default for MockHttpResponse {
128    fn default() -> Self {
129        Self::Success(Bytes::new())
130    }
131}
132
133/// An [`HttpClientExt`] implementation that records unary requests and returns
134/// a fixed response.
135#[derive(Clone, Debug, Default)]
136pub struct RecordingHttpClient {
137    requests: Arc<Mutex<Vec<CapturedHttpRequest>>>,
138    response: Arc<Mutex<MockHttpResponse>>,
139}
140
141impl RecordingHttpClient {
142    /// Create a client that returns `response_body` for unary requests.
143    pub fn new(response_body: impl Into<Bytes>) -> Self {
144        Self {
145            requests: Arc::new(Mutex::new(Vec::new())),
146            response: Arc::new(Mutex::new(MockHttpResponse::success(response_body))),
147        }
148    }
149
150    /// Create a client that returns an HTTP status error for unary requests.
151    pub fn with_error(status: http::StatusCode, message: impl Into<String>) -> Self {
152        Self {
153            requests: Arc::new(Mutex::new(Vec::new())),
154            response: Arc::new(Mutex::new(MockHttpResponse::error(status, message))),
155        }
156    }
157
158    /// Create a client that returns a non-success HTTP response (status and body)
159    /// for unary requests, instead of a transport-level error.
160    pub fn with_error_response(status: http::StatusCode, body: impl Into<Bytes>) -> Self {
161        Self {
162            requests: Arc::new(Mutex::new(Vec::new())),
163            response: Arc::new(Mutex::new(MockHttpResponse::ErrorResponse(
164                status,
165                body.into(),
166            ))),
167        }
168    }
169
170    /// Create a client whose transport reports a non-success status *and*
171    /// preserves the failed response's headers, as the bundled reqwest client
172    /// does (rig#2210).
173    pub fn with_error_headers(
174        status: http::StatusCode,
175        message: impl Into<String>,
176        headers: http::HeaderMap,
177    ) -> Self {
178        Self {
179            requests: Arc::new(Mutex::new(Vec::new())),
180            response: Arc::new(Mutex::new(MockHttpResponse::error_with_headers(
181                status, message, headers,
182            ))),
183        }
184    }
185
186    /// Create a client that hands back a non-success HTTP response carrying
187    /// `headers`, instead of a transport-level error.
188    pub fn with_error_response_headers(
189        status: http::StatusCode,
190        body: impl Into<Bytes>,
191        headers: http::HeaderMap,
192    ) -> Self {
193        Self {
194            requests: Arc::new(Mutex::new(Vec::new())),
195            response: Arc::new(Mutex::new(MockHttpResponse::ErrorResponseWithHeaders(
196                status,
197                body.into(),
198                headers,
199            ))),
200        }
201    }
202
203    /// Return the requests captured so far.
204    pub fn requests(&self) -> Vec<CapturedHttpRequest> {
205        self.requests_guard().clone()
206    }
207
208    /// Replace the scripted unary response.
209    pub fn set_response(&self, response: MockHttpResponse) {
210        *self.response_guard() = response;
211    }
212
213    fn requests_guard(&self) -> MutexGuard<'_, Vec<CapturedHttpRequest>> {
214        match self.requests.lock() {
215            Ok(guard) => guard,
216            Err(poisoned) => poisoned.into_inner(),
217        }
218    }
219
220    fn response_guard(&self) -> MutexGuard<'_, MockHttpResponse> {
221        match self.response.lock() {
222            Ok(guard) => guard,
223            Err(poisoned) => poisoned.into_inner(),
224        }
225    }
226
227    fn record_request(&self, uri: String, headers: http::HeaderMap, body: Bytes) {
228        self.requests_guard()
229            .push(CapturedHttpRequest { uri, headers, body });
230    }
231
232    fn build_unary_response<U>(
233        response: MockHttpResponse,
234    ) -> http_client::Result<Response<LazyBody<U>>>
235    where
236        U: From<Bytes> + WasmCompatSend + 'static,
237    {
238        let (status, response_body, response_headers) = match response {
239            MockHttpResponse::Success(response_body) => (http::StatusCode::OK, response_body, None),
240            MockHttpResponse::SuccessWithHeaders(response_body, headers) => {
241                (http::StatusCode::OK, response_body, Some(headers))
242            }
243            MockHttpResponse::ErrorWithHeaders(status, body, headers) => {
244                return Err(http_client::Error::InvalidStatusCodeWithDetails {
245                    status,
246                    body,
247                    headers,
248                });
249            }
250            MockHttpResponse::ErrorResponse(status, response_body) => (status, response_body, None),
251            MockHttpResponse::ErrorResponseWithHeaders(status, response_body, headers) => {
252                (status, response_body, Some(headers))
253            }
254        };
255        let body: LazyBody<U> = Box::pin(async move { Ok(U::from(response_body)) });
256        let mut builder = Response::builder().status(status);
257        if let Some(headers) = response_headers
258            && let Some(slot) = builder.headers_mut()
259        {
260            *slot = headers;
261        }
262        builder.body(body).map_err(http_client::Error::Protocol)
263    }
264}
265
266impl HttpClientExt for RecordingHttpClient {
267    fn send<T, U>(
268        &self,
269        req: Request<T>,
270    ) -> impl Future<Output = http_client::Result<Response<LazyBody<U>>>> + WasmCompatSend + 'static
271    where
272        T: Into<Bytes> + WasmCompatSend,
273        U: From<Bytes> + WasmCompatSend + 'static,
274    {
275        let response = self.response_guard().clone();
276        let (parts, body) = req.into_parts();
277        self.record_request(parts.uri.to_string(), parts.headers, body.into());
278
279        async move { Self::build_unary_response(response) }
280    }
281
282    fn send_multipart<U>(
283        &self,
284        req: Request<MultipartForm>,
285    ) -> impl Future<Output = http_client::Result<Response<LazyBody<U>>>> + WasmCompatSend + 'static
286    where
287        U: From<Bytes> + WasmCompatSend + 'static,
288    {
289        let response = self.response_guard().clone();
290        let (parts, body) = req.into_parts();
291        let (_, body) = body.boundary("recording-http-client").encode();
292        self.record_request(parts.uri.to_string(), parts.headers, body);
293
294        async move { Self::build_unary_response(response) }
295    }
296
297    fn send_streaming<T>(
298        &self,
299        _req: Request<T>,
300    ) -> impl Future<Output = http_client::Result<StreamingResponse>> + WasmCompatSend
301    where
302        T: Into<Bytes> + WasmCompatSend,
303    {
304        future::ready(Err(not_implemented()))
305    }
306}
307
308/// An [`HttpClientExt`] implementation that records requests and returns one
309/// scripted response per request. A streamed request's reply body arrives as
310/// one chunk.
311///
312/// This is useful for testing retry and recovery paths through real provider
313/// request/response conversion without live credentials.
314#[derive(Clone, Debug, Default)]
315pub struct SequencedHttpClient {
316    requests: Arc<Mutex<Vec<CapturedHttpRequest>>>,
317    responses: Arc<Mutex<VecDeque<MockHttpResponse>>>,
318}
319
320impl SequencedHttpClient {
321    /// Create a client that returns the supplied responses in order.
322    pub fn new(responses: impl IntoIterator<Item = MockHttpResponse>) -> Self {
323        Self {
324            requests: Arc::new(Mutex::new(Vec::new())),
325            responses: Arc::new(Mutex::new(responses.into_iter().collect())),
326        }
327    }
328
329    /// Return the requests captured so far.
330    pub fn requests(&self) -> Vec<CapturedHttpRequest> {
331        match self.requests.lock() {
332            Ok(guard) => guard.clone(),
333            Err(poisoned) => poisoned.into_inner().clone(),
334        }
335    }
336
337    /// Return the number of scripted responses that have not been consumed.
338    pub fn remaining_responses(&self) -> usize {
339        match self.responses.lock() {
340            Ok(guard) => guard.len(),
341            Err(poisoned) => poisoned.into_inner().len(),
342        }
343    }
344
345    fn record_request(&self, uri: String, headers: http::HeaderMap, body: Bytes) {
346        let request = CapturedHttpRequest { uri, headers, body };
347        match self.requests.lock() {
348            Ok(mut guard) => guard.push(request),
349            Err(poisoned) => poisoned.into_inner().push(request),
350        }
351    }
352
353    fn next_response(&self) -> Option<MockHttpResponse> {
354        match self.responses.lock() {
355            Ok(mut guard) => guard.pop_front(),
356            Err(poisoned) => poisoned.into_inner().pop_front(),
357        }
358    }
359}
360
361impl HttpClientExt for SequencedHttpClient {
362    fn send<T, U>(
363        &self,
364        req: Request<T>,
365    ) -> impl Future<Output = http_client::Result<Response<LazyBody<U>>>> + WasmCompatSend + 'static
366    where
367        T: Into<Bytes> + WasmCompatSend,
368        U: From<Bytes> + WasmCompatSend + 'static,
369    {
370        let response = self.next_response();
371        let (parts, body) = req.into_parts();
372        self.record_request(parts.uri.to_string(), parts.headers, body.into());
373
374        async move {
375            match response {
376                Some(response) => RecordingHttpClient::build_unary_response(response),
377                None => Err(not_implemented()),
378            }
379        }
380    }
381
382    fn send_multipart<U>(
383        &self,
384        req: Request<MultipartForm>,
385    ) -> impl Future<Output = http_client::Result<Response<LazyBody<U>>>> + WasmCompatSend + 'static
386    where
387        U: From<Bytes> + WasmCompatSend + 'static,
388    {
389        let response = self.next_response();
390        let (parts, _body) = req.into_parts();
391        self.record_request(parts.uri.to_string(), parts.headers, Bytes::new());
392
393        async move {
394            match response {
395                Some(response) => RecordingHttpClient::build_unary_response(response),
396                None => Err(not_implemented()),
397            }
398        }
399    }
400
401    fn send_streaming<T>(
402        &self,
403        req: Request<T>,
404    ) -> impl Future<Output = http_client::Result<StreamingResponse>> + WasmCompatSend
405    where
406        T: Into<Bytes> + WasmCompatSend,
407    {
408        let response = self.next_response();
409        let (parts, body) = req.into_parts();
410        self.record_request(parts.uri.to_string(), parts.headers, body.into());
411
412        future::ready(match response {
413            Some(response) => streaming_response(response),
414            None => Err(not_implemented()),
415        })
416    }
417}
418
419/// `response` as a streamed reply whose body is one chunk.
420fn streaming_response(response: MockHttpResponse) -> http_client::Result<StreamingResponse> {
421    let (status, body, headers) = match response {
422        MockHttpResponse::Success(body) => (http::StatusCode::OK, body, None),
423        MockHttpResponse::SuccessWithHeaders(body, headers) => {
424            (http::StatusCode::OK, body, Some(headers))
425        }
426        MockHttpResponse::ErrorWithHeaders(status, body, headers) => {
427            return Err(http_client::Error::InvalidStatusCodeWithDetails {
428                status,
429                body,
430                headers,
431            });
432        }
433        MockHttpResponse::ErrorResponse(status, body) => (status, body, None),
434        MockHttpResponse::ErrorResponseWithHeaders(status, body, headers) => {
435            (status, body, Some(headers))
436        }
437    };
438    let chunks: http_client::BoxedStream = Box::pin(futures::stream::iter([Ok::<
439        Bytes,
440        http_client::Error,
441    >(body)]));
442    let mut response = Response::builder()
443        .status(status)
444        .body(chunks)
445        .map_err(http_client::Error::Protocol)?;
446    if let Some(headers) = headers {
447        *response.headers_mut() = headers;
448    }
449    Ok(response)
450}
451
452/// A mock HTTP client that returns pre-built SSE bytes from `send_streaming`.
453///
454/// `send` and `send_multipart` always return `NOT_IMPLEMENTED`.
455#[derive(Clone, Debug, Default)]
456pub struct MockStreamingClient {
457    /// Bytes returned as a single streaming response chunk.
458    pub sse_bytes: Bytes,
459}
460
461impl HttpClientExt for MockStreamingClient {
462    no_unary_surface!();
463
464    fn send_streaming<T>(
465        &self,
466        _req: Request<T>,
467    ) -> impl Future<Output = http_client::Result<StreamingResponse>> + WasmCompatSend
468    where
469        T: Into<Bytes> + WasmCompatSend,
470    {
471        let sse_bytes = self.sse_bytes.clone();
472        async move {
473            let byte_stream =
474                futures::stream::iter(vec![Ok::<Bytes, http_client::Error>(sse_bytes)]);
475            let boxed_stream: http_client::BoxedStream = Box::pin(byte_stream);
476
477            Response::builder()
478                .status(http::StatusCode::OK)
479                .header(http::header::CONTENT_TYPE, "text/event-stream")
480                .body(boxed_stream)
481                .map_err(http_client::Error::Protocol)
482        }
483    }
484}
485
486/// An [`HttpClientExt`] implementation whose `send_streaming` fails immediately
487/// with a non-success HTTP status and response body.
488#[derive(Debug, Clone)]
489pub struct HttpErrorStreamingClient {
490    pub status: http::StatusCode,
491    pub body: String,
492}
493
494impl HttpErrorStreamingClient {
495    /// Create a streaming client that fails `send_streaming` with the given status and body.
496    pub fn new(status: http::StatusCode, body: impl Into<String>) -> Self {
497        Self {
498            status,
499            body: body.into(),
500        }
501    }
502}
503
504impl Default for HttpErrorStreamingClient {
505    /// The completion-model client bound requires `H: Default`; this lets the
506    /// streaming error client back a real model in tests.
507    fn default() -> Self {
508        Self::new(http::StatusCode::INTERNAL_SERVER_ERROR, String::new())
509    }
510}
511
512impl HttpClientExt for HttpErrorStreamingClient {
513    no_unary_surface!();
514
515    fn send_streaming<T>(
516        &self,
517        _req: Request<T>,
518    ) -> impl Future<Output = http_client::Result<StreamingResponse>> + WasmCompatSend
519    where
520        T: Into<Bytes> + WasmCompatSend,
521    {
522        let status = self.status;
523        let body = self.body.clone();
524        async move {
525            Err(http_client::Error::non_success_with_details(
526                status,
527                http::HeaderMap::new(),
528                body,
529            ))
530        }
531    }
532}
533
534/// An [`HttpClientExt`] whose `send_streaming` hands a non-success reply back
535/// as an `Ok` response — the fourth non-success cell: a custom transport that
536/// does not reject on status, so the driver must.
537#[derive(Debug, Clone)]
538pub struct NonSuccessStreamingClient {
539    pub status: http::StatusCode,
540    pub headers: http::HeaderMap,
541    pub body: Bytes,
542}
543
544impl HttpClientExt for NonSuccessStreamingClient {
545    no_unary_surface!();
546
547    fn send_streaming<T>(
548        &self,
549        _req: Request<T>,
550    ) -> impl Future<Output = http_client::Result<StreamingResponse>> + WasmCompatSend
551    where
552        T: Into<Bytes> + WasmCompatSend,
553    {
554        let status = self.status;
555        let headers = self.headers.clone();
556        let body = self.body.clone();
557        async move {
558            let byte_stream = futures::stream::iter(vec![Ok::<Bytes, http_client::Error>(body)]);
559            let boxed_stream: http_client::BoxedStream = Box::pin(byte_stream);
560            let mut response = Response::builder()
561                .status(status)
562                .body(boxed_stream)
563                .map_err(http_client::Error::Protocol)?;
564            *response.headers_mut() = headers;
565            Ok(response)
566        }
567    }
568}
569
570/// An [`HttpClientExt`] implementation that returns one scripted stream of byte
571/// chunks from `send_streaming`.
572#[derive(Debug, Clone, Default)]
573pub struct SequencedStreamingHttpClient {
574    chunks: Arc<Mutex<Option<Vec<http_client::Result<Bytes>>>>>,
575}
576
577impl SequencedStreamingHttpClient {
578    /// Create a streaming client from the chunks it should yield.
579    pub fn new(chunks: Vec<http_client::Result<Bytes>>) -> Self {
580        Self {
581            chunks: Arc::new(Mutex::new(Some(chunks))),
582        }
583    }
584}
585
586impl HttpClientExt for SequencedStreamingHttpClient {
587    no_unary_surface!();
588
589    fn send_streaming<T>(
590        &self,
591        _req: Request<T>,
592    ) -> impl Future<Output = http_client::Result<StreamingResponse>> + WasmCompatSend
593    where
594        T: Into<Bytes> + WasmCompatSend,
595    {
596        let chunks = match self.chunks.lock() {
597            Ok(mut guard) => guard.take(),
598            Err(poisoned) => poisoned.into_inner().take(),
599        };
600
601        async move {
602            let Some(chunks) = chunks else {
603                return Err(http_client::Error::non_success_with_details(
604                    http::StatusCode::INTERNAL_SERVER_ERROR,
605                    http::HeaderMap::new(),
606                    "streaming chunks should only be consumed once".to_string(),
607                ));
608            };
609
610            let byte_stream = futures::stream::iter(chunks);
611            let boxed_stream: http_client::BoxedStream = Box::pin(byte_stream);
612
613            Response::builder()
614                .status(http::StatusCode::OK)
615                .header(http::header::CONTENT_TYPE, "text/event-stream")
616                .body(boxed_stream)
617                .map_err(http_client::Error::Protocol)
618        }
619    }
620}