Skip to main content

rlx_moshi/stream/
tokio_impl.rs

1use super::duplex::{DuplexStreamEngine, StreamStepOutput};
2use super::sync::{StreamCommand, StreamEvent, StreamStats};
3use crate::session::{GenerationConfig, MoshiSession};
4use anyhow::Result;
5
6pub type TokioStreamEvent = StreamEvent;
7pub type TokioStreamCommand = StreamCommand;
8
9/// Tokio mpsc handle for duplex streaming.
10pub struct TokioStreamHandle {
11    pub cmd_tx: tokio::sync::mpsc::Sender<StreamCommand>,
12    pub event_rx: tokio::sync::mpsc::Receiver<StreamEvent>,
13    join: Option<std::thread::JoinHandle<()>>,
14}
15
16impl TokioStreamHandle {
17    pub fn stop(mut self) {
18        let _ = self.cmd_tx.blocking_send(StreamCommand::Stop);
19        if let Some(j) = self.join.take() {
20            let _ = j.join();
21        }
22    }
23}
24
25/// Worker on std thread; command/event channels are tokio mpsc (blocking_send/recv).
26pub fn spawn_duplex_tokio(
27    session: MoshiSession,
28    prompt: &str,
29    run_cfg: GenerationConfig,
30    channel_capacity: usize,
31) -> Result<TokioStreamHandle> {
32    let (cmd_tx, mut cmd_rx) = tokio::sync::mpsc::channel::<StreamCommand>(channel_capacity);
33    let (event_tx, event_rx) = tokio::sync::mpsc::channel::<StreamEvent>(channel_capacity);
34    let prompt = prompt.to_string();
35    let join = std::thread::spawn(move || {
36        let worker = || -> Result<()> {
37            let _ = event_tx.blocking_send(StreamEvent::Ready);
38            let mut engine = DuplexStreamEngine::from_session(session, &prompt, &run_cfg)?;
39            while let Some(cmd) = cmd_rx.blocking_recv() {
40                match cmd {
41                    StreamCommand::Pcm(pcm) => {
42                        for step in engine.feed_pcm(&pcm)? {
43                            emit_step_events(&event_tx, step)?;
44                        }
45                    }
46                    StreamCommand::Finish => {
47                        for step in engine.finish()? {
48                            emit_step_events(&event_tx, step)?;
49                        }
50                        let _ = event_tx.blocking_send(StreamEvent::Finished(StreamStats {
51                            steps: engine.steps_done(),
52                            output_samples: 0,
53                            device: engine.device(),
54                        }));
55                        break;
56                    }
57                    StreamCommand::Stop => break,
58                }
59            }
60            Ok(())
61        };
62        if let Err(e) = worker() {
63            let _ = event_tx.blocking_send(StreamEvent::Error(e.to_string()));
64        }
65    });
66    Ok(TokioStreamHandle {
67        cmd_tx,
68        event_rx,
69        join: Some(join),
70    })
71}
72
73fn emit_step_events(
74    tx: &tokio::sync::mpsc::Sender<StreamEvent>,
75    step: StreamStepOutput,
76) -> Result<()> {
77    if let Some(text) = step.transcript_delta.clone() {
78        let _ = tx.blocking_send(StreamEvent::Text {
79            step: step.step,
80            text,
81        });
82    }
83    if !step.moshi_pcm.is_empty() {
84        let _ = tx.blocking_send(StreamEvent::OutputPcm {
85            step: step.step,
86            samples: step.moshi_pcm.clone(),
87        });
88    }
89    let _ = tx.blocking_send(StreamEvent::Step(step));
90    Ok(())
91}