1use std::time::Duration;
5
6use async_stream::try_stream;
7use bytes::Bytes;
8use futures::Stream;
9use reqwest::{RequestBuilder, Response};
10use serde::de::DeserializeOwned;
11use tokio::time::{interval, sleep, timeout, MissedTickBehavior};
12use tracing::{debug_span, error, info, warn, Instrument};
13
14use super::{errors::FeedError, failures::FailureTracker, publisher::Publisher};
15
16#[derive(Clone, Debug)]
21pub struct HttpFeedConfig {
22 pub poll_interval: Duration,
26 pub request_timeout: Duration,
28 pub max_backoff_exp: u32,
34 pub max_consecutive_failures: Option<u32>,
37 pub max_snapshot_age: Option<Duration>,
45}
46
47#[expect(
51 dead_code,
52 reason = "nothing in this crate implements a book feed yet; the layer is exercised by its own tests"
53)]
54pub(crate) fn default_http_feed_config() -> HttpFeedConfig {
55 HttpFeedConfig {
56 poll_interval: Duration::from_secs(5),
57 request_timeout: Duration::from_secs(10),
58 max_backoff_exp: 3,
59 max_consecutive_failures: None,
60 max_snapshot_age: Some(Duration::from_secs(30)),
61 }
62}
63
64pub(crate) async fn run_http_poll_feed<S: HttpSource>(
77 config: HttpFeedConfig,
78 publisher: Publisher<S::Snapshot>,
79 source: S,
80) -> Result<(), FeedError> {
81 let HttpFeedConfig {
85 poll_interval,
86 request_timeout,
87 max_backoff_exp,
88 max_consecutive_failures,
89 max_snapshot_age,
90 } = config;
91
92 if poll_interval.is_zero() {
93 return Err(FeedError::InvalidInput("poll_interval must not be zero".to_string()));
94 }
95 if let Some(age) = max_snapshot_age {
96 if age.is_zero() {
97 return Err(FeedError::InvalidInput("max_snapshot_age must not be zero".to_string()));
98 }
99 if age <= poll_interval {
100 warn!(
103 max_snapshot_age = ?age,
104 ?poll_interval,
105 "max_snapshot_age does not exceed poll_interval: this feed will serve each \
106 snapshot for max_snapshot_age and nothing until the next poll"
107 );
108 }
109 }
110
111 let snapshots = polled_snapshots(
112 poll_interval,
113 request_timeout,
114 max_consecutive_failures,
115 max_backoff_exp,
116 source,
117 );
118 publisher
119 .publishing(max_snapshot_age, snapshots)
120 .await
121}
122
123pub(crate) trait HttpSource: Send + Sync {
126 type Snapshot: Send + Sync + 'static;
127
128 async fn fetch(&self) -> Result<Self::Snapshot, FeedError>;
129}
130
131fn polled_snapshots<S: HttpSource>(
141 poll_interval: Duration,
142 request_timeout: Duration,
143 max_consecutive_failures: Option<u32>,
144 max_backoff_exp: u32,
145 source: S,
146) -> impl Stream<Item = Result<S::Snapshot, FeedError>> {
147 let mut failures =
150 FailureTracker::new(max_consecutive_failures, poll_interval, max_backoff_exp);
151
152 let mut ticker = interval(poll_interval);
153 ticker.set_missed_tick_behavior(MissedTickBehavior::Delay);
159
160 let mut backoff = None;
163
164 info!(?poll_interval, "starting polling");
165
166 try_stream! {
167 loop {
168 match backoff.take() {
171 Some(backoff) => {
172 sleep(backoff).await;
173 ticker.reset();
177 }
178 None => {
179 ticker.tick().await;
180 }
181 }
182 let polled = timeout(request_timeout, source.fetch())
183 .instrument(debug_span!("poll"))
184 .await;
185
186 let error = match polled {
187 Ok(Ok(snapshot)) => {
188 let cleared = failures.record_success();
189 if cleared > 0 {
190 info!(after_failures = cleared, "poll succeeded again");
191 }
192 yield snapshot;
193 continue;
194 }
195 Ok(Err(e)) if e.is_fatal() => Err(e)?,
196 Ok(Err(e)) => e,
197 Err(_) => FeedError::Connection(format!(
198 "poll timed out after {request_timeout:?}"
199 )),
200 };
201
202 let retry = failures
203 .record_failure()
204 .inspect(|retry| {
205 warn!(
206 consecutive = retry.consecutive,
207 backoff = ?retry.backoff,
208 %error,
209 "poll failed, backing off"
210 )
211 })
212 .map_err(|consecutive| {
217 error!(consecutive, "giving up on the feed");
218 error
219 })?;
220 backoff = Some(retry.backoff);
221 }
222 }
223}
224
225pub(crate) async fn fetch_bytes(request: RequestBuilder, what: &str) -> Result<Bytes, FeedError> {
232 let response = request
233 .send()
234 .await
235 .map_err(|e| FeedError::Connection(format!("Failed to fetch {what}: {e}")))?;
236 read_bytes(response, what).await
237}
238
239async fn read_bytes(response: Response, what: &str) -> Result<Bytes, FeedError> {
242 let status = response.status();
245 let body = response
246 .bytes()
247 .await
248 .map_err(|e| FeedError::Connection(format!("Failed to read {what} response: {e}")))?;
249
250 if !status.is_success() {
251 return Err(FeedError::Connection(format!(
252 "{what} HTTP error {status}: {}",
253 String::from_utf8_lossy(&body)
254 )));
255 }
256 Ok(body)
257}
258
259pub(crate) async fn fetch_json<T: DeserializeOwned>(
263 request: RequestBuilder,
264 what: &str,
265) -> Result<T, FeedError> {
266 parse_json(&fetch_bytes(request, what).await?, what)
267}
268
269#[expect(
272 dead_code,
273 reason = "nothing in this crate implements a book feed yet; the layer is exercised by its own tests"
274)]
275pub(crate) async fn read_json<T: DeserializeOwned>(
276 response: Response,
277 what: &str,
278) -> Result<T, FeedError> {
279 parse_json(&read_bytes(response, what).await?, what)
280}
281
282fn parse_json<T: DeserializeOwned>(body: &[u8], what: &str) -> Result<T, FeedError> {
283 serde_json::from_slice(body)
284 .map_err(|e| FeedError::Parsing(format!("Failed to parse {what} response: {e}")))
285}
286
287#[cfg(test)]
290pub(crate) mod test_support {
291 use std::net::SocketAddr;
292
293 use tokio::{
294 io::{AsyncBufReadExt, AsyncWriteExt, BufReader},
295 net::TcpListener,
296 };
297
298 pub struct MockHttpServer {
300 address: SocketAddr,
301 }
302
303 impl MockHttpServer {
304 pub fn url(&self) -> String {
305 format!("http://{}", self.address)
306 }
307 }
308
309 pub async fn spawn_http_server(
310 respond: impl Fn() -> (&'static str, String) + Send + 'static,
311 ) -> MockHttpServer {
312 let listener = TcpListener::bind("127.0.0.1:0")
313 .await
314 .unwrap();
315 let address = listener.local_addr().unwrap();
316 tokio::spawn(async move {
317 while let Ok((stream, _)) = listener.accept().await {
318 let mut reader = BufReader::new(stream);
319 let mut request_line = String::new();
320 if reader
321 .read_line(&mut request_line)
322 .await
323 .is_err()
324 {
325 continue;
326 }
327 let (status, body) = respond();
328 let payload = format!(
329 "HTTP/1.1 {status}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
330 body.len()
331 );
332 let mut stream = reader.into_inner();
333 let _ = stream
334 .write_all(payload.as_bytes())
335 .await;
336 let _ = stream.shutdown().await;
337 }
338 });
339 MockHttpServer { address }
340 }
341}
342
343#[cfg(test)]
344mod tests {
345 use std::{
346 future::Future,
347 sync::{
348 atomic::{AtomicBool, AtomicU32, Ordering},
349 Arc,
350 },
351 };
352
353 use futures::StreamExt as _;
354 use tokio::sync::watch;
355 use tokio_stream::wrappers::WatchStream;
356
357 use super::{super::expect_to_finish, *};
358
359 impl<F, Fut, T> HttpSource for F
361 where
362 F: Fn() -> Fut + Send + Sync,
363 Fut: Future<Output = Result<T, FeedError>> + Send,
364 T: Send + Sync + 'static,
365 {
366 type Snapshot = T;
367
368 async fn fetch(&self) -> Result<T, FeedError> {
369 self().await
370 }
371 }
372
373 fn snapshot() -> u32 {
375 1
376 }
377
378 fn config(max_consecutive_failures: Option<u32>) -> HttpFeedConfig {
379 HttpFeedConfig {
380 poll_interval: Duration::from_millis(5),
381 request_timeout: Duration::from_millis(50),
382 max_backoff_exp: 0,
385 max_consecutive_failures,
386 max_snapshot_age: None,
387 }
388 }
389
390 #[tokio::test]
391 async fn publishes_a_snapshot_and_success_resets_failures() {
392 let polls = Arc::new(AtomicU32::new(0));
395 let polls_clone = Arc::clone(&polls);
396 let (publisher, mut rx) = Publisher::channel();
397 let feed = run_http_poll_feed(config(Some(2)), publisher, move || {
398 let n = polls_clone.fetch_add(1, Ordering::SeqCst);
399 async move {
400 if n == 0 {
401 Err(FeedError::Connection("boom".to_string()))
402 } else {
403 Ok(snapshot())
404 }
405 }
406 });
407 let feed = tokio::spawn(feed);
408
409 expect_to_finish("no snapshot published", async {
410 loop {
411 rx.changed().await.expect("feed ended");
412 if rx.borrow_and_update().is_some() {
413 return;
414 }
415 }
416 })
417 .await;
418
419 tokio::time::sleep(Duration::from_millis(50)).await;
421 assert!(!feed.is_finished());
422 }
423
424 #[tokio::test]
425 async fn withdraws_a_snapshot_that_goes_unrefreshed_and_restores_on_next_success() {
426 let polls = Arc::new(AtomicU32::new(0));
430 let polls_clone = Arc::clone(&polls);
431 let recovered = Arc::new(AtomicBool::new(false));
432 let recovered_clone = Arc::clone(&recovered);
433 let (publisher, mut rx) = Publisher::channel();
434 let config =
435 HttpFeedConfig { max_snapshot_age: Some(Duration::from_millis(20)), ..config(None) };
436 let feed = run_http_poll_feed(config, publisher, move || {
437 let n = polls_clone.fetch_add(1, Ordering::SeqCst);
438 let recovered = recovered_clone.load(Ordering::SeqCst);
439 async move {
440 if n == 0 || recovered {
441 Ok(snapshot())
442 } else {
443 Err(FeedError::Connection("boom".to_string()))
444 }
445 }
446 });
447 let _feed = tokio::spawn(feed);
448
449 let wait_for = |rx: &mut watch::Receiver<Option<u32>>, want_some: bool| {
450 let mut rx = rx.clone();
451 async move {
452 expect_to_finish("watch did not reach the expected state", async {
453 loop {
454 if rx.borrow_and_update().is_some() == want_some {
455 return;
456 }
457 rx.changed().await.expect("feed ended");
458 }
459 })
460 .await;
461 }
462 };
463 wait_for(&mut rx, true).await;
464 wait_for(&mut rx, false).await;
465 recovered.store(true, Ordering::SeqCst);
466 wait_for(&mut rx, true).await;
467 }
468
469 #[tokio::test]
470 async fn rejects_zero_poll_interval() {
471 let (publisher, _rx) = Publisher::<u32>::channel();
472 let config = HttpFeedConfig { poll_interval: Duration::ZERO, ..config(None) };
473 let result = run_http_poll_feed(config, publisher, || async { Ok(snapshot()) }).await;
474 assert!(matches!(result, Err(FeedError::InvalidInput(_))));
475 }
476
477 #[tokio::test]
478 async fn rejects_a_zero_max_snapshot_age() {
479 let (publisher, _rx) = Publisher::<u32>::channel();
480 let config = HttpFeedConfig { max_snapshot_age: Some(Duration::ZERO), ..config(None) };
481 let result = run_http_poll_feed(config, publisher, || async { Ok(snapshot()) }).await;
482 assert!(matches!(result, Err(FeedError::InvalidInput(_))));
483 }
484
485 #[tokio::test(start_paused = true)]
490 async fn a_max_snapshot_age_below_the_poll_interval_is_a_duty_cycle() {
491 let (publisher, rx) = Publisher::<u32>::channel();
492 let mut changes = WatchStream::from_changes(rx);
493 let config = HttpFeedConfig {
494 poll_interval: Duration::from_millis(100),
495 max_snapshot_age: Some(Duration::from_millis(20)),
496 ..config(None)
497 };
498 let _feed =
499 tokio::spawn(run_http_poll_feed(config, publisher, || async { Ok(snapshot()) }));
500
501 let start = tokio::time::Instant::now();
502 let mut seen = Vec::new();
503 for _ in 0..4 {
504 let change = expect_to_finish("the feed stopped changing", changes.next()).await;
505 seen.push((change.expect("the feed is still running"), start.elapsed()));
506 }
507
508 assert_eq!(
509 seen,
510 [
511 (Some(snapshot()), Duration::ZERO),
514 (None, Duration::from_millis(20)),
515 (Some(snapshot()), Duration::from_millis(100)),
516 (None, Duration::from_millis(120)),
517 ]
518 );
519 }
520
521 #[tokio::test]
522 async fn gives_up_after_max_consecutive_failed_polls() {
523 let (publisher, _rx) = Publisher::<u32>::channel();
526 let feed = run_http_poll_feed(config(Some(3)), publisher, || async {
527 Err(FeedError::Parsing("not a snapshot".to_string()))
528 });
529 let feed = tokio::spawn(feed);
530
531 let result = expect_to_finish("feed did not give up", feed)
532 .await
533 .unwrap();
534 assert!(matches!(result, Err(FeedError::Parsing(msg)) if msg == "not a snapshot"));
535 }
536
537 #[tokio::test]
538 async fn gives_up_on_a_fatal_poll_error_despite_an_unlimited_failure_budget() {
539 let (publisher, mut rx) = Publisher::channel();
542 let polls = Arc::new(AtomicU32::new(0));
543 let polls_clone = Arc::clone(&polls);
544 let feed = run_http_poll_feed(config(None), publisher, move || {
545 let n = polls_clone.fetch_add(1, Ordering::SeqCst);
546 async move {
547 if n == 0 {
548 Ok(snapshot())
549 } else {
550 Err(FeedError::Fatal("api key rejected".to_string()))
551 }
552 }
553 });
554 let feed = tokio::spawn(feed);
555
556 let result = expect_to_finish("feed did not give up", feed)
557 .await
558 .unwrap();
559
560 assert!(matches!(result, Err(FeedError::Fatal(_))));
561 assert_eq!(polls.load(Ordering::SeqCst), 2, "the feed polled again after the fatal");
562 assert!(rx.borrow_and_update().is_none(), "the snapshot must be withdrawn");
563 }
564
565 #[tokio::test]
566 async fn withdraws_the_published_snapshot_when_it_gives_up() {
567 let polls = Arc::new(AtomicU32::new(0));
570 let polls_clone = Arc::clone(&polls);
571 let (publisher, mut rx) = Publisher::channel();
572 let feed = run_http_poll_feed(config(Some(2)), publisher, move || {
573 let n = polls_clone.fetch_add(1, Ordering::SeqCst);
574 async move {
575 if n == 0 {
576 Ok(snapshot())
577 } else {
578 Err(FeedError::Connection("boom".to_string()))
579 }
580 }
581 });
582 let feed = tokio::spawn(feed);
583
584 expect_to_finish("no snapshot published", async {
585 loop {
586 rx.changed()
587 .await
588 .expect("feed ended before publishing");
589 if rx.borrow_and_update().is_some() {
590 return;
591 }
592 }
593 })
594 .await;
595
596 let result = expect_to_finish("feed did not give up", feed)
597 .await
598 .unwrap();
599 assert!(matches!(result, Err(FeedError::Connection(_))));
600 assert!(rx.borrow().is_none(), "the snapshot must be withdrawn when the feed gives up");
601 }
602
603 #[tokio::test]
604 async fn withdraws_the_published_snapshot_when_the_feed_is_dropped() {
605 let (publisher, mut rx) = Publisher::channel();
608 let feed = run_http_poll_feed(config(None), publisher, || async { Ok(snapshot()) });
609 let feed = tokio::spawn(feed);
610
611 expect_to_finish("no snapshot published", async {
612 loop {
613 rx.changed()
614 .await
615 .expect("feed ended before publishing");
616 if rx.borrow_and_update().is_some() {
617 return;
618 }
619 }
620 })
621 .await;
622
623 feed.abort();
624 expect_to_finish("the snapshot was not withdrawn", rx.changed())
625 .await
626 .expect("the withdrawal must reach the receiver before the sender drops");
627 assert!(rx.borrow().is_none());
628 }
629
630 #[tokio::test(start_paused = true)]
631 async fn failing_polls_stretch_the_polling_cadence() {
632 let attempts = Arc::new(std::sync::Mutex::new(Vec::new()));
636 let attempts_clone = Arc::clone(&attempts);
637 let start = tokio::time::Instant::now();
638 let (publisher, _rx) = Publisher::<u32>::channel();
639 let feed = tokio::spawn(run_http_poll_feed(
640 HttpFeedConfig {
641 poll_interval: Duration::from_millis(50),
642 max_backoff_exp: 3,
643 ..config(None)
644 },
645 publisher,
646 move || {
647 attempts_clone
648 .lock()
649 .unwrap()
650 .push(start.elapsed().as_millis());
651 async { Err(FeedError::Connection("boom".to_string())) }
652 },
653 ));
654
655 tokio::time::sleep(Duration::from_secs(1)).await;
656 feed.abort();
657
658 assert_eq!(*attempts.lock().unwrap(), vec![0, 50, 150, 350, 750]);
659 }
660
661 #[tokio::test(start_paused = true)]
662 async fn a_recovered_poll_resumes_the_cadence_instead_of_catching_up() {
663 let attempts = Arc::new(std::sync::Mutex::new(Vec::new()));
667 let attempts_clone = Arc::clone(&attempts);
668 let start = tokio::time::Instant::now();
669 let (publisher, _rx) = Publisher::<u32>::channel();
670 let feed = tokio::spawn(run_http_poll_feed(
671 HttpFeedConfig {
672 poll_interval: Duration::from_millis(50),
673 max_backoff_exp: 3,
674 ..config(None)
675 },
676 publisher,
677 move || {
678 let mut attempts = attempts_clone.lock().unwrap();
679 attempts.push(start.elapsed().as_millis());
680 let failing = attempts.len() <= 2;
681 async move {
682 if failing {
683 Err(FeedError::Connection("boom".to_string()))
684 } else {
685 Ok(snapshot())
686 }
687 }
688 },
689 ));
690
691 tokio::time::sleep(Duration::from_millis(300)).await;
692 feed.abort();
693
694 assert_eq!(*attempts.lock().unwrap(), vec![0, 50, 150, 200, 250]);
695 }
696
697 #[tokio::test(start_paused = true)]
698 async fn a_slow_poll_never_shortens_the_gap_to_the_next_one() {
699 let attempts = Arc::new(std::sync::Mutex::new(Vec::new()));
704 let attempts_clone = Arc::clone(&attempts);
705 let start = tokio::time::Instant::now();
706 let (publisher, _rx) = Publisher::<u32>::channel();
707 let feed = tokio::spawn(run_http_poll_feed(
708 HttpFeedConfig {
709 poll_interval: Duration::from_millis(50),
710 request_timeout: Duration::from_millis(500),
711 ..config(None)
712 },
713 publisher,
714 move || {
715 let mut attempts = attempts_clone.lock().unwrap();
716 attempts.push(start.elapsed().as_millis());
717 let slow = attempts.len() == 3;
718 async move {
719 if slow {
720 tokio::time::sleep(Duration::from_millis(120)).await;
721 }
722 Ok(snapshot())
723 }
724 },
725 ));
726
727 tokio::time::sleep(Duration::from_millis(350)).await;
728 feed.abort();
729
730 assert_eq!(*attempts.lock().unwrap(), vec![0, 50, 100, 220, 270, 320]);
731 }
732
733 #[tokio::test]
734 async fn retries_forever_without_limit() {
735 let polls = Arc::new(AtomicU32::new(0));
736 let polls_clone = Arc::clone(&polls);
737 let (publisher, _rx) = Publisher::<u32>::channel();
738 let feed = run_http_poll_feed(config(None), publisher, move || {
739 polls_clone.fetch_add(1, Ordering::SeqCst);
740 async { Err(FeedError::Connection("boom".to_string())) }
741 });
742 let feed = tokio::spawn(feed);
743
744 tokio::time::sleep(Duration::from_millis(100)).await;
745 assert!(polls.load(Ordering::SeqCst) > 5, "should keep polling");
746 assert!(!feed.is_finished());
747 }
748
749 #[tokio::test]
750 async fn hung_poll_counts_as_failure() {
751 let (publisher, _rx) = Publisher::<u32>::channel();
752 let feed = run_http_poll_feed(
753 HttpFeedConfig { request_timeout: Duration::from_millis(5), ..config(Some(2)) },
754 publisher,
755 || async {
756 tokio::time::sleep(Duration::from_secs(3600)).await;
757 Ok(snapshot())
758 },
759 );
760 let feed = tokio::spawn(feed);
761
762 let result = expect_to_finish("feed did not give up", feed)
763 .await
764 .unwrap();
765 assert!(matches!(result, Err(FeedError::Connection(msg)) if msg.contains("timed out")));
766 }
767
768 mod fetch_json {
769
770 use rstest::rstest;
771 use serde::Deserialize;
772
773 use super::{test_support::spawn_http_server, *};
774
775 #[derive(Debug, Deserialize, PartialEq)]
776 struct Payload {
777 value: u32,
778 }
779
780 #[tokio::test]
781 async fn parses_a_successful_json_body() {
782 let server = spawn_http_server(|| ("200 OK", r#"{"value":7}"#.to_string())).await;
783
784 let payload: Payload =
785 fetch_json(reqwest::Client::new().get(format!("{}/x", server.url())), "thing")
786 .await
787 .unwrap();
788
789 assert_eq!(payload, Payload { value: 7 });
790 }
791
792 #[rstest]
796 #[case::non_success_status_with_body("503 Service Unavailable", "maintenance", false)]
797 #[case::unparseable_success_body("200 OK", "<html>", true)]
798 #[tokio::test]
799 async fn classifies_status_and_body_failures(
800 #[case] status: &'static str,
801 #[case] body: &'static str,
802 #[case] expect_parsing_error: bool,
803 ) {
804 let server = spawn_http_server(move || (status, body.to_string())).await;
805
806 let result: Result<Payload, FeedError> =
807 fetch_json(reqwest::Client::new().get(format!("{}/x", server.url())), "thing")
808 .await;
809
810 match result {
811 Err(FeedError::Parsing(msg)) if expect_parsing_error => {
812 assert!(msg.contains("thing"), "{msg}")
813 }
814 Err(FeedError::Connection(msg)) if !expect_parsing_error => {
815 assert!(
816 msg.contains("thing") && msg.contains("503") && msg.contains(body),
817 "{msg}"
818 )
819 }
820 other => panic!("unexpected result: {other:?}"),
821 }
822 }
823
824 #[tokio::test]
825 async fn refused_connection_is_a_connection_error() {
826 let address = tokio::net::TcpListener::bind("127.0.0.1:0")
828 .await
829 .unwrap()
830 .local_addr()
831 .unwrap();
832
833 let result: Result<Payload, FeedError> =
834 fetch_json(reqwest::Client::new().get(format!("http://{address}/x")), "thing")
835 .await;
836
837 assert!(matches!(result, Err(FeedError::Connection(msg)) if msg.contains("thing")));
838 }
839 }
840}