use std::future::Future;
use crossfire::oneshot::{RxOneshot, TxOneshot, oneshot};
use parking_lot::Mutex;
#[derive(Debug, thiserror::Error)]
#[error("group commit pipeline broken")]
pub struct Broken;
pub type CommitRx = RxOneshot<Result<u64, Broken>>;
struct Waiter {
target: u64,
tx: Option<TxOneshot<Result<u64, Broken>>>,
}
struct Waiters {
leading: bool,
waiters: Vec<Waiter>,
}
pub enum Enter {
Done(u64),
Follow(CommitRx),
Lead,
}
pub struct GroupCommitPipeline {
waiters: Mutex<Waiters>,
}
impl Default for GroupCommitPipeline {
fn default() -> Self {
Self::new()
}
}
impl GroupCommitPipeline {
pub const fn new() -> Self {
Self {
waiters: Mutex::new(Waiters {
leading: false,
waiters: Vec::new(),
}),
}
}
pub fn enter(&self, target: u64, watermark: impl FnOnce() -> u64) -> Enter {
let mut lock = self.waiters.lock();
let committed = watermark();
if target <= committed {
return Enter::Done(committed);
}
if lock.leading {
let (tx, rx) = oneshot::<Result<u64, Broken>>();
lock.waiters.push(Waiter {
target,
tx: Some(tx),
});
return Enter::Follow(rx);
}
lock.leading = true;
Enter::Lead
}
pub async fn wait(
&self,
rx: CommitRx,
target: u64,
watermark: impl Fn() -> u64,
) -> Result<u64, Broken> {
match rx.await {
Ok(res) => res,
Err(_) => {
let committed = watermark();
if target <= committed {
Ok(committed)
} else {
Err(Broken)
}
}
}
}
pub async fn run_leader<S: GroupCommitStep>(&self, step: S) -> Result<u64, S::Error> {
let mut last_committed;
loop {
let batch_target = {
let lock = self.waiters.lock();
let max_waiter_target = lock.waiters.iter().map(|w| w.target).max().unwrap_or(0);
step.tail().max(max_waiter_target)
};
let watermark = step.watermark();
let step_res = if batch_target > watermark {
step.step(batch_target).await
} else {
Ok(watermark)
};
match step_res {
Ok(committed) => {
last_committed = committed;
let mut lock = self.waiters.lock();
lock.waiters.retain_mut(|waiter| {
if waiter.target <= committed {
if let Some(tx) = waiter.tx.take() {
tx.send(Ok(committed));
}
false
} else {
true
}
});
let has_lagging_waiters = lock.waiters.iter().any(|w| w.target > last_committed);
let has_new_tail = step.tail() > last_committed;
if !has_lagging_waiters && !has_new_tail {
lock.leading = false;
return Ok(last_committed);
}
}
Err(e) => {
let mut lock = self.waiters.lock();
lock.leading = false;
for mut waiter in lock.waiters.drain(..) {
if let Some(tx) = waiter.tx.take() {
tx.send(Err(Broken));
}
}
return Err(e);
}
}
}
}
}
pub trait GroupCommitStep {
type Error;
fn tail(&self) -> u64;
fn watermark(&self) -> u64;
fn step(&self, target: u64) -> impl Future<Output = Result<u64, Self::Error>>;
}
#[cfg(test)]
mod tests {
use std::{
sync::{
Arc,
atomic::{AtomicU64, Ordering},
},
thread,
};
use super::*;
use crate::future::block_on;
#[derive(Default)]
struct FakeStep {
watermark: AtomicU64,
steps: AtomicU64,
}
impl GroupCommitStep for Arc<FakeStep> {
type Error = Broken;
fn tail(&self) -> u64 {
self.watermark.load(Ordering::Acquire)
}
fn watermark(&self) -> u64 {
self.watermark.load(Ordering::Acquire)
}
async fn step(&self, target: u64) -> Result<u64, Broken> {
self.steps.fetch_add(1, Ordering::AcqRel);
self.watermark.fetch_max(target, Ordering::AcqRel);
Ok(self.watermark.load(Ordering::Acquire))
}
}
#[test]
fn test_enter_short_circuit_and_leadership() {
let pipeline = GroupCommitPipeline::new();
let watermark = AtomicU64::new(100);
match pipeline.enter(80, || watermark.load(Ordering::Acquire)) {
Enter::Done(v) => assert_eq!(v, 100),
_ => panic!("须命中快速短路"),
}
assert!(matches!(
pipeline.enter(120, || watermark.load(Ordering::Acquire)),
Enter::Lead
));
}
#[test]
fn test_batch_wake_and_site_consistency() {
let pipeline = Arc::new(GroupCommitPipeline::new());
let step = Arc::new(FakeStep::default());
assert!(matches!(pipeline.enter(300, || 0), Enter::Lead));
let mut handles = Vec::new();
for target in [100u64, 200, 250, 300] {
let Enter::Follow(rx) = pipeline.enter(target, || step.watermark()) else {
panic!("Leader 在位时须登记为 Follower");
};
let pipeline_bg = Arc::clone(&pipeline);
let step_bg = Arc::clone(&step);
handles.push(thread::spawn(move || {
block_on(pipeline_bg.wait(rx, target, || step_bg.watermark()))
}));
}
let last = block_on(pipeline.run_leader(Arc::clone(&step))).unwrap();
assert_eq!(last, 300);
assert_eq!(step.steps.load(Ordering::Acquire), 1);
for handle in handles {
assert_eq!(handle.join().unwrap().unwrap(), 300);
}
assert!(matches!(
pipeline.enter(301, || step.watermark()),
Enter::Lead
));
}
#[test]
fn test_error_broadcast_releases_leadership() {
struct FailingStep;
impl GroupCommitStep for FailingStep {
type Error = Broken;
fn tail(&self) -> u64 {
500
}
fn watermark(&self) -> u64 {
0
}
async fn step(&self, _target: u64) -> Result<u64, Broken> {
Err(Broken)
}
}
let pipeline = Arc::new(GroupCommitPipeline::new());
assert!(matches!(pipeline.enter(500, || 0), Enter::Lead));
let Enter::Follow(rx) = pipeline.enter(500, || 0) else {
panic!("须登记为 Follower");
};
let pipeline_bg = Arc::clone(&pipeline);
let handle = thread::spawn(move || block_on(pipeline_bg.wait(rx, 500, || 0)).unwrap_err());
assert!(block_on(pipeline.run_leader(FailingStep)).is_err());
let broken = handle.join().unwrap();
assert_eq!(broken.to_string(), "group commit pipeline broken");
assert!(matches!(pipeline.enter(1, || 0), Enter::Lead));
}
}