use std::sync::atomic::{AtomicUsize, Ordering};
use anyhow::bail;
use tokio::sync::Semaphore;
use super::session::OUTPUT_CHUNK;
use crate::rpc::ErrorResponse;
const MIN_OUTPUT_WINDOW: usize = 2 * OUTPUT_CHUNK;
pub(super) const STDIN_WINDOW: u32 = 256 * 1024;
pub(super) struct Flow {
output: Option<Semaphore>,
output_limit: usize,
stdin_queued: AtomicUsize,
}
impl Flow {
pub(super) fn new(output_window: u32) -> Result<Self, ErrorResponse> {
let output_limit = output_window as usize;
if output_limit != 0 && output_limit < MIN_OUTPUT_WINDOW {
return Err(ErrorResponse::new(
400,
format!("output_window must be 0 or at least {MIN_OUTPUT_WINDOW} bytes"),
));
}
Ok(Self {
output: (output_limit != 0).then(|| Semaphore::new(output_limit)),
output_limit,
stdin_queued: AtomicUsize::new(0),
})
}
pub(super) fn initial_stdin_window(&self) -> Option<u32> {
self.output.as_ref().map(|_| STDIN_WINDOW)
}
pub(super) async fn reserve_output(&self, len: usize) {
if let Some(window) = &self.output {
if let Ok(permit) = window.acquire_many(len as u32).await {
permit.forget();
}
}
}
pub(super) fn return_output(&self, bytes: u32) -> anyhow::Result<()> {
let Some(window) = &self.output else {
bail!("output window returned to a session without flow control");
};
if window.available_permits() + bytes as usize > self.output_limit {
bail!("host returned more output window than it was given");
}
window.add_permits(bytes as usize);
Ok(())
}
pub(super) fn admit_stdin(&self, len: usize) -> anyhow::Result<()> {
let queued = self.stdin_queued.fetch_add(len, Ordering::AcqRel) + len;
if self.output.is_some() && queued > STDIN_WINDOW as usize {
bail!("host sent {queued} bytes of stdin into a {STDIN_WINDOW}-byte window");
}
Ok(())
}
pub(super) fn stdin_delivered(&self, len: usize) -> Option<u32> {
self.stdin_queued.fetch_sub(len, Ordering::AcqRel);
self.output.as_ref().map(|_| len as u32)
}
}
#[cfg(test)]
mod tests {
use super::*;
const WINDOW: u32 = MIN_OUTPUT_WINDOW as u32;
#[test]
fn rejects_a_window_that_cannot_hold_a_full_output_frame() {
assert!(Flow::new(WINDOW - 1).is_err());
assert!(Flow::new(WINDOW).is_ok());
assert!(Flow::new(0).is_ok());
}
#[tokio::test]
async fn output_waits_for_the_window_the_host_returns() {
let flow = Flow::new(WINDOW).unwrap();
flow.reserve_output(WINDOW as usize).await;
let blocked =
tokio::time::timeout(std::time::Duration::from_millis(20), flow.reserve_output(1))
.await;
assert!(blocked.is_err(), "the window is used up");
flow.return_output(WINDOW).unwrap();
flow.reserve_output(WINDOW as usize).await;
}
#[test]
fn a_host_cannot_return_more_window_than_it_gave() {
let flow = Flow::new(WINDOW).unwrap();
assert!(flow.return_output(1).is_err());
assert!(Flow::new(0).unwrap().return_output(1).is_err());
}
#[test]
fn stdin_beyond_the_granted_window_is_refused() {
let flow = Flow::new(WINDOW).unwrap();
let window = STDIN_WINDOW as usize;
flow.admit_stdin(window).unwrap();
assert!(flow.admit_stdin(1).is_err());
let flow = Flow::new(WINDOW).unwrap();
flow.admit_stdin(window).unwrap();
assert_eq!(flow.stdin_delivered(window), Some(STDIN_WINDOW));
flow.admit_stdin(window).unwrap();
}
#[test]
fn without_flow_control_stdin_is_unlimited_and_never_acknowledged() {
let flow = Flow::new(0).unwrap();
assert_eq!(flow.initial_stdin_window(), None);
flow.admit_stdin(STDIN_WINDOW as usize + 1).unwrap();
assert_eq!(flow.stdin_delivered(1), None);
}
}