Skip to main content

rlx_moshi/stream/
sync.rs

1use super::duplex::{DuplexStreamEngine, StreamStepOutput};
2use crate::session::{GenerationConfig, MoshiSession};
3use anyhow::Result;
4use std::sync::mpsc::{self, Receiver, Sender};
5use std::thread::{self, JoinHandle};
6
7/// Events emitted by the duplex worker thread.
8#[derive(Debug, Clone)]
9pub enum StreamEvent {
10    Ready,
11    Step(StreamStepOutput),
12    OutputPcm { step: usize, samples: Vec<f32> },
13    Text { step: usize, text: String },
14    Finished(StreamStats),
15    Error(String),
16}
17
18#[derive(Debug, Clone)]
19pub struct StreamStats {
20    pub steps: usize,
21    pub output_samples: usize,
22    pub device: rlx_runtime::Device,
23}
24
25/// Commands sent to the duplex worker.
26#[derive(Debug)]
27pub enum StreamCommand {
28    Pcm(Vec<f32>),
29    Finish,
30    Stop,
31}
32
33/// Handle to a running duplex stream (std::sync::mpsc).
34pub struct StreamHandle {
35    pub cmd_tx: Sender<StreamCommand>,
36    pub event_rx: Receiver<StreamEvent>,
37    join: Option<JoinHandle<()>>,
38}
39
40impl StreamHandle {
41    pub fn stop(mut self) {
42        let _ = self.cmd_tx.send(StreamCommand::Stop);
43        if let Some(j) = self.join.take() {
44            let _ = j.join();
45        }
46    }
47}
48
49/// Spawn duplex streaming on a dedicated worker thread (LM + Mimi).
50pub fn spawn_duplex_stream(
51    session: MoshiSession,
52    prompt: &str,
53    run_cfg: GenerationConfig,
54) -> Result<StreamHandle> {
55    let (cmd_tx, cmd_rx) = mpsc::channel::<StreamCommand>();
56    let (event_tx, event_rx) = mpsc::channel::<StreamEvent>();
57    let prompt = prompt.to_string();
58    let join = thread::spawn(move || {
59        let worker = || -> Result<()> {
60            event_tx.send(StreamEvent::Ready)?;
61            let mut engine = DuplexStreamEngine::from_session(session, &prompt, &run_cfg)?;
62            loop {
63                match cmd_rx.recv() {
64                    Ok(StreamCommand::Pcm(pcm)) => {
65                        for step in engine.feed_pcm(&pcm)? {
66                            emit_step_events(&event_tx, step)?;
67                        }
68                    }
69                    Ok(StreamCommand::Finish) => {
70                        for step in engine.finish()? {
71                            emit_step_events(&event_tx, step)?;
72                        }
73                        let _ = event_tx.send(StreamEvent::Finished(StreamStats {
74                            steps: engine.steps_done(),
75                            output_samples: 0,
76                            device: engine.device(),
77                        }));
78                        break;
79                    }
80                    Ok(StreamCommand::Stop) | Err(_) => break,
81                }
82            }
83            Ok(())
84        };
85        if let Err(e) = worker() {
86            let _ = event_tx.send(StreamEvent::Error(e.to_string()));
87        }
88    });
89    Ok(StreamHandle {
90        cmd_tx,
91        event_rx,
92        join: Some(join),
93    })
94}
95
96pub(crate) fn emit_step_events(tx: &Sender<StreamEvent>, step: StreamStepOutput) -> Result<()> {
97    if let Some(text) = step.transcript_delta.clone() {
98        let _ = tx.send(StreamEvent::Text {
99            step: step.step,
100            text,
101        });
102    }
103    if !step.moshi_pcm.is_empty() {
104        let _ = tx.send(StreamEvent::OutputPcm {
105            step: step.step,
106            samples: step.moshi_pcm.clone(),
107        });
108    }
109    let _ = tx.send(StreamEvent::Step(step));
110    Ok(())
111}