use std::collections::HashMap;
use crate::outcome::{Outcome, UnavailableReason};
use crate::tier::Tier;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum RejectReason {
AlreadyClaimed,
Exhausted,
VersionConflict,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum StrongResult {
Committed,
Rejected(RejectReason),
Unavailable(UnavailableReason),
}
impl StrongResult {
pub fn is_committed(&self) -> bool {
matches!(self, StrongResult::Committed)
}
pub fn is_linearizable(&self) -> bool {
!matches!(self, StrongResult::Unavailable(_))
}
pub fn consistency(&self) -> Outcome {
match self {
StrongResult::Committed | StrongResult::Rejected(_) => Outcome::Committed(Tier::Strong),
StrongResult::Unavailable(reason) => Outcome::Unavailable(*reason),
}
}
}
#[derive(Clone, Debug)]
struct Member {
id: String,
online: bool,
store: HashMap<String, (Vec<u8>, u64)>,
}
pub struct QuorumGroup {
members: Vec<Member>,
}
impl QuorumGroup {
pub fn new<I, S>(member_ids: I) -> Self
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
let members = member_ids
.into_iter()
.map(|id| Member {
id: id.into(),
online: true,
store: HashMap::new(),
})
.collect();
Self { members }
}
pub fn size(&self) -> usize {
self.members.len()
}
pub fn majority(&self) -> usize {
self.members.len() / 2 + 1
}
pub fn online_count(&self) -> usize {
self.members.iter().filter(|m| m.online).count()
}
pub fn partition(&mut self, member_id: &str) -> bool {
self.set_online(member_id, false)
}
pub fn heal(&mut self, member_id: &str) -> bool {
self.set_online(member_id, true)
}
fn set_online(&mut self, member_id: &str, online: bool) -> bool {
match self.members.iter_mut().find(|m| m.id == member_id) {
Some(m) => {
m.online = online;
true
}
None => false,
}
}
fn online_indices(&self) -> Vec<usize> {
(0..self.members.len())
.filter(|&i| self.members[i].online)
.collect()
}
pub fn read(&self, key: &str) -> Result<(Option<Vec<u8>>, u64), UnavailableReason> {
let online = self.online_indices();
if online.len() < self.majority() {
return Err(UnavailableReason::QuorumUnreachable);
}
let best = online
.iter()
.filter_map(|&i| self.members[i].store.get(key))
.max_by_key(|(_, version)| *version);
Ok(match best {
Some((value, version)) => (Some(value.clone()), *version),
None => (None, 0),
})
}
pub fn cas(&mut self, key: &str, expected_version: u64, new_value: Vec<u8>) -> StrongResult {
let online = self.online_indices();
if online.len() < self.majority() {
return StrongResult::Unavailable(UnavailableReason::QuorumUnreachable);
}
let current = online
.iter()
.filter_map(|&i| self.members[i].store.get(key).map(|(_, v)| *v))
.max()
.unwrap_or(0);
if current != expected_version {
return StrongResult::Rejected(RejectReason::VersionConflict);
}
let new_version = expected_version + 1;
for &i in &online {
self.members[i]
.store
.insert(key.to_string(), (new_value.clone(), new_version));
}
StrongResult::Committed
}
pub fn claim_unique(&mut self, key: &str, owner: &[u8]) -> StrongResult {
let (current, version) = match self.read(key) {
Ok(read) => read,
Err(reason) => return StrongResult::Unavailable(reason),
};
if current.is_some() {
return StrongResult::Rejected(RejectReason::AlreadyClaimed);
}
self.cas(key, version, owner.to_vec())
}
pub fn init_seats(&mut self, key: &str, count: u64) -> StrongResult {
let (_, version) = match self.read(key) {
Ok(read) => read,
Err(reason) => return StrongResult::Unavailable(reason),
};
self.cas(key, version, count.to_le_bytes().to_vec())
}
pub fn acquire_seat(&mut self, key: &str) -> StrongResult {
let (current, version) = match self.read(key) {
Ok(read) => read,
Err(reason) => return StrongResult::Unavailable(reason),
};
let remaining = decode_count(current.as_deref());
if remaining == 0 {
return StrongResult::Rejected(RejectReason::Exhausted);
}
self.cas(key, version, (remaining - 1).to_le_bytes().to_vec())
}
pub fn seats_remaining(&self, key: &str) -> Result<u64, UnavailableReason> {
let (current, _) = self.read(key)?;
Ok(decode_count(current.as_deref()))
}
}
fn decode_count(bytes: Option<&[u8]>) -> u64 {
match bytes {
Some(b) if b.len() == 8 => u64::from_le_bytes(b.try_into().unwrap()),
_ => 0,
}
}