Files
Robovoice/rvsttd/src/main.rs
T

765 lines
25 KiB
Rust
Raw Normal View History

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<f32>,
log_lines: Vec<String>,
}
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(mut self) {
// Apply 5ms fade-out to the end to avoid click at the trailing edge
let fade_samples = (SAMPLE_RATE as usize * 5) / 1000; // 5ms
if self.audio.len() > fade_samples {
let start = self.audio.len() - fade_samples;
for i in 0..fade_samples {
let t = 1.0 - (i as f32 / fade_samples as f32);
self.audio[start + i] *= t;
}
}
// 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<TcpStream>,
session_id: u64,
transcriber_handle: i32,
debug_dir: Option<PathBuf>,
}
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<Shared>,
stop_signal: Arc<AtomicBool>,
aborted: Arc<AtomicBool>,
transcriber: thread::JoinHandle<()>,
cpal_stream: Option<cpal::Stream>,
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<Shared>) -> Option<Session> {
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::<Vec<f32>>();
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<Shared>,
rx: mpsc::Receiver<Vec<f32>>,
stop_signal: Arc<AtomicBool>,
aborted: Arc<AtomicBool>,
stream_handle: i32,
) {
let handle = shared.transcriber_handle;
let mut sent_ids: HashSet<u64> = 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.
let drain_deadline = std::time::Instant::now() + Duration::from_millis(100);
while std::time::Instant::now() < drain_deadline {
match rx.recv_timeout(drain_deadline - std::time::Instant::now()) {
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,
);
}
}
Err(mpsc::RecvTimeoutError::Timeout) => break,
Err(mpsc::RecvTimeoutError::Disconnected) => break,
}
}
// 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<String> = 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<u64>,
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<Vec<f32>>,
) -> Result<cpal::Stream> {
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<ManifestGroup>,
}
#[derive(Deserialize)]
struct ManifestGroup {
#[allow(dead_code)]
base_url: String,
files: Vec<ManifestFile>,
}
#[derive(Deserialize)]
struct ManifestFile {
name: String,
url: String,
size: Option<u64>,
}
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<Session> = 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();
2026-08-13 10:50:45 +00:00
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 <session>
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 <session>
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(())
}