Skip to main content

eventsource_client/
client.rs

1use base64::prelude::*;
2
3use futures::{ready, Stream};
4use http::{HeaderMap, HeaderName, HeaderValue, Request, Uri};
5use log::{debug, info, trace, warn};
6use pin_project::pin_project;
7use std::{
8    boxed,
9    fmt::{self, Debug, Formatter},
10    future::Future,
11    io::ErrorKind,
12    pin::Pin,
13    str::FromStr,
14    sync::Arc,
15    task::{Context, Poll},
16    time::{Duration, Instant},
17};
18
19use tokio::sync::watch;
20use tokio::time::Sleep;
21
22use crate::{
23    config::ReconnectOptions,
24    response::{ErrorBody, Response},
25};
26use crate::{
27    error::{Error, Result},
28    event_parser::ConnectionDetails,
29};
30use launchdarkly_sdk_transport::{ByteStream, HttpTransport, ResponseFuture};
31
32use crate::event_parser::EventParser;
33use crate::event_parser::SSE;
34
35use crate::retry::{BackoffRetry, RetryStrategy};
36use std::error::Error as StdError;
37
38/// Represents a [`Pin`]'d [`Send`] + [`Sync`] stream, returned by [`Client`]'s stream method.
39pub type BoxStream<T> = Pin<boxed::Box<dyn Stream<Item = T> + Send + Sync>>;
40
41/// Client is the Server-Sent-Events interface.
42/// This trait is sealed and cannot be implemented for types outside this crate.
43pub trait Client: Send + Sync + private::Sealed {
44    fn stream(&self) -> BoxStream<Result<SSE>>;
45}
46
47/*
48 * TODO remove debug output
49 * TODO specify list of stati to not retry (e.g. 204)
50 */
51
52/// Maximum amount of redirects that the client will follow before
53/// giving up, if not overridden via [ClientBuilder::redirect_limit].
54pub const DEFAULT_REDIRECT_LIMIT: u32 = 16;
55
56/// ClientBuilder provides a series of builder methods to easily construct a [`Client`].
57pub struct ClientBuilder {
58    url: Uri,
59    headers: HeaderMap,
60    reconnect_opts: ReconnectOptions,
61    last_event_id: Option<String>,
62    method: String,
63    body: Option<String>,
64    max_redirects: Option<u32>,
65    dynamic_url: Option<watch::Receiver<Uri>>,
66}
67
68impl ClientBuilder {
69    /// Create a builder for a given URL.
70    pub fn for_url(url: &str) -> Result<ClientBuilder> {
71        let url = url
72            .parse()
73            .map_err(|e| Error::InvalidParameter(Box::new(e)))?;
74
75        let mut header_map = HeaderMap::new();
76        header_map.insert("Accept", HeaderValue::from_static("text/event-stream"));
77        header_map.insert("Cache-Control", HeaderValue::from_static("no-cache"));
78
79        Ok(ClientBuilder {
80            url,
81            headers: header_map,
82            reconnect_opts: ReconnectOptions::default(),
83            last_event_id: None,
84            method: String::from("GET"),
85            max_redirects: None,
86            body: None,
87            dynamic_url: None,
88        })
89    }
90
91    /// Watch the given receiver for the url to use when attempting to connect
92    /// or reconnect. Overrides [`for_url`] if both are present.
93    pub fn dynamic_url(mut self, uri: watch::Receiver<Uri>) -> ClientBuilder {
94        self.dynamic_url = Some(uri);
95        self
96    }
97
98    /// Set the request method used for the initial connection to the SSE endpoint.
99    pub fn method(mut self, method: String) -> ClientBuilder {
100        self.method = method;
101        self
102    }
103
104    /// Set the request body used for the initial connection to the SSE endpoint.
105    pub fn body(mut self, body: String) -> ClientBuilder {
106        self.body = Some(body);
107        self
108    }
109
110    /// Set the last event id for a stream when it is created. If it is set, it will be sent to the
111    /// server in case it can replay missed events.
112    pub fn last_event_id(mut self, last_event_id: String) -> ClientBuilder {
113        self.last_event_id = Some(last_event_id);
114        self
115    }
116
117    /// Set a HTTP header on the SSE request.
118    pub fn header(mut self, name: &str, value: &str) -> Result<ClientBuilder> {
119        let name = HeaderName::from_str(name).map_err(|e| Error::InvalidParameter(Box::new(e)))?;
120
121        let value =
122            HeaderValue::from_str(value).map_err(|e| Error::InvalidParameter(Box::new(e)))?;
123
124        self.headers.insert(name, value);
125        Ok(self)
126    }
127
128    /// Set the Authorization header with the calculated basic authentication value.
129    pub fn basic_auth(self, username: &str, password: &str) -> Result<ClientBuilder> {
130        let auth = format!("{username}:{password}");
131        let encoded = BASE64_STANDARD.encode(auth);
132        let value = format!("Basic {encoded}");
133
134        self.header("Authorization", &value)
135    }
136
137    /// Configure the client's reconnect behaviour according to the supplied
138    /// [`ReconnectOptions`].
139    ///
140    /// [`ReconnectOptions`]: struct.ReconnectOptions.html
141    pub fn reconnect(mut self, opts: ReconnectOptions) -> ClientBuilder {
142        self.reconnect_opts = opts;
143        self
144    }
145
146    /// Customize the client's following behavior when served a redirect.
147    /// To disable following redirects, pass `0`.
148    /// By default, the limit is [`DEFAULT_REDIRECT_LIMIT`].
149    pub fn redirect_limit(mut self, limit: u32) -> ClientBuilder {
150        self.max_redirects = Some(limit);
151        self
152    }
153
154    /// Build a client with a custom HTTP transport implementation.
155    ///
156    /// # Arguments
157    ///
158    /// * `transport` - An implementation of the [`HttpTransport`] trait that will handle
159    ///   HTTP requests. See the `examples/` directory for reference implementations.
160    ///
161    /// # Example
162    ///
163    /// ```ignore
164    /// use eventsource_client::ClientBuilder;
165    ///
166    /// let transport = MyTransport::new();
167    /// let client = ClientBuilder::for_url("https://live-test-scores.herokuapp.com/scores")
168    ///     .expect("failed to create client builder")
169    ///     .build_with_transport(transport);
170    /// ```
171    pub fn build_with_transport<T>(self, transport: T) -> impl Client
172    where
173        T: HttpTransport,
174    {
175        ClientImpl {
176            transport: Arc::new(transport),
177            request_props: RequestProps {
178                url: self.url,
179                headers: self.headers,
180                method: self.method,
181                body: self.body,
182                reconnect_opts: self.reconnect_opts,
183                max_redirects: self.max_redirects.unwrap_or(DEFAULT_REDIRECT_LIMIT),
184                dynamic_url: self.dynamic_url,
185            },
186            last_event_id: self.last_event_id,
187        }
188    }
189}
190
191#[derive(Clone)]
192struct RequestProps {
193    url: Uri,
194    headers: HeaderMap,
195    method: String,
196    body: Option<String>,
197    reconnect_opts: ReconnectOptions,
198    max_redirects: u32,
199    dynamic_url: Option<watch::Receiver<Uri>>,
200}
201
202impl RequestProps {
203    fn resolve_url(&self) -> Uri {
204        self.dynamic_url
205            .as_ref()
206            .map(|rx| rx.borrow().clone())
207            .unwrap_or_else(|| self.url.clone())
208    }
209}
210
211/// A client implementation that connects to a server using the Server-Sent Events protocol
212/// and consumes the event stream indefinitely.
213struct ClientImpl<T: HttpTransport> {
214    transport: Arc<T>,
215    request_props: RequestProps,
216    last_event_id: Option<String>,
217}
218
219impl<T: HttpTransport> Client for ClientImpl<T> {
220    /// Connect to the server and begin consuming the stream. Produces a
221    /// [`Stream`] of [`Event`](crate::Event)s wrapped in [`Result`].
222    ///
223    /// Errors yielded by the stream are not terminal: keep polling.
224    /// When [`ReconnectOptions::reconnect`] is enabled (the default),
225    /// the stream schedules a reconnect on retryable errors and the
226    /// next poll resumes from a fresh connection.
227    ///
228    /// The stream is exhausted only when [`Stream::poll_next`] returns
229    /// [`Poll::Ready(None)`]. That happens when the underlying state
230    /// machine reaches `StreamClosed` (e.g. a redirect-limit overrun,
231    /// a malformed `Location` header, or an error during initial
232    /// connection while [`ReconnectOptions::retry_initial`] is
233    /// disabled), or after any error when reconnect is disabled.
234    ///
235    /// [`Poll::Ready(None)`]: std::task::Poll::Ready
236    /// [`Stream::poll_next`]: futures::Stream::poll_next
237    fn stream(&self) -> BoxStream<Result<SSE>> {
238        Box::pin(ReconnectingRequest::new(
239            Arc::clone(&self.transport),
240            self.request_props.clone(),
241            self.last_event_id.clone(),
242        ))
243    }
244}
245
246#[allow(clippy::large_enum_variant)] // false positive
247#[pin_project(project = StateProj)]
248enum State {
249    New,
250    Connecting {
251        retry: bool,
252        redirect_count: u32,
253        #[pin]
254        resp: ResponseFuture,
255    },
256    Connected(#[pin] ByteStream),
257    WaitingToReconnect(#[pin] Sleep),
258    FollowingRedirect {
259        header: Option<HeaderValue>,
260        redirect_count: u32,
261    },
262    StreamClosed,
263}
264
265impl State {
266    fn name(&self) -> &'static str {
267        match self {
268            State::New => "new",
269            State::Connecting { retry: false, .. } => "connecting(no-retry)",
270            State::Connecting { retry: true, .. } => "connecting(retry)",
271            State::Connected(_) => "connected",
272            State::WaitingToReconnect(_) => "waiting-to-reconnect",
273            State::FollowingRedirect { .. } => "following-redirect",
274            State::StreamClosed => "closed",
275        }
276    }
277}
278
279impl Debug for State {
280    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
281        write!(f, "{}", self.name())
282    }
283}
284
285#[must_use = "streams do nothing unless polled"]
286#[pin_project]
287pub struct ReconnectingRequest<T: HttpTransport> {
288    transport: Arc<T>,
289    props: RequestProps,
290    #[pin]
291    state: State,
292    retry_strategy: Box<dyn RetryStrategy + Send + Sync>,
293    event_parser: EventParser,
294    last_event_id: Option<String>,
295    #[pin]
296    initial_connection: bool,
297}
298
299impl<T: HttpTransport> ReconnectingRequest<T> {
300    fn new(
301        transport: Arc<T>,
302        props: RequestProps,
303        last_event_id: Option<String>,
304    ) -> ReconnectingRequest<T> {
305        let reconnect_delay = props.reconnect_opts.delay;
306        let delay_max = props.reconnect_opts.delay_max;
307        let backoff_factor = props.reconnect_opts.backoff_factor;
308
309        ReconnectingRequest {
310            props,
311            transport,
312            state: State::New,
313            retry_strategy: Box::new(BackoffRetry::new(
314                reconnect_delay,
315                delay_max,
316                backoff_factor,
317                true,
318            )),
319            event_parser: EventParser::new(),
320            last_event_id,
321            initial_connection: true,
322        }
323    }
324
325    fn send_request(&self, url: &Uri) -> Result<ResponseFuture> {
326        let mut request_builder = Request::builder()
327            .method(self.props.method.as_str())
328            .uri(url);
329
330        for (name, value) in &self.props.headers {
331            request_builder = request_builder.header(name, value);
332        }
333
334        if let Some(id) = self.last_event_id.as_ref() {
335            if !id.is_empty() {
336                let id_as_header =
337                    HeaderValue::from_str(id).map_err(|e| Error::InvalidParameter(Box::new(e)))?;
338
339                request_builder = request_builder.header("last-event-id", id_as_header);
340            }
341        }
342
343        // Include the request body if set. Most SSE requests use GET and will have None,
344        // but some implementations (e.g., using REPORT method) may include a body.
345        let request = request_builder
346            .body(self.props.body.clone().map(|b| b.into()))
347            .map_err(|e| Error::InvalidParameter(Box::new(e)))?;
348
349        Ok(self.transport.request(request))
350    }
351}
352
353impl<T: HttpTransport> Stream for ReconnectingRequest<T> {
354    type Item = Result<SSE>;
355
356    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
357        trace!("ReconnectingRequest::poll({:?})", self.state);
358
359        loop {
360            let this = self.as_mut().project();
361            if let Some(event) = this.event_parser.get_event() {
362                return match event {
363                    SSE::Connected(_) => Poll::Ready(Some(Ok(event))),
364                    SSE::Event(ref evt) => {
365                        this.last_event_id.clone_from(&evt.id);
366
367                        if let Some(retry) = evt.retry {
368                            this.retry_strategy
369                                .change_base_delay(Duration::from_millis(retry));
370                        }
371                        Poll::Ready(Some(Ok(event)))
372                    }
373                    SSE::Comment(_) => Poll::Ready(Some(Ok(event))),
374                };
375            }
376
377            trace!("ReconnectingRequest::poll loop({:?})", this.state);
378
379            let state = this.state.project();
380            match state {
381                StateProj::StreamClosed => return Poll::Ready(None),
382                // New immediately transitions to Connecting, and exists only
383                // to ensure that we only connect when polled.
384                StateProj::New => {
385                    *self.as_mut().project().event_parser = EventParser::new();
386                    let url = self.props.resolve_url();
387                    match self.send_request(&url) {
388                        Ok(resp) => {
389                            let retry = if self.initial_connection {
390                                self.props.reconnect_opts.retry_initial
391                            } else {
392                                self.props.reconnect_opts.reconnect
393                            };
394                            self.as_mut().project().state.set(State::Connecting {
395                                resp,
396                                retry,
397                                redirect_count: 0,
398                            })
399                        }
400                        Err(e) => {
401                            // This error seems to be unrecoverable. So we should just shut down the
402                            // stream.
403                            self.as_mut().project().state.set(State::StreamClosed);
404                            return Poll::Ready(Some(Err(e)));
405                        }
406                    }
407                }
408                StateProj::Connecting {
409                    retry,
410                    redirect_count,
411                    resp,
412                } => match ready!(resp.poll(cx)) {
413                    Ok(resp) => {
414                        debug!(
415                            "HTTP response status: {}, headers: {:?}",
416                            resp.status(),
417                            resp.headers()
418                        );
419
420                        if resp.status().is_success() {
421                            self.as_mut().project().retry_strategy.reset(Instant::now());
422
423                            let status = resp.status();
424                            let headers = resp.headers().clone();
425
426                            self.as_mut()
427                                .project()
428                                .state
429                                .set(State::Connected(resp.into_body()));
430                            self.as_mut().project().initial_connection.set(false);
431
432                            return Poll::Ready(Some(Ok(SSE::Connected(ConnectionDetails::new(
433                                Response::new(status, headers),
434                            )))));
435                        }
436
437                        if resp.status() == 301 || resp.status() == 307 {
438                            debug!("got redirected ({})", resp.status());
439
440                            let next_count = *redirect_count + 1;
441                            if next_count > self.props.max_redirects {
442                                debug!("redirect limit reached ({})", self.props.max_redirects);
443
444                                self.as_mut().project().state.set(State::StreamClosed);
445                                return Poll::Ready(Some(Err(Error::MaxRedirectLimitReached(
446                                    self.props.max_redirects,
447                                ))));
448                            } else {
449                                debug!("following redirect {}", next_count);
450
451                                self.as_mut().project().state.set(State::FollowingRedirect {
452                                    header: resp.headers().get("location").cloned(),
453                                    redirect_count: next_count,
454                                });
455                                continue;
456                            }
457                        }
458
459                        let status = resp.status();
460                        let headers = resp.headers().clone();
461                        let body = resp.into_body();
462
463                        let error = Error::UnexpectedResponse(
464                            Response::new(status, headers),
465                            ErrorBody::new(body),
466                        );
467
468                        if !*retry {
469                            self.as_mut().project().state.set(State::StreamClosed);
470                            return Poll::Ready(Some(Err(error)));
471                        }
472
473                        let duration = self
474                            .as_mut()
475                            .project()
476                            .retry_strategy
477                            .next_delay(Instant::now());
478
479                        self.as_mut()
480                            .project()
481                            .state
482                            .set(State::WaitingToReconnect(delay(duration, "retrying")));
483
484                        return Poll::Ready(Some(Err(error)));
485                    }
486                    Err(e) => {
487                        // This happens when the server is unreachable, e.g. connection refused.
488                        warn!("request returned an error: {e}");
489                        if !*retry {
490                            self.as_mut().project().state.set(State::StreamClosed);
491                            return Poll::Ready(Some(Err(Error::Transport(e))));
492                        }
493
494                        let duration = self
495                            .as_mut()
496                            .project()
497                            .retry_strategy
498                            .next_delay(Instant::now());
499
500                        self.as_mut()
501                            .project()
502                            .state
503                            .set(State::WaitingToReconnect(delay(duration, "retrying")));
504                    }
505                },
506                StateProj::FollowingRedirect {
507                    header,
508                    redirect_count,
509                } => match uri_from_header(header) {
510                    Ok(uri) => {
511                        let count = *redirect_count;
512                        match self.send_request(&uri) {
513                            Ok(resp) => {
514                                let retry = if self.initial_connection {
515                                    self.props.reconnect_opts.retry_initial
516                                } else {
517                                    self.props.reconnect_opts.reconnect
518                                };
519                                self.as_mut().project().state.set(State::Connecting {
520                                    resp,
521                                    retry,
522                                    redirect_count: count,
523                                });
524                            }
525                            Err(e) => {
526                                self.as_mut().project().state.set(State::StreamClosed);
527                                return Poll::Ready(Some(Err(e)));
528                            }
529                        }
530                    }
531                    Err(e) => {
532                        self.as_mut().project().state.set(State::StreamClosed);
533                        return Poll::Ready(Some(Err(e)));
534                    }
535                },
536                StateProj::Connected(mut body) => match ready!(body.as_mut().poll_next(cx)) {
537                    Some(Ok(result)) => {
538                        if let Err(e) = this.event_parser.process_bytes(result) {
539                            // The current response body is unusable. Either
540                            // schedule a reconnect or close the stream so a
541                            // caller that disabled reconnect doesn't keep
542                            // reading from a poisoned parser.
543                            if self.props.reconnect_opts.reconnect {
544                                let duration = self
545                                    .as_mut()
546                                    .project()
547                                    .retry_strategy
548                                    .next_delay(Instant::now());
549                                self.as_mut().project().state.set(State::WaitingToReconnect(
550                                    delay(duration, "reconnecting"),
551                                ));
552                            } else {
553                                self.as_mut().project().state.set(State::StreamClosed);
554                            }
555                            return Poll::Ready(Some(Err(e)));
556                        }
557                        continue;
558                    }
559                    Some(Err(e)) => {
560                        if self.props.reconnect_opts.reconnect {
561                            let duration = self
562                                .as_mut()
563                                .project()
564                                .retry_strategy
565                                .next_delay(Instant::now());
566                            self.as_mut()
567                                .project()
568                                .state
569                                .set(State::WaitingToReconnect(delay(duration, "reconnecting")));
570                        }
571
572                        // Check if the underlying error is a timeout
573                        if let Some(cause) = e.source() {
574                            if let Some(downcast) = cause.downcast_ref::<std::io::Error>() {
575                                if let std::io::ErrorKind::TimedOut = downcast.kind() {
576                                    return Poll::Ready(Some(Err(Error::TimedOut)));
577                                }
578                            }
579                        }
580
581                        return Poll::Ready(Some(Err(Error::Transport(e))));
582                    }
583                    None => {
584                        if self.props.reconnect_opts.reconnect {
585                            let duration = self
586                                .as_mut()
587                                .project()
588                                .retry_strategy
589                                .next_delay(Instant::now());
590                            self.as_mut()
591                                .project()
592                                .state
593                                .set(State::WaitingToReconnect(delay(duration, "retrying")));
594                        } else {
595                            self.as_mut().project().state.set(State::StreamClosed);
596                        }
597
598                        if self.event_parser.was_processing() {
599                            return Poll::Ready(Some(Err(Error::UnexpectedEof)));
600                        }
601                        return Poll::Ready(Some(Err(Error::Eof)));
602                    }
603                },
604                StateProj::WaitingToReconnect(delay) => {
605                    ready!(delay.poll(cx));
606                    info!("Reconnecting");
607                    self.as_mut().project().state.set(State::New);
608                }
609            };
610        }
611    }
612}
613
614fn uri_from_header(maybe_header: &Option<HeaderValue>) -> Result<Uri> {
615    let header = maybe_header.as_ref().ok_or_else(|| {
616        Error::MalformedLocationHeader(Box::new(std::io::Error::new(
617            ErrorKind::NotFound,
618            "missing Location header",
619        )))
620    })?;
621
622    let header_string = header
623        .to_str()
624        .map_err(|e| Error::MalformedLocationHeader(Box::new(e)))?;
625
626    header_string
627        .parse::<Uri>()
628        .map_err(|e| Error::MalformedLocationHeader(Box::new(e)))
629}
630
631fn delay(dur: Duration, description: &str) -> Sleep {
632    info!("Waiting {dur:?} before {description}");
633    tokio::time::sleep(dur)
634}
635
636mod private {
637    use crate::client::ClientImpl;
638    use launchdarkly_sdk_transport::HttpTransport;
639
640    pub trait Sealed {}
641    impl<T: HttpTransport> Sealed for ClientImpl<T> {}
642}
643
644#[cfg(test)]
645mod tests {
646    use crate::ClientBuilder;
647    use http::HeaderValue;
648    use test_case::test_case;
649
650    #[test_case("user", "pass", "dXNlcjpwYXNz")]
651    #[test_case("user1", "password123", "dXNlcjE6cGFzc3dvcmQxMjM=")]
652    #[test_case("user2", "", "dXNlcjI6")]
653    #[test_case("user@name", "pass#word!", "dXNlckBuYW1lOnBhc3Mjd29yZCE=")]
654    #[test_case("user3", "my pass", "dXNlcjM6bXkgcGFzcw==")]
655    #[test_case(
656        "weird@-/:stuff",
657        "goes@-/:here",
658        "d2VpcmRALS86c3R1ZmY6Z29lc0AtLzpoZXJl"
659    )]
660    fn basic_auth_generates_correct_headers(username: &str, password: &str, expected: &str) {
661        let builder = ClientBuilder::for_url("http://example.com")
662            .expect("failed to build client")
663            .basic_auth(username, password)
664            .expect("failed to add authentication");
665
666        let actual = builder.headers.get("Authorization");
667        let expected = HeaderValue::from_str(format!("Basic {expected}").as_str())
668            .expect("unable to create expected header");
669
670        assert_eq!(Some(&expected), actual);
671    }
672
673    use std::{
674        pin::pin,
675        sync::{Arc, Mutex},
676        time::Duration,
677    };
678
679    use bytes::Bytes;
680    use futures::{stream, TryStreamExt};
681    use http::HeaderMap;
682    use tokio::time::timeout;
683
684    use crate::{
685        client::{RequestProps, State},
686        ReconnectOptionsBuilder, ReconnectingRequest,
687    };
688    use launchdarkly_sdk_transport::{ByteStream, HttpTransport, ResponseFuture, TransportError};
689
690    // Mock transport for testing
691    #[derive(Clone)]
692    struct MockTransport {
693        fail_request: bool,
694    }
695
696    impl MockTransport {
697        fn new(_url: String, fail_request: bool) -> Self {
698            Self { fail_request }
699        }
700    }
701
702    impl HttpTransport for MockTransport {
703        fn request(&self, _request: http::Request<Option<Bytes>>) -> ResponseFuture {
704            if self.fail_request {
705                // Simulate a connection error
706                Box::pin(async {
707                    Err(TransportError::new(std::io::Error::new(
708                        std::io::ErrorKind::ConnectionRefused,
709                        "connection refused",
710                    )))
711                })
712            } else {
713                // Return a 404 response
714                Box::pin(async {
715                    let byte_stream: ByteStream =
716                        Box::pin(stream::iter(vec![Ok(Bytes::from("not found"))]));
717                    let response = http::Response::builder()
718                        .status(404)
719                        .body(byte_stream)
720                        .unwrap();
721                    Ok(response)
722                })
723            }
724        }
725    }
726
727    const INVALID_URI: &str = "http://mycrazyunexsistenturl.invaliddomainext";
728
729    #[test_case(INVALID_URI, false, |state| matches!(state, State::StreamClosed))]
730    #[test_case(INVALID_URI, true, |state| matches!(state, State::WaitingToReconnect(_)))]
731    #[tokio::test]
732    async fn initial_connection(uri: &str, retry_initial: bool, expected: fn(&State) -> bool) {
733        let reconnect_opts = ReconnectOptionsBuilder::new(false)
734            .backoff_factor(1)
735            .delay(Duration::from_secs(1))
736            .retry_initial(retry_initial)
737            .build();
738
739        let transport = Arc::new(MockTransport::new(uri.to_string(), true));
740        let req_props = RequestProps {
741            url: uri.parse().unwrap(),
742            headers: HeaderMap::new(),
743            method: "GET".to_string(),
744            body: None,
745            reconnect_opts,
746            max_redirects: 10,
747            dynamic_url: None,
748        };
749
750        let mut reconnecting_request = ReconnectingRequest::new(transport.clone(), req_props, None);
751
752        // sets initial state with a failing request
753        let resp = transport.request(http::Request::builder().uri(uri).body(None).unwrap());
754
755        reconnecting_request.state = State::Connecting {
756            retry: reconnecting_request.props.reconnect_opts.retry_initial,
757            redirect_count: 0,
758            resp,
759        };
760
761        let mut reconnecting_request = pin!(reconnecting_request);
762
763        timeout(Duration::from_millis(500), reconnecting_request.try_next())
764            .await
765            .ok();
766
767        assert!(expected(&reconnecting_request.state));
768    }
769
770    #[test_case(false, |state| matches!(state, State::StreamClosed))]
771    #[test_case(true, |state| matches!(state, State::WaitingToReconnect(_)))]
772    #[tokio::test]
773    async fn initial_connection_mocked_server(retry_initial: bool, expected: fn(&State) -> bool) {
774        let mut mock_server = mockito::Server::new_async().await;
775        let _mock = mock_server
776            .mock("GET", "/")
777            .with_status(404)
778            .create_async()
779            .await;
780
781        initial_connection(&mock_server.url(), retry_initial, expected).await;
782    }
783
784    #[derive(Clone)]
785    struct CapturingTransport {
786        captured_uris: Arc<Mutex<Vec<http::Uri>>>,
787    }
788
789    impl HttpTransport for CapturingTransport {
790        fn request(&self, request: http::Request<Option<Bytes>>) -> ResponseFuture {
791            self.captured_uris
792                .lock()
793                .unwrap()
794                .push(request.uri().clone());
795            Box::pin(async {
796                Err(TransportError::new(std::io::Error::new(
797                    std::io::ErrorKind::ConnectionRefused,
798                    "test",
799                )))
800            })
801        }
802    }
803
804    fn props_with_dynamic(
805        static_url: &str,
806        rx: tokio::sync::watch::Receiver<http::Uri>,
807    ) -> RequestProps {
808        RequestProps {
809            url: static_url.parse().unwrap(),
810            headers: HeaderMap::new(),
811            method: "GET".to_string(),
812            body: None,
813            reconnect_opts: ReconnectOptionsBuilder::new(false).build(),
814            max_redirects: 10,
815            dynamic_url: Some(rx),
816        }
817    }
818
819    #[tokio::test]
820    async fn dynamic_url_is_used_on_initial_connect() {
821        let (_tx, rx) = tokio::sync::watch::channel(
822            "http://dynamic.example.com/".parse::<http::Uri>().unwrap(),
823        );
824        let captured = Arc::new(Mutex::new(Vec::new()));
825        let transport = CapturingTransport {
826            captured_uris: captured.clone(),
827        };
828        let props = props_with_dynamic("http://static.example.com/", rx);
829        let req = ReconnectingRequest::new(Arc::new(transport), props, None);
830
831        let _ = req.send_request(&req.props.resolve_url());
832
833        let uris = captured.lock().unwrap();
834        assert_eq!(uris.len(), 1);
835        assert_eq!(uris[0].to_string(), "http://dynamic.example.com/");
836    }
837
838    #[derive(Clone)]
839    struct RedirectTransport {
840        location: String,
841    }
842
843    impl HttpTransport for RedirectTransport {
844        fn request(&self, _request: http::Request<Option<Bytes>>) -> ResponseFuture {
845            let location = self.location.clone();
846            Box::pin(async move {
847                let byte_stream: ByteStream = Box::pin(stream::iter(Vec::<
848                    std::result::Result<Bytes, TransportError>,
849                >::new()));
850                Ok(http::Response::builder()
851                    .status(301)
852                    .header("Location", location)
853                    .body(byte_stream)
854                    .unwrap())
855            })
856        }
857    }
858
859    #[derive(Clone)]
860    struct RedirectOnceTransport {
861        location: String,
862        captured_uris: Arc<Mutex<Vec<http::Uri>>>,
863    }
864
865    impl HttpTransport for RedirectOnceTransport {
866        fn request(&self, request: http::Request<Option<Bytes>>) -> ResponseFuture {
867            let mut uris = self.captured_uris.lock().unwrap();
868            let is_first = uris.is_empty();
869            uris.push(request.uri().clone());
870            drop(uris);
871            let location = self.location.clone();
872            Box::pin(async move {
873                if is_first {
874                    let byte_stream: ByteStream = Box::pin(stream::iter(Vec::<
875                        std::result::Result<Bytes, TransportError>,
876                    >::new(
877                    )));
878                    Ok(http::Response::builder()
879                        .status(301)
880                        .header("Location", location)
881                        .body(byte_stream)
882                        .unwrap())
883                } else {
884                    Err(TransportError::new(std::io::Error::new(
885                        std::io::ErrorKind::ConnectionRefused,
886                        "stop",
887                    )))
888                }
889            })
890        }
891    }
892
893    #[tokio::test]
894    async fn connecting_sees_301_follows_to_location() {
895        let captured = Arc::new(Mutex::new(Vec::new()));
896        let transport = Arc::new(RedirectOnceTransport {
897            location: "http://redirect.example.com/".to_string(),
898            captured_uris: captured.clone(),
899        });
900        let props = RequestProps {
901            url: "http://start.example.com/".parse().unwrap(),
902            headers: HeaderMap::new(),
903            method: "GET".to_string(),
904            body: None,
905            reconnect_opts: ReconnectOptionsBuilder::new(false).build(),
906            max_redirects: 3,
907            dynamic_url: None,
908        };
909        let req = ReconnectingRequest::new(transport, props, None);
910        let mut req = pin!(req);
911        timeout(Duration::from_millis(500), req.try_next())
912            .await
913            .ok();
914
915        let uris = captured.lock().unwrap();
916        assert_eq!(uris.len(), 2);
917        assert_eq!(uris[0].to_string(), "http://start.example.com/");
918        assert_eq!(uris[1].to_string(), "http://redirect.example.com/");
919    }
920
921    #[tokio::test]
922    async fn connecting_sees_301_at_redirect_limit_closes_stream() {
923        let transport = Arc::new(RedirectTransport {
924            location: "http://redirect.example.com/".to_string(),
925        });
926        let props = RequestProps {
927            url: "http://start.example.com/".parse().unwrap(),
928            headers: HeaderMap::new(),
929            method: "GET".to_string(),
930            body: None,
931            reconnect_opts: ReconnectOptionsBuilder::new(false).build(),
932            max_redirects: 3,
933            dynamic_url: None,
934        };
935        let mut req = ReconnectingRequest::new(transport.clone(), props, None);
936
937        let resp = transport.request(
938            http::Request::builder()
939                .uri("http://start.example.com/")
940                .body(None)
941                .unwrap(),
942        );
943        // Already at max_redirects=3, so the next redirect (would be #4) should fail.
944        req.state = State::Connecting {
945            retry: true,
946            redirect_count: 3,
947            resp,
948        };
949
950        let mut req = pin!(req);
951        let result = timeout(Duration::from_millis(500), req.try_next()).await;
952
953        assert!(matches!(&req.state, State::StreamClosed));
954        assert!(matches!(
955            result,
956            Ok(Err(crate::Error::MaxRedirectLimitReached(3)))
957        ));
958    }
959
960    #[tokio::test]
961    async fn redirect_target_overrides_dynamic_url() {
962        let (_tx, rx) = tokio::sync::watch::channel(
963            "http://dynamic.example.com/".parse::<http::Uri>().unwrap(),
964        );
965        let captured = Arc::new(Mutex::new(Vec::new()));
966        let transport = CapturingTransport {
967            captured_uris: captured.clone(),
968        };
969        let props = props_with_dynamic("http://static.example.com/", rx);
970        let mut req = ReconnectingRequest::new(Arc::new(transport), props, None);
971
972        // Jump straight to FollowingRedirect. The poll loop should parse the
973        // location header and call send_request with the redirect target,
974        // not with the dynamic-uri watch value.
975        req.state = State::FollowingRedirect {
976            header: Some(http::HeaderValue::from_static(
977                "http://redirect.example.com/",
978            )),
979            redirect_count: 1,
980        };
981
982        let mut req = pin!(req);
983        timeout(Duration::from_millis(500), req.try_next())
984            .await
985            .ok();
986
987        let uris = captured.lock().unwrap();
988        assert_eq!(uris.len(), 1);
989        assert_eq!(uris[0].to_string(), "http://redirect.example.com/");
990    }
991
992    #[tokio::test]
993    async fn updated_dynamic_url_is_used_on_next_send_request() {
994        let (tx, rx) =
995            tokio::sync::watch::channel("http://v1.example.com/".parse::<http::Uri>().unwrap());
996        let captured = Arc::new(Mutex::new(Vec::new()));
997        let transport = CapturingTransport {
998            captured_uris: captured.clone(),
999        };
1000        let props = props_with_dynamic("http://static.example.com/", rx);
1001        let req = ReconnectingRequest::new(Arc::new(transport), props, None);
1002
1003        let _ = req.send_request(&req.props.resolve_url());
1004        tx.send("http://v2.example.com/".parse().unwrap()).unwrap();
1005        let _ = req.send_request(&req.props.resolve_url());
1006
1007        let uris = captured.lock().unwrap();
1008        assert_eq!(uris.len(), 2);
1009        assert_eq!(uris[0].to_string(), "http://v1.example.com/");
1010        assert_eq!(uris[1].to_string(), "http://v2.example.com/");
1011    }
1012
1013    #[cfg(feature = "hyper")]
1014    #[tokio::test(flavor = "multi_thread")]
1015    async fn unknown_field_block_does_not_interrupt_stream() {
1016        use crate::{Client, ClientBuilder, Error, Event, ReconnectOptionsBuilder, SSE};
1017        use futures::StreamExt;
1018        use launchdarkly_sdk_transport::HyperTransport;
1019
1020        let mut server = mockito::Server::new_async().await;
1021        let mock = server
1022            .mock("GET", "/")
1023            .with_status(200)
1024            .with_body("unknown: value\n\ndata: hello\n\n")
1025            .expect(1)
1026            .create_async()
1027            .await;
1028
1029        let transport = HyperTransport::new().expect("failed to build transport");
1030        let client = ClientBuilder::for_url(&server.url())
1031            .unwrap()
1032            .reconnect(ReconnectOptionsBuilder::new(false).build())
1033            .build_with_transport(transport);
1034
1035        let items =
1036            tokio::time::timeout(Duration::from_secs(2), client.stream().collect::<Vec<_>>())
1037                .await
1038                .expect("timed out waiting for the stream to end");
1039
1040        assert_eq!(items.len(), 3);
1041        assert!(matches!(items.first(), Some(Ok(SSE::Connected(_)))));
1042        assert_eq!(
1043            items[1].as_ref().unwrap(),
1044            &SSE::Event(Event {
1045                event_type: "message".into(),
1046                data: "hello".into(),
1047                id: None,
1048                retry: None,
1049            })
1050        );
1051        assert!(matches!(items.last(), Some(Err(Error::Eof))));
1052        mock.assert_async().await;
1053    }
1054
1055    // When a parse error happens during streaming and reconnect is
1056    // enabled, the next stream item should be a fresh `Connected` from
1057    // the reconnect, not another error from continuing to drain the
1058    // broken response body.
1059    #[cfg(feature = "hyper")]
1060    #[tokio::test(flavor = "multi_thread")]
1061    async fn parser_error_schedules_reconnect_immediately() {
1062        use crate::{Client, ClientBuilder, Error, ReconnectOptionsBuilder, SSE};
1063        use futures::StreamExt;
1064        use launchdarkly_sdk_transport::HyperTransport;
1065
1066        let mut server = mockito::Server::new_async().await;
1067        let _mock = server
1068            .mock("GET", "/")
1069            .with_status(200)
1070            .with_body(b"\xff\xfe:bad\n\n".as_ref())
1071            .create_async()
1072            .await;
1073
1074        let transport = HyperTransport::new().expect("failed to build transport");
1075        let client = ClientBuilder::for_url(&server.url())
1076            .unwrap()
1077            .reconnect(
1078                ReconnectOptionsBuilder::new(true)
1079                    .delay(Duration::from_millis(10))
1080                    .delay_max(Duration::from_millis(10))
1081                    .retry_initial(true)
1082                    .build(),
1083            )
1084            .build_with_transport(transport);
1085
1086        let mut stream = client.stream();
1087
1088        // Expected order: Connected, parse error, Connected (reconnect).
1089        let mut items = Vec::new();
1090        tokio::time::timeout(Duration::from_secs(2), async {
1091            while items.len() < 3 {
1092                match stream.next().await {
1093                    Some(item) => items.push(item),
1094                    None => break,
1095                }
1096            }
1097        })
1098        .await
1099        .expect("timed out waiting for parse error and reconnect");
1100
1101        assert!(
1102            matches!(items.first(), Some(Ok(SSE::Connected(_)))),
1103            "expected initial Connected, got {:?}",
1104            items.first()
1105        );
1106        assert!(
1107            matches!(items.get(1), Some(Err(Error::InvalidLine(_)))),
1108            "expected InvalidLine error after first connection, got {:?}",
1109            items.get(1)
1110        );
1111        assert!(
1112            matches!(items.get(2), Some(Ok(SSE::Connected(_)))),
1113            "expected reconnect (Connected) immediately after parse error, got {:?}",
1114            items.get(2)
1115        );
1116    }
1117
1118    // With reconnect disabled, a parse error should close the stream so the
1119    // next poll returns `None` rather than continuing to read from a poisoned
1120    // parser or reconnecting via the EOF arm.
1121    #[cfg(feature = "hyper")]
1122    #[tokio::test(flavor = "multi_thread")]
1123    async fn parser_error_closes_stream_when_reconnect_disabled() {
1124        use crate::{Client, ClientBuilder, Error, ReconnectOptionsBuilder, SSE};
1125        use futures::StreamExt;
1126        use launchdarkly_sdk_transport::HyperTransport;
1127
1128        let mut server = mockito::Server::new_async().await;
1129        let _mock = server
1130            .mock("GET", "/")
1131            .with_status(200)
1132            .with_body(b"\xff\xfe:bad\n\n".as_ref())
1133            .create_async()
1134            .await;
1135
1136        let transport = HyperTransport::new().expect("failed to build transport");
1137        let client = ClientBuilder::for_url(&server.url())
1138            .unwrap()
1139            .reconnect(
1140                ReconnectOptionsBuilder::new(false)
1141                    .retry_initial(true)
1142                    .build(),
1143            )
1144            .build_with_transport(transport);
1145
1146        let mut stream = client.stream();
1147
1148        let mut items = Vec::new();
1149        tokio::time::timeout(Duration::from_secs(2), async {
1150            while items.len() < 3 {
1151                match stream.next().await {
1152                    Some(item) => items.push(item),
1153                    None => {
1154                        items.push(Ok(SSE::Comment("__stream_ended__".into())));
1155                        break;
1156                    }
1157                }
1158            }
1159        })
1160        .await
1161        .expect("timed out waiting for stream to close");
1162
1163        assert!(
1164            matches!(items.first(), Some(Ok(SSE::Connected(_)))),
1165            "expected initial Connected, got {:?}",
1166            items.first()
1167        );
1168        assert!(
1169            matches!(items.get(1), Some(Err(Error::InvalidLine(_)))),
1170            "expected InvalidLine error, got {:?}",
1171            items.get(1)
1172        );
1173        assert!(
1174            matches!(
1175                items.get(2),
1176                Some(Ok(SSE::Comment(s))) if s == "__stream_ended__"
1177            ),
1178            "expected stream to end (None) after parse error with reconnect disabled, got {:?}",
1179            items.get(2)
1180        );
1181    }
1182
1183    // With reconnect disabled, a clean end-of-body should close the stream
1184    // rather than scheduling a reconnect.
1185    #[cfg(feature = "hyper")]
1186    #[tokio::test(flavor = "multi_thread")]
1187    async fn eof_closes_stream_when_reconnect_disabled() {
1188        use crate::{Client, ClientBuilder, Error, ReconnectOptionsBuilder, SSE};
1189        use futures::StreamExt;
1190        use launchdarkly_sdk_transport::HyperTransport;
1191
1192        let mut server = mockito::Server::new_async().await;
1193        let _mock = server
1194            .mock("GET", "/")
1195            .with_status(200)
1196            .with_body("event: hello\ndata: world\n\n")
1197            .create_async()
1198            .await;
1199
1200        let transport = HyperTransport::new().expect("failed to build transport");
1201        let client = ClientBuilder::for_url(&server.url())
1202            .unwrap()
1203            .reconnect(
1204                ReconnectOptionsBuilder::new(false)
1205                    .retry_initial(true)
1206                    .build(),
1207            )
1208            .build_with_transport(transport);
1209
1210        let mut stream = client.stream();
1211
1212        let mut items: Vec<Option<crate::Result<SSE>>> = Vec::new();
1213        tokio::time::timeout(Duration::from_secs(2), async {
1214            for _ in 0..4 {
1215                let item = stream.next().await;
1216                let is_terminal = item.is_none();
1217                items.push(item);
1218                if is_terminal {
1219                    break;
1220                }
1221            }
1222        })
1223        .await
1224        .expect("timed out waiting for stream to close");
1225
1226        assert!(
1227            matches!(items.first(), Some(Some(Ok(SSE::Connected(_))))),
1228            "expected initial Connected, got {:?}",
1229            items.first()
1230        );
1231        assert!(
1232            matches!(items.get(1), Some(Some(Ok(SSE::Event(e)))) if e.event_type == "hello"),
1233            "expected hello event, got {:?}",
1234            items.get(1)
1235        );
1236        assert!(
1237            matches!(items.get(2), Some(Some(Err(Error::Eof)))),
1238            "expected Eof error after body ends, got {:?}",
1239            items.get(2)
1240        );
1241        assert!(
1242            matches!(items.get(3), Some(None)),
1243            "expected stream to end (None) after EOF with reconnect disabled, got {:?}",
1244            items.get(3)
1245        );
1246    }
1247}