use anyhow::{anyhow, Result}; use cpal::traits::{DeviceTrait, HostTrait, StreamTrait}; use cpal::{SampleFormat, SampleRate}; use serde::Deserialize; use std::collections::HashSet; use std::ffi::CStr; use std::io::{BufRead, BufReader, Write}; use std::net::{TcpListener, TcpStream}; use std::path::PathBuf; use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::mpsc; use std::sync::{Arc, Mutex}; use std::thread; use std::time::{Duration, SystemTime, UNIX_EPOCH}; include!(concat!(env!("OUT_DIR"), "/moonshine_bindings.rs")); const SAMPLE_RATE: i32 = 16000; const HEADER_VERSION: i32 = 30000; const ARCH: u32 = 5; // MOONSHINE_MODEL_ARCH_MEDIUM_STREAMING const BIND_ADDR: &str = "127.0.0.1:6996"; const MAX_TEXT_BYTES: usize = 1380; // ─── helpers ────────────────────────────────────────────────────────────── fn ts() -> String { let now = SystemTime::now() .duration_since(UNIX_EPOCH) .unwrap_or_default(); let secs = now.as_secs() % 86400; let h = secs / 3600; let m = (secs % 3600) / 60; let s = secs % 60; let ms = now.subsec_millis(); format!("{:02}:{:02}:{:02}.{:03}", h, m, s, ms) } fn log(msg: &str) { eprintln!("[{}] {}", ts(), msg); } fn err_str(code: i32) -> String { unsafe { let s = moonshine_error_to_string(code); if s.is_null() { format!("error {}", code) } else { CStr::from_ptr(s).to_string_lossy().into_owned() } } } fn truncate_to_word(text: &str, max_bytes: usize) -> &str { if text.len() <= max_bytes { return text; } let cut = &text[..max_bytes.min(text.len())]; match cut.rfind(' ') { Some(pos) => &text[..pos], None => cut, } } fn line_text(line: &transcript_line_t) -> String { if line.text.is_null() { return String::new(); } unsafe { CStr::from_ptr(line.text) } .to_string_lossy() .into_owned() } fn unix_now() -> u64 { SystemTime::now() .duration_since(UNIX_EPOCH) .unwrap_or_default() .as_secs() } // ─── debug session recorder ────────────────────────────────────────────── struct DebugRecorder { dir: PathBuf, audio: Vec, log_lines: Vec, } impl DebugRecorder { fn new(base_dir: &PathBuf) -> Self { let dir = base_dir.join(unix_now().to_string()); std::fs::create_dir_all(&dir).ok(); Self { dir, audio: Vec::new(), log_lines: Vec::new(), } } fn log_event(&mut self, event: &str) { self.log_lines.push(format!("[{}] {}", ts(), event)); } fn log_segment(&mut self, line: &transcript_line_t, prefix: &str) { let text = line_text(line); let event = format!( "SEGMENT id={} prefix={} is_complete={} start={:.3}s duration={:.3}s text=\"{}\"", line.id, prefix, line.is_complete, line.start_time, line.duration, text ); self.log_event(&event); } fn add_audio(&mut self, samples: &[f32]) { self.audio.extend_from_slice(samples); } fn save(self) { // Write log let log_path = self.dir.join("session.log"); match std::fs::File::create(&log_path) { Ok(mut f) => { for line in &self.log_lines { writeln!(f, "{}", line).ok(); } } Err(e) => log(&format!("debug: failed to write log: {}", e)), } // Write WAV let wav_path = self.dir.join("audio.wav"); match write_wav(&wav_path, &self.audio, SAMPLE_RATE as u32) { Ok(()) => { let secs = self.audio.len() as f64 / SAMPLE_RATE as f64; log(&format!("debug: saved {} ({:.1}s, {} samples)", self.dir.display(), secs, self.audio.len())); } Err(e) => log(&format!("debug: failed to write wav: {}", e)), } } } fn write_wav(path: &PathBuf, samples: &[f32], sample_rate: u32) -> Result<()> { let num_samples = samples.len() as u32; let data_size = num_samples * 4; // f32 = 4 bytes let file_size = 44 + data_size; let mut f = std::fs::File::create(path)?; // RIFF header f.write_all(b"RIFF")?; f.write_all(&(file_size - 8).to_le_bytes())?; f.write_all(b"WAVE")?; // fmt chunk f.write_all(b"fmt ")?; f.write_all(&16u32.to_le_bytes())?; // chunk size f.write_all(&3u16.to_le_bytes())?; // IEEE float f.write_all(&1u16.to_le_bytes())?; // mono f.write_all(&sample_rate.to_le_bytes())?; f.write_all(&(sample_rate * 4).to_le_bytes())?; // byte rate f.write_all(&4u16.to_le_bytes())?; // block align f.write_all(&32u16.to_le_bytes())?; // bits per sample // data chunk f.write_all(b"data")?; f.write_all(&data_size.to_le_bytes())?; // Convert f32 samples to little-endian bytes let mut bytes = Vec::with_capacity(data_size as usize); for &s in samples { bytes.extend_from_slice(&s.to_le_bytes()); } f.write_all(&bytes)?; Ok(()) } // ─── shared state ───────────────────────────────────────────────────────── struct Shared { writer: Mutex, session_id: u64, transcriber_handle: i32, debug_dir: Option, } impl Shared { fn send_msg(&self, prefix: &str, text: &str) { let text = truncate_to_word(text, MAX_TEXT_BYTES); let line = if text.is_empty() { format!("{} {}\n", prefix, self.session_id) } else { format!("{} {} {}\n", prefix, self.session_id, text) }; let mut writer = self.writer.lock().unwrap(); match writer.write_all(line.as_bytes()) { Ok(_) => log(&format!("TX {} {} {}", prefix, self.session_id, text)), Err(e) => log(&format!("TX failed: {}", e)), } } } // ─── session ────────────────────────────────────────────────────────────── struct Session { shared: Arc, stop_signal: Arc, aborted: Arc, transcriber: thread::JoinHandle<()>, cpal_stream: Option, stream_handle: i32, } impl Session { fn stop(mut self) { // Signal transcriber to exit main loop, then wait for it to drain // trailing audio + final flush. cpal stream stays alive during drain. self.stop_signal.store(true, Ordering::SeqCst); self.transcriber.join().ok(); // Now safe to kill ALSA — transcriber is done self.cpal_stream.take(); unsafe { moonshine_free_stream(self.shared.transcriber_handle, self.stream_handle) }; } fn abort(mut self) { self.aborted.store(true, Ordering::SeqCst); self.stop_signal.store(true, Ordering::SeqCst); self.transcriber.join().ok(); self.cpal_stream.take(); unsafe { moonshine_free_stream(self.shared.transcriber_handle, self.stream_handle) }; } } fn start_session(shared: Arc) -> Option { let stream_handle = unsafe { moonshine_create_stream(shared.transcriber_handle, 0) }; if stream_handle < 0 { log(&format!("create_stream failed: {}", err_str(stream_handle))); return None; } let rc = unsafe { moonshine_start_stream(shared.transcriber_handle, stream_handle) }; if rc != 0 { log(&format!("start_stream failed: {}", err_str(rc))); unsafe { moonshine_free_stream(shared.transcriber_handle, stream_handle) }; return None; } let (tx, rx) = mpsc::channel::>(); let stop_signal = Arc::new(AtomicBool::new(false)); let aborted = Arc::new(AtomicBool::new(false)); let cpal_stream = match start_cpal(tx) { Ok(s) => Some(s), Err(e) => { log(&format!("cpal failed: {}", e)); unsafe { moonshine_free_stream(shared.transcriber_handle, stream_handle) }; return None; } }; let shared_clone = shared.clone(); let stop_signal_clone = stop_signal.clone(); let aborted_clone = aborted.clone(); let transcriber = thread::spawn(move || { transcriber_loop(shared_clone, rx, stop_signal_clone, aborted_clone, stream_handle); }); Some(Session { shared, stop_signal, aborted, transcriber, cpal_stream, stream_handle, }) } fn transcriber_loop( shared: Arc, rx: mpsc::Receiver>, stop_signal: Arc, aborted: Arc, stream_handle: i32, ) { let handle = shared.transcriber_handle; let mut sent_ids: HashSet = HashSet::new(); // Debug recorder (if enabled) let mut recorder = shared.debug_dir.as_ref().map(|_| DebugRecorder::new( &shared.debug_dir.as_ref().unwrap().join(shared.session_id.to_string()), )); if let Some(ref mut r) = recorder { r.log_event(&format!("session {} started", shared.session_id)); } // Main loop: process audio until stop_signal while !stop_signal.load(Ordering::SeqCst) { match rx.recv_timeout(Duration::from_millis(100)) { Ok(chunk) => { if let Some(ref mut r) = recorder { r.add_audio(&chunk); } unsafe { moonshine_transcribe_add_audio_to_stream( handle, stream_handle, chunk.as_ptr(), chunk.len() as u64, SAMPLE_RATE, 0, ); } let mut t_ptr: *mut transcript_t = std::ptr::null_mut(); let rc = unsafe { moonshine_transcribe_stream(handle, stream_handle, 0, &mut t_ptr) }; if rc != 0 || t_ptr.is_null() { continue; } if let Some(ref mut r) = recorder { log_transcript_lines(r, t_ptr); } send_new_segments(&shared, t_ptr, &mut sent_ids, "P"); } Err(mpsc::RecvTimeoutError::Timeout) => continue, Err(mpsc::RecvTimeoutError::Disconnected) => break, } } // Drain trailing audio from ALSA buffer (cpal stream still alive). // Fixed 100ms window — cpal delivers every ~50ms (800 samples @ 16kHz), // so this captures 1-2 more callbacks worth of trailing audio. // Hold one chunk back so we can apply a fade-out to the very last one. let drain_deadline = std::time::Instant::now() + Duration::from_millis(100); let mut held_chunk: Option> = None; while std::time::Instant::now() < drain_deadline { match rx.recv_timeout(drain_deadline - std::time::Instant::now()) { Ok(chunk) => { // Feed previously held chunk to Moonshine + recorder (no fade) if let Some(prev) = held_chunk.take() { if let Some(ref mut r) = recorder { r.add_audio(&prev); } unsafe { moonshine_transcribe_add_audio_to_stream( handle, stream_handle, prev.as_ptr(), prev.len() as u64, SAMPLE_RATE, 0, ); } } held_chunk = Some(chunk); } Err(mpsc::RecvTimeoutError::Timeout) => break, Err(mpsc::RecvTimeoutError::Disconnected) => break, } } // Apply 5ms fade-out to the last chunk, then feed to both Moonshine and recorder if let Some(mut chunk) = held_chunk.take() { let fade_samples = (SAMPLE_RATE as usize * 25) / 1000; // 25ms if chunk.len() > fade_samples { let start = chunk.len() - fade_samples; for i in 0..fade_samples { let t = 1.0 - (i as f32 / fade_samples as f32); chunk[start + i] *= t; } } if let Some(ref mut r) = recorder { r.add_audio(&chunk); } unsafe { moonshine_transcribe_add_audio_to_stream( handle, stream_handle, chunk.as_ptr(), chunk.len() as u64, SAMPLE_RATE, 0, ); } } // If aborted (new session took over), skip final flush entirely if aborted.load(Ordering::SeqCst) { unsafe { moonshine_stop_stream(handle, stream_handle) }; log(&format!("Session {} aborted, skipping final flush", shared.session_id)); if let Some(mut r) = recorder { r.log_event("aborted (new session took over)"); r.save(); } return; } // Final flush unsafe { moonshine_stop_stream(handle, stream_handle) }; let mut t_ptr: *mut transcript_t = std::ptr::null_mut(); let rc = unsafe { moonshine_transcribe_stream(handle, stream_handle, 0, &mut t_ptr) }; if rc == 0 && !t_ptr.is_null() { let t = unsafe { &*t_ptr }; let mut new_segments: Vec = Vec::new(); if let Some(ref mut r) = recorder { log_transcript_lines(r, t_ptr); } for i in 0..t.line_count as usize { let line = unsafe { &*t.lines.add(i) }; if line.text.is_null() || line.is_complete == 0 { continue; } if !sent_ids.insert(line.id) { continue; } let text = line_text(line); if !text.is_empty() { new_segments.push(text); } } if new_segments.is_empty() { shared.send_msg("F", ""); } else { let last = new_segments.len() - 1; for (i, text) in new_segments.iter().enumerate() { let prefix = if i == last { "F" } else { "P" }; shared.send_msg(prefix, text); if let Some(ref mut r) = recorder { r.log_event(&format!("TX {} \"{}\"", prefix, text)); } } } } else { shared.send_msg("F", ""); if let Some(ref mut r) = recorder { r.log_event("TX F (empty)"); } } if let Some(mut r) = recorder { r.log_event("session ended"); r.save(); } } fn log_transcript_lines(recorder: &mut DebugRecorder, t_ptr: *const transcript_t) { let t = unsafe { &*t_ptr }; for i in 0..t.line_count as usize { let line = unsafe { &*t.lines.add(i) }; if line.text.is_null() { continue; } recorder.log_segment(line, if line.is_complete != 0 { "complete" } else { "partial" }); } } fn send_new_segments( shared: &Shared, t_ptr: *const transcript_t, sent_ids: &mut HashSet, prefix: &str, ) { let t = unsafe { &*t_ptr }; for i in 0..t.line_count as usize { let line = unsafe { &*t.lines.add(i) }; if line.text.is_null() || line.is_complete == 0 { continue; } if !sent_ids.insert(line.id) { continue; } let text = line_text(line); if text.is_empty() { continue; } shared.send_msg(prefix, &text); } } // ─── cpal ───────────────────────────────────────────────────────────────── fn start_cpal( tx: mpsc::Sender>, ) -> Result { let host = cpal::default_host(); let dev = host .default_input_device() .ok_or_else(|| anyhow!("no input device"))?; let supported = dev .supported_input_configs()? .filter(|c| c.channels() <= 2 && c.min_sample_rate().0 <= 16000) .min_by_key(|c| match c.sample_format() { SampleFormat::F32 => 0, SampleFormat::I16 => 1, SampleFormat::U8 => 2, _ => 99, }) .ok_or_else(|| anyhow!("no suitable input config"))?; let fmt = supported.sample_format(); let mut config = supported.with_max_sample_rate().config(); if config.channels > 1 { config.channels = 1; } config.sample_rate = SampleRate(16000); config.buffer_size = cpal::BufferSize::Fixed(800); let err_fn = |e: cpal::StreamError| log(&format!("cpal error: {}", e)); let stream = match fmt { SampleFormat::F32 => dev.build_input_stream( &config, move |data: &[f32], _: &_| { let _ = tx.send(data.to_vec()); }, err_fn, None, )?, SampleFormat::I16 => dev.build_input_stream( &config, move |data: &[i16], _: &_| { let _ = tx.send(data.iter().map(|&x| x as f32 / 32768.0).collect()); }, err_fn, None, )?, SampleFormat::U8 => dev.build_input_stream( &config, move |data: &[u8], _: &_| { let _ = tx.send(data.iter().map(|&x| (x as f32 - 128.0) / 128.0).collect()); }, err_fn, None, )?, _ => return Err(anyhow!("unsupported sample format {:?}", fmt)), }; stream.play()?; Ok(stream) } // ─── model fetch ────────────────────────────────────────────────────────── #[derive(Deserialize)] struct Manifest { groups: Vec, } #[derive(Deserialize)] struct ManifestGroup { #[allow(dead_code)] base_url: String, files: Vec, } #[derive(Deserialize)] struct ManifestFile { name: String, url: String, size: Option, } fn fetch_model(model_dir: &str) -> Result<()> { let dest = PathBuf::from(model_dir); log(&format!("Fetching medium-streaming-en model to {}", dest.display())); let lang = std::ffi::CString::new("en").unwrap(); let opt_name = std::ffi::CString::new("model_arch").unwrap(); let opt_value = std::ffi::CString::new("5").unwrap(); let mut options = [moonshine_option_t { name: opt_name.as_ptr(), value: opt_value.as_ptr(), }]; let mut json_ptr: *mut i8 = std::ptr::null_mut(); let rc = unsafe { moonshine_get_stt_dependencies( lang.as_ptr(), options.as_mut_ptr(), options.len() as u64, &mut json_ptr, ) }; if rc != 0 || json_ptr.is_null() { return Err(anyhow!("moonshine_get_stt_dependencies failed: {}", err_str(rc))); } let json_str = unsafe { CStr::from_ptr(json_ptr) } .to_string_lossy() .into_owned(); unsafe { moonshine_free_buffer(json_ptr as *mut std::ffi::c_void) }; let manifest: Manifest = serde_json::from_str(&json_str)?; std::fs::create_dir_all(&dest)?; let mut total_files = 0; let mut total_bytes: u64 = 0; for group in &manifest.groups { for file in &group.files { let dest_path = dest.join(&file.name); if dest_path.exists() { log(&format!(" SKIP {} (already exists)", file.name)); continue; } if let Some(parent) = dest_path.parent() { std::fs::create_dir_all(parent)?; } log(&format!(" GET {}", file.url)); if let Some(expected) = file.size { log(&format!(" {} bytes", expected)); total_bytes += expected; } let status = std::process::Command::new("curl") .arg("-sSL") .arg("-o") .arg(&dest_path) .arg(&file.url) .status()?; if !status.success() { return Err(anyhow!("curl failed for {}", file.url)); } total_files += 1; } } log(&format!("Done: {} files, ~{} MB", total_files, total_bytes / (1024 * 1024))); log(&format!("Model directory: {}", dest.display())); Ok(()) } // ─── main ───────────────────────────────────────────────────────────────── fn main() -> Result<()> { let mut args = std::env::args().skip(1); let home = std::env::var("HOME").unwrap_or_else(|_| "/root".to_string()); let default_model = format!("{}/.rvsttd/model", home); let first = args.next(); if first.as_deref() == Some("fetch") { let model_dir = args.next().unwrap_or_else(|| default_model.clone()); return fetch_model(&model_dir); } let mut args = first.into_iter().chain(args); let mut model_dir = default_model; let mut debug = false; while let Some(a) = args.next() { match a.as_str() { "--model-dir" | "-m" => { model_dir = args.next().unwrap_or(model_dir); } "--debug" => { debug = true; } "--help" | "-h" => { println!("Usage: rvsttd [--model-dir DIR] [--debug]"); println!(" rvsttd fetch [DIR]"); println!("Listens on TCP {}", BIND_ADDR); println!("Model: medium-streaming (Moonshine)"); println!("--debug: save session audio + transcript log to ~/.rvsttd/debug/"); println!("'rvsttd fetch' downloads the English medium-streaming model"); return Ok(()); } _ => return Err(anyhow!("unknown arg: {}", a)), } } let debug_dir = if debug { let d = PathBuf::from(&home).join(".rvsttd").join("debug"); std::fs::create_dir_all(&d).ok(); log(&format!("Debug mode enabled — sessions saved to {}", d.display())); Some(d) } else { None }; let model_path = std::fs::canonicalize(&model_dir) .unwrap_or_else(|_| std::path::PathBuf::from(&model_dir)); log(&format!("Loading model from {}...", model_path.display())); let c_dir = std::ffi::CString::new(model_path.to_str().unwrap()).unwrap(); let transcriber_handle = unsafe { moonshine_load_transcriber_from_files( c_dir.as_ptr(), ARCH, std::ptr::null(), 0, HEADER_VERSION, ) }; if transcriber_handle < 0 { return Err(anyhow!("failed to load model: {}", err_str(transcriber_handle))); } log(&format!("Model loaded (handle {})", transcriber_handle)); let listener = TcpListener::bind(BIND_ADDR)?; log(&format!("STT server listening on TCP {}", BIND_ADDR)); let mut current_session: Option = None; for stream in listener.incoming() { let stream = match stream { Ok(s) => s, Err(e) => { log(&format!("accept failed: {}", e)); continue; } }; stream.set_nodelay(true).ok(); log(&format!("Client connected: {}", stream.peer_addr().map(|a| a.to_string()).unwrap_or_else(|_| "?".to_string()))); let writer_stream = stream.try_clone()?; let reader = BufReader::new(stream); for line in reader.lines() { let line = match line { Ok(l) => l, Err(_) => break, }; let line = line.trim(); log(&format!("RX {}", line)); // ON if let Some(rest) = line.strip_prefix("ON ") { let new_session_id: u64 = rest.parse().unwrap_or(0); if let Some(s) = current_session.take() { log(&format!("Aborting session {} for new session {}", s.shared.session_id, new_session_id)); s.abort(); } let shared = Arc::new(Shared { writer: Mutex::new(writer_stream.try_clone()?), session_id: new_session_id, transcriber_handle, debug_dir: debug_dir.clone(), }); log(&format!("PTT on session {}", new_session_id)); match start_session(shared) { Some(s) => current_session = Some(s), None => log("Failed to start session"), } } // OFF else if let Some(rest) = line.strip_prefix("OFF ") { let off_session: u64 = rest.parse().unwrap_or(0); if let Some(s) = current_session.as_ref() { if s.shared.session_id == off_session { log(&format!("OFF session {}", off_session)); if let Some(s) = current_session.take() { s.stop(); } } else { log(&format!("OFF session {} (stale, current={}), ignoring", off_session, s.shared.session_id)); } } else { log(&format!("OFF session {} (no active session), ignoring", off_session)); } } else { log(&format!("Unknown command: {}", line)); } } log("Client disconnected"); if let Some(s) = current_session.take() { s.abort(); } } Ok(()) }