use std::{
future::Future,
pin::Pin,
sync::atomic::{AtomicBool, Ordering},
task::{Context, Poll, Waker},
};
use parking_lot::Mutex;
use crate::types::RespCommand;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ObserverStatus {
WaitingForResult,
ResultSet,
SessionDisposed,
}
#[derive(Debug)]
struct ObserverState {
status: ObserverStatus,
result: CollectionItemResult,
}
#[derive(Debug, Default)]
pub(crate) struct Wakeup {
fired: AtomicBool,
waiters: Mutex<Vec<Waker>>,
}
impl Wakeup {
#[inline]
pub fn new() -> Self {
Self::default()
}
pub fn notify_one(&self) {
self.fired.store(true, Ordering::SeqCst);
if let Some(w) = self.waiters.lock().pop() {
w.wake();
}
}
pub fn poll_wait(&self, cx: &mut Context<'_>) -> Poll<()> {
if self.fired.swap(false, Ordering::SeqCst) {
return Poll::Ready(());
}
let mut waiters = self.waiters.lock();
if self.fired.load(Ordering::SeqCst) {
return Poll::Ready(());
}
if !waiters.iter().any(|w| w.will_wake(cx.waker())) {
waiters.push(cx.waker().clone());
}
Poll::Pending
}
}
pub(crate) struct WakeupFuture<'a> {
wakeup: &'a Wakeup,
}
impl Future for WakeupFuture<'_> {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<()> {
self.wakeup.poll_wait(cx)
}
}
impl Wakeup {
pub fn wait(&self) -> WakeupFuture<'_> {
WakeupFuture { wakeup: self }
}
}
#[derive(Debug)]
pub struct CollectionItemObserver {
pub session_id: usize,
pub command: RespCommand,
pub command_args: Vec<Vec<u8>>,
state: Mutex<ObserverState>,
result_found: Wakeup,
}
impl CollectionItemObserver {
pub fn new(session_id: usize, command: RespCommand, command_args: Vec<Vec<u8>>) -> Self {
Self {
session_id,
command,
command_args,
state: Mutex::new(ObserverState {
status: ObserverStatus::WaitingForResult,
result: CollectionItemResult::empty(),
}),
result_found: Wakeup::new(),
}
}
#[inline]
pub fn status(&self) -> ObserverStatus {
self.state.lock().status
}
#[inline]
pub fn result(&self) -> CollectionItemResult {
self.state.lock().result.clone()
}
pub async fn wait_result(&self) {
self.result_found.wait().await;
}
pub fn poll_wait(&self, cx: &mut Context<'_>) -> Poll<()> {
self.result_found.poll_wait(cx)
}
pub fn handle_set_result(&self, result: CollectionItemResult) {
let mut state = self.state.lock();
if state.status != ObserverStatus::WaitingForResult {
return;
}
state.result = result;
state.status = ObserverStatus::ResultSet;
drop(state);
self.result_found.notify_one();
}
pub fn try_force_unblock(&self, throw_error: bool) -> bool {
let mut state = self.state.lock();
if state.status != ObserverStatus::WaitingForResult {
return false;
}
state.result = if throw_error {
CollectionItemResult::force_unblocked()
} else {
CollectionItemResult::empty()
};
state.status = ObserverStatus::ResultSet;
drop(state);
self.result_found.notify_one();
true
}
pub fn handle_session_disposed(&self) {
let mut state = self.state.lock();
state.status = ObserverStatus::SessionDisposed;
drop(state);
self.result_found.notify_one();
}
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct CollectionItemResult {
pub key: Option<Vec<u8>>,
pub item: Option<Vec<u8>>,
pub score: Option<f64>,
pub items: Option<Vec<Vec<u8>>>,
pub scores: Option<Vec<f64>>,
pub is_force_unblocked: bool,
pub is_type_mismatch: bool,
}
impl CollectionItemResult {
pub fn empty() -> Self {
Self::default()
}
pub fn single(key: Vec<u8>, item: Vec<u8>) -> Self {
Self {
key: Some(key),
item: Some(item),
..Self::default()
}
}
pub fn single_with_score(key: Vec<u8>, score: f64, item: Vec<u8>) -> Self {
Self {
key: Some(key),
item: Some(item),
score: Some(score),
..Self::default()
}
}
pub fn multiple(key: Vec<u8>, items: Vec<Vec<u8>>) -> Self {
Self {
key: Some(key),
items: Some(items),
..Self::default()
}
}
pub fn multiple_with_scores(key: Vec<u8>, scores: Vec<f64>, items: Vec<Vec<u8>>) -> Self {
Self {
key: Some(key),
scores: Some(scores),
items: Some(items),
..Self::default()
}
}
#[inline]
pub fn found(&self) -> bool {
self.key.is_some()
}
pub fn force_unblocked() -> Self {
Self {
is_force_unblocked: true,
..Self::default()
}
}
pub fn type_mismatch() -> Self {
Self {
is_type_mismatch: true,
..Self::default()
}
}
pub const EMPTY: CollectionItemResult = CollectionItemResult {
key: None,
item: None,
score: None,
items: None,
scores: None,
is_force_unblocked: false,
is_type_mismatch: false,
};
pub const FORCE_UNBLOCKED: CollectionItemResult = CollectionItemResult {
is_force_unblocked: true,
..CollectionItemResult::EMPTY
};
pub const TYPE_MISMATCH: CollectionItemResult = CollectionItemResult {
is_type_mismatch: true,
..CollectionItemResult::EMPTY
};
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn set_result_once_only() {
let obs = CollectionItemObserver::new(7, RespCommand::Blpop, vec![]);
assert_eq!(obs.status(), ObserverStatus::WaitingForResult);
obs.handle_set_result(CollectionItemResult::single(b"k".to_vec(), b"v".to_vec()));
assert_eq!(obs.status(), ObserverStatus::ResultSet);
assert!(obs.result().found());
obs.handle_set_result(CollectionItemResult::empty());
assert!(obs.result().found());
}
#[test]
fn force_unblock_and_dispose() {
let obs = CollectionItemObserver::new(1, RespCommand::Bzpopmin, vec![]);
assert!(obs.try_force_unblock(false)); assert!(!obs.try_force_unblock(false)); assert!(!obs.result().found());
let obs2 = CollectionItemObserver::new(2, RespCommand::Blpop, vec![]);
assert!(obs2.try_force_unblock(true));
assert!(obs2.result().is_force_unblocked);
let obs3 = CollectionItemObserver::new(3, RespCommand::Blpop, vec![]);
obs3.handle_session_disposed();
assert_eq!(obs3.status(), ObserverStatus::SessionDisposed);
obs3.handle_set_result(CollectionItemResult::single(b"k".to_vec(), b"v".to_vec()));
assert!(!obs3.result().found());
}
#[test]
fn result_shapes() {
let multi = CollectionItemResult::multiple_with_scores(
b"z".to_vec(),
vec![1.0, 2.0],
vec![b"x".to_vec(), b"y".to_vec()],
);
assert!(multi.found());
assert_eq!(multi.scores.as_ref().unwrap().len(), 2);
assert!(!CollectionItemResult::EMPTY.found());
const {
assert!(CollectionItemResult::TYPE_MISMATCH.is_type_mismatch);
}
assert!(!CollectionItemResult::FORCE_UNBLOCKED.found());
}
}