use super::Stream;
use parking_lot::Mutex;
use std::collections::VecDeque;
use std::fmt;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll, Waker};
const PARTITION_COOPERATIVE_BUDGET: usize = 1024;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Lane {
Matched,
Unmatched,
}
impl Lane {
#[inline]
const fn other(self) -> Self {
match self {
Self::Matched => Self::Unmatched,
Self::Unmatched => Self::Matched,
}
}
}
struct PartitionInner<S: Stream, P> {
stream: S,
predicate: P,
matched: VecDeque<S::Item>,
unmatched: VecDeque<S::Item>,
capacity: usize,
done: bool,
matched_waker: Option<Waker>,
unmatched_waker: Option<Waker>,
matched_dropped: bool,
unmatched_dropped: bool,
}
impl<S: Stream, P> PartitionInner<S, P> {
#[inline]
fn queue(&mut self, lane: Lane) -> &mut VecDeque<S::Item> {
match lane {
Lane::Matched => &mut self.matched,
Lane::Unmatched => &mut self.unmatched,
}
}
#[inline]
fn len_of(&self, lane: Lane) -> usize {
match lane {
Lane::Matched => self.matched.len(),
Lane::Unmatched => self.unmatched.len(),
}
}
#[inline]
fn is_dropped(&self, lane: Lane) -> bool {
match lane {
Lane::Matched => self.matched_dropped,
Lane::Unmatched => self.unmatched_dropped,
}
}
#[inline]
fn take_waker(&mut self, lane: Lane) -> Option<Waker> {
match lane {
Lane::Matched => self.matched_waker.take(),
Lane::Unmatched => self.unmatched_waker.take(),
}
}
#[inline]
fn register_waker(&mut self, lane: Lane, waker: &Waker) {
let slot = match lane {
Lane::Matched => &mut self.matched_waker,
Lane::Unmatched => &mut self.unmatched_waker,
};
match slot {
Some(existing) if existing.will_wake(waker) => {}
_ => *slot = Some(waker.clone()),
}
}
}
#[must_use = "streams do nothing unless polled"]
pub struct Partition<S: Stream, P> {
inner: Arc<Mutex<PartitionInner<S, P>>>,
lane: Lane,
}
impl<S: Stream, P> fmt::Debug for Partition<S, P> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let inner = self.inner.lock();
f.debug_struct("Partition")
.field("lane", &self.lane)
.field("buffered", &inner.len_of(self.lane))
.field("peer_buffered", &inner.len_of(self.lane.other()))
.field("capacity", &inner.capacity)
.field("done", &inner.done)
.finish_non_exhaustive()
}
}
impl<S: Stream, P> Partition<S, P> {
#[must_use]
pub fn buffered_len(&self) -> usize {
self.inner.lock().len_of(self.lane)
}
#[must_use]
pub fn peer_buffered_len(&self) -> usize {
self.inner.lock().len_of(self.lane.other())
}
}
impl<S: Stream, P> Drop for Partition<S, P> {
fn drop(&mut self) {
let mut inner = self.inner.lock();
match self.lane {
Lane::Matched => {
inner.matched_dropped = true;
inner.matched.clear();
}
Lane::Unmatched => {
inner.unmatched_dropped = true;
inner.unmatched.clear();
}
}
let peer = inner.take_waker(self.lane.other());
drop(inner);
if let Some(waker) = peer {
waker.wake();
}
}
}
impl<S, P> Stream for Partition<S, P>
where
S: Stream + Unpin,
P: FnMut(&S::Item) -> bool,
{
type Item = S::Item;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let lane = self.lane;
let peer = lane.other();
let mut inner = self.inner.lock();
let mut routed_this_poll = 0usize;
loop {
let own = inner.queue(lane).pop_front();
if let Some(item) = own {
let peer_waker = inner.take_waker(peer);
drop(inner);
if let Some(waker) = peer_waker {
waker.wake();
}
return Poll::Ready(Some(item));
}
if inner.done {
return Poll::Ready(None);
}
if !inner.is_dropped(peer) && inner.len_of(peer) >= inner.capacity {
inner.register_waker(lane, cx.waker());
let peer_waker = inner.take_waker(peer);
drop(inner);
if let Some(waker) = peer_waker {
waker.wake();
}
return Poll::Pending;
}
if routed_this_poll >= PARTITION_COOPERATIVE_BUDGET {
inner.register_waker(lane, cx.waker());
drop(inner);
cx.waker().wake_by_ref();
return Poll::Pending;
}
let polled = {
let inner = &mut *inner;
Pin::new(&mut inner.stream).poll_next(cx)
};
match polled {
Poll::Ready(Some(item)) => {
routed_this_poll += 1;
let target = {
let inner = &mut *inner;
if (inner.predicate)(&item) {
Lane::Matched
} else {
Lane::Unmatched
}
};
if target == lane {
let peer_waker = inner.take_waker(peer);
drop(inner);
if let Some(waker) = peer_waker {
waker.wake();
}
return Poll::Ready(Some(item));
}
if inner.is_dropped(target) {
continue;
}
inner.queue(target).push_back(item);
if let Some(waker) = inner.take_waker(target) {
waker.wake();
}
}
Poll::Ready(None) => {
inner.done = true;
let peer_waker = inner.take_waker(peer);
drop(inner);
if let Some(waker) = peer_waker {
waker.wake();
}
return Poll::Ready(None);
}
Poll::Pending => {
inner.register_waker(lane, cx.waker());
return Poll::Pending;
}
}
}
}
}
pub fn partition<S, P>(
stream: S,
predicate: P,
lane_capacity: usize,
) -> (Partition<S, P>, Partition<S, P>)
where
S: Stream + Unpin,
P: FnMut(&S::Item) -> bool,
{
assert!(
lane_capacity > 0,
"partition lane_capacity must be non-zero; a zero-capacity lane cannot buffer a peer-bound item"
);
let inner = Arc::new(Mutex::new(PartitionInner {
stream,
predicate,
matched: VecDeque::new(),
unmatched: VecDeque::new(),
capacity: lane_capacity,
done: false,
matched_waker: None,
unmatched_waker: None,
matched_dropped: false,
unmatched_dropped: false,
}));
(
Partition {
inner: Arc::clone(&inner),
lane: Lane::Matched,
},
Partition {
inner,
lane: Lane::Unmatched,
},
)
}
#[cfg(test)]
mod tests {
#![allow(
clippy::pedantic,
clippy::nursery,
clippy::expect_fun_call,
clippy::future_not_send
)]
use super::*;
use crate::stream::iter;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::task::Wake;
fn init_test(name: &str) {
crate::test_utils::init_test_logging();
crate::test_phase!(name);
}
fn noop_waker() -> Waker {
std::task::Waker::noop().clone()
}
struct TrackWaker(Arc<AtomicBool>);
impl Wake for TrackWaker {
fn wake(self: Arc<Self>) {
self.0.store(true, Ordering::SeqCst);
}
fn wake_by_ref(self: &Arc<Self>) {
self.0.store(true, Ordering::SeqCst);
}
}
fn poll_once<S: Stream + Unpin>(stream: &mut S) -> Poll<Option<S::Item>> {
let waker = noop_waker();
let mut cx = Context::from_waker(&waker);
Pin::new(stream).poll_next(&mut cx)
}
fn drain_half<S: Stream + Unpin>(stream: &mut S) -> Vec<S::Item> {
let mut out = Vec::new();
let mut idle = 0usize;
loop {
match poll_once(stream) {
Poll::Ready(Some(item)) => {
out.push(item);
idle = 0;
}
Poll::Ready(None) => return out,
Poll::Pending => {
idle += 1;
assert!(idle <= 64, "partition half stalled for {idle} polls");
}
}
}
}
#[test]
fn partition_routes_every_item_to_exactly_one_half() {
init_test("partition_routes_every_item_to_exactly_one_half");
let (mut evens, mut odds) =
partition(iter(vec![0, 1, 2, 3, 4, 5]), |x: &i32| x % 2 == 0, 8);
let even_items = drain_half(&mut evens);
crate::assert_with_log!(
even_items == vec![0, 2, 4],
"matched half receives predicate-true items in order",
vec![0, 2, 4],
even_items
);
let odd_items = drain_half(&mut odds);
crate::assert_with_log!(
odd_items == vec![1, 3, 5],
"unmatched half receives the rest, in order",
vec![1, 3, 5],
odd_items
);
crate::test_complete!("partition_routes_every_item_to_exactly_one_half");
}
#[test]
fn partition_stalls_at_peer_capacity() {
init_test("partition_stalls_at_peer_capacity");
let (mut evens, _odds) = partition(iter(vec![1, 3, 5, 7, 9]), |x: &i32| x % 2 == 0, 2);
let woke = Arc::new(AtomicBool::new(false));
let waker = Waker::from(Arc::new(TrackWaker(woke.clone())));
let mut cx = Context::from_waker(&waker);
let polled = Pin::new(&mut evens).poll_next(&mut cx);
crate::assert_with_log!(
matches!(polled, Poll::Pending),
"matched half yields once the peer buffer is full",
"Poll::Pending",
polled
);
crate::assert_with_log!(
evens.peer_buffered_len() == 2,
"peer buffer bounded by lane_capacity, not by source length",
2,
evens.peer_buffered_len()
);
crate::test_complete!("partition_stalls_at_peer_capacity");
}
#[test]
fn partition_peer_consumption_releases_the_stall() {
init_test("partition_peer_consumption_releases_the_stall");
let (mut evens, mut odds) = partition(iter(vec![1, 3, 2, 5, 7]), |x: &i32| x % 2 == 0, 2);
crate::assert_with_log!(
matches!(poll_once(&mut evens), Poll::Pending),
"even half stalls at capacity",
"Poll::Pending",
"Poll::Pending"
);
let first_odd = poll_once(&mut odds);
crate::assert_with_log!(
matches!(first_odd, Poll::Ready(Some(1))),
"odd half drains its buffer",
"Poll::Ready(Some(1))",
first_odd
);
let resumed = poll_once(&mut evens);
crate::assert_with_log!(
matches!(resumed, Poll::Ready(Some(2))),
"even half resumes once the peer made room",
"Poll::Ready(Some(2))",
resumed
);
let stalled_again = poll_once(&mut evens);
crate::assert_with_log!(
matches!(stalled_again, Poll::Pending),
"even half re-stalls once the peer buffer refills",
"Poll::Pending",
stalled_again
);
crate::assert_with_log!(
evens.peer_buffered_len() == 2,
"peer buffer back at capacity",
2,
evens.peer_buffered_len()
);
crate::test_complete!("partition_peer_consumption_releases_the_stall");
}
#[test]
fn partition_delivers_every_item_when_both_halves_are_consumed() {
init_test("partition_delivers_every_item_when_both_halves_are_consumed");
let (mut evens, mut odds) =
partition(iter(vec![1, 3, 2, 5, 7, 4, 9, 6]), |x: &i32| x % 2 == 0, 2);
let mut even_items = Vec::new();
let mut odd_items = Vec::new();
let mut even_done = false;
let mut odd_done = false;
let mut idle = 0usize;
while !(even_done && odd_done) {
let mut progressed = false;
if !even_done {
match poll_once(&mut evens) {
Poll::Ready(Some(item)) => {
even_items.push(item);
progressed = true;
}
Poll::Ready(None) => {
even_done = true;
progressed = true;
}
Poll::Pending => {}
}
}
if !odd_done {
match poll_once(&mut odds) {
Poll::Ready(Some(item)) => {
odd_items.push(item);
progressed = true;
}
Poll::Ready(None) => {
odd_done = true;
progressed = true;
}
Poll::Pending => {}
}
}
if progressed {
idle = 0;
} else {
idle += 1;
assert!(idle <= 16, "both halves stalled together for {idle} rounds");
}
}
crate::assert_with_log!(
even_items == vec![2, 4, 6],
"matched items delivered in order",
vec![2, 4, 6],
even_items
);
crate::assert_with_log!(
odd_items == vec![1, 3, 5, 7, 9],
"unmatched items delivered in order",
vec![1, 3, 5, 7, 9],
odd_items
);
crate::test_complete!("partition_delivers_every_item_when_both_halves_are_consumed");
}
#[test]
fn partition_dropped_peer_lifts_backpressure() {
init_test("partition_dropped_peer_lifts_backpressure");
let (mut evens, odds) =
partition(iter(vec![1, 3, 2, 5, 7, 9, 11]), |x: &i32| x % 2 == 0, 1);
drop(odds);
let even_items = drain_half(&mut evens);
crate::assert_with_log!(
even_items == vec![2],
"surviving half runs to completion, peer items discarded",
vec![2],
even_items
);
crate::test_complete!("partition_dropped_peer_lifts_backpressure");
}
#[test]
fn partition_end_of_source_terminates_both_halves() {
init_test("partition_end_of_source_terminates_both_halves");
let (mut evens, mut odds) = partition(iter(vec![2, 4]), |x: &i32| x % 2 == 0, 4);
let even_items = drain_half(&mut evens);
crate::assert_with_log!(
even_items == vec![2, 4],
"matched half drains the source",
vec![2, 4],
even_items
);
crate::assert_with_log!(
matches!(poll_once(&mut odds), Poll::Ready(None)),
"unmatched half observes end-of-source with an empty buffer",
"Poll::Ready(None)",
"Poll::Ready(None)"
);
crate::test_complete!("partition_end_of_source_terminates_both_halves");
}
#[test]
#[should_panic(expected = "partition lane_capacity must be non-zero")]
fn partition_rejects_zero_capacity() {
let _ = partition(iter(Vec::<i32>::new()), |_: &i32| true, 0);
}
mod partition_conformance {
use super::*;
#[test]
fn peer_buffer_never_exceeds_capacity() {
init_test("partition_conformance::peer_buffer_never_exceeds_capacity");
let source: Vec<i32> = (0..64).map(|i| i * 2 + 1).collect(); let (mut evens, _odds) = partition(iter(source), |x: &i32| x % 2 == 0, 3);
for _ in 0..32 {
let _ = poll_once(&mut evens);
crate::assert_with_log!(
evens.peer_buffered_len() <= 3,
"peer buffer stays within lane_capacity across repeated polls",
"<= 3",
evens.peer_buffered_len()
);
}
crate::test_complete!("partition_conformance::peer_buffer_never_exceeds_capacity");
}
#[test]
fn routing_yields_after_cooperative_budget() {
init_test("partition_conformance::routing_yields_after_cooperative_budget");
let source: Vec<i32> = (0..PARTITION_COOPERATIVE_BUDGET + 16)
.map(|i| (i as i32) * 2 + 1)
.collect();
let (mut evens, _odds) = partition(
iter(source),
|x: &i32| x % 2 == 0,
PARTITION_COOPERATIVE_BUDGET * 4,
);
let woke = Arc::new(AtomicBool::new(false));
let waker = Waker::from(Arc::new(TrackWaker(woke.clone())));
let mut cx = Context::from_waker(&waker);
let polled = Pin::new(&mut evens).poll_next(&mut cx);
crate::assert_with_log!(
matches!(polled, Poll::Pending),
"routing yields at the cooperative budget",
"Poll::Pending",
polled
);
crate::assert_with_log!(
evens.peer_buffered_len() == PARTITION_COOPERATIVE_BUDGET,
"routed exactly the cooperative budget in one poll",
PARTITION_COOPERATIVE_BUDGET,
evens.peer_buffered_len()
);
crate::assert_with_log!(
woke.load(Ordering::SeqCst),
"self-wake requested so the yield is not a stall",
true,
woke.load(Ordering::SeqCst)
);
crate::test_complete!("partition_conformance::routing_yields_after_cooperative_budget");
}
}
}