//! Standalone HTTP server for the rtx-csm pipeline. //! //! Loads CSM-1B + (optionally) AudioSeal + (optionally) WavLM-SV at startup //! and exposes: //! //! GET /health → "ok" //! POST /v1/tts → audio/wav (24 kHz mono) //! POST /v1/detect [multipart audio] → JSON { mean_presence, message } //! POST /v1/speaker_embed [multipart audio] → JSON { embedding: [512 floats] } //! POST /v1/speaker_compare [multipart a + b] → JSON { cosine } //! //! Inference is sync + GPU-heavy → spawn_blocking for every request. The //! Generator + speaker scorer are wrapped in `std::sync::Mutex` so only //! one inference runs per process; AudioSeal embed/detect are stateless //! and run concurrently. Run multiple processes behind a load balancer //! for higher throughput. //! //! Usage (full pipeline): //! ```bash //! cargo run -p rtx-csm --release --features metal --example tts_server -- \ //! --bind 127.0.0.1:8080 \ //! --audioseal-generator /tmp/audioseal_generator.safetensors \ //! --audioseal-detector /tmp/audioseal_detector.safetensors \ //! --audioseal-message 0xCAFE \ //! --wavlm-sv /tmp/wavlm_sv.safetensors //! //! curl -X POST http://127.0.0.1:8080/v1/tts \ //! -H 'content-type: application/json' \ //! -d '{"text":"Hello.","speaker":0}' --output /tmp/out.wav //! curl -X POST http://127.0.0.1:8080/v1/detect \ //! -F audio=@/tmp/out.wav -F source_rate=24000 //! curl -X POST http://127.0.0.1:8080/v1/speaker_embed \ //! -F audio=@/tmp/out.wav -F source_rate=24000 //! curl -X POST http://127.0.0.1:8080/v1/speaker_compare \ //! -F a=@/tmp/a.wav -F b=@/tmp/b.wav -F source_rate=24000 //! ``` use anyhow::Result; use axum::{ Json, Router, body::{Body, Bytes}, extract::{Multipart, State}, http::{StatusCode, header}, response::IntoResponse, routing::{get, post}, }; use candle_core::DType; use clap::Parser; use futures_util::stream::Stream; use rtx_csm::{ GenerateOptions, Generator, PostProcess, audio_io, audioseal::AudioSealWatermarker, speaker_sim::WavLmSimilarity, watermark::ResampledWatermarker, }; use serde::{Deserialize, Serialize}; use std::io::Cursor; use std::net::SocketAddr; use std::path::PathBuf; use std::sync::{Arc, Mutex}; const AUDIOSEAL_RATE: u32 = 16_000; #[derive(Debug, Parser)] struct Cli { /// Bind address (e.g. 127.0.0.1:8080 or 0.0.0.0:8080). #[arg(long, default_value = "127.0.0.1:8080")] bind: SocketAddr, /// Optional path to a quantized GGUF (output of `examples/quantize`). #[arg(long)] quantized_gguf: Option, /// Force CPU even if metal/cuda features are enabled. #[arg(long)] cpu: bool, /// AudioSeal generator safetensors. When BOTH this and /// --audioseal-detector are set, the inline watermarker is installed /// on the Generator (every TTS response is watermarked) and /// /v1/detect is enabled. #[arg(long)] audioseal_generator: Option, #[arg(long)] audioseal_detector: Option, /// 16-bit watermark message (decimal or 0xHEX). #[arg(long, default_value = "0")] audioseal_message: String, /// WavLM-SV safetensors (output of `wavlm_sv_convert`). When set, /// /v1/speaker_embed and /v1/speaker_compare are enabled. #[arg(long)] wavlm_sv: Option, /// LoRA voice adapter (safetensors, output of `lora_train` or /// `lora_train_emotional`). Loaded into the backbone (FP or quantized /// — combines with --quantized-gguf per Phase 12.6) before serving so /// every TTS request uses the trained voice. #[arg(long)] lora: Option, /// LoRA rank override. When omitted, auto-detected from the adapter's /// embedded metadata (Phase 12.5); falls back to 8 for older adapters. #[arg(long)] lora_rank: Option, /// LoRA alpha override. Auto-detected from metadata when omitted. #[arg(long)] lora_alpha: Option, /// Force extended LoRA coverage (q+k+v+output_proj + MLP). Auto- /// detected from metadata when omitted; setting it on a classic q+v /// adapter just allocates extra unused B=0 slots. #[arg(long, default_value_t = false)] extended_lora: bool, } fn parse_message(s: &str) -> Result { let s = s.trim(); if let Some(rest) = s.strip_prefix("0x").or_else(|| s.strip_prefix("0X")) { Ok(u16::from_str_radix(rest, 16)?) } else { Ok(s.parse::()?) } } struct AppState { generator: Mutex, /// Held by the generator's set_watermarker; we keep a separate /// detector instance for /v1/detect (mmap'd weights are shared). audioseal_detector: Option, wavlm: Option>, } #[derive(Debug, Deserialize)] struct TtsRequest { text: String, #[serde(default)] speaker: u32, #[serde(default = "default_max_audio_ms")] max_audio_ms: u32, #[serde(default = "default_temperature")] temperature: f64, #[serde(default = "default_top_k")] top_k: usize, #[serde(default = "default_top_p")] top_p: f64, #[serde(default = "default_seed")] seed: u64, /// Optional per-request 16-bit watermark message override. If set /// AND the server has AudioSeal loaded, the watermark on this /// response carries this message instead of the server-startup /// default. Useful for tagging each generation with a unique ID /// (e.g. clawsample's job_id mod 0x10000) for audit trails. /// Accepts decimal or 0xHEX as a string for safety with JSON. #[serde(default)] watermark_message: Option, } fn default_max_audio_ms() -> u32 { 10_000 } fn default_temperature() -> f64 { 0.9 } fn default_top_k() -> usize { 50 } fn default_top_p() -> f64 { 0.9 } fn default_seed() -> u64 { 42 } #[derive(Debug, Serialize)] struct DetectResponse { mean_presence: f32, message: u16, message_hex: String, } #[derive(Debug, Serialize)] struct EmbedResponse { embedding: Vec, } #[derive(Debug, Serialize)] struct CompareResponse { cosine: f32, } async fn health() -> &'static str { "ok" } async fn tts( State(state): State>, Json(req): Json, ) -> Result { if req.text.trim().is_empty() { return Err((StatusCode::BAD_REQUEST, "text is required".to_string())); } let opts = GenerateOptions { max_audio_ms: req.max_audio_ms, temperature: req.temperature, top_k: req.top_k, top_p: req.top_p, seed: req.seed, ..GenerateOptions::default() }; let state2 = state.clone(); let text = req.text.clone(); let speaker = req.speaker; // Per-request message override; None falls back to startup default. let req_msg = match req.watermark_message.as_deref() { Some(s) => Some( parse_message(s) .map_err(|e| (StatusCode::BAD_REQUEST, format!("watermark_message: {e}")))?, ), None => None, }; let result: Result, String> = tokio::task::spawn_blocking(move || { let mut g = state2.generator.lock().expect("generator mutex poisoned"); let mut pcm = g .generate(&text, speaker, &[], opts) .map_err(|e| format!("inference: {e}"))?; PostProcess::default() .apply(&mut pcm, g.config.sample_rate) .map_err(|e| format!("post: {e}"))?; if let Some(wm) = g.watermarker.as_ref() { pcm = match req_msg { Some(m) => wm.embed_with_message(&pcm, m), None => wm.embed(&pcm), } .map_err(|e| format!("watermark: {e}"))?; } Ok(pcm) }) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("join: {e}")))?; let pcm = result.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e))?; let mut buf: Vec = Vec::new(); { let cursor = Cursor::new(&mut buf); let spec = hound::WavSpec { channels: 1, sample_rate: 24_000, bits_per_sample: 16, sample_format: hound::SampleFormat::Int, }; let mut w = hound::WavWriter::new(cursor, spec).map_err(|e| { ( StatusCode::INTERNAL_SERVER_ERROR, format!("wav writer: {e}"), ) })?; for &s in pcm.iter() { let v = (s.clamp(-1.0, 1.0) * i16::MAX as f32) as i16; w.write_sample(v) .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("wav write: {e}")))?; } w.finalize().map_err(|e| { ( StatusCode::INTERNAL_SERVER_ERROR, format!("wav finalize: {e}"), ) })?; } let _ = audio_io::TARGET_SAMPLE_RATE; Ok(( StatusCode::OK, [(header::CONTENT_TYPE, "audio/wav")], Bytes::from(buf), )) } /// Streaming TTS: chunks of raw 16-bit little-endian PCM (24 kHz mono) /// are emitted as soon as Mimi produces them. Content-Type is /// `audio/L16; rate=24000; channels=1` per RFC 2586. First chunk arrives /// in roughly `chunk_frames × 80ms + per-frame compute × chunk_frames` /// — typically ~600 ms with `chunk_frames=4` on Metal. /// /// Clients: `curl --output - http://.../v1/tts_stream | aplay -f S16_LE -r 24000` /// or pipe directly into a player. Note: this endpoint does NOT apply /// post-processing or the inline watermarker; streaming watermarking /// requires a chunked AudioSeal port (deferred). Use /v1/tts for the /// post-processed + watermarked path. async fn tts_stream( State(state): State>, Json(req): Json, ) -> Result { if req.text.trim().is_empty() { return Err((StatusCode::BAD_REQUEST, "text is required".to_string())); } let opts = GenerateOptions { max_audio_ms: req.max_audio_ms, temperature: req.temperature, top_k: req.top_k, top_p: req.top_p, seed: req.seed, ..GenerateOptions::default() }; // mpsc channel: blocking inference task pushes chunks, async response // consumer yields them as the body stream. let (tx, rx) = tokio::sync::mpsc::channel::>(8); let state2 = state.clone(); let text = req.text.clone(); let speaker = req.speaker; tokio::task::spawn_blocking(move || { let mut g = state2.generator.lock().expect("generator mutex poisoned"); // chunk_frames=4 → ~320ms of audio per emitted chunk; first chunk // arrives after ~4 frames × per-frame-compute. Tunable. let chunk_frames = 4usize; let res = g.generate_streaming( &text, speaker, &[], opts, chunk_frames, |samples: &[f32]| { // Encode chunk as 16-bit LE PCM. let mut buf = Vec::with_capacity(samples.len() * 2); for &s in samples { let v = (s.clamp(-1.0, 1.0) * i16::MAX as f32) as i16; buf.extend_from_slice(&v.to_le_bytes()); } // blocking_send is correct here — we're inside spawn_blocking. let _ = tx.blocking_send(Ok(Bytes::from(buf))); Ok(()) }, ); if let Err(e) = res { let _ = tx.blocking_send(Err(std::io::Error::new( std::io::ErrorKind::Other, format!("stream: {e}"), ))); } // dropping tx closes the channel }); let stream = ReceiverStream { rx }; let body = Body::from_stream(stream); Ok(( StatusCode::OK, [(header::CONTENT_TYPE, "audio/L16; rate=24000; channels=1")], body, )) } /// Adapter: tokio mpsc Receiver → futures Stream. struct ReceiverStream { rx: tokio::sync::mpsc::Receiver>, } impl Stream for ReceiverStream { type Item = Result; fn poll_next( mut self: std::pin::Pin<&mut Self>, cx: &mut std::task::Context<'_>, ) -> std::task::Poll> { self.rx.poll_recv(cx) } } /// Read a multipart `audio` field as raw WAV bytes, decode at the given /// `source_rate` (default 24 kHz, CSM-1B native), then resample to 16 kHz /// for AudioSeal/WavLM-SV consumption. async fn read_audio_at( multipart: &mut Multipart, field_name: &str, source_rate: u32, target_rate: u32, ) -> Result, (StatusCode, String)> { while let Some(field) = multipart .next_field() .await .map_err(|e| (StatusCode::BAD_REQUEST, format!("multipart: {e}")))? { let name = field.name().unwrap_or_default().to_string(); if name != field_name { continue; } let bytes = field .bytes() .await .map_err(|e| (StatusCode::BAD_REQUEST, format!("read field: {e}")))?; let tmp = std::env::temp_dir().join(format!( "tts_server_in_{}_{:x}.wav", field_name, std::process::id() as u64 ^ rand::random::() )); std::fs::write(&tmp, &bytes).map_err(|e| { ( StatusCode::INTERNAL_SERVER_ERROR, format!("temp write: {e}"), ) })?; let raw = audio_io::load_mono_at_rate(&tmp, source_rate) .map_err(|e| (StatusCode::BAD_REQUEST, format!("decode: {e}")))?; let _ = std::fs::remove_file(&tmp); let resampled = audio_io::resample(&raw, source_rate, target_rate) .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("resample: {e}")))?; return Ok(resampled); } Err(( StatusCode::BAD_REQUEST, format!("missing field: {field_name}"), )) } async fn detect( State(state): State>, mut multipart: Multipart, ) -> Result, (StatusCode, String)> { let det = state.audioseal_detector.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, "AudioSeal not loaded — pass --audioseal-generator + --audioseal-detector at startup" .to_string(), ))?; // Multipart fields are consumed once; read source_rate first if present, // then pull audio. Most clients send audio first; fall back to default 24k. let source_rate = 24_000u32; let samples = read_audio_at(&mut multipart, "audio", source_rate, AUDIOSEAL_RATE).await?; let result = det .detect(&samples) .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("detect: {e}")))?; let message = result.message.unwrap_or(0); Ok(Json(DetectResponse { mean_presence: result.mean_presence, message, message_hex: format!("0x{:04X}", message), })) } async fn speaker_embed( State(state): State>, mut multipart: Multipart, ) -> Result, (StatusCode, String)> { let scorer_mu = state.wavlm.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, "WavLM-SV not loaded — pass --wavlm-sv at startup".to_string(), ))?; let source_rate = 24_000u32; let samples = read_audio_at(&mut multipart, "audio", source_rate, AUDIOSEAL_RATE).await?; let scorer = scorer_mu.lock().expect("wavlm mutex poisoned"); let embedding = scorer .embed(&samples) .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("embed: {e}")))?; Ok(Json(EmbedResponse { embedding })) } async fn speaker_compare( State(state): State>, mut multipart: Multipart, ) -> Result, (StatusCode, String)> { let scorer_mu = state.wavlm.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, "WavLM-SV not loaded — pass --wavlm-sv at startup".to_string(), ))?; let source_rate = 24_000u32; // Read both `a` and `b` audio fields. We can't seek over the multipart // stream, so collect any audio fields we see in order. let mut a: Option> = None; let mut b: Option> = None; while let Some(field) = multipart .next_field() .await .map_err(|e| (StatusCode::BAD_REQUEST, format!("multipart: {e}")))? { let name = field.name().unwrap_or_default().to_string(); match name.as_str() { "a" | "b" => { let bytes = field .bytes() .await .map_err(|e| (StatusCode::BAD_REQUEST, format!("read field: {e}")))?; let tmp = std::env::temp_dir().join(format!( "tts_server_cmp_{name}_{:x}.wav", std::process::id() as u64 ^ rand::random::() )); std::fs::write(&tmp, &bytes).map_err(|e| { ( StatusCode::INTERNAL_SERVER_ERROR, format!("temp write: {e}"), ) })?; let raw = audio_io::load_mono_at_rate(&tmp, source_rate) .map_err(|e| (StatusCode::BAD_REQUEST, format!("decode {name}: {e}")))?; let _ = std::fs::remove_file(&tmp); let resampled = audio_io::resample(&raw, source_rate, AUDIOSEAL_RATE) .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("resample: {e}")))?; if name == "a" { a = Some(resampled); } else { b = Some(resampled); } } _ => {} } } let a = a.ok_or((StatusCode::BAD_REQUEST, "missing 'a' audio".to_string()))?; let b = b.ok_or((StatusCode::BAD_REQUEST, "missing 'b' audio".to_string()))?; let scorer = scorer_mu.lock().expect("wavlm mutex poisoned"); use rtx_csm::speaker_sim::SpeakerSimilarity; let cosine = scorer .score(&a, &b) .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("score: {e}")))?; Ok(Json(CompareResponse { cosine })) } #[tokio::main] async fn main() -> Result<()> { tracing_subscriber::fmt().init(); let cli = Cli::parse(); let device = if cli.cpu { candle_core::Device::Cpu } else { Generator::default_device()? }; tracing::info!("device: {device:?}"); let mut generator = if let Some(gguf) = cli.quantized_gguf.as_ref() { tracing::info!("loading quantized CSM from {}", gguf.display()); Generator::load_csm_1b_quantized(gguf, &device, false)? } else { Generator::load_csm_1b(&device)? }; tracing::info!("CSM-1B loaded"); // Optional LoRA adapter (Phase 12.4 + 12.5 self-describing metadata). if let Some(lora_path) = cli.lora.as_ref() { let extended_override = if cli.extended_lora { Some(true) } else { None }; rtx_csm::training::apply_lora_adapter( &mut generator, lora_path, cli.lora_rank, cli.lora_alpha, extended_override, &device, )?; } // Optional AudioSeal — install inline + keep a detector instance. let audioseal_message = parse_message(&cli.audioseal_message)?; let audioseal_detector = if let (Some(g), Some(d)) = ( cli.audioseal_generator.as_ref(), cli.audioseal_detector.as_ref(), ) { let g_vb = unsafe { candle_nn::VarBuilder::from_mmaped_safetensors(&[g], DType::F32, &device) }?; let d_vb = unsafe { candle_nn::VarBuilder::from_mmaped_safetensors(&[d], DType::F32, &device) }?; let inner = AudioSealWatermarker::from_var_builders(g_vb, d_vb, device.clone(), audioseal_message)?; let wm = ResampledWatermarker::new(inner, generator.config.sample_rate, AUDIOSEAL_RATE); generator.set_watermarker(Box::new(wm)); // Build a second detector instance for /v1/detect (mmap'd weights shared). let g_vb2 = unsafe { candle_nn::VarBuilder::from_mmaped_safetensors(&[g], DType::F32, &device) }?; let d_vb2 = unsafe { candle_nn::VarBuilder::from_mmaped_safetensors(&[d], DType::F32, &device) }?; let det = AudioSealWatermarker::from_var_builders( g_vb2, d_vb2, device.clone(), audioseal_message, )?; tracing::info!( "AudioSeal installed (message=0x{:04X}); /v1/tts auto-watermarks, /v1/detect enabled", audioseal_message ); Some(det) } else { None }; let wavlm = if let Some(p) = cli.wavlm_sv.as_ref() { let scorer = WavLmSimilarity::load(p, &device)?; tracing::info!("WavLM-SV loaded; /v1/speaker_embed and /v1/speaker_compare enabled"); Some(Mutex::new(scorer)) } else { None }; let _ = audioseal_message; // captured into the watermarker; field not needed on AppState. let state = Arc::new(AppState { generator: Mutex::new(generator), audioseal_detector, wavlm, }); let app = Router::new() .route("/health", get(health)) .route("/v1/tts", post(tts)) .route("/v1/tts_stream", post(tts_stream)) .route("/v1/detect", post(detect)) .route("/v1/speaker_embed", post(speaker_embed)) .route("/v1/speaker_compare", post(speaker_compare)) .with_state(state); let listener = tokio::net::TcpListener::bind(&cli.bind).await?; tracing::info!("listening on http://{}", cli.bind); axum::serve(listener, app).await?; Ok(()) }