use std::time::Duration;
use async_stream::try_stream;
use bytes::Bytes;
use futures::Stream;
use reqwest::{RequestBuilder, Response};
use serde::de::DeserializeOwned;
use tokio::time::{interval, sleep, timeout, MissedTickBehavior};
use tracing::{debug_span, error, info, warn, Instrument};
use super::{errors::FeedError, failures::FailureTracker, publisher::Publisher};
#[derive(Clone, Debug)]
pub struct HttpFeedConfig {
pub poll_interval: Duration,
pub request_timeout: Duration,
pub max_backoff_exp: u32,
pub max_consecutive_failures: Option<u32>,
pub max_snapshot_age: Option<Duration>,
}
#[expect(
dead_code,
reason = "nothing in this crate implements a book feed yet; the layer is exercised by its own tests"
)]
pub(crate) fn default_http_feed_config() -> HttpFeedConfig {
HttpFeedConfig {
poll_interval: Duration::from_secs(5),
request_timeout: Duration::from_secs(10),
max_backoff_exp: 3,
max_consecutive_failures: None,
max_snapshot_age: Some(Duration::from_secs(30)),
}
}
pub(crate) async fn run_http_poll_feed<S: HttpSource>(
config: HttpFeedConfig,
publisher: Publisher<S::Snapshot>,
source: S,
) -> Result<(), FeedError> {
let HttpFeedConfig {
poll_interval,
request_timeout,
max_backoff_exp,
max_consecutive_failures,
max_snapshot_age,
} = config;
if poll_interval.is_zero() {
return Err(FeedError::InvalidInput("poll_interval must not be zero".to_string()));
}
if let Some(age) = max_snapshot_age {
if age.is_zero() {
return Err(FeedError::InvalidInput("max_snapshot_age must not be zero".to_string()));
}
if age <= poll_interval {
warn!(
max_snapshot_age = ?age,
?poll_interval,
"max_snapshot_age does not exceed poll_interval: this feed will serve each \
snapshot for max_snapshot_age and nothing until the next poll"
);
}
}
let snapshots = polled_snapshots(
poll_interval,
request_timeout,
max_consecutive_failures,
max_backoff_exp,
source,
);
publisher
.publishing(max_snapshot_age, snapshots)
.await
}
pub(crate) trait HttpSource: Send + Sync {
type Snapshot: Send + Sync + 'static;
async fn fetch(&self) -> Result<Self::Snapshot, FeedError>;
}
fn polled_snapshots<S: HttpSource>(
poll_interval: Duration,
request_timeout: Duration,
max_consecutive_failures: Option<u32>,
max_backoff_exp: u32,
source: S,
) -> impl Stream<Item = Result<S::Snapshot, FeedError>> {
let mut failures =
FailureTracker::new(max_consecutive_failures, poll_interval, max_backoff_exp);
let mut ticker = interval(poll_interval);
ticker.set_missed_tick_behavior(MissedTickBehavior::Delay);
let mut backoff = None;
info!(?poll_interval, "starting polling");
try_stream! {
loop {
match backoff.take() {
Some(backoff) => {
sleep(backoff).await;
ticker.reset();
}
None => {
ticker.tick().await;
}
}
let polled = timeout(request_timeout, source.fetch())
.instrument(debug_span!("poll"))
.await;
let error = match polled {
Ok(Ok(snapshot)) => {
let cleared = failures.record_success();
if cleared > 0 {
info!(after_failures = cleared, "poll succeeded again");
}
yield snapshot;
continue;
}
Ok(Err(e)) if e.is_fatal() => Err(e)?,
Ok(Err(e)) => e,
Err(_) => FeedError::Connection(format!(
"poll timed out after {request_timeout:?}"
)),
};
let retry = failures
.record_failure()
.inspect(|retry| {
warn!(
consecutive = retry.consecutive,
backoff = ?retry.backoff,
%error,
"poll failed, backing off"
)
})
.map_err(|consecutive| {
error!(consecutive, "giving up on the feed");
error
})?;
backoff = Some(retry.backoff);
}
}
}
pub(crate) async fn fetch_bytes(request: RequestBuilder, what: &str) -> Result<Bytes, FeedError> {
let response = request
.send()
.await
.map_err(|e| FeedError::Connection(format!("Failed to fetch {what}: {e}")))?;
read_bytes(response, what).await
}
async fn read_bytes(response: Response, what: &str) -> Result<Bytes, FeedError> {
let status = response.status();
let body = response
.bytes()
.await
.map_err(|e| FeedError::Connection(format!("Failed to read {what} response: {e}")))?;
if !status.is_success() {
return Err(FeedError::Connection(format!(
"{what} HTTP error {status}: {}",
String::from_utf8_lossy(&body)
)));
}
Ok(body)
}
pub(crate) async fn fetch_json<T: DeserializeOwned>(
request: RequestBuilder,
what: &str,
) -> Result<T, FeedError> {
parse_json(&fetch_bytes(request, what).await?, what)
}
#[expect(
dead_code,
reason = "nothing in this crate implements a book feed yet; the layer is exercised by its own tests"
)]
pub(crate) async fn read_json<T: DeserializeOwned>(
response: Response,
what: &str,
) -> Result<T, FeedError> {
parse_json(&read_bytes(response, what).await?, what)
}
fn parse_json<T: DeserializeOwned>(body: &[u8], what: &str) -> Result<T, FeedError> {
serde_json::from_slice(body)
.map_err(|e| FeedError::Parsing(format!("Failed to parse {what} response: {e}")))
}
#[cfg(test)]
pub(crate) mod test_support {
use std::net::SocketAddr;
use tokio::{
io::{AsyncBufReadExt, AsyncWriteExt, BufReader},
net::TcpListener,
};
pub struct MockHttpServer {
address: SocketAddr,
}
impl MockHttpServer {
pub fn url(&self) -> String {
format!("http://{}", self.address)
}
}
pub async fn spawn_http_server(
respond: impl Fn() -> (&'static str, String) + Send + 'static,
) -> MockHttpServer {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.unwrap();
let address = listener.local_addr().unwrap();
tokio::spawn(async move {
while let Ok((stream, _)) = listener.accept().await {
let mut reader = BufReader::new(stream);
let mut request_line = String::new();
if reader
.read_line(&mut request_line)
.await
.is_err()
{
continue;
}
let (status, body) = respond();
let payload = format!(
"HTTP/1.1 {status}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
body.len()
);
let mut stream = reader.into_inner();
let _ = stream
.write_all(payload.as_bytes())
.await;
let _ = stream.shutdown().await;
}
});
MockHttpServer { address }
}
}
#[cfg(test)]
mod tests {
use std::{
future::Future,
sync::{
atomic::{AtomicBool, AtomicU32, Ordering},
Arc,
},
};
use futures::StreamExt as _;
use tokio::sync::watch;
use tokio_stream::wrappers::WatchStream;
use super::{super::expect_to_finish, *};
impl<F, Fut, T> HttpSource for F
where
F: Fn() -> Fut + Send + Sync,
Fut: Future<Output = Result<T, FeedError>> + Send,
T: Send + Sync + 'static,
{
type Snapshot = T;
async fn fetch(&self) -> Result<T, FeedError> {
self().await
}
}
fn snapshot() -> u32 {
1
}
fn config(max_consecutive_failures: Option<u32>) -> HttpFeedConfig {
HttpFeedConfig {
poll_interval: Duration::from_millis(5),
request_timeout: Duration::from_millis(50),
max_backoff_exp: 0,
max_consecutive_failures,
max_snapshot_age: None,
}
}
#[tokio::test]
async fn publishes_a_snapshot_and_success_resets_failures() {
let polls = Arc::new(AtomicU32::new(0));
let polls_clone = Arc::clone(&polls);
let (publisher, mut rx) = Publisher::channel();
let feed = run_http_poll_feed(config(Some(2)), publisher, move || {
let n = polls_clone.fetch_add(1, Ordering::SeqCst);
async move {
if n == 0 {
Err(FeedError::Connection("boom".to_string()))
} else {
Ok(snapshot())
}
}
});
let feed = tokio::spawn(feed);
expect_to_finish("no snapshot published", async {
loop {
rx.changed().await.expect("feed ended");
if rx.borrow_and_update().is_some() {
return;
}
}
})
.await;
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(!feed.is_finished());
}
#[tokio::test]
async fn withdraws_a_snapshot_that_goes_unrefreshed_and_restores_on_next_success() {
let polls = Arc::new(AtomicU32::new(0));
let polls_clone = Arc::clone(&polls);
let recovered = Arc::new(AtomicBool::new(false));
let recovered_clone = Arc::clone(&recovered);
let (publisher, mut rx) = Publisher::channel();
let config =
HttpFeedConfig { max_snapshot_age: Some(Duration::from_millis(20)), ..config(None) };
let feed = run_http_poll_feed(config, publisher, move || {
let n = polls_clone.fetch_add(1, Ordering::SeqCst);
let recovered = recovered_clone.load(Ordering::SeqCst);
async move {
if n == 0 || recovered {
Ok(snapshot())
} else {
Err(FeedError::Connection("boom".to_string()))
}
}
});
let _feed = tokio::spawn(feed);
let wait_for = |rx: &mut watch::Receiver<Option<u32>>, want_some: bool| {
let mut rx = rx.clone();
async move {
expect_to_finish("watch did not reach the expected state", async {
loop {
if rx.borrow_and_update().is_some() == want_some {
return;
}
rx.changed().await.expect("feed ended");
}
})
.await;
}
};
wait_for(&mut rx, true).await;
wait_for(&mut rx, false).await;
recovered.store(true, Ordering::SeqCst);
wait_for(&mut rx, true).await;
}
#[tokio::test]
async fn rejects_zero_poll_interval() {
let (publisher, _rx) = Publisher::<u32>::channel();
let config = HttpFeedConfig { poll_interval: Duration::ZERO, ..config(None) };
let result = run_http_poll_feed(config, publisher, || async { Ok(snapshot()) }).await;
assert!(matches!(result, Err(FeedError::InvalidInput(_))));
}
#[tokio::test]
async fn rejects_a_zero_max_snapshot_age() {
let (publisher, _rx) = Publisher::<u32>::channel();
let config = HttpFeedConfig { max_snapshot_age: Some(Duration::ZERO), ..config(None) };
let result = run_http_poll_feed(config, publisher, || async { Ok(snapshot()) }).await;
assert!(matches!(result, Err(FeedError::InvalidInput(_))));
}
#[tokio::test(start_paused = true)]
async fn a_max_snapshot_age_below_the_poll_interval_is_a_duty_cycle() {
let (publisher, rx) = Publisher::<u32>::channel();
let mut changes = WatchStream::from_changes(rx);
let config = HttpFeedConfig {
poll_interval: Duration::from_millis(100),
max_snapshot_age: Some(Duration::from_millis(20)),
..config(None)
};
let _feed =
tokio::spawn(run_http_poll_feed(config, publisher, || async { Ok(snapshot()) }));
let start = tokio::time::Instant::now();
let mut seen = Vec::new();
for _ in 0..4 {
let change = expect_to_finish("the feed stopped changing", changes.next()).await;
seen.push((change.expect("the feed is still running"), start.elapsed()));
}
assert_eq!(
seen,
[
(Some(snapshot()), Duration::ZERO),
(None, Duration::from_millis(20)),
(Some(snapshot()), Duration::from_millis(100)),
(None, Duration::from_millis(120)),
]
);
}
#[tokio::test]
async fn gives_up_after_max_consecutive_failed_polls() {
let (publisher, _rx) = Publisher::<u32>::channel();
let feed = run_http_poll_feed(config(Some(3)), publisher, || async {
Err(FeedError::Parsing("not a snapshot".to_string()))
});
let feed = tokio::spawn(feed);
let result = expect_to_finish("feed did not give up", feed)
.await
.unwrap();
assert!(matches!(result, Err(FeedError::Parsing(msg)) if msg == "not a snapshot"));
}
#[tokio::test]
async fn gives_up_on_a_fatal_poll_error_despite_an_unlimited_failure_budget() {
let (publisher, mut rx) = Publisher::channel();
let polls = Arc::new(AtomicU32::new(0));
let polls_clone = Arc::clone(&polls);
let feed = run_http_poll_feed(config(None), publisher, move || {
let n = polls_clone.fetch_add(1, Ordering::SeqCst);
async move {
if n == 0 {
Ok(snapshot())
} else {
Err(FeedError::Fatal("api key rejected".to_string()))
}
}
});
let feed = tokio::spawn(feed);
let result = expect_to_finish("feed did not give up", feed)
.await
.unwrap();
assert!(matches!(result, Err(FeedError::Fatal(_))));
assert_eq!(polls.load(Ordering::SeqCst), 2, "the feed polled again after the fatal");
assert!(rx.borrow_and_update().is_none(), "the snapshot must be withdrawn");
}
#[tokio::test]
async fn withdraws_the_published_snapshot_when_it_gives_up() {
let polls = Arc::new(AtomicU32::new(0));
let polls_clone = Arc::clone(&polls);
let (publisher, mut rx) = Publisher::channel();
let feed = run_http_poll_feed(config(Some(2)), publisher, move || {
let n = polls_clone.fetch_add(1, Ordering::SeqCst);
async move {
if n == 0 {
Ok(snapshot())
} else {
Err(FeedError::Connection("boom".to_string()))
}
}
});
let feed = tokio::spawn(feed);
expect_to_finish("no snapshot published", async {
loop {
rx.changed()
.await
.expect("feed ended before publishing");
if rx.borrow_and_update().is_some() {
return;
}
}
})
.await;
let result = expect_to_finish("feed did not give up", feed)
.await
.unwrap();
assert!(matches!(result, Err(FeedError::Connection(_))));
assert!(rx.borrow().is_none(), "the snapshot must be withdrawn when the feed gives up");
}
#[tokio::test]
async fn withdraws_the_published_snapshot_when_the_feed_is_dropped() {
let (publisher, mut rx) = Publisher::channel();
let feed = run_http_poll_feed(config(None), publisher, || async { Ok(snapshot()) });
let feed = tokio::spawn(feed);
expect_to_finish("no snapshot published", async {
loop {
rx.changed()
.await
.expect("feed ended before publishing");
if rx.borrow_and_update().is_some() {
return;
}
}
})
.await;
feed.abort();
expect_to_finish("the snapshot was not withdrawn", rx.changed())
.await
.expect("the withdrawal must reach the receiver before the sender drops");
assert!(rx.borrow().is_none());
}
#[tokio::test(start_paused = true)]
async fn failing_polls_stretch_the_polling_cadence() {
let attempts = Arc::new(std::sync::Mutex::new(Vec::new()));
let attempts_clone = Arc::clone(&attempts);
let start = tokio::time::Instant::now();
let (publisher, _rx) = Publisher::<u32>::channel();
let feed = tokio::spawn(run_http_poll_feed(
HttpFeedConfig {
poll_interval: Duration::from_millis(50),
max_backoff_exp: 3,
..config(None)
},
publisher,
move || {
attempts_clone
.lock()
.unwrap()
.push(start.elapsed().as_millis());
async { Err(FeedError::Connection("boom".to_string())) }
},
));
tokio::time::sleep(Duration::from_secs(1)).await;
feed.abort();
assert_eq!(*attempts.lock().unwrap(), vec![0, 50, 150, 350, 750]);
}
#[tokio::test(start_paused = true)]
async fn a_recovered_poll_resumes_the_cadence_instead_of_catching_up() {
let attempts = Arc::new(std::sync::Mutex::new(Vec::new()));
let attempts_clone = Arc::clone(&attempts);
let start = tokio::time::Instant::now();
let (publisher, _rx) = Publisher::<u32>::channel();
let feed = tokio::spawn(run_http_poll_feed(
HttpFeedConfig {
poll_interval: Duration::from_millis(50),
max_backoff_exp: 3,
..config(None)
},
publisher,
move || {
let mut attempts = attempts_clone.lock().unwrap();
attempts.push(start.elapsed().as_millis());
let failing = attempts.len() <= 2;
async move {
if failing {
Err(FeedError::Connection("boom".to_string()))
} else {
Ok(snapshot())
}
}
},
));
tokio::time::sleep(Duration::from_millis(300)).await;
feed.abort();
assert_eq!(*attempts.lock().unwrap(), vec![0, 50, 150, 200, 250]);
}
#[tokio::test(start_paused = true)]
async fn a_slow_poll_never_shortens_the_gap_to_the_next_one() {
let attempts = Arc::new(std::sync::Mutex::new(Vec::new()));
let attempts_clone = Arc::clone(&attempts);
let start = tokio::time::Instant::now();
let (publisher, _rx) = Publisher::<u32>::channel();
let feed = tokio::spawn(run_http_poll_feed(
HttpFeedConfig {
poll_interval: Duration::from_millis(50),
request_timeout: Duration::from_millis(500),
..config(None)
},
publisher,
move || {
let mut attempts = attempts_clone.lock().unwrap();
attempts.push(start.elapsed().as_millis());
let slow = attempts.len() == 3;
async move {
if slow {
tokio::time::sleep(Duration::from_millis(120)).await;
}
Ok(snapshot())
}
},
));
tokio::time::sleep(Duration::from_millis(350)).await;
feed.abort();
assert_eq!(*attempts.lock().unwrap(), vec![0, 50, 100, 220, 270, 320]);
}
#[tokio::test]
async fn retries_forever_without_limit() {
let polls = Arc::new(AtomicU32::new(0));
let polls_clone = Arc::clone(&polls);
let (publisher, _rx) = Publisher::<u32>::channel();
let feed = run_http_poll_feed(config(None), publisher, move || {
polls_clone.fetch_add(1, Ordering::SeqCst);
async { Err(FeedError::Connection("boom".to_string())) }
});
let feed = tokio::spawn(feed);
tokio::time::sleep(Duration::from_millis(100)).await;
assert!(polls.load(Ordering::SeqCst) > 5, "should keep polling");
assert!(!feed.is_finished());
}
#[tokio::test]
async fn hung_poll_counts_as_failure() {
let (publisher, _rx) = Publisher::<u32>::channel();
let feed = run_http_poll_feed(
HttpFeedConfig { request_timeout: Duration::from_millis(5), ..config(Some(2)) },
publisher,
|| async {
tokio::time::sleep(Duration::from_secs(3600)).await;
Ok(snapshot())
},
);
let feed = tokio::spawn(feed);
let result = expect_to_finish("feed did not give up", feed)
.await
.unwrap();
assert!(matches!(result, Err(FeedError::Connection(msg)) if msg.contains("timed out")));
}
mod fetch_json {
use rstest::rstest;
use serde::Deserialize;
use super::{test_support::spawn_http_server, *};
#[derive(Debug, Deserialize, PartialEq)]
struct Payload {
value: u32,
}
#[tokio::test]
async fn parses_a_successful_json_body() {
let server = spawn_http_server(|| ("200 OK", r#"{"value":7}"#.to_string())).await;
let payload: Payload =
fetch_json(reqwest::Client::new().get(format!("{}/x", server.url())), "thing")
.await
.unwrap();
assert_eq!(payload, Payload { value: 7 });
}
#[rstest]
#[case::non_success_status_with_body("503 Service Unavailable", "maintenance", false)]
#[case::unparseable_success_body("200 OK", "<html>", true)]
#[tokio::test]
async fn classifies_status_and_body_failures(
#[case] status: &'static str,
#[case] body: &'static str,
#[case] expect_parsing_error: bool,
) {
let server = spawn_http_server(move || (status, body.to_string())).await;
let result: Result<Payload, FeedError> =
fetch_json(reqwest::Client::new().get(format!("{}/x", server.url())), "thing")
.await;
match result {
Err(FeedError::Parsing(msg)) if expect_parsing_error => {
assert!(msg.contains("thing"), "{msg}")
}
Err(FeedError::Connection(msg)) if !expect_parsing_error => {
assert!(
msg.contains("thing") && msg.contains("503") && msg.contains(body),
"{msg}"
)
}
other => panic!("unexpected result: {other:?}"),
}
}
#[tokio::test]
async fn refused_connection_is_a_connection_error() {
let address = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.unwrap()
.local_addr()
.unwrap();
let result: Result<Payload, FeedError> =
fetch_json(reqwest::Client::new().get(format!("http://{address}/x")), "thing")
.await;
assert!(matches!(result, Err(FeedError::Connection(msg)) if msg.contains("thing")));
}
}
}