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
38pub type BoxStream<T> = Pin<boxed::Box<dyn Stream<Item = T> + Send + Sync>>;
40
41pub trait Client: Send + Sync + private::Sealed {
44 fn stream(&self) -> BoxStream<Result<SSE>>;
45}
46
47pub const DEFAULT_REDIRECT_LIMIT: u32 = 16;
55
56pub 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 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 pub fn dynamic_url(mut self, uri: watch::Receiver<Uri>) -> ClientBuilder {
94 self.dynamic_url = Some(uri);
95 self
96 }
97
98 pub fn method(mut self, method: String) -> ClientBuilder {
100 self.method = method;
101 self
102 }
103
104 pub fn body(mut self, body: String) -> ClientBuilder {
106 self.body = Some(body);
107 self
108 }
109
110 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 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 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 pub fn reconnect(mut self, opts: ReconnectOptions) -> ClientBuilder {
142 self.reconnect_opts = opts;
143 self
144 }
145
146 pub fn redirect_limit(mut self, limit: u32) -> ClientBuilder {
150 self.max_redirects = Some(limit);
151 self
152 }
153
154 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
211struct 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 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)] #[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 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 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 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 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 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 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 #[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 Box::pin(async {
707 Err(TransportError::new(std::io::Error::new(
708 std::io::ErrorKind::ConnectionRefused,
709 "connection refused",
710 )))
711 })
712 } else {
713 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 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 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 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")]
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 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 #[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 #[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}