use anyhow::{Context, Result};
use skippy_protocol::{
StageActivationCodec, StageActivationCodecPolicy,
binary::{StageWireMessage, read_stage_message_for_codec_policy},
};
use std::io;
use std::net::{Shutdown, TcpStream};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
use std::sync::mpsc;
use std::thread;
use std::time::Duration;
use super::ConnectionWorkerControl;
use super::stale_discard::StaleDiscardRegistry;
static BINARY_SESSION_COUNTER: AtomicU64 = AtomicU64::new(1);
struct StoppableRead<'a> {
stream: &'a TcpStream,
stopped: &'a AtomicBool,
worker_control: &'a ConnectionWorkerControl,
}
impl io::Read for StoppableRead<'_> {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
loop {
match io::Read::read(&mut &*self.stream, buf) {
Err(error)
if matches!(
error.kind(),
io::ErrorKind::TimedOut | io::ErrorKind::WouldBlock
) =>
{
if self.stopped.load(Ordering::Acquire)
|| self.worker_control.is_shutting_down()
{
return Err(io::Error::new(
io::ErrorKind::ConnectionAborted,
"inbound message reader stopped mid-frame",
));
}
}
other => return other,
}
}
}
}
pub(super) fn next_connection_session_id() -> u64 {
BINARY_SESSION_COUNTER.fetch_add(1, Ordering::Relaxed)
}
pub(super) const INBOUND_LOOKAHEAD_MESSAGES: usize =
2 * skippy_protocol::MAX_VERIFY_WINDOW_PIPELINE_DEPTH;
pub(super) const INBOUND_LOOKAHEAD_BYTES: usize = 32 * 1024 * 1024;
pub(super) struct InboundMessageReader {
receiver: Option<mpsc::Receiver<io::Result<StageWireMessage>>>,
stream: Arc<TcpStream>,
thread: Option<thread::JoinHandle<()>>,
queued_bytes: Arc<AtomicUsize>,
stopped: Arc<AtomicBool>,
}
impl Drop for InboundMessageReader {
fn drop(&mut self) {
self.stopped.store(true, Ordering::Release);
drop(self.receiver.take());
let _ = self.stream.shutdown(Shutdown::Both);
if let Some(thread) = self.thread.take() {
let _ = thread.join();
}
}
}
pub(super) fn spawn_message_reader(
upstream: &TcpStream,
activation_width: i32,
activation_codec: StageActivationCodec,
activation_codec_policy: StageActivationCodecPolicy,
capacity: usize,
registry: Arc<StaleDiscardRegistry>,
worker_control: Arc<ConnectionWorkerControl>,
) -> Result<InboundMessageReader> {
let stream = Arc::new(
upstream
.try_clone()
.context("clone upstream stream for inbound message reader")?,
);
let reader = stream.clone();
let (sender, receiver) = mpsc::sync_channel(capacity.max(INBOUND_LOOKAHEAD_MESSAGES));
let queued_bytes = Arc::new(AtomicUsize::new(0));
let reader_queued_bytes = queued_bytes.clone();
let stopped = Arc::new(AtomicBool::new(false));
let reader_stopped = stopped.clone();
let thread = thread::spawn(move || {
loop {
while reader_queued_bytes.load(Ordering::Acquire) >= INBOUND_LOOKAHEAD_BYTES {
if reader_stopped.load(Ordering::Acquire) || worker_control.is_shutting_down() {
return;
}
thread::sleep(Duration::from_millis(1));
}
match worker_control.wait_for_readable(&reader) {
Ok(true) => {}
Ok(false) => return,
Err(error) => {
let _ = sender.send(Err(error));
return;
}
}
#[cfg(windows)]
let _ = reader.set_read_timeout(Some(super::WORKER_SHUTDOWN_POLL));
let mut source = StoppableRead {
stream: &reader,
stopped: &reader_stopped,
worker_control: &worker_control,
};
match read_stage_message_for_codec_policy(
&mut source,
activation_width,
activation_codec,
activation_codec_policy,
) {
Ok(message) => {
if message.kind.is_stale_window_discard() {
registry.record_message(&message);
}
let message_bytes = message.estimated_wire_bytes();
reader_queued_bytes.fetch_add(message_bytes, Ordering::AcqRel);
if sender.send(Ok(message)).is_err() {
return;
}
}
Err(error) => {
let _ = sender.send(Err(error));
return;
}
}
}
});
Ok(InboundMessageReader {
receiver: Some(receiver),
stream,
thread: Some(thread),
queued_bytes,
stopped,
})
}
impl InboundMessageReader {
pub(super) fn next(
&self,
first_message: Option<StageWireMessage>,
pending_prefill_replies: usize,
observed_message_count: usize,
) -> Result<Option<StageWireMessage>> {
if first_message.is_some() {
return Ok(first_message);
}
let receiver = self
.receiver
.as_ref()
.expect("inbound receiver present until drop");
match receiver.recv() {
Ok(Ok(message)) => {
let bytes = message.estimated_wire_bytes();
self.queued_bytes
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |queued| {
Some(queued.saturating_sub(bytes))
})
.ok();
Ok(Some(message))
}
Ok(Err(error))
if error.kind() == io::ErrorKind::UnexpectedEof
&& pending_prefill_replies == 0
&& observed_message_count == 0 =>
{
Ok(None)
}
Ok(Err(error)) => Err(error).context("read binary stage message"),
Err(_) => Ok(None),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use skippy_protocol::binary::{StageStateHeader, WireMessageKind, write_stage_message};
use std::io::Write;
use std::net::TcpListener;
use std::time::{Duration, Instant};
fn test_worker_control() -> Arc<ConnectionWorkerControl> {
Arc::new(ConnectionWorkerControl::default())
}
fn spawn_test_message_reader(
upstream: &TcpStream,
capacity: usize,
registry: Arc<StaleDiscardRegistry>,
) -> Result<InboundMessageReader> {
spawn_message_reader(
upstream,
4,
StageActivationCodec::default(),
StageActivationCodecPolicy::default(),
capacity,
registry,
test_worker_control(),
)
}
fn control_message(kind: WireMessageKind, tokens: Vec<i32>) -> StageWireMessage {
StageWireMessage {
kind,
pos_start: 0,
token_count: 0,
state: StageStateHeader::new(kind),
request_id: 7,
session_id: 9,
sampling: None,
chat_sampling_metadata: None,
tokens,
positions: Vec::new(),
activation: Vec::new(),
raw_bytes: Vec::new(),
}
}
fn connected_pair() -> (TcpStream, TcpStream) {
let listener = TcpListener::bind("127.0.0.1:0").expect("bind");
let address = listener.local_addr().expect("local addr");
let client = TcpStream::connect(address).expect("connect");
let (server, _) = listener.accept().expect("accept");
(client, server)
}
#[test]
fn discard_is_recorded_behind_a_backlog_larger_than_the_execution_queue() {
let (mut peer, upstream) = connected_pair();
let registry = Arc::new(StaleDiscardRegistry::default());
let reader = spawn_test_message_reader(&upstream, 1, registry.clone()).expect("spawn");
for _ in 0..skippy_protocol::MAX_VERIFY_WINDOW_PIPELINE_DEPTH {
let stale = control_message(WireMessageKind::Stop, Vec::new());
write_stage_message(&mut peer, &stale).expect("write");
}
let discard = control_message(WireMessageKind::DiscardStaleWindows, vec![3, 9]);
write_stage_message(&mut peer, &discard).expect("write");
let deadline = Instant::now() + Duration::from_secs(5);
while !registry.is_discarded(7, 9, 5) {
assert!(
Instant::now() < deadline,
"discard must be recorded while the stale backlog is still queued"
);
thread::sleep(Duration::from_millis(5));
}
drop(reader);
}
#[test]
fn dropping_the_reader_completes_while_it_is_parked_on_the_byte_ceiling() {
let (mut peer, upstream) = connected_pair();
let registry = Arc::new(StaleDiscardRegistry::default());
let reader = spawn_test_message_reader(&upstream, 1, registry).expect("spawn");
reader
.queued_bytes
.store(INBOUND_LOOKAHEAD_BYTES, Ordering::Release);
write_stage_message(
&mut peer,
&control_message(WireMessageKind::Stop, Vec::new()),
)
.expect("write");
thread::sleep(Duration::from_millis(50));
let (done, dropped) = mpsc::channel();
thread::spawn(move || {
drop(reader);
let _ = done.send(());
});
dropped
.recv_timeout(Duration::from_secs(5))
.expect("a reader parked on the byte ceiling must still be released on drop");
drop(peer);
}
#[test]
fn dropping_the_reader_completes_while_the_lookahead_channel_is_full() {
let (mut peer, upstream) = connected_pair();
let registry = Arc::new(StaleDiscardRegistry::default());
let reader = spawn_test_message_reader(&upstream, 1, registry).expect("spawn");
for _ in 0..(INBOUND_LOOKAHEAD_MESSAGES + 4) {
let message = control_message(WireMessageKind::Stop, Vec::new());
write_stage_message(&mut peer, &message).expect("write");
}
thread::sleep(Duration::from_millis(100));
let (done, dropped) = mpsc::channel();
thread::spawn(move || {
drop(reader);
let _ = done.send(());
});
dropped
.recv_timeout(Duration::from_secs(5))
.expect("dropping the receiver must unblock the queued send before the join");
drop(peer);
}
#[test]
fn dropping_the_reader_completes_while_a_read_is_stalled_mid_frame() {
let (mut peer, upstream) = connected_pair();
let registry = Arc::new(StaleDiscardRegistry::default());
let reader = spawn_test_message_reader(&upstream, 1, registry).expect("spawn");
let mut frame = Vec::new();
write_stage_message(
&mut frame,
&control_message(WireMessageKind::Stop, Vec::new()),
)
.expect("encode");
peer.write_all(&frame[..4]).expect("write frame prefix");
peer.flush().expect("flush");
thread::sleep(Duration::from_millis(100));
let (done, dropped) = mpsc::channel();
let dropper = thread::spawn(move || {
drop(reader);
let _ = done.send(());
});
dropped
.recv_timeout(Duration::from_secs(5))
.expect("drop must interrupt a read stalled mid-frame, not block in join");
dropper.join().expect("dropper thread");
drop(peer);
}
#[test]
fn a_frame_slower_than_the_read_timeout_still_arrives_whole() {
let (mut peer, upstream) = connected_pair();
let registry = Arc::new(StaleDiscardRegistry::default());
let reader = spawn_test_message_reader(&upstream, 1, registry).expect("spawn");
let mut frame = Vec::new();
write_stage_message(
&mut frame,
&control_message(WireMessageKind::DiscardStaleWindows, vec![3, 9]),
)
.expect("encode");
peer.write_all(&frame[..4]).expect("write frame prefix");
peer.flush().expect("flush");
thread::sleep(super::super::WORKER_SHUTDOWN_POLL * 3);
peer.write_all(&frame[4..]).expect("write frame rest");
peer.flush().expect("flush");
let message = reader
.next(None, 0, 0)
.expect("a slow frame must not fail the read")
.expect("a message, not a closed connection");
assert_eq!(message.kind, WireMessageKind::DiscardStaleWindows);
assert_eq!((message.request_id, message.session_id), (7, 9));
drop(reader);
}
#[test]
fn dropping_the_reader_joins_the_thread_while_the_peer_stays_open() {
let (peer, upstream) = connected_pair();
let registry = Arc::new(StaleDiscardRegistry::default());
let reader = spawn_test_message_reader(&upstream, 1, registry).expect("spawn");
let (done, dropped) = mpsc::channel();
thread::spawn(move || {
drop(reader);
let _ = done.send(());
});
dropped
.recv_timeout(Duration::from_secs(5))
.expect("drop must shut the socket down and join the blocked reader thread");
drop(peer);
}
}