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    // When a parse error happens during streaming and reconnect is
1014    // enabled, the next stream item should be a fresh `Connected` from
1015    // the reconnect, not another error from continuing to drain the
1016    // broken response body.
1017    #[cfg(feature = "hyper")]
1018    #[tokio::test(flavor = "multi_thread")]
1019    async fn parser_error_schedules_reconnect_immediately() {
1020        use crate::{Client, ClientBuilder, Error, ReconnectOptionsBuilder, SSE};
1021        use futures::StreamExt;
1022        use launchdarkly_sdk_transport::HyperTransport;
1023
1024        let mut server = mockito::Server::new_async().await;
1025        let _mock = server
1026            .mock("GET", "/")
1027            .with_status(200)
1028            .with_body(b"\xff\xfe:bad\n\n".as_ref())
1029            .create_async()
1030            .await;
1031
1032        let transport = HyperTransport::new().expect("failed to build transport");
1033        let client = ClientBuilder::for_url(&server.url())
1034            .unwrap()
1035            .reconnect(
1036                ReconnectOptionsBuilder::new(true)
1037                    .delay(Duration::from_millis(10))
1038                    .delay_max(Duration::from_millis(10))
1039                    .retry_initial(true)
1040                    .build(),
1041            )
1042            .build_with_transport(transport);
1043
1044        let mut stream = client.stream();
1045
1046        // Expected order: Connected, parse error, Connected (reconnect).
1047        let mut items = Vec::new();
1048        tokio::time::timeout(Duration::from_secs(2), async {
1049            while items.len() < 3 {
1050                match stream.next().await {
1051                    Some(item) => items.push(item),
1052                    None => break,
1053                }
1054            }
1055        })
1056        .await
1057        .expect("timed out waiting for parse error and reconnect");
1058
1059        assert!(
1060            matches!(items.first(), Some(Ok(SSE::Connected(_)))),
1061            "expected initial Connected, got {:?}",
1062            items.first()
1063        );
1064        assert!(
1065            matches!(items.get(1), Some(Err(Error::InvalidLine(_)))),
1066            "expected InvalidLine error after first connection, got {:?}",
1067            items.get(1)
1068        );
1069        assert!(
1070            matches!(items.get(2), Some(Ok(SSE::Connected(_)))),
1071            "expected reconnect (Connected) immediately after parse error, got {:?}",
1072            items.get(2)
1073        );
1074    }
1075
1076    // With reconnect disabled, a parse error should close the stream so the
1077    // next poll returns `None` rather than continuing to read from a poisoned
1078    // parser or reconnecting via the EOF arm.
1079    #[cfg(feature = "hyper")]
1080    #[tokio::test(flavor = "multi_thread")]
1081    async fn parser_error_closes_stream_when_reconnect_disabled() {
1082        use crate::{Client, ClientBuilder, Error, ReconnectOptionsBuilder, SSE};
1083        use futures::StreamExt;
1084        use launchdarkly_sdk_transport::HyperTransport;
1085
1086        let mut server = mockito::Server::new_async().await;
1087        let _mock = server
1088            .mock("GET", "/")
1089            .with_status(200)
1090            .with_body(b"\xff\xfe:bad\n\n".as_ref())
1091            .create_async()
1092            .await;
1093
1094        let transport = HyperTransport::new().expect("failed to build transport");
1095        let client = ClientBuilder::for_url(&server.url())
1096            .unwrap()
1097            .reconnect(
1098                ReconnectOptionsBuilder::new(false)
1099                    .retry_initial(true)
1100                    .build(),
1101            )
1102            .build_with_transport(transport);
1103
1104        let mut stream = client.stream();
1105
1106        let mut items = Vec::new();
1107        tokio::time::timeout(Duration::from_secs(2), async {
1108            while items.len() < 3 {
1109                match stream.next().await {
1110                    Some(item) => items.push(item),
1111                    None => {
1112                        items.push(Ok(SSE::Comment("__stream_ended__".into())));
1113                        break;
1114                    }
1115                }
1116            }
1117        })
1118        .await
1119        .expect("timed out waiting for stream to close");
1120
1121        assert!(
1122            matches!(items.first(), Some(Ok(SSE::Connected(_)))),
1123            "expected initial Connected, got {:?}",
1124            items.first()
1125        );
1126        assert!(
1127            matches!(items.get(1), Some(Err(Error::InvalidLine(_)))),
1128            "expected InvalidLine error, got {:?}",
1129            items.get(1)
1130        );
1131        assert!(
1132            matches!(
1133                items.get(2),
1134                Some(Ok(SSE::Comment(s))) if s == "__stream_ended__"
1135            ),
1136            "expected stream to end (None) after parse error with reconnect disabled, got {:?}",
1137            items.get(2)
1138        );
1139    }
1140
1141    // With reconnect disabled, a clean end-of-body should close the stream
1142    // rather than scheduling a reconnect.
1143    #[cfg(feature = "hyper")]
1144    #[tokio::test(flavor = "multi_thread")]
1145    async fn eof_closes_stream_when_reconnect_disabled() {
1146        use crate::{Client, ClientBuilder, Error, ReconnectOptionsBuilder, SSE};
1147        use futures::StreamExt;
1148        use launchdarkly_sdk_transport::HyperTransport;
1149
1150        let mut server = mockito::Server::new_async().await;
1151        let _mock = server
1152            .mock("GET", "/")
1153            .with_status(200)
1154            .with_body("event: hello\ndata: world\n\n")
1155            .create_async()
1156            .await;
1157
1158        let transport = HyperTransport::new().expect("failed to build transport");
1159        let client = ClientBuilder::for_url(&server.url())
1160            .unwrap()
1161            .reconnect(
1162                ReconnectOptionsBuilder::new(false)
1163                    .retry_initial(true)
1164                    .build(),
1165            )
1166            .build_with_transport(transport);
1167
1168        let mut stream = client.stream();
1169
1170        let mut items: Vec<Option<crate::Result<SSE>>> = Vec::new();
1171        tokio::time::timeout(Duration::from_secs(2), async {
1172            for _ in 0..4 {
1173                let item = stream.next().await;
1174                let is_terminal = item.is_none();
1175                items.push(item);
1176                if is_terminal {
1177                    break;
1178                }
1179            }
1180        })
1181        .await
1182        .expect("timed out waiting for stream to close");
1183
1184        assert!(
1185            matches!(items.first(), Some(Some(Ok(SSE::Connected(_))))),
1186            "expected initial Connected, got {:?}",
1187            items.first()
1188        );
1189        assert!(
1190            matches!(items.get(1), Some(Some(Ok(SSE::Event(e)))) if e.event_type == "hello"),
1191            "expected hello event, got {:?}",
1192            items.get(1)
1193        );
1194        assert!(
1195            matches!(items.get(2), Some(Some(Err(Error::Eof)))),
1196            "expected Eof error after body ends, got {:?}",
1197            items.get(2)
1198        );
1199        assert!(
1200            matches!(items.get(3), Some(None)),
1201            "expected stream to end (None) after EOF with reconnect disabled, got {:?}",
1202            items.get(3)
1203        );
1204    }
1205}