use std::io::{self, Read};
use std::time::{Duration, Instant};
use tokio::io::AsyncWriteExt as _;
use super::{
lex::{LexOutput, LexWriter},
syntax as s,
};
use crate::support::async_io::ServerIo;
pub enum OutputEvent {
ResponseLine {
line: s::ResponseLine<'static>,
ctl: OutputControl,
},
ContinuationLine {
prompt: &'static str,
},
Flush,
EnableUnicode,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum OutputControl {
Buffer,
Flush,
EnableCompression,
Disconnect,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum OutputDisconnect {
ByControl,
InputClosed,
}
pub async fn write_responses(
mut io: ServerIo,
mut outputs: tokio::sync::mpsc::Receiver<OutputEvent>,
) -> io::Result<OutputDisconnect> {
let mut state = State::new();
while let Some(evt) = outputs.recv().await {
if state.text.is_empty() {
state.last_flush = Instant::now();
}
let ctl = match evt {
OutputEvent::ResponseLine { mut line, ctl } => {
let unicode = state.unicode;
line.write_to(&mut LexWriter::new(&mut state, unicode, false))?;
state.text.extend_from_slice(b"\r\n");
ctl
},
OutputEvent::ContinuationLine { prompt } => {
state.text.extend_from_slice(b"+ ");
state.text.extend_from_slice(prompt.as_bytes());
state.text.extend_from_slice(b"\r\n");
OutputControl::Flush
},
OutputEvent::Flush => OutputControl::Flush,
OutputEvent::EnableUnicode => {
state.unicode = true;
continue;
},
};
match ctl {
OutputControl::Buffer => {
let flush_due_to_size = state.text.len() >= TEXT_FLUSH_THRESH
|| state.splices.len() >= SPLICE_FLUSH_THRESH;
let flush_due_to_time =
state.last_flush.elapsed() >= Duration::from_secs(3);
if flush_due_to_size || flush_due_to_time {
let flush_compress = if flush_due_to_time {
flate2::FlushCompress::Sync
} else {
flate2::FlushCompress::None
};
state.flush(&mut io, flush_compress).await?;
}
},
OutputControl::Flush => {
state.flush(&mut io, flate2::FlushCompress::Sync).await?;
},
OutputControl::EnableCompression => {
assert!(state.compress.is_none());
state.flush(&mut io, flate2::FlushCompress::None).await?;
state.compress = Some(flate2::Compress::new(
flate2::Compression::new(3),
false,
));
state.compressed = vec![0u8; TEXT_FLUSH_THRESH];
},
OutputControl::Disconnect => {
state.flush(&mut io, flate2::FlushCompress::Finish).await?;
return Ok(OutputDisconnect::ByControl);
},
}
}
state.flush(&mut io, flate2::FlushCompress::Finish).await?;
Ok(OutputDisconnect::InputClosed)
}
const TEXT_FLUSH_THRESH: usize = 4096;
const SPLICE_FLUSH_THRESH: usize = 4;
struct State {
text: Vec<u8>,
splices: Vec<LiteralSplice>,
splice_read: Vec<u8>,
compress: Option<flate2::Compress>,
compressed: Vec<u8>,
last_flush: Instant,
unicode: bool,
}
struct LiteralSplice {
offset: usize,
data: Box<dyn Read>,
}
impl State {
fn new() -> Self {
Self {
text: Vec::with_capacity(TEXT_FLUSH_THRESH * 5 / 4),
splices: Vec::with_capacity(SPLICE_FLUSH_THRESH * 2),
splice_read: vec![0; 4096],
compress: None,
compressed: Vec::new(),
last_flush: Instant::now(),
unicode: false,
}
}
async fn flush(
&mut self,
io: &mut ServerIo,
flush_compress: flate2::FlushCompress,
) -> io::Result<()> {
tokio::time::timeout(
Duration::from_secs(60),
self.flush_impl(io, flush_compress),
)
.await
.unwrap_or_else(|_| Err(io::ErrorKind::TimedOut.into()))
}
async fn flush_impl(
&mut self,
io: &mut ServerIo,
flush_compress: flate2::FlushCompress,
) -> io::Result<()> {
#[allow(clippy::collapsible_else_if)] async fn do_write(
io: &mut ServerIo,
compress: Option<&mut flate2::Compress>,
compressed: &mut [u8],
mut data: &[u8],
) -> io::Result<()> {
if let Some(compress) = compress {
while !data.is_empty() {
let before_in = compress.total_in();
let before_out = compress.total_out();
compress
.compress(data, compressed, flate2::FlushCompress::None)
.map_err(io::Error::other)?;
let after_in = compress.total_in();
let after_out = compress.total_out();
data = &data[(after_in - before_in) as usize..];
if after_out != before_out {
io.write_all(
&compressed[..(after_out - before_out) as usize],
)
.await?;
}
}
} else {
if !data.is_empty() {
io.write_all(data).await?;
}
}
Ok(())
}
let mut offset = 0usize;
for mut splice in self.splices.drain(..) {
if splice.offset > offset {
do_write(
io,
self.compress.as_mut(),
&mut self.compressed,
&self.text[offset..splice.offset],
)
.await?;
offset = splice.offset;
}
loop {
let nread = splice.data.read(&mut self.splice_read)?;
if 0 == nread {
break;
}
do_write(
io,
self.compress.as_mut(),
&mut self.compressed,
&self.splice_read[..nread],
)
.await?;
}
}
if offset < self.text.len() {
do_write(
io,
self.compress.as_mut(),
&mut self.compressed,
&self.text[offset..],
)
.await?;
}
if flate2::FlushCompress::None != flush_compress {
if let Some(ref mut compress) = self.compress {
loop {
let before_out = compress.total_out();
compress
.compress(&[], &mut self.compressed, flush_compress)
.map_err(io::Error::other)?;
let after_out = compress.total_out();
if after_out == before_out {
break;
}
io.write_all(
&self.compressed[..(after_out - before_out) as usize],
)
.await?;
}
}
}
self.text.clear();
self.last_flush = Instant::now();
Ok(())
}
}
impl io::Write for &mut State {
fn write(&mut self, data: &[u8]) -> io::Result<usize> {
self.text.extend_from_slice(data);
Ok(data.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
impl LexOutput for &mut State {
fn splice<R: Read + 'static>(&mut self, data: R) -> io::Result<()> {
self.splices.push(LiteralSplice {
offset: self.text.len(),
data: Box::new(data),
});
Ok(())
}
}