Skip to main content

studio_worker/stt_stream/
server.rs

1//! The LAN streaming listener: `ws://<host>:<port>/transcribe?token=...`.
2//! A stream token (minted on the loopback API) names the model; the
3//! session runs on that loaded model's lane until it finalises, the
4//! client leaves, or the model is unloaded.
5//!
6//! Binds the LAN on purpose (a phone streams to it); the stream token is the
7//! guard.  One session per model at a time (`try_with_lane`).
8
9use super::session::{ClientFrame, Next, ServerFrame, StreamSession};
10pub use super::tokens::StreamTokens;
11use crate::host::{Lane, LoadedModel, ModelHost};
12use crate::job_run::JobRun;
13use crate::runtime::{CurrentJob, JobOutcome, JobSource, WorkerObservers};
14use crate::types::TaskKind;
15use chrono::Utc;
16use std::net::{SocketAddr, TcpListener, TcpStream};
17use std::sync::atomic::{AtomicBool, Ordering};
18use std::sync::Arc;
19use std::time::{Duration, Instant};
20use tungstenite::handshake::server::{Callback, ErrorResponse, Request, Response};
21use tungstenite::{Message, WebSocket};
22
23/// Port the listener binds when none is configured (next to the local
24/// API's 4787).  Safe range: any free port.
25pub const DEFAULT_STREAM_PORT: u16 = 4798;
26/// The one path served.
27pub const STREAM_PATH: &str = "/transcribe";
28const TRACE_TARGET: &str = "studio_worker::stt_stream";
29/// Read timeout while streaming: how often an idle session checks for an
30/// unload.  Safe range 20..=500 ms.
31const POLL: Duration = Duration::from_millis(100);
32/// Handshake read timeout: a client that connects and says nothing.
33const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10);
34/// How long a closing session waits for the client's close frame.
35const CLOSE_GRACE: Duration = Duration::from_secs(1);
36
37type Socket = WebSocket<TcpStream>;
38
39/// The listener, bound but not yet serving.
40pub struct StreamServer {
41    listener: TcpListener,
42    addr: SocketAddr,
43    host: ModelHost,
44    tokens: Arc<StreamTokens>,
45    observers: WorkerObservers,
46}
47
48impl StreamServer {
49    pub fn bind(
50        addr: &str,
51        host: ModelHost,
52        tokens: Arc<StreamTokens>,
53        observers: WorkerObservers,
54    ) -> anyhow::Result<Self> {
55        let listener = TcpListener::bind(addr)
56            .map_err(|e| anyhow::anyhow!("stream listener bind {addr}: {e}"))?;
57        listener.set_nonblocking(true)?;
58        let addr = listener.local_addr()?;
59        Ok(Self {
60            listener,
61            addr,
62            host,
63            tokens,
64            observers,
65        })
66    }
67
68    pub fn local_addr(&self) -> SocketAddr {
69        self.addr
70    }
71
72    /// Accept sessions until `stop` is set; one thread per session.
73    pub fn serve(&self, stop: &AtomicBool) {
74        while !stop.load(Ordering::Relaxed) {
75            match self.listener.accept() {
76                Ok((stream, peer)) => {
77                    let (host, tokens, observers) = (
78                        self.host.clone(),
79                        self.tokens.clone(),
80                        self.observers.clone(),
81                    );
82                    std::thread::spawn(move || session(stream, peer, &host, &tokens, &observers));
83                }
84                Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {
85                    std::thread::sleep(POLL / 2)
86                }
87                Err(e) => {
88                    tracing::warn!(target: TRACE_TARGET, op = "accept", error = %e, "stream accept failed");
89                    std::thread::sleep(POLL);
90                }
91            }
92        }
93    }
94}
95
96fn query_param<'a>(query: Option<&'a str>, name: &str) -> Option<&'a str> {
97    query?
98        .split('&')
99        .filter_map(|pair| pair.split_once('='))
100        .find(|(k, _)| *k == name)
101        .map(|(_, v)| v)
102}
103
104fn reject(status: u16, message: &str) -> ErrorResponse {
105    let mut resp = ErrorResponse::new(Some(message.to_string()));
106    *resp.status_mut() = tungstenite::http::StatusCode::from_u16(status)
107        .unwrap_or(tungstenite::http::StatusCode::BAD_REQUEST);
108    resp
109}
110
111/// Checks the path and the stream token during the WebSocket handshake,
112/// remembering which model the token opens.
113struct Handshake<'a> {
114    tokens: &'a StreamTokens,
115    model: &'a mut Option<String>,
116}
117
118impl Callback for Handshake<'_> {
119    fn on_request(self, req: &Request, resp: Response) -> Result<Response, ErrorResponse> {
120        if req.uri().path() != STREAM_PATH {
121            return Err(reject(404, "not found; stream to /transcribe"));
122        }
123        match query_param(req.uri().query(), "token").map(|t| self.tokens.check(t, Utc::now())) {
124            Some(Ok(m)) => {
125                *self.model = Some(m);
126                Ok(resp)
127            }
128            Some(Err(rejection)) => Err(reject(401, &rejection.to_string())),
129            None => Err(reject(401, "missing stream token")),
130        }
131    }
132}
133
134/// What a session did, for the log and the local-jobs ring.
135#[derive(Default)]
136struct Summary {
137    audio_bytes: usize,
138    final_text: Option<String>,
139    error: Option<String>,
140}
141
142fn session(
143    stream: TcpStream,
144    peer: SocketAddr,
145    host: &ModelHost,
146    tokens: &StreamTokens,
147    observers: &WorkerObservers,
148) {
149    let started = Instant::now();
150    let started_at = Utc::now();
151    if stream.set_nonblocking(false).is_err()
152        || stream.set_read_timeout(Some(HANDSHAKE_TIMEOUT)).is_err()
153    {
154        return;
155    }
156    let mut model: Option<String> = None;
157    let handshake = Handshake {
158        tokens,
159        model: &mut model,
160    };
161    let mut ws = match tungstenite::accept_hdr(stream, handshake) {
162        Ok(ws) => ws,
163        Err(e) => {
164            tracing::info!(target: TRACE_TARGET, op = "handshake", %peer, error = %e, "stream refused");
165            return;
166        }
167    };
168    let Some(model) = model else { return };
169    if ws.get_ref().set_read_timeout(Some(POLL)).is_err() {
170        return;
171    }
172    tracing::info!(target: TRACE_TARGET, op = "stream", %peer, model = %model, "stream opened");
173    let served = host.try_with_lane(&model, |loaded, lane| {
174        let job = JobRun::begin(
175            observers,
176            CurrentJob {
177                job_id: crate::local::next_job_id(),
178                kind: TaskKind::AudioStt,
179                model: model.clone(),
180                prompt: String::new(),
181                started_at,
182                source: JobSource::Stream,
183            },
184        );
185        let summary = job.span().in_scope(|| run(&mut ws, loaded, lane));
186        (job, summary)
187    });
188    match served {
189        Ok((mut job, summary)) => {
190            let outcome = match &summary.error {
191                Some(reason) => JobOutcome::Failed {
192                    reason: reason.clone(),
193                },
194                None => JobOutcome::Completed,
195            };
196            job.span().in_scope(|| {
197                tracing::info!(
198                    target: TRACE_TARGET,
199                    op = "stream",
200                    %peer,
201                    model = %model,
202                    audio_ms = summary.audio_bytes / 32,
203                    final_chars = summary.final_text.as_ref().map_or(0, String::len),
204                    error = summary.error.as_deref().unwrap_or(""),
205                    elapsed_ms = started.elapsed().as_millis() as u64,
206                    "stream closed"
207                );
208            });
209            job.set_prompt(summary.final_text.as_deref().unwrap_or(""));
210            job.finish(outcome);
211        }
212        Err(err) => {
213            tracing::info!(target: TRACE_TARGET, op = "stream", %peer, model = %model, error = %err, "stream refused");
214            send(&mut ws, &ServerFrame::Error(err.to_string()));
215            close(&mut ws);
216        }
217    }
218}
219
220/// Serve one session on a loaded model's lane.
221fn run(ws: &mut Socket, loaded: &dyn LoadedModel, lane: &Lane) -> Summary {
222    let summary = Summary::default();
223    let Some(streaming) = loaded.as_stream() else {
224        return fail(ws, summary, "model is not a streaming speech model".into());
225    };
226    let mut transcriber = match streaming.open() {
227        Ok(t) => t,
228        Err(e) => return fail(ws, summary, format!("could not open a stream: {e:#}")),
229    };
230    let mut session = StreamSession::new(transcriber.as_mut());
231    let mut summary = summary;
232    loop {
233        if lane.cancelled() {
234            return fail(ws, summary, "model unloaded".into());
235        }
236        let frame = match ws.read() {
237            Ok(Message::Binary(bytes)) => {
238                summary.audio_bytes += bytes.len();
239                ClientFrame::Audio(bytes.to_vec())
240            }
241            Ok(Message::Text(text)) => match ClientFrame::from_text(&text) {
242                Some(frame) => frame,
243                None => {
244                    send(
245                        ws,
246                        &ServerFrame::Error(format!(
247                            "unknown frame {:?}; send audio, end or cancel",
248                            text.as_str()
249                        )),
250                    );
251                    continue;
252                }
253            },
254            Ok(Message::Close(_)) => {
255                summary.error = Some("client left before the final".into());
256                return summary;
257            }
258            Ok(_) => continue,
259            Err(tungstenite::Error::Io(e)) if is_timeout(&e) => continue,
260            Err(e) => {
261                summary.error = Some(format!("stream broke: {e}"));
262                return summary;
263            }
264        };
265        let (frames, next) = session.handle(frame);
266        for frame in &frames {
267            match frame {
268                ServerFrame::Final(text) => summary.final_text = Some(text.clone()),
269                ServerFrame::Error(e) => summary.error = Some(e.clone()),
270                ServerFrame::Partial(_) => {}
271            }
272            send(ws, frame);
273        }
274        if next == Next::Close {
275            close(ws);
276            return summary;
277        }
278    }
279}
280
281fn is_timeout(e: &std::io::Error) -> bool {
282    matches!(
283        e.kind(),
284        std::io::ErrorKind::WouldBlock | std::io::ErrorKind::TimedOut
285    )
286}
287
288fn fail(ws: &mut Socket, mut summary: Summary, error: String) -> Summary {
289    send(ws, &ServerFrame::Error(error.clone()));
290    close(ws);
291    summary.error = Some(error);
292    summary
293}
294
295fn send(ws: &mut Socket, frame: &ServerFrame) {
296    if let Err(e) = ws.send(Message::Text(frame.to_json().to_string().into())) {
297        tracing::debug!(target: TRACE_TARGET, op = "send", error = %e, "stream send failed");
298    }
299}
300
301/// Start the close handshake and give the client a moment to answer.
302fn close(ws: &mut Socket) {
303    let _ = ws.close(None);
304    let deadline = Instant::now() + CLOSE_GRACE;
305    while Instant::now() < deadline {
306        match ws.read() {
307            Ok(_) => {}
308            Err(tungstenite::Error::Io(e)) if is_timeout(&e) => {}
309            Err(_) => return,
310        }
311    }
312}
313
314#[cfg(test)]
315mod tests {
316    use super::*;
317    use crate::catalog::{Catalog, CatalogModel};
318    use crate::host::ModelHost;
319    use crate::lifecycle::ModelState;
320    use crate::runtime::WorkerObservers;
321    use crate::types::{ModelEngine, ModelSource, TaskKind};
322    use chrono::Duration as ChronoDuration;
323    use parking_lot::Mutex;
324    use std::sync::atomic::AtomicBool;
325    use std::sync::Arc;
326    use std::time::Duration;
327    use tungstenite::Message;
328
329    const WAIT: Duration = Duration::from_secs(5);
330
331    fn stt(id: &str) -> CatalogModel {
332        CatalogModel {
333            id: id.into(),
334            display_name: id.into(),
335            kind: TaskKind::AudioStt,
336            vram_gb_estimate: 1.0,
337            description: None,
338            source: ModelSource {
339                engine: ModelEngine::Parakeet,
340                files: vec![],
341                cli_defaults: Default::default(),
342            },
343            enabled: true,
344            origin: "local".into(),
345            exclusive_group: None,
346        }
347    }
348
349    struct Harness {
350        host: ModelHost,
351        tokens: Arc<StreamTokens>,
352        observers: WorkerObservers,
353        addr: std::net::SocketAddr,
354        stop: Arc<AtomicBool>,
355        handle: Option<std::thread::JoinHandle<()>>,
356    }
357
358    impl Drop for Harness {
359        fn drop(&mut self) {
360            self.stop.store(true, std::sync::atomic::Ordering::SeqCst);
361            if let Some(h) = self.handle.take() {
362                let _ = h.join();
363            }
364        }
365    }
366
367    fn start(loaded: bool) -> Harness {
368        let catalog = Arc::new(Mutex::new(Catalog {
369            models: vec![stt("stt-a")],
370            ..Default::default()
371        }));
372        let host = ModelHost::new(
373            catalog,
374            Arc::new(crate::test_support::InstantRuntime),
375            Arc::new(crate::test_support::FixedProbe(20.0)),
376            crate::residency::Residency::load_for_serving(None),
377        );
378        if loaded {
379            host.load("stt-a").unwrap();
380            host.wait_for("stt-a", ModelState::serves, WAIT).unwrap();
381        }
382        let tokens = Arc::new(StreamTokens::default());
383        let observers = WorkerObservers::default();
384        let server = StreamServer::bind(
385            "127.0.0.1:0",
386            host.clone(),
387            tokens.clone(),
388            observers.clone(),
389        )
390        .unwrap();
391        let addr = server.local_addr();
392        let stop = Arc::new(AtomicBool::new(false));
393        let s = stop.clone();
394        let handle = std::thread::spawn(move || server.serve(&s));
395        Harness {
396            host,
397            tokens,
398            observers,
399            addr,
400            stop,
401            handle: Some(handle),
402        }
403    }
404
405    type Client = tungstenite::WebSocket<tungstenite::stream::MaybeTlsStream<std::net::TcpStream>>;
406
407    /// Open a session, or the HTTP status the handshake was refused with.
408    fn connect_to(h: &Harness, path: &str, token: &str) -> Result<Client, u16> {
409        let url = format!("ws://{}{path}?token={token}", h.addr);
410        match tungstenite::connect(url) {
411            Ok((ws, _)) => Ok(ws),
412            Err(tungstenite::Error::Http(resp)) => Err(resp.status().as_u16()),
413            Err(other) => panic!("unexpected connect error: {other}"),
414        }
415    }
416
417    fn connect(h: &Harness, token: &str) -> Result<Client, u16> {
418        connect_to(h, "/transcribe", token)
419    }
420
421    fn token(h: &Harness) -> String {
422        h.tokens
423            .mint("stt-a", ChronoDuration::minutes(5), chrono::Utc::now())
424            .token
425    }
426
427    fn pcm(ms: usize, level: i16) -> Vec<u8> {
428        (0..ms * 16)
429            .flat_map(|i| (if i % 2 == 0 { level } else { -level }).to_le_bytes())
430            .collect()
431    }
432
433    /// Every JSON text frame until the server closes.
434    fn drain(ws: &mut Client) -> Vec<serde_json::Value> {
435        let mut out = Vec::new();
436        loop {
437            match ws.read() {
438                Ok(Message::Text(t)) => out.push(serde_json::from_str(&t).unwrap()),
439                Ok(Message::Close(_)) | Err(_) => return out,
440                Ok(_) => {}
441            }
442        }
443    }
444
445    #[test]
446    fn a_session_streams_partials_then_the_final() {
447        let h = start(true);
448        let mut ws = connect(&h, &token(&h)).unwrap();
449        ws.send(Message::Binary(pcm(200, 8000).into())).unwrap();
450        ws.send(Message::Text("end".into())).unwrap();
451        let frames = drain(&mut ws);
452        assert_eq!(
453            frames[0],
454            serde_json::json!({ "partial": true, "text": "w1" })
455        );
456        assert_eq!(
457            frames[1],
458            serde_json::json!({ "partial": true, "text": "w1 w2" })
459        );
460        assert_eq!(
461            frames.last().unwrap(),
462            &serde_json::json!({ "final": true, "text": "w1 w2" })
463        );
464    }
465
466    #[test]
467    fn a_finished_session_is_recorded_as_a_local_job() {
468        let h = start(true);
469        let mut ws = connect(&h, &token(&h)).unwrap();
470        ws.send(Message::Binary(pcm(100, 8000).into())).unwrap();
471        ws.send(Message::Text("end".into())).unwrap();
472        drain(&mut ws);
473        let deadline = std::time::Instant::now() + WAIT;
474        loop {
475            if let Some(job) = h.observers.local_jobs.lock().front().cloned() {
476                assert_eq!(job.kind, TaskKind::AudioStt);
477                assert_eq!(job.model, "stt-a");
478                assert_eq!(job.prompt, "w1");
479                assert_eq!(job.source, crate::runtime::JobSource::Stream);
480                assert!(h.observers.active_jobs.lock().is_empty());
481                break;
482            }
483            assert!(std::time::Instant::now() < deadline, "never recorded");
484            std::thread::sleep(Duration::from_millis(10));
485        }
486    }
487
488    #[test]
489    fn an_unknown_token_is_refused_at_the_handshake() {
490        let h = start(true);
491        assert_eq!(connect(&h, "nope").err(), Some(401));
492    }
493
494    #[test]
495    fn an_expired_token_is_refused_at_the_handshake() {
496        let h = start(true);
497        let old = h.tokens.mint(
498            "stt-a",
499            ChronoDuration::minutes(1),
500            chrono::Utc::now() - ChronoDuration::hours(1),
501        );
502        assert_eq!(connect(&h, &old.token).err(), Some(401));
503    }
504
505    #[test]
506    fn only_the_transcribe_path_is_served() {
507        let h = start(true);
508        assert_eq!(connect_to(&h, "/elsewhere", &token(&h)).err(), Some(404));
509    }
510
511    #[test]
512    fn a_model_that_is_not_loaded_answers_with_an_error() {
513        let h = start(false);
514        let mut ws = connect(&h, &token(&h)).unwrap();
515        let frames = drain(&mut ws);
516        assert_eq!(
517            frames,
518            [serde_json::json!({ "error": "model stt-a is not loaded (unloaded)" })]
519        );
520    }
521
522    #[test]
523    fn a_second_stream_on_the_same_model_is_told_it_is_busy() {
524        let h = start(true);
525        let mut first = connect(&h, &token(&h)).unwrap();
526        first.send(Message::Binary(pcm(100, 8000).into())).unwrap();
527        // Wait until the first session is serving (its partial arrives).
528        assert!(matches!(first.read(), Ok(Message::Text(_))));
529        let mut second = connect(&h, &token(&h)).unwrap();
530        let frames = drain(&mut second);
531        assert_eq!(
532            frames,
533            [serde_json::json!({ "error": "model stt-a is busy serving another request" })]
534        );
535        first.send(Message::Text("cancel".into())).unwrap();
536    }
537
538    #[test]
539    fn unloading_the_model_ends_the_stream() {
540        let h = start(true);
541        let mut ws = connect(&h, &token(&h)).unwrap();
542        ws.send(Message::Binary(pcm(100, 8000).into())).unwrap();
543        assert!(matches!(ws.read(), Ok(Message::Text(_))));
544        h.host.unload("stt-a").unwrap();
545        let frames = drain(&mut ws);
546        assert_eq!(frames, [serde_json::json!({ "error": "model unloaded" })]);
547        h.host
548            .wait_for("stt-a", |s| *s == ModelState::Unloaded, WAIT)
549            .unwrap();
550    }
551
552    #[test]
553    fn an_unknown_text_frame_is_reported_and_the_session_continues() {
554        let h = start(true);
555        let mut ws = connect(&h, &token(&h)).unwrap();
556        ws.send(Message::Text("hello?".into())).unwrap();
557        ws.send(Message::Binary(pcm(100, 8000).into())).unwrap();
558        ws.send(Message::Text("end".into())).unwrap();
559        let frames = drain(&mut ws);
560        assert_eq!(
561            frames[0],
562            serde_json::json!({ "error": "unknown frame \"hello?\"; send audio, end or cancel" })
563        );
564        assert_eq!(
565            frames.last().unwrap(),
566            &serde_json::json!({ "final": true, "text": "w1" })
567        );
568    }
569}