use super::types::{
ConnectionCommand, CurrentRequest, HandleKind, RemoteCommand, RemoteConnectionHandle,
RemoteError, RemoteRenderHandle, RemoteRenderRequest, RemoteStream, RemoteUpdate,
RenderCommand, StreamEventFrame, host_from_address,
};
use crate::settings::WorkerSettings;
use indicatrix_net::{
client::{Accumulator, ApplyOutcome},
framing::{self, LEN_PREFIX_BYTES, MAX_FRAME_LEN},
messages::{NetError, RenderRequest, StreamEvent},
};
use socket2::{SockRef, TcpKeepalive};
use std::{
io::Read,
net::{TcpStream, ToSocketAddrs},
sync::{
Arc, Mutex, PoisonError,
mpsc::{self, TryRecvError},
},
thread,
time::{Duration, Instant},
};
const POLL_INTERVAL: Duration = Duration::from_millis(100);
const CONNECT_TIMEOUT: Duration = Duration::from_secs(5);
const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10);
const WRITE_TIMEOUT: Duration = Duration::from_secs(10);
const KEEPALIVE_TIME: Duration = Duration::from_secs(30);
const KEEPALIVE_INTERVAL: Duration = Duration::from_secs(10);
const LIVENESS_TIMEOUT: Duration = Duration::from_secs(8);
const FIRST_EVENT_TIMEOUT: Duration = Duration::from_secs(30);
const FRAME_REMAINDER_TIMEOUT: Duration = Duration::from_secs(5);
const fn liveness_deadline(seen_first_event: bool, liveness_timeout: Duration) -> Duration {
if seen_first_event {
liveness_timeout
} else {
FIRST_EVENT_TIMEOUT
}
}
pub fn spawn_remote_render(
request: RemoteRenderRequest,
accumulator: Arc<Mutex<Accumulator>>,
mut on_update: impl FnMut(super::types::RemoteUpdate) + Send + 'static,
) -> RemoteRenderHandle {
let (tx, rx) = mpsc::channel();
thread::spawn(move || {
let request_id = request.request_id;
let result = run(&request, &accumulator, &rx, &mut on_update);
if let Err(e) = result {
on_update(super::types::RemoteUpdate::Failed {
request_id,
message: e.to_string(),
});
}
});
RemoteRenderHandle(HandleKind::OneShot(tx))
}
#[must_use]
pub fn spawn_remote_connection(worker: WorkerSettings) -> RemoteConnectionHandle {
let (tx, rx) = mpsc::channel();
let worker_for_thread = worker.clone();
thread::spawn(move || run_connection(&worker_for_thread, &rx, LIVENESS_TIMEOUT));
RemoteConnectionHandle {
worker,
commands: tx,
}
}
fn open_tcp(address: &str) -> std::io::Result<TcpStream> {
let addr = address.to_socket_addrs()?.next().ok_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::InvalidInput,
format!("no address found for {address:?}"),
)
})?;
let tcp = TcpStream::connect_timeout(&addr, CONNECT_TIMEOUT)?;
tcp.set_nodelay(true)?;
let keepalive = TcpKeepalive::new()
.with_time(KEEPALIVE_TIME)
.with_interval(KEEPALIVE_INTERVAL);
SockRef::from(&tcp).set_tcp_keepalive(&keepalive)?;
Ok(tcp)
}
pub fn connect_and_handshake(
worker: &WorkerSettings,
) -> Result<(RemoteStream, indicatrix_net::messages::Welcome), RemoteError> {
let ca = indicatrix_net::tls::load_ca(&worker.ca_path())?;
let cert_chain = indicatrix_net::tls::load_certs(&worker.client_cert_path())?;
let key = indicatrix_net::tls::load_private_key(&worker.client_key_path())?;
let config = indicatrix_net::tls::client_config(ca, cert_chain, key)?;
let tcp = open_tcp(&worker.address)?;
let host = host_from_address(&worker.address);
let server_name = rustls::pki_types::ServerName::try_from(host.to_string())
.map_err(|_| RemoteError::InvalidServerName(host.to_string()))?;
let conn = rustls::ClientConnection::new(config, server_name)
.map_err(indicatrix_net::tls::TlsError::Rustls)?;
let mut stream = rustls::StreamOwned::new(conn, tcp);
stream
.sock
.set_read_timeout(Some(HANDSHAKE_TIMEOUT))
.map_err(RemoteError::Io)?;
stream
.sock
.set_write_timeout(Some(WRITE_TIMEOUT))
.map_err(RemoteError::Io)?;
stream
.conn
.complete_io(&mut stream.sock)
.map_err(RemoteError::Io)?;
let welcome = indicatrix_net::client::handshake::handshake(&mut stream)?;
stream
.sock
.set_read_timeout(None)
.map_err(RemoteError::Io)?;
Ok((stream, welcome))
}
pub fn test_connection(
worker: &WorkerSettings,
) -> Result<indicatrix_net::client::ConnectionInfo, RemoteError> {
let (_stream, welcome) = connect_and_handshake(worker)?;
Ok(welcome.into())
}
fn run(
request: &RemoteRenderRequest,
accumulator: &Arc<Mutex<Accumulator>>,
commands: &mpsc::Receiver<RemoteCommand>,
on_update: &mut dyn FnMut(super::types::RemoteUpdate),
) -> Result<(), RemoteError> {
run_with_liveness_timeout(request, accumulator, commands, on_update, LIVENESS_TIMEOUT)
}
fn run_with_liveness_timeout(
request: &RemoteRenderRequest,
accumulator: &Arc<Mutex<Accumulator>>,
commands: &mpsc::Receiver<RemoteCommand>,
on_update: &mut dyn FnMut(super::types::RemoteUpdate),
liveness_timeout: Duration,
) -> Result<(), RemoteError> {
let request_id = request.request_id;
let (mut stream, welcome) = connect_and_handshake(&request.worker)?;
on_update(RemoteUpdate::Connected {
request_id,
info: welcome.clone().into(),
});
let Some(capability) = welcome.render.as_ref() else {
return Err(RemoteError::NoRenderCapacity);
};
let render_request = RenderRequest {
request_id,
scene: request.scene.clone(),
first_sample: request.first_sample,
samples: request.samples,
stream: request
.worker
.export_stream_config(capability.min_cadence_ms),
};
{
let mut acc = accumulator.lock().unwrap_or_else(PoisonError::into_inner);
acc.begin_request(request_id);
}
indicatrix_net::client::send_render_request(&mut stream, &render_request)?;
let mut last_event = Instant::now();
let mut seen_first_event = false;
loop {
match commands.try_recv() {
Ok(RemoteCommand::Cancel) => {
indicatrix_net::client::send_cancel(&mut stream, request_id)?;
}
Err(TryRecvError::Empty | TryRecvError::Disconnected) => {
}
}
let Some((event, payload)) = try_read_stream_event(&mut stream, POLL_INTERVAL)? else {
let deadline = liveness_deadline(seen_first_event, liveness_timeout);
if last_event.elapsed() > deadline {
return Err(RemoteError::WorkerSilent(last_event.elapsed()));
}
continue;
};
last_event = Instant::now();
seen_first_event = true;
let outcome = {
let mut acc = accumulator.lock().unwrap_or_else(PoisonError::into_inner);
acc.apply(&event, payload.as_deref())
.map_err(|e| RemoteError::Client(e.into()))?
};
let is_terminal = matches!(
outcome,
ApplyOutcome::Done { .. } | ApplyOutcome::WorkerError
);
if let Some(update) = to_remote_update(request_id, &event, outcome) {
on_update(update);
}
if is_terminal {
return Ok(());
}
}
}
fn run_connection(
worker: &WorkerSettings,
commands: &mpsc::Receiver<ConnectionCommand>,
liveness_timeout: Duration,
) {
let mut stream: Option<RemoteStream> = None;
let mut welcome: Option<indicatrix_net::messages::Welcome> = None;
let mut current: Option<CurrentRequest> = None;
let mut last_event = Instant::now();
let mut seen_first_event = false;
loop {
match commands.try_recv() {
Ok(ConnectionCommand::Render(cmd)) => {
dispatch_render(worker, &mut stream, &mut welcome, &mut current, *cmd);
last_event = Instant::now();
seen_first_event = false;
}
Ok(ConnectionCommand::Cancel { request_id }) => {
if current.as_ref().is_some_and(|c| c.request_id == request_id)
&& let Some(s) = stream.as_mut()
&& indicatrix_net::client::send_cancel(s, request_id).is_err()
{
stream = None;
welcome = None;
}
}
Err(TryRecvError::Empty) => {}
Err(TryRecvError::Disconnected) => return,
}
let Some(s) = stream.as_mut() else {
thread::sleep(POLL_INTERVAL);
continue;
};
match try_read_stream_event(s, POLL_INTERVAL) {
Ok(Some((event, payload))) => {
last_event = Instant::now();
seen_first_event = true;
route_event(&mut current, &event, payload.as_deref());
}
Ok(None) => {
let deadline = liveness_deadline(seen_first_event, liveness_timeout);
if check_liveness(&mut current, last_event, deadline) {
stream = None;
welcome = None;
}
}
Err(e) => {
if let Some(mut cur) = current.take() {
(cur.on_update)(RemoteUpdate::Failed {
request_id: cur.request_id,
message: RemoteError::from(e).to_string(),
});
}
stream = None;
welcome = None;
}
}
}
}
fn check_liveness(
current: &mut Option<CurrentRequest>,
last_event: Instant,
timeout: Duration,
) -> bool {
if current.is_none() {
return false;
}
let elapsed = last_event.elapsed();
if elapsed <= timeout {
return false;
}
let mut cur = current.take().expect("just checked Some above");
(cur.on_update)(RemoteUpdate::Failed {
request_id: cur.request_id,
message: RemoteError::WorkerSilent(elapsed).to_string(),
});
true
}
fn dispatch_render(
worker: &WorkerSettings,
stream: &mut Option<RemoteStream>,
welcome: &mut Option<indicatrix_net::messages::Welcome>,
current: &mut Option<CurrentRequest>,
cmd: RenderCommand,
) {
let RenderCommand {
request,
accumulator,
mut on_update,
} = cmd;
let request_id = request.request_id;
if let Some(previous) = current.take()
&& let Some(s) = stream.as_mut()
{
let _ = indicatrix_net::client::send_cancel(s, previous.request_id);
}
if stream.is_none() {
match connect_and_handshake(worker) {
Ok((s, w)) => {
*stream = Some(s);
*welcome = Some(w);
}
Err(e) => {
on_update(RemoteUpdate::Failed {
request_id,
message: e.to_string(),
});
return; }
}
}
let w = welcome
.as_ref()
.expect("stream is Some at this point, and the two are always set together");
on_update(RemoteUpdate::Connected {
request_id,
info: w.clone().into(),
});
let Some(capability) = w.render.as_ref() else {
on_update(RemoteUpdate::Failed {
request_id,
message: RemoteError::NoRenderCapacity.to_string(),
});
return;
};
let render_request = RenderRequest {
request_id,
scene: request.scene,
first_sample: request.first_sample,
samples: request.samples,
stream: worker.stream_config(capability.min_cadence_ms, request.width, request.height),
};
{
let mut acc = accumulator.lock().unwrap_or_else(PoisonError::into_inner);
acc.begin_request(request_id);
}
let s = stream
.as_mut()
.expect("just connected above if this was None");
match indicatrix_net::client::send_render_request(s, &render_request) {
Ok(()) => {
*current = Some(CurrentRequest {
request_id,
accumulator,
on_update,
});
}
Err(e) => {
*stream = None;
*welcome = None;
on_update(RemoteUpdate::Failed {
request_id,
message: RemoteError::from(e).to_string(),
});
}
}
}
fn route_event(current: &mut Option<CurrentRequest>, event: &StreamEvent, payload: Option<&[u8]>) {
let Some(cur) = current.as_mut() else {
return;
};
let outcome = {
let mut acc = cur
.accumulator
.lock()
.unwrap_or_else(PoisonError::into_inner);
acc.apply(event, payload)
};
let request_id = cur.request_id;
let outcome = match outcome {
Ok(outcome) => outcome,
Err(e) => {
(cur.on_update)(RemoteUpdate::Failed {
request_id,
message: RemoteError::Client(e.into()).to_string(),
});
*current = None;
return;
}
};
let is_terminal = matches!(
outcome,
ApplyOutcome::Done { .. } | ApplyOutcome::WorkerError
);
if let Some(update) = to_remote_update(request_id, event, outcome) {
(cur.on_update)(update);
}
if is_terminal {
*current = None;
}
}
fn to_remote_update(
request_id: u32,
event: &StreamEvent,
outcome: ApplyOutcome,
) -> Option<super::types::RemoteUpdate> {
match outcome {
ApplyOutcome::FrameSummed { samples_done } => Some(RemoteUpdate::Frame {
request_id,
samples_done,
}),
ApplyOutcome::PreviewReplaced => Some(RemoteUpdate::Preview { request_id }),
ApplyOutcome::Progress { samples_done } => Some(RemoteUpdate::Progress {
request_id,
samples_done,
}),
ApplyOutcome::Done { cancelled } => Some(RemoteUpdate::Done {
request_id,
cancelled,
}),
ApplyOutcome::WorkerError => {
let StreamEvent::Error(e) = event else {
unreachable!("WorkerError only ever comes from applying an Error event")
};
Some(RemoteUpdate::Failed {
request_id,
message: e.message.clone(),
})
}
ApplyOutcome::StaleDropped => None,
}
}
fn try_read_one_frame(
stream: &mut RemoteStream,
poll_timeout: Duration,
) -> Result<Option<Vec<u8>>, NetError> {
stream
.sock
.set_read_timeout(Some(poll_timeout))
.map_err(|e| NetError::Framing(framing::FramingError::Io(e)))?;
let mut len_bytes = [0u8; LEN_PREFIX_BYTES];
let n = match stream.read(&mut len_bytes) {
Ok(0) => {
return Err(NetError::Framing(framing::FramingError::Io(
std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
"worker closed the connection",
),
)));
}
Ok(n) => n,
Err(e)
if e.kind() == std::io::ErrorKind::WouldBlock
|| e.kind() == std::io::ErrorKind::TimedOut =>
{
return Ok(None);
}
Err(e) => return Err(NetError::Framing(framing::FramingError::Io(e))),
};
stream
.sock
.set_read_timeout(Some(FRAME_REMAINDER_TIMEOUT))
.map_err(|e| NetError::Framing(framing::FramingError::Io(e)))?;
if n < len_bytes.len() {
stream
.read_exact(&mut len_bytes[n..])
.map_err(|e| NetError::Framing(framing::FramingError::Io(e)))?;
}
let len = u32::from_le_bytes(len_bytes);
if len > MAX_FRAME_LEN {
return Err(NetError::Framing(framing::FramingError::FrameTooLarge {
len,
max: MAX_FRAME_LEN,
}));
}
let mut payload = vec![0u8; len as usize];
stream
.read_exact(&mut payload)
.map_err(|e| NetError::Framing(framing::FramingError::Io(e)))?;
Ok(Some(payload))
}
fn try_read_stream_event(
stream: &mut RemoteStream,
poll_timeout: Duration,
) -> Result<Option<StreamEventFrame>, NetError> {
let Some(header_bytes) = try_read_one_frame(stream, poll_timeout)? else {
return Ok(None);
};
let event: StreamEvent = postcard::from_bytes(&header_bytes)?;
let expected_len = match &event {
StreamEvent::Frame(h) => Some(h.payload_len),
StreamEvent::Preview(h) => Some(h.payload_len),
StreamEvent::Progress(_) | StreamEvent::Done(_) | StreamEvent::Error(_) => None,
};
let payload = match expected_len {
Some(expected) => {
let bytes = framing::read_frame(stream).map_err(NetError::Framing)?;
if bytes.len() as u32 != expected {
return Err(NetError::FramePayloadLenMismatch {
declared: expected,
actual: bytes.len(),
});
}
Some(bytes)
}
None => None,
};
Ok(Some((event, payload)))
}
#[cfg(test)]
mod tests;