Skip to main content

cherenkov/server/
output.rs

1//! Slow clients consume bounded writer queues, never the GPU worker thread.
2
3use super::{
4    ApiKind, registry::Ticket, request::Request, response::Response, sessions::Commit,
5    tool_call::WireToolCall,
6};
7use crate::control::State;
8use anyhow::{Result, bail};
9use serde_json::Value;
10use std::{
11    net::TcpStream,
12    sync::{Arc, mpsc},
13    time::{SystemTime, UNIX_EPOCH},
14};
15
16pub(super) enum Frame {
17    Text(String),
18    Finish {
19        text: String,
20        reason: &'static str,
21        usage: Value,
22        turn: Option<Box<Commit>>,
23        tool_calls: Option<Vec<WireToolCall>>,
24    },
25    Error(String),
26}
27
28pub(super) struct Output {
29    sender: mpsc::SyncSender<Frame>,
30    ticket: Arc<Ticket>,
31}
32
33impl Output {
34    pub(super) fn new(
35        mut stream: TcpStream,
36        kind: ApiKind,
37        request: &Request,
38        ticket: Arc<Ticket>,
39        state: Arc<State>,
40    ) -> Self {
41        let (sender, receiver) = mpsc::sync_channel(16);
42        let writer_ticket = ticket.clone();
43        let (streaming, include_usage) = (request.stream, request.include_usage);
44
45        std::thread::spawn(move || {
46            let created = SystemTime::now()
47                .duration_since(UNIX_EPOCH)
48                .unwrap_or_default()
49                .as_secs();
50            let mut response = Response::for_writer(
51                &mut stream,
52                kind,
53                &writer_ticket.id,
54                created,
55                streaming,
56                include_usage,
57            );
58            let result = write_frames(&mut response, receiver, &writer_ticket);
59
60            if result.is_err() {
61                writer_ticket.cancel();
62            }
63
64            state.update(|s| {
65                if writer_ticket.cancelled() {
66                    s.cancelled_requests += 1;
67                } else if matches!(result, Ok(true)) {
68                    s.completed_requests += 1;
69                } else {
70                    s.failed_requests += 1;
71                }
72            });
73        });
74
75        Self { sender, ticket }
76    }
77
78    #[cfg(test)]
79    pub(super) fn for_test(ticket: Arc<Ticket>) -> (Self, mpsc::Receiver<Frame>) {
80        let (sender, receiver) = mpsc::sync_channel(16);
81
82        (Self { sender, ticket }, receiver)
83    }
84
85    pub(super) fn send(&self, frame: Frame) -> Result<()> {
86        if self.sender.try_send(frame).is_err() {
87            self.ticket.cancel();
88            bail!("client disconnected or output queue is full");
89        }
90
91        Ok(())
92    }
93}
94
95pub(super) fn write_frames(
96    response: &mut Response<'_, impl std::io::Write>,
97    receiver: mpsc::Receiver<Frame>,
98    ticket: &Ticket,
99) -> Result<bool> {
100    response.start()?;
101
102    for frame in receiver {
103        match frame {
104            Frame::Text(text) => response.text(&text)?,
105            Frame::Error(message) => {
106                response.fail(&message)?;
107
108                return Ok(false);
109            }
110            Frame::Finish {
111                text,
112                reason,
113                usage,
114                turn,
115                tool_calls,
116            } => {
117                // Cancellation must win before publication. Dropping an unpublished
118                // turn rolls it back before the final response is written.
119                let completed = ticket.complete();
120
121                if let Some(turn) = turn.filter(|_| completed) {
122                    turn.publish();
123                }
124
125                response.finish(
126                    &text,
127                    tool_calls.as_deref(),
128                    if completed { reason } else { "cancelled" },
129                    usage,
130                )?;
131
132                return Ok(completed);
133            }
134        }
135    }
136
137    bail!("output ended without a terminal response")
138}
139
140#[cfg(test)]
141#[path = "../../tests/unit/server/output.rs"]
142mod tests;