use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex, mpsc};
use std::time::{Duration, Instant};
use tau_proto::{Event, ProviderResponseStats};
use super::RendererCmd;
use super::cold_attach_stager::RendererPresentation;
use super::delivery_memory::{DeliveryMemoryCut, DeliveryMemoryTracker};
#[cfg(test)]
#[path = "renderer_scheduler/tests.rs"]
mod tests;
pub(super) struct RemoteRendererSender {
tx: Option<mpsc::SyncSender<RendererCmd>>,
wake: tau_blocking_notify_channel::Sender,
}
impl RemoteRendererSender {
pub(super) fn channel(
capacity: usize,
wake: tau_blocking_notify_channel::Sender,
) -> (Self, mpsc::Receiver<RendererCmd>) {
let (tx, rx) = mpsc::sync_channel(capacity);
(Self { tx: Some(tx), wake }, rx)
}
pub(super) fn send(&self, cmd: RendererCmd) -> Result<(), mpsc::SendError<RendererCmd>> {
let result = self.tx.as_ref().expect("live remote sender").send(cmd);
if result.is_ok() {
self.wake.notify();
}
result
}
}
impl Drop for RemoteRendererSender {
fn drop(&mut self) {
drop(self.tx.take());
self.wake.notify();
}
}
struct LocalRendererCmd {
cmd: RendererCmd,
remote_watermark: u64,
}
pub(super) struct LocalRendererReceiver {
rx: mpsc::Receiver<LocalRendererCmd>,
}
impl LocalRendererReceiver {
#[cfg(test)]
pub(super) fn try_recv(&self) -> Result<RendererCmd, mpsc::TryRecvError> {
self.rx.try_recv().map(|local| local.cmd)
}
}
pub(super) struct LocalRendererSender {
tx: Option<mpsc::Sender<LocalRendererCmd>>,
remote_admitted: Arc<AtomicU64>,
arbiter: Arc<Mutex<()>>,
wake: tau_blocking_notify_channel::Sender,
}
impl LocalRendererSender {
pub(super) fn channel(
remote_admitted: Arc<AtomicU64>,
arbiter: Arc<Mutex<()>>,
wake: tau_blocking_notify_channel::Sender,
) -> (Self, LocalRendererReceiver) {
let (tx, rx) = mpsc::channel();
(
Self {
tx: Some(tx),
remote_admitted,
arbiter,
wake,
},
LocalRendererReceiver { rx },
)
}
pub(super) fn send(&self, cmd: RendererCmd) -> Result<(), mpsc::SendError<RendererCmd>> {
let _guard = self
.arbiter
.lock()
.expect("renderer arbiter mutex poisoned");
let local = LocalRendererCmd {
cmd,
remote_watermark: self.remote_admitted.load(Ordering::Acquire),
};
let result = self
.tx
.as_ref()
.expect("live local sender")
.send(local)
.map_err(|error| mpsc::SendError(error.0.cmd));
if result.is_ok() {
self.wake.notify();
}
result
}
}
impl Clone for LocalRendererSender {
fn clone(&self) -> Self {
Self {
tx: self.tx.clone(),
remote_admitted: self.remote_admitted.clone(),
arbiter: self.arbiter.clone(),
wake: self.wake.clone(),
}
}
}
impl Drop for LocalRendererSender {
fn drop(&mut self) {
drop(self.tx.take());
self.wake.notify();
}
}
pub(super) struct RendererCommandScheduler {
remote_rx: mpsc::Receiver<RendererCmd>,
remote_admitted: Arc<AtomicU64>,
local_rx: LocalRendererReceiver,
remote_closed: bool,
remote_processed: u64,
pending_remote: Option<RendererCmd>,
pending_local: Option<LocalRendererCmd>,
arbiter: Arc<Mutex<()>>,
wake: tau_blocking_notify_channel::Receiver,
delivery_memory: Option<Arc<DeliveryMemoryTracker>>,
}
impl RendererCommandScheduler {
pub(super) fn new(
remote_rx: mpsc::Receiver<RendererCmd>,
local_rx: LocalRendererReceiver,
remote_admitted: Arc<AtomicU64>,
arbiter: Arc<Mutex<()>>,
wake: tau_blocking_notify_channel::Receiver,
delivery_memory: Option<Arc<DeliveryMemoryTracker>>,
) -> Self {
Self {
remote_rx,
remote_admitted,
local_rx,
remote_closed: false,
remote_processed: 0,
pending_remote: None,
pending_local: None,
arbiter,
wake,
delivery_memory,
}
}
#[cfg(test)]
pub(super) fn remote_closed(&self) -> bool {
self.remote_closed
}
pub(super) fn recv_timeout(
&mut self,
timeout: Duration,
) -> Result<RendererCmd, mpsc::RecvTimeoutError> {
self.recv_timeout_inner(timeout, None, None, false)
}
#[cfg(test)]
pub(super) fn recv_timeout_after_local_check(
&mut self,
timeout: Duration,
hook: &mut dyn FnMut(),
) -> Result<RendererCmd, mpsc::RecvTimeoutError> {
self.recv_timeout_inner(timeout, Some(hook), None, false)
}
#[cfg(test)]
pub(super) fn recv_timeout_before_wait(
&mut self,
timeout: Duration,
hook: &mut dyn FnMut(),
) -> Result<RendererCmd, mpsc::RecvTimeoutError> {
self.recv_timeout_inner(timeout, None, Some(hook), false)
}
#[cfg(test)]
pub(super) fn recv_timeout_before_each_wait(
&mut self,
timeout: Duration,
hook: &mut dyn FnMut(),
) -> Result<RendererCmd, mpsc::RecvTimeoutError> {
self.recv_timeout_inner(timeout, None, Some(hook), true)
}
fn recv_timeout_inner(
&mut self,
timeout: Duration,
mut after_local_check: Option<&mut dyn FnMut()>,
mut before_wait: Option<&mut dyn FnMut()>,
repeat_before_wait: bool,
) -> Result<RendererCmd, mpsc::RecvTimeoutError> {
let started_at = Instant::now();
let mut deadline_elapsed = false;
loop {
let arbiter = self.arbiter.clone();
let guard = arbiter.lock().expect("renderer arbiter mutex poisoned");
if self.pending_local.is_none() {
match self.local_rx.rx.try_recv() {
Ok(local) => self.pending_local = Some(local),
Err(mpsc::TryRecvError::Disconnected) if self.remote_closed => {
return Err(mpsc::RecvTimeoutError::Disconnected);
}
Err(mpsc::TryRecvError::Disconnected | mpsc::TryRecvError::Empty) => {}
}
}
if self.pending_local.as_ref().is_some_and(|local| {
self.remote_closed || local.remote_watermark <= self.remote_processed
}) {
return Ok(self
.pending_local
.take()
.expect("pending local command")
.cmd);
}
if self.pending_local.is_some() {
match self.try_recv_remote() {
Ok(cmd) => return Ok(self.dequeue_remote(cmd)),
Err(mpsc::TryRecvError::Disconnected) => {
self.remote_closed = true;
continue;
}
Err(mpsc::TryRecvError::Empty) => {}
}
}
if self.remote_closed
&& self.pending_local.is_none()
&& matches!(
self.local_rx.rx.try_recv(),
Err(mpsc::TryRecvError::Disconnected)
)
{
return Err(mpsc::RecvTimeoutError::Disconnected);
}
if let Some(hook) = after_local_check.take() {
hook();
}
if !self.remote_closed {
match self.try_recv_remote() {
Ok(cmd) => return Ok(self.dequeue_remote(cmd)),
Err(mpsc::TryRecvError::Disconnected) => {
self.remote_closed = true;
continue;
}
Err(mpsc::TryRecvError::Empty) => {}
}
}
drop(guard);
if repeat_before_wait {
if let Some(hook) = before_wait.as_deref_mut() {
hook();
}
} else if let Some(hook) = before_wait.take() {
hook();
}
if deadline_elapsed || started_at.elapsed() >= timeout {
return Err(mpsc::RecvTimeoutError::Timeout);
}
let remaining = timeout.saturating_sub(started_at.elapsed());
match self.wake.recv_timeout(remaining) {
Ok(()) | Err(tau_blocking_notify_channel::RecvTimeoutError::Disconnected) => {}
Err(tau_blocking_notify_channel::RecvTimeoutError::Timeout) => {
deadline_elapsed = true;
}
}
}
}
fn try_recv_remote(&mut self) -> Result<RendererCmd, mpsc::TryRecvError> {
let result = self
.pending_remote
.take()
.map_or_else(|| self.remote_rx.try_recv(), Ok);
if let Ok(
RendererCmd::Remote { delivery_id, .. }
| RendererCmd::RemoteDisconnect { delivery_id, .. },
) = &result
&& let Some(memory) = &self.delivery_memory
{
memory.transition(*delivery_id, DeliveryMemoryCut::Scheduler);
}
result
}
fn dequeue_remote(&mut self, mut cmd: RendererCmd) -> RendererCmd {
self.remote_processed = self.remote_processed.saturating_add(1);
let captured = self.remote_admitted.load(Ordering::Acquire);
let fold_before = self
.pending_local
.as_ref()
.map_or(captured, |local| captured.min(local.remote_watermark));
while self.remote_processed < fold_before && is_pure_provider_update(&cmd) {
let next = match self.try_recv_remote() {
Ok(next) => next,
Err(mpsc::TryRecvError::Empty) => break,
Err(mpsc::TryRecvError::Disconnected) => {
self.remote_closed = true;
break;
}
};
match fold_provider_update(cmd, next) {
(folded, None) => {
cmd = folded;
self.remote_processed = self.remote_processed.saturating_add(1);
}
(current, Some(barrier)) => {
cmd = current;
self.pending_remote = Some(barrier);
break;
}
}
}
cmd
}
}
fn is_pure_provider_update(cmd: &RendererCmd) -> bool {
matches!(
cmd,
RendererCmd::Remote {
event,
presentation: RendererPresentation::Ordinary,
abandoned_shell_starts,
..
} if abandoned_shell_starts.is_empty()
&& matches!(event.as_ref(), Event::ProviderResponseUpdated(update)
if update.status.is_none() && update.compaction.is_none())
)
}
fn fold_provider_update(
mut current: RendererCmd,
next: RendererCmd,
) -> (RendererCmd, Option<RendererCmd>) {
let RendererCmd::Remote {
event: next_event,
presentation: RendererPresentation::Ordinary,
abandoned_shell_starts: next_abandoned,
folded_frames: next_folded,
..
} = &next
else {
return (current, Some(next));
};
if !next_abandoned.is_empty() || !next_folded.is_empty() {
return (current, Some(next));
}
let Event::ProviderResponseUpdated(next_update) = next_event.as_ref() else {
return (current, Some(next));
};
if next_update.status.is_some() || next_update.compaction.is_some() {
return (current, Some(next));
}
let RendererCmd::Remote { event, .. } = ¤t else {
return (current, Some(next));
};
let Event::ProviderResponseUpdated(update) = event.as_ref() else {
return (current, Some(next));
};
if update.agent_id != next_update.agent_id
|| update.agent_prompt_id != next_update.agent_prompt_id
|| update.originator != next_update.originator
{
return (current, Some(next));
}
let RendererCmd::Remote {
event: next_event,
delivery_id,
queue_bytes,
enqueued_at,
..
} = next
else {
unreachable!("validated remote update");
};
let Event::ProviderResponseUpdated(mut next_update) = *next_event else {
unreachable!("validated provider response update");
};
let RendererCmd::Remote {
event,
folded_frames,
..
} = &mut current
else {
unreachable!("validated current remote update");
};
let Event::ProviderResponseUpdated(update) = event.as_mut() else {
unreachable!("validated current provider response update");
};
update.deltas.append(&mut next_update.deltas);
fold_response_stats(&mut update.response_stats, next_update.response_stats);
folded_frames.push(super::RendererQueueFrame {
delivery_id,
queue_bytes,
enqueued_at,
});
(current, None)
}
fn fold_response_stats(
current: &mut Option<ProviderResponseStats>,
next: Option<ProviderResponseStats>,
) {
match (current.as_mut(), next) {
(Some(current), Some(next)) => {
current.current = next.current;
current.first_semantic_output_elapsed_micros = current
.first_semantic_output_elapsed_micros
.or(next.first_semantic_output_elapsed_micros);
}
(None, Some(next)) => *current = Some(next),
(_, None) => {}
}
}