diff --git a/rs/moq-auth/src/client.rs b/rs/moq-auth/src/client.rs index 54ebc845ce..020358d65c 100644 --- a/rs/moq-auth/src/client.rs +++ b/rs/moq-auth/src/client.rs @@ -250,7 +250,7 @@ fn backoff(failures: u32, cadence: Duration) -> Duration { mod tests { use super::*; use moq_pattern::Patterns; - use std::sync::{Arc, Mutex}; + use std::task::Poll; use std::time::SystemTime; use wiremock::matchers::{method, path}; use wiremock::{Mock, MockServer, Request as Received, ResponseTemplate}; @@ -267,21 +267,38 @@ mod tests { /// Records every request body the server saw, in order. #[derive(Clone, Default)] - struct Log(Arc>>); + struct Log(kio::Shared>); impl Log { fn events(&self) -> Vec { - self.0.lock().unwrap().iter().map(|r| r.event.clone()).collect() + self.0.read().iter().map(|r| r.event.clone()).collect() } - fn last(&self) -> Request { - self.0.lock().unwrap().last().cloned().unwrap() + fn revalidates(&self) -> usize { + self.events() + .iter() + .filter(|event| **event == Event::Revalidate) + .count() + } + + /// Wait until the requests seen so far satisfy `done`. + async fn until(&self, mut done: impl FnMut(&[Request]) -> bool + Unpin) { + self.0 + .wait(|log| if done(log) { Poll::Ready(()) } else { Poll::Pending }) + .await; + } + + /// Wait for the background `end` POST and return it. + async fn end(&self) -> Request { + let is_end = |r: &Request| matches!(r.event, Event::End { .. }); + self.until(|log| log.iter().any(is_end)).await; + self.0.read().iter().find(|r| is_end(r)).cloned().unwrap() } } impl wiremock::Match for Log { fn matches(&self, received: &Received) -> bool { - self.0.lock().unwrap().push(received.body_json().unwrap()); + self.0.lock().push(received.body_json().unwrap()); true } } @@ -310,10 +327,6 @@ mod tests { grant } - async fn settle() { - tokio::time::sleep(Duration::from_millis(50)).await; - } - #[tokio::test] async fn connect_admits_and_end_follows_the_close() { let log = Log::default(); @@ -327,9 +340,8 @@ mod tests { assert_eq!(log.events(), [Event::Connect]); consumer.close("disconnected", Bytes { sent: 7, received: 11 }); - settle().await; - let end = log.last(); + let end = log.end().await; assert_eq!(end.id, "0123"); match end.event { Event::End { reason, bytes, .. } => { @@ -350,9 +362,8 @@ mod tests { let consumer = client(&server).connect(request()).await.unwrap(); drop(consumer); - settle().await; - match log.last().event { + match log.end().await.event { Event::End { reason: Reason::Dropped, bytes, @@ -484,10 +495,7 @@ mod tests { let consumer = client(&server).connect(request()).await.unwrap(); tokio::time::sleep(Duration::from_millis(1500)).await; - assert!( - log.events().iter().filter(|e| **e == Event::Revalidate).count() >= 1, - "re-checks happened" - ); + assert!(log.revalidates() >= 1, "re-checks happened"); assert_eq!( consumer.grant().publish, patterns(&["**"]), @@ -498,9 +506,8 @@ mod tests { .await .expect("expired"); assert_eq!(reason, Reason::Expired); - settle().await; assert!(matches!( - log.last().event, + log.end().await.event, Event::End { reason: Reason::Expired, .. @@ -543,9 +550,8 @@ mod tests { tokio::time::sleep(Duration::from_millis(1500)).await; assert!(log.events().contains(&Event::Revalidate), "the re-check is in flight"); consumer.close("disconnected", Bytes::default()); - settle().await; assert!(matches!( - log.last().event, + log.end().await.event, Event::End { reason: Reason::Session(_), .. @@ -571,6 +577,8 @@ mod tests { assert!(Client::new("unix:///run/moq-auth.sock".parse().unwrap(), None).is_ok()); } + /// Backoff leaves the driver in this same state (nothing in flight, a timer armed), + /// so this covers a nudge during backoff too. #[tokio::test] async fn a_nudge_while_idle_posts_at_once() { let log = Log::default(); @@ -583,56 +591,14 @@ mod tests { let consumer = client(&server).connect(request()).await.unwrap(); assert_eq!(log.events(), [Event::Connect]); consumer.revalidate(); - tokio::time::timeout(Duration::from_secs(2), async { - loop { - if log.events().contains(&Event::Revalidate) { - break; - } - tokio::time::sleep(Duration::from_millis(10)).await; - } - }) + tokio::time::timeout( + Duration::from_secs(2), + log.until(|log| log.iter().any(|r| r.event == Event::Revalidate)), + ) .await .expect("a nudge while idle POSTs at once"); } - #[tokio::test] - async fn a_nudge_during_backoff_posts_at_once() { - let log = Log::default(); - let server = server(log.clone(), |request| match request.event { - Event::Connect => ResponseTemplate::new(200) - .set_body_json(grant(Some(Duration::from_secs(3600)), Some(Duration::from_secs(3600)))), - _ => ResponseTemplate::new(503), - }) - .await; - - let consumer = client(&server).connect(request()).await.unwrap(); - consumer.revalidate(); - tokio::time::timeout(Duration::from_secs(2), async { - loop { - if log.events().contains(&Event::Revalidate) { - break; - } - tokio::time::sleep(Duration::from_millis(10)).await; - } - }) - .await - .expect("the first re-check ran"); - settle().await; - let before = log.events().iter().filter(|event| **event == Event::Revalidate).count(); - - consumer.revalidate(); - tokio::time::timeout(Duration::from_secs(2), async { - loop { - if log.events().iter().filter(|event| **event == Event::Revalidate).count() > before { - break; - } - tokio::time::sleep(Duration::from_millis(10)).await; - } - }) - .await - .expect("a nudge during backoff POSTs at once"); - } - #[tokio::test] async fn a_nudge_during_inflight_posts_once_more_when_the_reply_lands() { let log = Log::default(); @@ -648,32 +614,27 @@ mod tests { let consumer = client(&server).connect(request()).await.unwrap(); consumer.revalidate(); - tokio::time::timeout(Duration::from_secs(2), async { - loop { - if log.events().contains(&Event::Revalidate) { - break; - } - tokio::time::sleep(Duration::from_millis(10)).await; - } - }) + tokio::time::timeout( + Duration::from_secs(2), + log.until(|log| log.iter().any(|r| r.event == Event::Revalidate)), + ) .await .expect("the first re-check is in flight"); consumer.revalidate(); consumer.revalidate(); - tokio::time::timeout(Duration::from_secs(3), async { - loop { - if log.events().iter().filter(|event| **event == Event::Revalidate).count() >= 2 { - break; - } - tokio::time::sleep(Duration::from_millis(10)).await; - } - }) + // `changed` jumps to the latest grant epoch, so two replies can arrive as one + // observation. The request log does not coalesce. + tokio::time::timeout( + Duration::from_secs(3), + log.until(|log| log.iter().filter(|r| r.event == Event::Revalidate).count() >= 2), + ) .await .expect("the in-flight nudge POSTs once more when the reply lands"); - settle().await; + consumer.close("disconnected", Bytes::default()); + log.end().await; assert_eq!( - log.events().iter().filter(|event| **event == Event::Revalidate).count(), + log.revalidates(), 2, "a burst during an in-flight re-check is one extra POST" );