1use 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
23pub const DEFAULT_STREAM_PORT: u16 = 4798;
26pub const STREAM_PATH: &str = "/transcribe";
28const TRACE_TARGET: &str = "studio_worker::stt_stream";
29const POLL: Duration = Duration::from_millis(100);
32const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10);
34const CLOSE_GRACE: Duration = Duration::from_secs(1);
36
37type Socket = WebSocket<TcpStream>;
38
39pub 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 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
111struct 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#[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
220fn 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
301fn 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 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 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 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}