use std::task::Poll;
use serde::{Deserialize, Serialize};
use crate::{Bytes, Grant};
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum Reason {
Dropped,
Expired,
Refused,
Invalid,
Session(String),
}
impl Reason {
pub fn as_str(&self) -> &str {
match self {
Self::Dropped => "dropped",
Self::Expired => "expired",
Self::Refused => "refused",
Self::Invalid => "invalid",
Self::Session(reason) => reason,
}
}
}
impl std::fmt::Display for Reason {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
impl From<&str> for Reason {
fn from(reason: &str) -> Self {
match reason {
"dropped" => Self::Dropped,
"expired" => Self::Expired,
"refused" => Self::Refused,
"invalid" => Self::Invalid,
other => Self::Session(other.to_string()),
}
}
}
impl From<String> for Reason {
fn from(reason: String) -> Self {
Self::from(reason.as_str())
}
}
impl Serialize for Reason {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
serializer.serialize_str(self.as_str())
}
}
impl<'de> Deserialize<'de> for Reason {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
Ok(String::deserialize(deserializer)?.into())
}
}
#[derive(Debug)]
struct State {
grant: Grant,
epoch: u64,
closed: Option<(Reason, Bytes)>,
revalidate: u64,
}
#[derive(Debug)]
pub struct Producer {
state: kio::Shared<State>,
}
impl Producer {
pub fn new(grant: Grant) -> (Self, Consumer) {
let state = kio::Shared::new(State {
grant,
epoch: 0,
closed: None,
revalidate: 0,
});
(Self { state: state.clone() }, Consumer { state, seen: 0 })
}
pub fn update(&self, grant: Grant) {
let mut state = self.state.lock();
if state.closed.is_some() {
return;
}
state.grant = grant;
state.epoch += 1;
}
pub fn revoke(self, reason: Reason) -> Reason {
self.finish(reason, Bytes::default()).0
}
pub fn poll_closed(&self, waiter: &kio::Waiter) -> Poll<(Reason, Bytes)> {
let state = std::task::ready!(self.state.poll(waiter, |state| ready_if(state.closed.is_some())));
Poll::Ready(state.closed.clone().expect("waited for a close"))
}
pub async fn closed(&self) -> (Reason, Bytes) {
kio::wait(|waiter| self.poll_closed(waiter)).await
}
pub fn poll_revalidate(&self, waiter: &kio::Waiter) -> Poll<()> {
let mut state = std::task::ready!(self.state.poll(waiter, |state| ready_if(state.revalidate > 0)));
state.revalidate = 0;
Poll::Ready(())
}
pub async fn revalidate_requested(&self) {
kio::wait(|waiter| self.poll_revalidate(waiter)).await
}
pub(crate) fn finish(&self, reason: Reason, bytes: Bytes) -> (Reason, Bytes) {
self.state.lock().closed.get_or_insert((reason, bytes)).clone()
}
}
impl Drop for Producer {
fn drop(&mut self) {
self.finish(Reason::Dropped, Bytes::default());
}
}
#[derive(Debug)]
pub struct Consumer {
state: kio::Shared<State>,
seen: u64,
}
impl Consumer {
pub fn fixed(grant: Grant) -> Self {
let state = kio::Shared::new(State {
grant,
epoch: 0,
closed: None,
revalidate: 0,
});
Self { state, seen: 0 }
}
pub fn grant(&self) -> Grant {
self.state.read().grant.clone()
}
pub fn poll_changed(&mut self, waiter: &kio::Waiter) -> Poll<Result<Grant, Reason>> {
let seen = self.seen;
let state = std::task::ready!(
self.state
.poll(waiter, |state| ready_if(state.epoch > seen || state.closed.is_some()))
);
if let Some((reason, _)) = &state.closed {
return Poll::Ready(Err(reason.clone()));
}
self.seen = state.epoch;
Poll::Ready(Ok(state.grant.clone()))
}
pub async fn changed(&mut self) -> Result<Grant, Reason> {
kio::wait(|waiter| self.poll_changed(waiter)).await
}
pub fn poll_closed(&self, waiter: &kio::Waiter) -> Poll<Reason> {
let state = std::task::ready!(self.state.poll(waiter, |state| ready_if(state.closed.is_some())));
Poll::Ready(state.closed.as_ref().expect("waited for a close").0.clone())
}
pub async fn closed(&self) -> Reason {
kio::wait(|waiter| self.poll_closed(waiter)).await
}
pub fn revalidate(&self) {
self.state.lock().revalidate += 1;
}
pub fn close(self, reason: impl Into<Reason>, bytes: Bytes) -> Reason {
self.state.lock().closed.get_or_insert((reason.into(), bytes)).0.clone()
}
}
impl Drop for Consumer {
fn drop(&mut self) {
self.state
.lock()
.closed
.get_or_insert((Reason::Dropped, Bytes::default()));
}
}
fn ready_if(condition: bool) -> Poll<()> {
if condition { Poll::Ready(()) } else { Poll::Pending }
}
#[cfg(test)]
mod tests {
use super::*;
use std::future::Future;
use std::pin::pin;
use std::task::{Context, Waker};
fn grant(publish: &str) -> Grant {
Grant::new([publish.parse().unwrap()].into_iter().collect(), Default::default())
}
fn poll<F: Future>(future: F) -> Poll<F::Output> {
pin!(future).poll(&mut Context::from_waker(Waker::noop()))
}
#[test]
fn update_wakes_the_consumer_once_per_change() {
let (producer, mut consumer) = Producer::new(grant("a/**"));
assert_eq!(consumer.grant(), grant("a/**"));
assert!(poll(consumer.changed()).is_pending());
producer.update(grant("b/**"));
assert_eq!(poll(consumer.changed()), Poll::Ready(Ok(grant("b/**"))));
assert_eq!(consumer.grant(), grant("b/**"));
assert!(poll(consumer.changed()).is_pending());
}
#[test]
fn revoke_reaches_the_consumer() {
let (producer, mut consumer) = Producer::new(grant("a/**"));
assert!(poll(consumer.closed()).is_pending());
producer.revoke(Reason::Refused);
assert_eq!(poll(consumer.closed()), Poll::Ready(Reason::Refused));
assert_eq!(poll(consumer.changed()), Poll::Ready(Err(Reason::Refused)));
}
#[test]
fn the_session_close_hands_over_byte_totals() {
let (producer, consumer) = Producer::new(grant("a/**"));
consumer.close("done", Bytes { sent: 3, received: 5 });
assert_eq!(
poll(producer.closed()),
Poll::Ready((Reason::Session("done".into()), Bytes { sent: 3, received: 5 }))
);
}
#[test]
fn dropping_the_producer_revokes() {
let (producer, consumer) = Producer::new(grant("a/**"));
drop(producer);
assert_eq!(poll(consumer.closed()), Poll::Ready(Reason::Dropped));
}
#[test]
fn dropping_the_consumer_reports_zero_bytes() {
let (producer, consumer) = Producer::new(grant("a/**"));
drop(consumer);
assert_eq!(
poll(producer.closed()),
Poll::Ready((Reason::Dropped, Bytes::default()))
);
}
#[test]
fn the_session_close_reaches_the_producer_and_the_first_reason_wins() {
let (producer, consumer) = Producer::new(grant("a/**"));
assert!(poll(producer.closed()).is_pending());
let recorded = consumer.close("disconnected", Bytes { sent: 1, received: 2 });
assert_eq!(recorded, Reason::Session("disconnected".into()));
assert_eq!(
poll(producer.closed()),
Poll::Ready((Reason::Session("disconnected".into()), Bytes { sent: 1, received: 2 }))
);
assert_eq!(producer.revoke(Reason::Expired), Reason::Session("disconnected".into()));
}
#[test]
fn a_fixed_lease_only_ends_by_the_holder() {
let mut consumer = Consumer::fixed(grant("a/**"));
assert_eq!(consumer.grant(), grant("a/**"));
assert!(poll(consumer.changed()).is_pending());
assert!(poll(consumer.closed()).is_pending());
consumer.revalidate();
assert_eq!(consumer.grant(), grant("a/**"));
assert!(poll(consumer.changed()).is_pending());
assert!(poll(consumer.closed()).is_pending());
assert_eq!(consumer.close("done", Bytes::default()), Reason::Session("done".into()));
}
#[test]
fn n_nudges_wake_the_producer_once() {
let (producer, consumer) = Producer::new(grant("a/**"));
assert!(poll(producer.revalidate_requested()).is_pending());
for _ in 0..8 {
consumer.revalidate();
}
assert_eq!(poll(producer.revalidate_requested()), Poll::Ready(()));
assert!(poll(producer.revalidate_requested()).is_pending());
consumer.revalidate();
assert_eq!(poll(producer.revalidate_requested()), Poll::Ready(()));
}
#[test]
fn reason_round_trips_as_one_string() {
for (reason, text) in [
(Reason::Dropped, "\"dropped\""),
(Reason::Expired, "\"expired\""),
(Reason::Refused, "\"refused\""),
(Reason::Invalid, "\"invalid\""),
(Reason::Session("protocol error".into()), "\"protocol error\""),
] {
assert_eq!(serde_json::to_string(&reason).unwrap(), text);
assert_eq!(serde_json::from_str::<Reason>(text).unwrap(), reason);
}
}
}