use std::future::Future;
use std::pin::pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::task::Poll;
use arcbox_engine::agent_client::ExecSessionInput;
use tokio::sync::{Notify, mpsc};
const HIGH_WATER: usize = 1024 * 1024;
pub const LIMIT: usize = 64 * 1024 * 1024;
const PROCESS_CAPACITY: usize = 4;
#[derive(Debug)]
pub struct Overflow;
#[derive(Default)]
pub struct Outbound {
blocked: AtomicUsize,
blocked_one: Notify,
}
impl Outbound {
pub async fn send<F: Future>(&self, send: F) -> F::Output {
let mut send = pin!(send);
if let Poll::Ready(done) =
std::future::poll_fn(|cx| Poll::Ready(send.as_mut().poll(cx))).await
{
return done;
}
let _blocked = Blocked::new(self);
send.await
}
fn any_blocked(&self) -> bool {
self.blocked.load(Ordering::Acquire) > 0
}
}
struct Blocked<'a>(&'a Outbound);
impl<'a> Blocked<'a> {
fn new(outbound: &'a Outbound) -> Self {
outbound.blocked.fetch_add(1, Ordering::AcqRel);
outbound.blocked_one.notify_waiters();
Self(outbound)
}
}
impl Drop for Blocked<'_> {
fn drop(&mut self) {
self.0.blocked.fetch_sub(1, Ordering::AcqRel);
}
}
pub struct SessionInput {
queue: mpsc::UnboundedSender<ExecSessionInput>,
backlog: Arc<Backlog>,
}
#[derive(Default)]
struct Backlog {
bytes: AtomicUsize,
moved: Notify,
closed: AtomicBool,
}
impl SessionInput {
pub fn new() -> (Self, mpsc::Receiver<ExecSessionInput>) {
let (process, taken) = mpsc::channel(PROCESS_CAPACITY);
let (queue, pending) = mpsc::unbounded_channel();
let backlog = Arc::new(Backlog::default());
tokio::spawn(hand_on(pending, process, Arc::clone(&backlog)));
(Self { queue, backlog }, taken)
}
pub async fn send(&self, input: ExecSessionInput, outbound: &Outbound) -> Result<(), Overflow> {
let len = stdin_len(&input);
let queued = self.backlog.bytes.fetch_add(len, Ordering::AcqRel) + len;
let _ = self.queue.send(input);
if queued > LIMIT {
return Err(Overflow);
}
if len > 0 {
self.backlog.wait_for_room(outbound).await;
}
Ok(())
}
}
impl Backlog {
async fn wait_for_room(&self, outbound: &Outbound) {
loop {
let mut moved = pin!(self.moved.notified());
let mut blocked = pin!(outbound.blocked_one.notified());
moved.as_mut().enable();
blocked.as_mut().enable();
if self.bytes.load(Ordering::Acquire) <= HIGH_WATER
|| self.closed.load(Ordering::Acquire)
|| outbound.any_blocked()
{
return;
}
tokio::select! {
() = moved => {}
() = blocked => {}
}
}
}
}
async fn hand_on(
mut pending: mpsc::UnboundedReceiver<ExecSessionInput>,
process: mpsc::Sender<ExecSessionInput>,
backlog: Arc<Backlog>,
) {
while let Some(input) = pending.recv().await {
let len = stdin_len(&input);
let taken = process.send(input).await;
backlog.bytes.fetch_sub(len, Ordering::AcqRel);
backlog.moved.notify_waiters();
if taken.is_err() {
break;
}
}
backlog.closed.store(true, Ordering::Release);
backlog.moved.notify_waiters();
}
fn stdin_len(input: &ExecSessionInput) -> usize {
match input {
ExecSessionInput::Stdin(data) => data.len(),
ExecSessionInput::Resize { .. } | ExecSessionInput::Signal(_) => 0,
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::*;
const CHUNK: usize = 256 * 1024;
fn stdin() -> ExecSessionInput {
ExecSessionInput::Stdin(vec![0; CHUNK])
}
async fn backlogged(input: &SessionInput, outbound: &Outbound) {
while input.backlog.bytes.load(Ordering::Acquire) <= HIGH_WATER {
let queued =
tokio::time::timeout(Duration::from_millis(100), input.send(stdin(), outbound))
.await;
if queued.is_err() {
return;
}
}
}
async fn completes<F: Future>(future: F) -> bool {
tokio::time::timeout(Duration::from_millis(100), future)
.await
.is_ok()
}
#[tokio::test]
async fn stdin_past_the_high_water_mark_waits_for_the_process() {
let (input, mut taken) = SessionInput::new();
let outbound = Outbound::default();
backlogged(&input, &outbound).await;
let send = input.send(stdin(), &outbound);
let mut send = pin!(send);
assert!(
!completes(send.as_mut()).await,
"the queue is past the mark"
);
while input.backlog.bytes.load(Ordering::Acquire) > HIGH_WATER {
taken.recv().await.unwrap();
}
assert!(completes(send).await, "the process caught up");
}
#[tokio::test]
async fn control_input_never_waits() {
let (input, _taken) = SessionInput::new();
let outbound = Outbound::default();
backlogged(&input, &outbound).await;
let resize = ExecSessionInput::Resize {
width: 80,
height: 24,
};
assert!(completes(input.send(resize, &outbound)).await);
}
#[tokio::test]
async fn output_blocked_on_the_connection_releases_waiting_input() {
let (input, _taken) = SessionInput::new();
let outbound = Outbound::default();
backlogged(&input, &outbound).await;
let send = input.send(stdin(), &outbound);
let mut send = pin!(send);
assert!(!completes(send.as_mut()).await);
let never = outbound.send(std::future::pending::<()>());
let mut never = pin!(never);
assert!(!completes(never.as_mut()).await);
assert!(completes(send).await);
assert!(completes(input.send(stdin(), &outbound)).await);
}
#[tokio::test]
async fn output_that_completes_at_once_does_not_count_as_blocked() {
let outbound = Outbound::default();
outbound.send(std::future::ready(())).await;
assert!(!outbound.any_blocked());
}
#[tokio::test]
async fn a_session_past_the_limit_overflows() {
let (input, _taken) = SessionInput::new();
let outbound = Outbound::default();
let never = outbound.send(std::future::pending::<()>());
let mut never = pin!(never);
assert!(!completes(never.as_mut()).await);
let chunk = vec![0; 4 * 1024 * 1024];
let mut queued = 0;
loop {
queued += chunk.len();
let send = input.send(ExecSessionInput::Stdin(chunk.clone()), &outbound);
let sent = tokio::time::timeout(Duration::from_secs(1), send)
.await
.expect("input does not wait while output is blocked");
if sent.is_err() {
break;
}
}
assert!(queued > LIMIT);
}
#[tokio::test]
async fn a_process_that_is_gone_releases_waiting_input() {
let (input, taken) = SessionInput::new();
let outbound = Outbound::default();
backlogged(&input, &outbound).await;
let send = input.send(stdin(), &outbound);
let mut send = pin!(send);
assert!(!completes(send.as_mut()).await);
drop(taken);
assert!(completes(send).await);
}
}