Files
mudmouth/src/kokoro.rs
T
pleb 561ace9443
CI / macos-latest (push) Has been cancelled
CI / ubuntu-latest (push) Has been cancelled
CI / windows-latest (push) Has been cancelled
Limit voice selection to avoid awkward assignments
2026-07-28 18:08:10 -07:00

227 lines
6.2 KiB
Rust

use anyhow::{Context, Result};
use reqwest::Url;
use serde::{Deserialize, Serialize};
use std::time::Duration;
/// Curated English Kokoro voices used for NPC selection and the voice manager.
///
/// Keeping this allowlist here means a server update cannot unexpectedly add
/// voices to Mudmouth's random assignment pool.
const NPC_VOICE_ALLOWLIST: &[&str] = &[
"af_heart",
"af_bella",
"af_nicole",
"af_alloy",
"af_aoede",
"af_kore",
"af_nova",
"af_sarah",
"af_river",
"af_jessica",
"af_sky",
"am_michael",
"am_fenrir",
"am_eric",
"am_echo",
"am_liam",
"am_onyx",
"am_adam",
"am_puck",
"bf_alice",
"bf_emma",
"bf_isabella",
"bf_lily",
"bm_daniel",
"bm_fable",
"bm_george",
"bm_lewis",
];
#[derive(Clone)]
pub struct KokoroClient {
client: reqwest::Client,
base_url: Url,
}
#[derive(Debug, Deserialize)]
pub struct Voice {
pub id: String,
#[allow(dead_code)]
pub name: Option<String>,
}
#[derive(Debug, Deserialize)]
#[serde(untagged)]
enum VoiceResponse {
Wrapped { voices: Vec<Voice> },
Bare(Vec<Voice>),
}
impl VoiceResponse {
fn into_voices(self) -> Vec<Voice> {
match self {
Self::Wrapped { voices } | Self::Bare(voices) => voices,
}
}
}
#[derive(Serialize)]
struct SpeechRequest<'a> {
input: &'a str,
voice: &'a str,
model: &'static str,
response_format: &'static str,
speed: f32,
stream: bool,
}
impl KokoroClient {
pub fn new(base_url: &str) -> Result<Self> {
let normalized = format!("{}/", base_url.trim_end_matches('/'));
Ok(Self {
client: reqwest::Client::builder()
.connect_timeout(Duration::from_secs(2))
.timeout(Duration::from_secs(20))
.build()?,
base_url: Url::parse(&normalized).context("Kokoro endpoint is not a valid URL")?,
})
}
fn endpoint(&self, path: &str) -> Result<Url> {
self.base_url
.join(path)
.context("could not build Kokoro URL")
}
pub async fn health(&self) -> Result<()> {
self.client
.get(self.endpoint("health")?)
.send()
.await?
.error_for_status()?;
Ok(())
}
pub async fn voices(&self) -> Result<Vec<String>> {
let response = self
.client
.get(self.endpoint("v1/audio/voices")?)
.send()
.await?
.error_for_status()?;
let response: VoiceResponse = response.json().await?;
Ok(response
.into_voices()
.into_iter()
.map(|voice| voice.id)
.filter(|voice| NPC_VOICE_ALLOWLIST.contains(&voice.as_str()))
.collect())
}
pub async fn synthesize_wav(&self, text: &str, voice: &str, speed: f32) -> Result<Vec<u8>> {
let response = self
.client
.post(self.endpoint("v1/audio/speech")?)
.json(&SpeechRequest {
input: text,
voice,
model: "kokoro",
response_format: "wav",
speed,
stream: false,
})
.send()
.await?
.error_for_status()?;
Ok(response.bytes().await?.to_vec())
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::{
io::{Read, Write},
net::TcpListener,
thread,
};
fn serve_once(status: &str, content_type: &str, body: &'static [u8]) -> String {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let address = listener.local_addr().unwrap();
let status = status.to_owned();
let content_type = content_type.to_owned();
thread::spawn(move || {
let (mut stream, _) = listener.accept().unwrap();
let mut request = [0_u8; 4096];
let count = stream.read(&mut request).unwrap();
assert!(count > 0);
let response = format!(
"HTTP/1.1 {status}\r\nContent-Type: {content_type}\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
body.len()
);
stream.write_all(response.as_bytes()).unwrap();
stream.write_all(body).unwrap();
});
format!("http://{address}")
}
#[tokio::test]
async fn reads_voice_catalog_from_mock_server() {
let url = serve_once(
"200 OK",
"application/json",
br#"[{"id":"af_heart","name":"Heart"},{"id":"am_adam","name":"Adam"},{"id":"af_onyx","name":"Onyx"}]"#,
);
let voices = KokoroClient::new(&url).unwrap().voices().await.unwrap();
assert_eq!(voices, vec!["af_heart", "am_adam"]);
}
#[tokio::test]
async fn reads_wrapped_voice_catalog_from_mock_server() {
let url = serve_once(
"200 OK",
"application/json",
br#"{"voices":[{"id":"af_heart","name":"Heart"}]}"#,
);
let voices = KokoroClient::new(&url).unwrap().voices().await.unwrap();
assert_eq!(voices, vec!["af_heart"]);
}
#[tokio::test]
async fn returns_wav_body_from_mock_server() {
let url = serve_once("200 OK", "audio/wav", b"RIFFmock");
let wav = KokoroClient::new(&url)
.unwrap()
.synthesize_wav("hello", "af_heart", 1.0)
.await
.unwrap();
assert_eq!(wav, b"RIFFmock");
}
#[test]
fn speech_request_includes_speed() {
let request = SpeechRequest {
input: "hello",
voice: "af_heart",
model: "kokoro",
response_format: "wav",
speed: 1.5,
stream: false,
};
assert_eq!(serde_json::to_value(request).unwrap()["speed"], 1.5);
}
#[tokio::test]
async fn treats_non_success_as_failure() {
let url = serve_once("503 Service Unavailable", "application/json", b"{}");
assert!(KokoroClient::new(&url).unwrap().health().await.is_err());
}
#[tokio::test]
async fn rejects_malformed_voice_catalog() {
let url = serve_once("200 OK", "application/json", br#"{"voices":"not a list"}"#);
assert!(KokoroClient::new(&url).unwrap().voices().await.is_err());
}
}