use std::collections::VecDeque;
use std::time::Duration;
use tokio::io::BufReader;
use tokio::process::{ChildStdin, ChildStdout};
use tokio::sync::{mpsc, oneshot, watch};
use tokio::time::Instant;
use super::protocol::{Line, read_line, read_line_within, write_line};
use super::{BlockResult, Boundary, Delivery, Event};
use crate::Error;
use crate::internal::process::PersistentChild;
use crate::limits::ControlLimits;
const COMMAND_QUEUE: usize = 16;
pub(super) const EVENT_QUEUE: usize = 256;
pub(super) const HELD_WHILE_AWAITING: usize = EVENT_QUEUE * 8;
pub(super) struct OpenedConnection {
pub(super) commands: mpsc::Sender<Request>,
pub(super) events: mpsc::Receiver<Delivery>,
pub(super) stop: watch::Sender<()>,
pub(super) connection: tokio::task::JoinHandle<Result<(), Error>>,
}
pub(super) async fn open(
mut child: PersistentChild,
limits: ControlLimits,
timeout: Duration,
) -> Result<OpenedConnection, Error> {
let Some(stdin) = child.take_stdin() else {
let _ = child.terminate().await;
return Err(Error::control_mode_pipes());
};
let Some(stdout) = child.take_stdout() else {
let _ = child.terminate().await;
return Err(Error::control_mode_pipes());
};
let core_stopped = child.stopped();
let (commands, queue) = mpsc::channel(COMMAND_QUEUE);
let (events, received) = mpsc::channel(EVENT_QUEUE);
let (stop, stopped) = watch::channel(());
let actor = Connection {
child,
stdin,
stdout: BufReader::new(stdout),
limits,
timeout,
line: Vec::new(),
commands: queue,
events,
stopped,
core_stopped,
awaiting: ReplySlots::default(),
pending: VecDeque::new(),
};
let (ready, mut opened) = oneshot::channel();
let mut connection = tokio::spawn(actor.run(ready));
tokio::select! {
biased;
result = &mut opened => {
if result.is_err() {
return match connection.await {
Ok(Err(error)) => Err(error),
Ok(Ok(())) | Err(_) => Err(Error::control_mode_closed()),
};
}
}
result = &mut connection => {
return match result {
Ok(Err(error)) => Err(error),
Ok(Ok(())) | Err(_) => Err(Error::control_mode_closed()),
};
}
}
Ok(OpenedConnection {
commands,
events: received,
stop,
connection,
})
}
#[derive(Debug)]
pub(super) struct Request {
pub(super) line: String,
pub(super) deadline: Option<Instant>,
pub(super) result: oneshot::Sender<Result<BlockResult, Error>>,
pub(super) commit: oneshot::Sender<()>,
pub(super) boundary: Option<Boundary>,
}
#[derive(Debug)]
pub(super) struct CommittedRequest {
pub(super) line: String,
pub(super) deadline: Option<Instant>,
pub(super) result: oneshot::Sender<Result<BlockResult, Error>>,
pub(super) boundary: Option<Boundary>,
}
impl Request {
pub(super) fn commit(self) -> Option<CommittedRequest> {
let Self {
line,
deadline,
result,
commit,
boundary,
} = self;
commit.send(()).ok()?;
Some(CommittedRequest {
line,
deadline,
result,
boundary,
})
}
}
pub(super) fn admit_request(request: Request, pending_events: usize) -> Option<Request> {
if pending_events < HELD_WHILE_AWAITING {
return Some(request);
}
let _ = request.result.send(Err(Error::control_mode_unread()));
None
}
#[derive(Debug)]
pub(super) enum ReplySlot {
Live {
result: oneshot::Sender<Result<BlockResult, Error>>,
deadline: Option<Instant>,
boundary: Option<Boundary>,
},
Tombstone {
deadline: Option<Instant>,
boundary: Option<Boundary>,
},
}
impl ReplySlot {
const fn deadline(&self) -> Option<Instant> {
match self {
Self::Live { deadline, .. } | Self::Tombstone { deadline, .. } => *deadline,
}
}
const fn boundary(&self) -> Option<Boundary> {
match self {
Self::Live { boundary, .. } | Self::Tombstone { boundary, .. } => *boundary,
}
}
}
#[derive(Debug, Default)]
pub(super) struct ReplySlots {
pub(super) slots: VecDeque<ReplySlot>,
live: usize,
earliest: Option<Instant>,
}
impl ReplySlots {
#[cfg(test)]
pub(super) fn push(
&mut self,
result: oneshot::Sender<Result<BlockResult, Error>>,
deadline: Option<Instant>,
) {
self.push_ordered(result, deadline, None);
}
fn push_ordered(
&mut self,
result: oneshot::Sender<Result<BlockResult, Error>>,
deadline: Option<Instant>,
boundary: Option<Boundary>,
) {
self.earliest = earliest_deadline(self.earliest, deadline);
if result.is_closed() {
self.slots
.push_back(ReplySlot::Tombstone { deadline, boundary });
} else {
self.slots.push_back(ReplySlot::Live {
result,
deadline,
boundary,
});
self.live += 1;
}
}
fn front_boundary(&self) -> Option<Boundary> {
self.slots.front().and_then(ReplySlot::boundary)
}
pub(super) const fn has_live(&self) -> bool {
self.live != 0
}
fn has_slots(&self) -> bool {
!self.slots.is_empty()
}
pub(super) const fn earliest_deadline(&self) -> Option<Instant> {
self.earliest
}
fn block_deadline(&self, timeout: Duration) -> Option<Instant> {
if self.has_slots() {
self.earliest
} else {
Instant::now().checked_add(timeout)
}
}
pub(super) fn refuse_live(&mut self) {
for slot in &mut self.slots {
let deadline = slot.deadline();
let boundary = slot.boundary();
let ReplySlot::Live { result, .. } =
std::mem::replace(slot, ReplySlot::Tombstone { deadline, boundary })
else {
continue;
};
let _ = result.send(Err(Error::control_mode_unread()));
}
self.live = 0;
}
pub(super) fn complete(&mut self, block: BlockResult) {
let Some(slot) = self.slots.pop_front() else {
return;
};
let deadline = slot.deadline();
if let ReplySlot::Live { result, .. } = slot {
self.live -= 1;
let _ = result.send(Ok(block));
}
if deadline == self.earliest {
self.earliest = self.slots.iter().filter_map(ReplySlot::deadline).min();
}
}
fn fail_all(&mut self, mut reason: impl FnMut() -> Error) {
while let Some(slot) = self.slots.pop_front() {
if let ReplySlot::Live { result, .. } = slot {
let _ = result.send(Err(reason()));
}
}
self.live = 0;
self.earliest = None;
}
}
enum Step {
Deliver(bool),
Read(Result<Option<Line>, Error>),
Send(Option<Request>),
Unwatched {
asked: bool,
},
CoreStopped,
TimedOut,
}
enum BlockRead {
Complete(BlockResult),
Stopped,
}
#[derive(Clone, Copy)]
enum TerminalError {
Closed,
Frame(&'static str, usize),
Shutdown,
TimedOut,
}
impl TerminalError {
fn build(self, child: &PersistentChild) -> Error {
match self {
Self::Closed => Error::control_mode_closed(),
Self::Frame(frame, limit) => Error::control_mode_frame_too_large(frame, limit),
Self::Shutdown => child.shutdown_error(),
Self::TimedOut => Error::control_mode_timeout(),
}
}
}
struct Connection {
child: PersistentChild,
stdin: ChildStdin,
stdout: BufReader<ChildStdout>,
limits: ControlLimits,
timeout: Duration,
line: Vec<u8>,
commands: mpsc::Receiver<Request>,
events: mpsc::Sender<Delivery>,
stopped: watch::Receiver<()>,
core_stopped: watch::Receiver<bool>,
awaiting: ReplySlots,
pending: VecDeque<Delivery>,
}
impl Connection {
async fn run(mut self, mut ready: oneshot::Sender<()>) -> Result<(), Error> {
let opening_deadline = Instant::now().checked_add(self.timeout);
let (outcome, established) = match self
.discard_opening_block(&mut ready, opening_deadline)
.await
{
Ok(true) if ready.send(()).is_ok() => (self.serve().await, true),
Ok(_) => (Ok(()), false),
Err(error) => (Err(error), false),
};
let reason = match &outcome {
Err(Error::ControlModeFrameTooLarge { frame, limit }) => {
TerminalError::Frame(frame, *limit)
}
Err(Error::ControlMode {
kind: crate::ControlModeErrorKind::TimedOut,
..
}) => TerminalError::TimedOut,
Err(Error::ExecutorShutdown { .. }) => TerminalError::Shutdown,
_ => TerminalError::Closed,
};
let child = &self.child;
self.commands.close();
self.awaiting.fail_all(|| reason.build(child));
while let Ok(request) = self.commands.try_recv() {
let _ = request.result.send(Err(reason.build(child)));
}
let drained = if established {
self.drain_pending_events().await
} else {
Ok(())
};
drop(self.stdin);
let cleanup = self.child.terminate().await;
match outcome {
Err(error) => Err(error),
Ok(()) => drained.and(cleanup),
}
}
async fn serve(&mut self) -> Result<(), Error> {
let mut sending = true;
let mut watching = true;
while sending || watching {
if self.awaiting.has_live() && self.pending.len() >= HELD_WHILE_AWAITING {
self.awaiting.refuse_live();
}
let held_back = !self.awaiting.has_live() && self.pending.len() >= EVENT_QUEUE;
let reply_deadline = self.awaiting.earliest_deadline();
let step = tokio::select! {
line = read_line(&mut self.stdout, &mut self.line, self.limits.max_line_bytes),
if !held_back => Step::Read(line),
room = self.events.reserve(), if !self.pending.is_empty() => Step::Deliver(room.is_ok()),
request = self.commands.recv(), if sending => Step::Send(request),
asked = self.stopped.changed(), if watching => Step::Unwatched {
asked: asked.is_ok(),
},
() = cancellation_requested(&mut self.core_stopped) => Step::CoreStopped,
() = deadline_elapsed(reply_deadline), if self.awaiting.has_slots() => Step::TimedOut,
};
match step {
Step::Read(Err(error)) => return Err(error),
Step::Read(Ok(None)) => return Ok(()),
Step::Unwatched { asked: true } => {
return Ok(());
}
Step::CoreStopped => return Err(self.child.shutdown_error()),
Step::TimedOut => return Err(Error::control_mode_timeout()),
Step::Read(Ok(Some(line))) => {
if !self.dispatch(line, &mut watching).await? {
return Ok(());
}
}
Step::Send(Some(request)) => {
if request
.deadline
.is_some_and(|deadline| deadline <= Instant::now())
{
let _ = request
.result
.send(Err(Error::control_mode_dispatch_timeout()));
continue;
}
let Some(request) = admit_request(request, self.pending.len()) else {
continue;
};
let Some(request) = request.commit() else {
continue;
};
let write_deadline = earliest_deadline(reply_deadline, request.deadline);
let write = write_line(&mut self.stdin, &request.line);
tokio::pin!(write);
loop {
let result = tokio::select! {
biased;
() = cancellation_requested(&mut self.core_stopped) => {
let _ = request.result.send(Err(self.child.shutdown_error()));
return Err(self.child.shutdown_error());
}
changed = self.stopped.changed(), if watching => {
if changed.is_ok() {
let _ = request.result.send(Err(Error::control_mode_closed()));
return Ok(());
}
watching = false;
continue;
}
() = deadline_elapsed(write_deadline) => {
let _ = request.result.send(Err(Error::control_mode_timeout()));
return Err(Error::control_mode_timeout());
}
result = &mut write => result,
};
if let Err(error) = result {
let _ = request.result.send(Err(Error::control_mode_closed()));
return Err(error);
}
break;
}
self.awaiting
.push_ordered(request.result, request.deadline, request.boundary);
}
Step::Deliver(true) => {
if let Some(event) = self.pending.pop_front() {
let _ = self.events.try_send(event);
}
}
Step::Deliver(false) => self.pending.clear(),
Step::Send(None) => sending = false,
Step::Unwatched { asked: false } => watching = false,
}
}
Ok(())
}
async fn discard_opening_block(
&mut self,
ready: &mut oneshot::Sender<()>,
deadline: Option<Instant>,
) -> Result<bool, Error> {
loop {
let held_back = self.events.capacity() == 0;
let line = tokio::select! {
biased;
() = ready.closed() => return Ok(false),
() = cancellation_requested(&mut self.core_stopped) => {
return Err(self.child.shutdown_error());
}
() = deadline_elapsed(deadline) => {
return Err(Error::control_mode_timeout());
}
line = read_line(
&mut self.stdout,
&mut self.line,
self.limits.max_line_bytes,
), if !held_back => line?,
};
match line {
Some(Line::BlockStart(number)) => {
return match self.read_opening_block(number, ready, deadline).await? {
Some(true) => Ok(true),
Some(false) => Err(Error::control_mode_closed()),
None => Ok(false),
};
}
Some(Line::Event(exit @ Event::Exit { .. })) => {
let _ = self.events.try_send(Delivery::Event(exit));
return Err(Error::control_mode_closed());
}
Some(Line::Event(event)) => {
let _ = self.events.try_send(Delivery::Event(event));
}
Some(Line::Text(_) | Line::BlockEnd { .. }) => {}
None => return Err(Error::control_mode_closed()),
}
}
}
async fn read_opening_block(
&mut self,
number: u64,
ready: &mut oneshot::Sender<()>,
deadline: Option<Instant>,
) -> Result<Option<bool>, Error> {
let mut accumulated = 0usize;
loop {
let line = tokio::select! {
biased;
() = ready.closed() => return Ok(None),
() = cancellation_requested(&mut self.core_stopped) => {
return Err(self.child.shutdown_error());
}
() = deadline_elapsed(deadline) => {
return Err(Error::control_mode_timeout());
}
line = read_line_within(
&mut self.stdout,
&mut self.line,
self.limits.max_line_bytes,
Some(number),
) => line?,
};
match line {
Some(Line::BlockEnd {
number: end,
succeeded,
}) if end == number => return Ok(Some(succeeded)),
Some(Line::Text(text)) => {
accumulated = accumulated.saturating_add(text.as_bytes().len());
if accumulated > self.limits.max_block_bytes {
return Err(Error::control_mode_frame_too_large(
"block",
self.limits.max_block_bytes,
));
}
}
Some(Line::Event(_) | Line::BlockStart(_) | Line::BlockEnd { .. }) => {}
None => return Err(Error::control_mode_closed()),
}
}
}
async fn drain_pending_events(&mut self) -> Result<(), Error> {
while let Some(delivery) = self.pending.pop_front() {
let permit = tokio::select! {
biased;
() = cancellation_requested(&mut self.core_stopped) => {
return Err(self.child.shutdown_error());
}
_ = self.stopped.changed() => return Ok(()),
permit = self.events.reserve() => permit,
};
let Ok(permit) = permit else {
return Ok(());
};
permit.send(delivery);
}
Ok(())
}
async fn dispatch(&mut self, line: Line, watching: &mut bool) -> Result<bool, Error> {
match line {
Line::BlockStart(number) => {
let deadline = self.awaiting.block_deadline(self.timeout);
match self.read_block(number, deadline, watching).await? {
BlockRead::Complete(block) => {
if let Some(boundary) = self.awaiting.front_boundary() {
self.report(Delivery::Boundary(boundary));
}
self.awaiting.complete(block);
Ok(true)
}
BlockRead::Stopped => Ok(false),
}
}
Line::Event(exit @ Event::Exit { .. }) => {
self.report(Delivery::Event(exit));
Ok(false)
}
Line::Event(event) => {
self.report(Delivery::Event(event));
Ok(true)
}
Line::Text(_) | Line::BlockEnd { .. } => Ok(true),
}
}
fn report(&mut self, delivery: Delivery) {
if !self.pending.is_empty() {
self.pending.push_back(delivery);
return;
}
let Err(mpsc::error::TrySendError::Full(delivery)) = self.events.try_send(delivery) else {
return;
};
self.pending.push_back(delivery);
}
async fn read_block(
&mut self,
number: u64,
deadline: Option<Instant>,
watching: &mut bool,
) -> Result<BlockRead, Error> {
let mut output = Vec::new();
let mut accumulated = 0usize;
loop {
let line = tokio::select! {
biased;
() = cancellation_requested(&mut self.core_stopped) => {
return Err(self.child.shutdown_error());
}
changed = self.stopped.changed(), if *watching => {
if changed.is_ok() {
return Ok(BlockRead::Stopped);
}
*watching = false;
continue;
}
() = deadline_elapsed(deadline) => {
return Err(Error::control_mode_timeout());
}
line = read_line_within(
&mut self.stdout,
&mut self.line,
self.limits.max_line_bytes,
Some(number),
) => line?,
};
match line {
Some(Line::BlockEnd {
number: end,
succeeded,
}) if end == number => {
return Ok(BlockRead::Complete(BlockResult {
number,
succeeded,
output,
sensitive_input: false,
}));
}
Some(Line::Text(text)) => {
accumulated = accumulated.saturating_add(text.as_bytes().len());
if accumulated > self.limits.max_block_bytes {
return Err(Error::control_mode_frame_too_large(
"block",
self.limits.max_block_bytes,
));
}
output.push(text);
}
Some(Line::Event(_) | Line::BlockStart(_) | Line::BlockEnd { .. }) => {}
None => return Err(Error::control_mode_closed()),
}
}
}
}
async fn cancellation_requested(stopped: &mut watch::Receiver<bool>) {
loop {
if *stopped.borrow() {
return;
}
if stopped.changed().await.is_err() {
return;
}
}
}
pub(super) async fn deadline_elapsed(deadline: Option<Instant>) {
match deadline {
Some(deadline) => tokio::time::sleep_until(deadline).await,
None => std::future::pending().await,
}
}
fn earliest_deadline(left: Option<Instant>, right: Option<Instant>) -> Option<Instant> {
match (left, right) {
(Some(left), Some(right)) => Some(left.min(right)),
(Some(deadline), None) | (None, Some(deadline)) => Some(deadline),
(None, None) => None,
}
}