use crate::error::TalkError;
use directories::{ProjectDirs, UserDirs};
use serde::Deserialize;
use std::env;
use std::fs;
use std::path::{Path, PathBuf};
#[derive(Debug, Clone, Deserialize)]
pub struct Config {
pub output_dir: PathBuf,
pub providers: ProvidersConfig,
pub indicators: Option<IndicatorsConfig>,
pub transcription: Option<TranscriptionConfig>,
#[serde(default)]
pub speak: Option<SpeakConfig>,
#[serde(default)]
pub paste: Option<PasteConfig>,
#[serde(default)]
pub audio: Option<AudioSettings>,
#[serde(default)]
pub recording: Option<RecordingConfig>,
}
#[derive(Debug, Clone, Default, Deserialize)]
pub struct AudioSettings {
#[serde(default)]
pub bt_auto_switch: Option<bool>,
}
impl AudioSettings {
pub fn bt_auto_switch_enabled(&self) -> bool {
self.bt_auto_switch.unwrap_or(true)
}
}
pub const DEFAULT_RECORDING_SAMPLE_RATE: u32 = 48_000;
pub const DEFAULT_RECORDING_CHANNELS: u8 = 1;
pub const DEFAULT_RECORDING_BITRATE: u32 = 128_000;
#[derive(Debug, Clone, Default, Deserialize)]
pub struct RecordingConfig {
#[serde(default)]
pub sample_rate: Option<u32>,
#[serde(default)]
pub channels: Option<u8>,
#[serde(default)]
pub bitrate: Option<u32>,
}
impl RecordingConfig {
pub fn resolved(&self) -> AudioConfig {
AudioConfig {
sample_rate: self.sample_rate.unwrap_or(DEFAULT_RECORDING_SAMPLE_RATE),
channels: self.channels.unwrap_or(DEFAULT_RECORDING_CHANNELS),
bitrate: self.bitrate.unwrap_or(DEFAULT_RECORDING_BITRATE),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum Provider {
Mistral,
#[serde(alias = "openai")]
OpenAI,
Parakeet,
}
impl std::fmt::Display for Provider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Provider::Mistral => write!(f, "mistral"),
Provider::OpenAI => write!(f, "openai"),
Provider::Parakeet => write!(f, "parakeet"),
}
}
}
impl std::str::FromStr for Provider {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.to_lowercase().as_str() {
"mistral" => Ok(Provider::Mistral),
"openai" => Ok(Provider::OpenAI),
"parakeet" => Ok(Provider::Parakeet),
other => Err(format!(
"unknown provider '{}' (expected 'mistral', 'openai', or 'parakeet')",
other
)),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum SynthesisProvider {
Kokoro,
Mistral,
}
impl std::fmt::Display for SynthesisProvider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
SynthesisProvider::Kokoro => write!(f, "kokoro"),
SynthesisProvider::Mistral => write!(f, "mistral"),
}
}
}
impl std::str::FromStr for SynthesisProvider {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.to_lowercase().as_str() {
"kokoro" => Ok(SynthesisProvider::Kokoro),
"mistral" => Ok(SynthesisProvider::Mistral),
other => Err(format!(
"unknown synthesis provider '{}' (expected 'kokoro' or 'mistral')",
other
)),
}
}
}
#[derive(Debug, Clone, Deserialize)]
pub struct ProvidersConfig {
#[serde(default)]
pub mistral: Option<MistralConfig>,
#[serde(default)]
pub openai: Option<OpenAIConfig>,
#[serde(default)]
pub parakeet: Option<ParakeetConfig>,
#[serde(default)]
pub kokoro: Option<KokoroConfig>,
}
#[derive(Debug, Deserialize, Clone)]
pub struct MistralConfig {
pub api_key: String,
#[serde(default)]
pub url: Option<String>,
#[serde(default = "default_mistral_model")]
pub model: String,
#[serde(default)]
pub context_bias: Option<String>,
#[serde(default = "default_mistral_tts_model")]
pub tts_model: String,
#[serde(default)]
pub tts_voice: Option<String>,
#[serde(default)]
pub tts_voices: Option<std::collections::HashMap<String, String>>,
}
fn default_mistral_model() -> String {
"voxtral-mini-2507".to_string()
}
fn default_mistral_tts_model() -> String {
"voxtral-mini-tts-latest".to_string()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum OpenAIRealtimeDelay {
Minimal,
Low,
Medium,
High,
Xhigh,
}
impl std::fmt::Display for OpenAIRealtimeDelay {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let value = match self {
Self::Minimal => "minimal",
Self::Low => "low",
Self::Medium => "medium",
Self::High => "high",
Self::Xhigh => "xhigh",
};
write!(f, "{value}")
}
}
impl std::str::FromStr for OpenAIRealtimeDelay {
type Err = String;
fn from_str(value: &str) -> Result<Self, Self::Err> {
match value.trim().to_ascii_lowercase().as_str() {
"minimal" => Ok(Self::Minimal),
"low" => Ok(Self::Low),
"medium" => Ok(Self::Medium),
"high" => Ok(Self::High),
"xhigh" => Ok(Self::Xhigh),
_ => Err(format!(
"invalid OpenAI realtime delay '{value}' (expected minimal, low, medium, high, or xhigh)"
)),
}
}
}
#[derive(Debug, Deserialize, Clone)]
pub struct OpenAIConfig {
pub api_key: String,
#[serde(default)]
pub url: Option<String>,
#[serde(default = "default_openai_model")]
pub model: String,
#[serde(default = "default_openai_realtime_model")]
pub realtime_model: String,
#[serde(default)]
pub prompt: Option<String>,
#[serde(default)]
pub keywords: Option<Vec<String>>,
#[serde(default)]
pub languages: Option<Vec<String>>,
#[serde(default)]
pub realtime_delay: Option<OpenAIRealtimeDelay>,
}
fn default_openai_model() -> String {
"gpt-transcribe".to_string()
}
fn default_openai_realtime_model() -> String {
"gpt-live-transcribe".to_string()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum ParakeetVariant {
#[default]
Int8,
Fp32,
}
impl std::fmt::Display for ParakeetVariant {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ParakeetVariant::Int8 => write!(f, "int8"),
ParakeetVariant::Fp32 => write!(f, "fp32"),
}
}
}
impl std::str::FromStr for ParakeetVariant {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.to_lowercase().as_str() {
"int8" => Ok(ParakeetVariant::Int8),
"fp32" => Ok(ParakeetVariant::Fp32),
other => Err(format!(
"unknown parakeet variant '{}' (expected 'int8' or 'fp32')",
other
)),
}
}
}
#[derive(Debug, Clone, Default, Deserialize)]
pub struct ParakeetConfig {
#[serde(default)]
pub variant: ParakeetVariant,
#[serde(default)]
pub model_dir: Option<PathBuf>,
#[serde(default = "default_parakeet_threads")]
pub num_threads: i32,
#[serde(default)]
pub model: Option<String>,
}
fn default_parakeet_threads() -> i32 {
2
}
impl ParakeetConfig {
pub fn resolved_variant(&self) -> ParakeetVariant {
self.variant
}
pub fn resolved_model_dir(&self) -> Result<PathBuf, TalkError> {
if let Some(ref dir) = self.model_dir {
return Ok(dir.clone());
}
let data_dir = ProjectDirs::from("org", "kalysto", "talk-rs")
.map(|dirs| dirs.data_dir().to_path_buf())
.ok_or_else(|| {
TalkError::Config(
"Could not determine data directory for parakeet model_dir".to_string(),
)
})?;
Ok(data_dir
.join("models")
.join(format!("parakeet-tdt-0.6b-v3-{}", self.variant)))
}
pub fn resolved_model_name(&self) -> String {
if let Some(ref m) = self.model {
return m.clone();
}
format!("parakeet-tdt-0.6b-v3-{}", self.variant)
}
}
#[derive(Debug, Clone, Default, Deserialize)]
pub struct KokoroConfig {
#[serde(default)]
pub variant: Option<String>,
#[serde(default)]
pub model_dir: Option<PathBuf>,
#[serde(default)]
pub voice: Option<String>,
#[serde(default)]
pub num_threads: Option<usize>,
#[serde(default)]
pub lang: Option<String>,
}
impl KokoroConfig {
pub fn resolved_model_dir(&self) -> Result<PathBuf, TalkError> {
if let Some(ref dir) = self.model_dir {
return Ok(dir.clone());
}
let data_dir = ProjectDirs::from("org", "kalysto", "talk-rs")
.map(|dirs| dirs.data_dir().to_path_buf())
.ok_or_else(|| {
TalkError::Config(
"Could not determine data directory for kokoro model_dir".to_string(),
)
})?;
Ok(data_dir.join("models").join("kokoro-multi-lang-v1_0"))
}
pub fn resolved_num_threads(&self) -> usize {
self.num_threads.unwrap_or(4)
}
}
#[derive(Debug, Deserialize, Clone)]
pub struct TranscriptionConfig {
#[serde(default = "default_provider")]
pub default_provider: Provider,
}
fn default_provider() -> Provider {
Provider::Mistral
}
#[derive(Debug, Deserialize, Clone)]
pub struct SpeakConfig {
#[serde(default)]
pub default_provider: Option<SynthesisProvider>,
}
#[derive(Debug, Clone)]
pub struct AudioConfig {
pub sample_rate: u32,
pub channels: u8,
pub bitrate: u32,
}
impl Default for AudioConfig {
fn default() -> Self {
Self::new()
}
}
impl AudioConfig {
pub fn new() -> Self {
Self {
sample_rate: 16_000,
channels: 1,
bitrate: 32_000,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum VizMode {
Waterfall,
Amplitude,
Spectrum,
}
impl std::fmt::Display for VizMode {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
VizMode::Waterfall => write!(f, "waterfall"),
VizMode::Amplitude => write!(f, "amplitude"),
VizMode::Spectrum => write!(f, "spectrum"),
}
}
}
impl std::str::FromStr for VizMode {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.to_lowercase().as_str() {
"waterfall" => Ok(VizMode::Waterfall),
"amplitude" => Ok(VizMode::Amplitude),
"spectrum" => Ok(VizMode::Spectrum),
other => Err(format!(
"unknown visualizer mode '{}' (expected: waterfall, amplitude, spectrum)",
other
)),
}
}
}
#[derive(Debug, Clone, Deserialize)]
pub struct IndicatorsConfig {
pub boop_interval_ms: u64,
pub visual_overlay: bool,
pub viz: Option<VizMode>,
}
#[derive(Debug, Clone, Copy, Deserialize, Default, PartialEq, Eq)]
pub enum PasteShortcut {
#[default]
#[serde(rename = "ctrl_shift_v")]
CtrlShiftV,
#[serde(rename = "ctrl_v")]
CtrlV,
}
#[derive(Debug, Clone, Deserialize)]
pub struct FlatPasteConfig {
#[serde(default = "default_paste_chunk_chars")]
pub chunk_chars: usize,
#[serde(default)]
pub shortcut: PasteShortcut,
#[serde(default = "default_paste_restore_settle_ms")]
pub restore_settle_ms: u64,
#[serde(default = "default_paste_chunk_fetch_timeout_ms")]
pub chunk_fetch_timeout_ms: u64,
#[serde(default = "default_paste_target_fetch_retries")]
pub target_fetch_retries: u32,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(untagged)]
pub enum PasteConfig {
#[cfg(feature = "ui")]
Tree(crate::paste::PasteNodeConfig),
Flat(FlatPasteConfig),
}
impl PasteConfig {
#[cfg(feature = "ui")]
pub fn build_root(&self, no_chunk_paste: bool) -> Box<dyn crate::paste::PasteNode> {
let tree = self.to_tree();
crate::paste::build_root_from_config(&tree, no_chunk_paste)
}
#[cfg(feature = "ui")]
pub fn timing(&self) -> crate::paste::PasteTiming {
crate::paste::timing_from_root(&self.to_tree())
}
#[cfg(feature = "ui")]
pub fn to_tree(&self) -> crate::paste::PasteNodeConfig {
match self {
Self::Tree(t) => t.clone(),
Self::Flat(f) => flat_to_tree(f),
}
}
}
impl PasteConfig {
pub fn chunk_chars(&self) -> usize {
#[cfg(feature = "ui")]
fn find(cfg: &crate::paste::PasteNodeConfig) -> Option<usize> {
match cfg {
crate::paste::PasteNodeConfig::Chunk { chunk_chars, .. } => Some(*chunk_chars),
crate::paste::PasteNodeConfig::DetectDisplayServer { x11, .. } => find(x11),
crate::paste::PasteNodeConfig::MatchWmClass { default, .. } => find(default),
crate::paste::PasteNodeConfig::Clipboard { .. }
| crate::paste::PasteNodeConfig::XtestType {} => None,
}
}
match self {
Self::Flat(f) => f.chunk_chars,
#[cfg(feature = "ui")]
Self::Tree(t) => find(t).unwrap_or(0),
}
}
pub fn shortcut(&self) -> PasteShortcut {
#[cfg(feature = "ui")]
fn find(cfg: &crate::paste::PasteNodeConfig) -> Option<PasteShortcut> {
match cfg {
crate::paste::PasteNodeConfig::Clipboard { shortcut, .. } => Some(*shortcut),
crate::paste::PasteNodeConfig::Chunk { child, .. } => find(child),
crate::paste::PasteNodeConfig::DetectDisplayServer { x11, .. } => find(x11),
crate::paste::PasteNodeConfig::MatchWmClass { default, .. } => find(default),
crate::paste::PasteNodeConfig::XtestType {} => None,
}
}
match self {
Self::Flat(f) => f.shortcut,
#[cfg(feature = "ui")]
Self::Tree(t) => find(t).unwrap_or_default(),
}
}
pub fn restore_settle_ms(&self) -> u64 {
match self {
Self::Flat(f) => f.restore_settle_ms,
#[cfg(feature = "ui")]
Self::Tree(_) => self.timing().restore_settle_ms,
}
}
pub fn chunk_fetch_timeout_ms(&self) -> u64 {
match self {
Self::Flat(f) => f.chunk_fetch_timeout_ms,
#[cfg(feature = "ui")]
Self::Tree(_) => self.timing().chunk_fetch_timeout_ms,
}
}
}
#[cfg(feature = "ui")]
fn flat_to_tree(f: &FlatPasteConfig) -> crate::paste::PasteNodeConfig {
let clipboard = crate::paste::PasteNodeConfig::Clipboard {
shortcut: f.shortcut,
restore_settle_ms: f.restore_settle_ms,
chunk_fetch_timeout_ms: f.chunk_fetch_timeout_ms,
target_quiescence_ms: crate::paste::PasteTiming::default().target_quiescence_ms,
target_fetch_retries: f.target_fetch_retries,
};
if f.chunk_chars == 0 {
clipboard
} else {
crate::paste::PasteNodeConfig::Chunk {
chunk_chars: f.chunk_chars,
child: Box::new(clipboard),
}
}
}
fn default_paste_chunk_chars() -> usize {
150
}
fn default_paste_restore_settle_ms() -> u64 {
200
}
fn default_paste_chunk_fetch_timeout_ms() -> u64 {
500
}
fn default_paste_target_fetch_retries() -> u32 {
2
}
fn expand_tilde(path: &Path) -> Result<PathBuf, TalkError> {
let mut components = path.components();
let first = match components.next() {
Some(std::path::Component::Normal(part)) => part,
_ => return Ok(path.to_path_buf()),
};
if first != "~" {
return Ok(path.to_path_buf());
}
let home = UserDirs::new()
.map(|dirs| dirs.home_dir().to_path_buf())
.ok_or_else(|| {
TalkError::Config(
"Could not determine home directory to expand '~' in output_dir".to_string(),
)
})?;
let rest = components.as_path();
if rest.as_os_str().is_empty() {
Ok(home)
} else {
Ok(home.join(rest))
}
}
pub fn config_dir() -> Result<PathBuf, TalkError> {
ProjectDirs::from("org", "kalysto", "talk-rs")
.map(|dirs| dirs.config_dir().to_path_buf())
.ok_or_else(|| TalkError::Config("Could not determine config directory".to_string()))
}
pub fn config_path() -> Result<PathBuf, TalkError> {
Ok(config_dir()?.join("config.yaml"))
}
impl Config {
pub fn load(path: Option<&Path>) -> Result<Self, TalkError> {
let config_path = match path {
Some(path) => path.to_path_buf(),
None => config_path()?,
};
let content = fs::read_to_string(&config_path).map_err(|err| {
TalkError::Config(format!(
"Failed to read config file {}: {}",
config_path.display(),
err
))
})?;
let mut config: Config = serde_yaml::from_str(&content).map_err(|err| {
TalkError::Config(format!(
"Failed to parse config file {}: {}",
config_path.display(),
err
))
})?;
if let Some(value) = env_var_string("TALK_RS_OUTPUT_DIR")? {
config.output_dir = PathBuf::from(value);
}
if let Some(value) = env_var_string("TALK_RS_INDICATORS_VIZ")? {
let mode: VizMode = value.parse().map_err(TalkError::Config)?;
if let Some(ref mut ind) = config.indicators {
ind.viz = Some(mode);
}
}
if let Some(value) = env_var_string("TALK_RS_AUDIO_BT_AUTO_SWITCH")? {
let parsed = parse_bool_env(&value).ok_or_else(|| {
TalkError::Config(format!(
"TALK_RS_AUDIO_BT_AUTO_SWITCH must be true/false/1/0/yes/no, got '{}'",
value
))
})?;
let audio = config.audio.get_or_insert_with(AudioSettings::default);
audio.bt_auto_switch = Some(parsed);
}
if let Some(value) = env_var_string("TALK_RS_RECORDING_SAMPLE_RATE")? {
let parsed = parse_u32_env("TALK_RS_RECORDING_SAMPLE_RATE", &value)?;
let rec = config
.recording
.get_or_insert_with(RecordingConfig::default);
rec.sample_rate = Some(parsed);
}
if let Some(value) = env_var_string("TALK_RS_RECORDING_CHANNELS")? {
let parsed = parse_u32_env("TALK_RS_RECORDING_CHANNELS", &value)?;
let channels = u8::try_from(parsed).map_err(|_| {
TalkError::Config(format!(
"TALK_RS_RECORDING_CHANNELS must be 1 or 2, got '{}'",
value
))
})?;
let rec = config
.recording
.get_or_insert_with(RecordingConfig::default);
rec.channels = Some(channels);
}
if let Some(value) = env_var_string("TALK_RS_RECORDING_BITRATE")? {
let parsed = parse_u32_env("TALK_RS_RECORDING_BITRATE", &value)?;
let rec = config
.recording
.get_or_insert_with(RecordingConfig::default);
rec.bitrate = Some(parsed);
}
if let Some(ref mut mistral) = config.providers.mistral {
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_MISTRAL_API_KEY")? {
mistral.api_key = value;
}
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_MISTRAL_URL")? {
mistral.url = Some(value);
}
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_MISTRAL_MODEL")? {
mistral.model = value;
}
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_MISTRAL_CONTEXT_BIAS")? {
mistral.context_bias = Some(value);
}
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_MISTRAL_TTS_MODEL")? {
mistral.tts_model = value;
}
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_MISTRAL_TTS_VOICE")? {
mistral.tts_voice = Some(value);
}
} else {
if let Some(api_key) = env_var_string("TALK_RS_PROVIDERS_MISTRAL_API_KEY")? {
config.providers.mistral = Some(MistralConfig {
api_key,
url: env_var_string("TALK_RS_PROVIDERS_MISTRAL_URL")?,
model: env_var_string("TALK_RS_PROVIDERS_MISTRAL_MODEL")?
.unwrap_or_else(default_mistral_model),
context_bias: env_var_string("TALK_RS_PROVIDERS_MISTRAL_CONTEXT_BIAS")?,
tts_model: env_var_string("TALK_RS_PROVIDERS_MISTRAL_TTS_MODEL")?
.unwrap_or_else(default_mistral_tts_model),
tts_voice: env_var_string("TALK_RS_PROVIDERS_MISTRAL_TTS_VOICE")?,
tts_voices: None,
});
}
}
if let Some(ref mut openai) = config.providers.openai {
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_OPENAI_API_KEY")? {
openai.api_key = value;
}
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_OPENAI_URL")? {
openai.url = Some(value);
}
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_OPENAI_MODEL")? {
openai.model = value;
}
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_OPENAI_REALTIME_MODEL")? {
openai.realtime_model = value;
}
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_OPENAI_PROMPT")? {
openai.prompt = Some(value);
}
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_OPENAI_KEYWORDS")? {
openai.keywords = parse_string_list_env(&value);
}
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_OPENAI_LANGUAGES")? {
openai.languages = parse_string_list_env(&value);
}
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_OPENAI_REALTIME_DELAY")? {
openai.realtime_delay = Some(value.parse().map_err(TalkError::Config)?);
}
} else {
if let Some(api_key) = env_var_string("TALK_RS_PROVIDERS_OPENAI_API_KEY")? {
config.providers.openai = Some(OpenAIConfig {
api_key,
url: env_var_string("TALK_RS_PROVIDERS_OPENAI_URL")?,
model: env_var_string("TALK_RS_PROVIDERS_OPENAI_MODEL")?
.unwrap_or_else(default_openai_model),
realtime_model: env_var_string("TALK_RS_PROVIDERS_OPENAI_REALTIME_MODEL")?
.unwrap_or_else(default_openai_realtime_model),
prompt: env_var_string("TALK_RS_PROVIDERS_OPENAI_PROMPT")?,
keywords: env_var_string("TALK_RS_PROVIDERS_OPENAI_KEYWORDS")?
.and_then(|value| parse_string_list_env(&value)),
languages: env_var_string("TALK_RS_PROVIDERS_OPENAI_LANGUAGES")?
.and_then(|value| parse_string_list_env(&value)),
realtime_delay: env_var_string("TALK_RS_PROVIDERS_OPENAI_REALTIME_DELAY")?
.map(|value| value.parse().map_err(TalkError::Config))
.transpose()?,
});
}
}
if let Some(ref mut parakeet) = config.providers.parakeet {
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_PARAKEET_VARIANT")? {
let variant: ParakeetVariant = value.parse().map_err(TalkError::Config)?;
parakeet.variant = variant;
}
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_PARAKEET_MODEL_DIR")? {
parakeet.model_dir = Some(PathBuf::from(value));
}
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_PARAKEET_NUM_THREADS")? {
let parsed = parse_u32_env("TALK_RS_PROVIDERS_PARAKEET_NUM_THREADS", &value)?;
parakeet.num_threads = i32::try_from(parsed).map_err(|_| {
TalkError::Config(format!(
"TALK_RS_PROVIDERS_PARAKEET_NUM_THREADS must fit in i32, got '{}'",
value
))
})?;
}
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_PARAKEET_MODEL")? {
parakeet.model = Some(value);
}
} else {
let variant_env = env_var_string("TALK_RS_PROVIDERS_PARAKEET_VARIANT")?;
let model_dir_env = env_var_string("TALK_RS_PROVIDERS_PARAKEET_MODEL_DIR")?;
let num_threads_env = env_var_string("TALK_RS_PROVIDERS_PARAKEET_NUM_THREADS")?;
let model_env = env_var_string("TALK_RS_PROVIDERS_PARAKEET_MODEL")?;
if variant_env.is_some()
|| model_dir_env.is_some()
|| num_threads_env.is_some()
|| model_env.is_some()
{
let variant = match variant_env {
Some(v) => v.parse::<ParakeetVariant>().map_err(TalkError::Config)?,
None => ParakeetVariant::default(),
};
let num_threads = match num_threads_env {
Some(v) => {
let parsed = parse_u32_env("TALK_RS_PROVIDERS_PARAKEET_NUM_THREADS", &v)?;
i32::try_from(parsed).map_err(|_| {
TalkError::Config(format!(
"TALK_RS_PROVIDERS_PARAKEET_NUM_THREADS must fit in i32, got '{}'",
v
))
})?
}
None => default_parakeet_threads(),
};
config.providers.parakeet = Some(ParakeetConfig {
variant,
model_dir: model_dir_env.map(PathBuf::from),
num_threads,
model: model_env,
});
}
}
if let Some(ref mut kokoro) = config.providers.kokoro {
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_KOKORO_VARIANT")? {
kokoro.variant = Some(value);
}
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_KOKORO_MODEL_DIR")? {
kokoro.model_dir = Some(PathBuf::from(value));
}
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_KOKORO_VOICE")? {
kokoro.voice = Some(value);
}
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_KOKORO_NUM_THREADS")? {
let parsed = parse_u32_env("TALK_RS_PROVIDERS_KOKORO_NUM_THREADS", &value)?;
kokoro.num_threads = Some(parsed as usize);
}
if let Some(value) = env_var_string("TALK_RS_PROVIDERS_KOKORO_LANG")? {
kokoro.lang = Some(value);
}
} else {
let variant_env = env_var_string("TALK_RS_PROVIDERS_KOKORO_VARIANT")?;
let model_dir_env = env_var_string("TALK_RS_PROVIDERS_KOKORO_MODEL_DIR")?;
let voice_env = env_var_string("TALK_RS_PROVIDERS_KOKORO_VOICE")?;
let num_threads_env = env_var_string("TALK_RS_PROVIDERS_KOKORO_NUM_THREADS")?;
let lang_env = env_var_string("TALK_RS_PROVIDERS_KOKORO_LANG")?;
if variant_env.is_some()
|| model_dir_env.is_some()
|| voice_env.is_some()
|| num_threads_env.is_some()
|| lang_env.is_some()
{
let num_threads = match num_threads_env {
Some(v) => {
let parsed = parse_u32_env("TALK_RS_PROVIDERS_KOKORO_NUM_THREADS", &v)?;
Some(parsed as usize)
}
None => None,
};
config.providers.kokoro = Some(KokoroConfig {
variant: variant_env,
model_dir: model_dir_env.map(PathBuf::from),
voice: voice_env,
num_threads,
lang: lang_env,
});
}
}
config.output_dir = expand_tilde(&config.output_dir)?;
validate_config(&config)?;
Ok(config)
}
}
fn env_var_string(key: &str) -> Result<Option<String>, TalkError> {
match env::var(key) {
Ok(value) => Ok(Some(value)),
Err(env::VarError::NotPresent) => Ok(None),
Err(env::VarError::NotUnicode(_)) => {
Err(TalkError::Config(format!("{} must be valid UTF-8", key)))
}
}
}
fn parse_bool_env(value: &str) -> Option<bool> {
match value.trim().to_ascii_lowercase().as_str() {
"true" | "yes" | "1" | "on" => Some(true),
"false" | "no" | "0" | "off" => Some(false),
_ => None,
}
}
fn parse_string_list_env(value: &str) -> Option<Vec<String>> {
let values = value
.split(',')
.map(str::trim)
.filter(|entry| !entry.is_empty())
.map(ToString::to_string)
.collect::<Vec<_>>();
if values.is_empty() {
None
} else {
Some(values)
}
}
fn parse_u32_env(key: &str, value: &str) -> Result<u32, TalkError> {
value.trim().parse::<u32>().map_err(|_| {
TalkError::Config(format!(
"{} must be a positive integer, got '{}'",
key, value
))
})
}
fn validate_config(config: &Config) -> Result<(), TalkError> {
if config.output_dir.as_os_str().is_empty() {
return Err(TalkError::Config("output_dir is required".to_string()));
}
if !config.output_dir.is_absolute() {
return Err(TalkError::Config(format!(
"output_dir must be an absolute path, got '{}'",
config.output_dir.display()
)));
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use indoc::indoc;
use std::error::Error;
use std::ffi::OsString;
use std::io::Write;
use std::sync::{Mutex, MutexGuard, OnceLock};
use tempfile::NamedTempFile;
struct EnvGuard {
key: String,
value: Option<OsString>,
}
impl EnvGuard {
fn set(key: &str, value: &str) -> Result<Self, Box<dyn Error>> {
let previous = env::var_os(key);
env::set_var(key, value);
Ok(Self {
key: key.to_string(),
value: previous,
})
}
fn clear(key: &str) -> Result<Self, Box<dyn Error>> {
let previous = env::var_os(key);
env::remove_var(key);
Ok(Self {
key: key.to_string(),
value: previous,
})
}
}
impl Drop for EnvGuard {
fn drop(&mut self) {
if let Some(value) = &self.value {
env::set_var(&self.key, value);
} else {
env::remove_var(&self.key);
}
}
}
fn env_lock() -> Result<MutexGuard<'static, ()>, Box<dyn Error>> {
static ENV_LOCK: OnceLock<Mutex<()>> = OnceLock::new();
let mutex = ENV_LOCK.get_or_init(|| Mutex::new(()));
mutex.lock().map_err(|_| "Env lock poisoned".into())
}
fn write_config(contents: &str) -> Result<NamedTempFile, Box<dyn Error>> {
let mut file = NamedTempFile::new()?;
file.write_all(contents.as_bytes())?;
Ok(file)
}
fn clear_all_provider_env_vars() -> Result<Vec<EnvGuard>, Box<dyn Error>> {
Ok(vec![
EnvGuard::clear("TALK_RS_OUTPUT_DIR")?,
EnvGuard::clear("TALK_RS_RECORDING_SAMPLE_RATE")?,
EnvGuard::clear("TALK_RS_RECORDING_CHANNELS")?,
EnvGuard::clear("TALK_RS_RECORDING_BITRATE")?,
EnvGuard::clear("TALK_RS_PROVIDERS_MISTRAL_API_KEY")?,
EnvGuard::clear("TALK_RS_PROVIDERS_MISTRAL_URL")?,
EnvGuard::clear("TALK_RS_PROVIDERS_MISTRAL_MODEL")?,
EnvGuard::clear("TALK_RS_PROVIDERS_MISTRAL_CONTEXT_BIAS")?,
EnvGuard::clear("TALK_RS_PROVIDERS_MISTRAL_TTS_MODEL")?,
EnvGuard::clear("TALK_RS_PROVIDERS_MISTRAL_TTS_VOICE")?,
EnvGuard::clear("TALK_RS_PROVIDERS_OPENAI_API_KEY")?,
EnvGuard::clear("TALK_RS_PROVIDERS_OPENAI_URL")?,
EnvGuard::clear("TALK_RS_PROVIDERS_OPENAI_MODEL")?,
EnvGuard::clear("TALK_RS_PROVIDERS_OPENAI_REALTIME_MODEL")?,
EnvGuard::clear("TALK_RS_PROVIDERS_OPENAI_PROMPT")?,
EnvGuard::clear("TALK_RS_PROVIDERS_OPENAI_KEYWORDS")?,
EnvGuard::clear("TALK_RS_PROVIDERS_OPENAI_LANGUAGES")?,
EnvGuard::clear("TALK_RS_PROVIDERS_OPENAI_REALTIME_DELAY")?,
EnvGuard::clear("TALK_RS_PROVIDERS_PARAKEET_VARIANT")?,
EnvGuard::clear("TALK_RS_PROVIDERS_PARAKEET_MODEL_DIR")?,
EnvGuard::clear("TALK_RS_PROVIDERS_PARAKEET_NUM_THREADS")?,
EnvGuard::clear("TALK_RS_PROVIDERS_PARAKEET_MODEL")?,
EnvGuard::clear("TALK_RS_PROVIDERS_KOKORO_VARIANT")?,
EnvGuard::clear("TALK_RS_PROVIDERS_KOKORO_MODEL_DIR")?,
EnvGuard::clear("TALK_RS_PROVIDERS_KOKORO_VOICE")?,
EnvGuard::clear("TALK_RS_PROVIDERS_KOKORO_NUM_THREADS")?,
EnvGuard::clear("TALK_RS_PROVIDERS_KOKORO_LANG")?,
])
}
#[test]
fn test_config_load_valid() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
assert_eq!(config.output_dir, PathBuf::from("/tmp/test-output"));
let m = config.providers.mistral.as_ref().expect("mistral present");
assert_eq!(m.api_key, "test-api-key");
assert_eq!(m.model, "voxtral-mini-2507");
assert!(m.context_bias.is_none());
assert!(config.providers.openai.is_none());
Ok(())
}
#[test]
fn test_config_env_override() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let _set_key = EnvGuard::set("TALK_RS_PROVIDERS_MISTRAL_API_KEY", "override-key")?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let m = config.providers.mistral.as_ref().expect("mistral present");
assert_eq!(m.api_key, "override-key");
Ok(())
}
#[test]
fn test_config_model_and_context_bias() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
model: voxtral-mini-2602
context_bias: "Kalysto,talk-rs,Voxtral"
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let m = config.providers.mistral.as_ref().expect("mistral present");
assert_eq!(m.model, "voxtral-mini-2602");
assert_eq!(m.context_bias.as_deref(), Some("Kalysto,talk-rs,Voxtral"));
Ok(())
}
#[test]
fn test_config_env_override_model_and_context_bias() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let _set_model = EnvGuard::set("TALK_RS_PROVIDERS_MISTRAL_MODEL", "voxtral-mini-2602")?;
let _set_bias = EnvGuard::set("TALK_RS_PROVIDERS_MISTRAL_CONTEXT_BIAS", "custom,terms")?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let m = config.providers.mistral.as_ref().expect("mistral present");
assert_eq!(m.model, "voxtral-mini-2602");
assert_eq!(m.context_bias.as_deref(), Some("custom,terms"));
Ok(())
}
#[test]
fn test_config_missing_required_field() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: ""
providers: {}
"#;
let file = write_config(yaml)?;
let result = Config::load(Some(file.path()));
match result {
Ok(_) => Err("Expected empty output_dir to fail".into()),
Err(err) => {
assert!(err.to_string().contains("output_dir is required"));
Ok(())
}
}
}
#[test]
fn test_config_relative_output_dir_rejected() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: relative/path
providers:
mistral:
api_key: test-api-key
"#;
let file = write_config(yaml)?;
let result = Config::load(Some(file.path()));
match result {
Ok(_) => Err("Expected relative output_dir to fail".into()),
Err(err) => {
assert!(
err.to_string().contains("absolute"),
"error should mention 'absolute', got: {}",
err
);
Ok(())
}
}
}
#[test]
fn test_config_relative_output_dir_via_env_rejected() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let _set_dir = EnvGuard::set("TALK_RS_OUTPUT_DIR", "relative/from/env")?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
"#;
let file = write_config(yaml)?;
let result = Config::load(Some(file.path()));
match result {
Ok(_) => Err("Expected relative output_dir from env to fail".into()),
Err(err) => {
assert!(
err.to_string().contains("absolute"),
"error should mention 'absolute', got: {}",
err
);
Ok(())
}
}
}
#[test]
fn test_config_tilde_output_dir_expanded() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let home = directories::UserDirs::new()
.map(|d| d.home_dir().to_path_buf())
.ok_or("home dir unavailable")?;
let yaml = r#"
output_dir: ~/talk-rs-output
providers:
mistral:
api_key: test-api-key
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
assert_eq!(config.output_dir, home.join("talk-rs-output"));
Ok(())
}
#[test]
fn test_config_bare_tilde_output_dir_expanded() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let home = directories::UserDirs::new()
.map(|d| d.home_dir().to_path_buf())
.ok_or("home dir unavailable")?;
let yaml = r#"
output_dir: ~
providers:
mistral:
api_key: test-api-key
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
assert_eq!(config.output_dir, home);
Ok(())
}
#[test]
fn test_config_tilde_output_dir_via_env_expanded() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let home = directories::UserDirs::new()
.map(|d| d.home_dir().to_path_buf())
.ok_or("home dir unavailable")?;
let _set_dir = EnvGuard::set("TALK_RS_OUTPUT_DIR", "~/from-env")?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
assert_eq!(config.output_dir, home.join("from-env"));
Ok(())
}
#[test]
fn test_config_non_leading_tilde_not_expanded() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/~/x
providers:
mistral:
api_key: test-api-key
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
assert_eq!(config.output_dir, PathBuf::from("/tmp/~/x"));
Ok(())
}
#[test]
fn test_config_openai_provider() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
openai:
api_key: sk-test-key
model: gpt-4o-transcribe
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
assert!(config.providers.mistral.is_none());
let o = config.providers.openai.as_ref().expect("openai present");
assert_eq!(o.api_key, "sk-test-key");
assert_eq!(o.model, "gpt-4o-transcribe");
Ok(())
}
#[test]
fn openai_omitted_migration_fields_preserve_defaults() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let file = write_config(indoc! {"
output_dir: /tmp/test-output
providers:
openai:
api_key: sk-test-key
"})?;
let config = Config::load(Some(file.path()))?;
let openai = config.providers.openai.as_ref().ok_or("openai missing")?;
assert_eq!(openai.model, "gpt-transcribe");
assert_eq!(openai.realtime_model, "gpt-live-transcribe");
assert!(openai.prompt.is_none());
assert!(openai.keywords.is_none());
assert!(openai.languages.is_none());
assert!(openai.realtime_delay.is_none());
Ok(())
}
#[test]
fn openai_yaml_parses_migration_hints() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let file = write_config(indoc! {"
output_dir: /tmp/test-output
providers:
openai:
api_key: sk-test-key
prompt: Preserve punctuation and casing.
keywords: [Kalysto, talk-rs]
languages: [fr, en]
realtime_delay: high
"})?;
let config = Config::load(Some(file.path()))?;
let openai = config.providers.openai.as_ref().ok_or("openai missing")?;
assert_eq!(
openai.prompt.as_deref(),
Some("Preserve punctuation and casing.")
);
assert_eq!(
openai.keywords.as_deref(),
Some(&["Kalysto".into(), "talk-rs".into()][..])
);
assert_eq!(
openai.languages.as_deref(),
Some(&["fr".into(), "en".into()][..])
);
assert_eq!(openai.realtime_delay, Some(OpenAIRealtimeDelay::High));
Ok(())
}
#[test]
fn openai_env_hints_override_yaml_with_trimmed_nonempty_lists() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let _prompt = EnvGuard::set("TALK_RS_PROVIDERS_OPENAI_PROMPT", "env prompt")?;
let _keywords = EnvGuard::set("TALK_RS_PROVIDERS_OPENAI_KEYWORDS", " Kalysto, ,talk-rs,")?;
let _languages = EnvGuard::set("TALK_RS_PROVIDERS_OPENAI_LANGUAGES", " fr, en ,, ")?;
let _delay = EnvGuard::set("TALK_RS_PROVIDERS_OPENAI_REALTIME_DELAY", "xhigh")?;
let file = write_config(indoc! {"
output_dir: /tmp/test-output
providers:
openai:
api_key: sk-test-key
prompt: yaml prompt
keywords: [yaml]
languages: [de]
realtime_delay: low
"})?;
let config = Config::load(Some(file.path()))?;
let openai = config.providers.openai.as_ref().ok_or("openai missing")?;
assert_eq!(openai.prompt.as_deref(), Some("env prompt"));
assert_eq!(
openai.keywords.as_deref(),
Some(&["Kalysto".into(), "talk-rs".into()][..])
);
assert_eq!(
openai.languages.as_deref(),
Some(&["fr".into(), "en".into()][..])
);
assert_eq!(openai.realtime_delay, Some(OpenAIRealtimeDelay::Xhigh));
Ok(())
}
#[test]
fn openai_env_created_section_carries_migration_hints() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let _key = EnvGuard::set("TALK_RS_PROVIDERS_OPENAI_API_KEY", "sk-env")?;
let _prompt = EnvGuard::set("TALK_RS_PROVIDERS_OPENAI_PROMPT", "env prompt")?;
let _keywords = EnvGuard::set("TALK_RS_PROVIDERS_OPENAI_KEYWORDS", "one,two")?;
let _languages = EnvGuard::set("TALK_RS_PROVIDERS_OPENAI_LANGUAGES", "fr,en")?;
let _delay = EnvGuard::set("TALK_RS_PROVIDERS_OPENAI_REALTIME_DELAY", "minimal")?;
let file = write_config(indoc! {"
output_dir: /tmp/test-output
providers: {}
"})?;
let config = Config::load(Some(file.path()))?;
let openai = config.providers.openai.as_ref().ok_or("openai missing")?;
assert_eq!(openai.prompt.as_deref(), Some("env prompt"));
assert_eq!(
openai.keywords.as_deref(),
Some(&["one".into(), "two".into()][..])
);
assert_eq!(
openai.languages.as_deref(),
Some(&["fr".into(), "en".into()][..])
);
assert_eq!(openai.realtime_delay, Some(OpenAIRealtimeDelay::Minimal));
Ok(())
}
#[test]
fn openai_empty_env_lists_clear_existing_yaml_hints() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let _keywords = EnvGuard::set("TALK_RS_PROVIDERS_OPENAI_KEYWORDS", " , , ")?;
let _languages = EnvGuard::set("TALK_RS_PROVIDERS_OPENAI_LANGUAGES", ",,")?;
let file = write_config(indoc! {"
output_dir: /tmp/test-output
providers:
openai:
api_key: sk-test-key
keywords: [yaml]
languages: [de]
"})?;
let config = Config::load(Some(file.path()))?;
let openai = config.providers.openai.as_ref().ok_or("openai missing")?;
assert!(openai.keywords.is_none());
assert!(openai.languages.is_none());
Ok(())
}
#[test]
fn openai_empty_env_lists_are_none_in_env_created_section() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let _key = EnvGuard::set("TALK_RS_PROVIDERS_OPENAI_API_KEY", "sk-env")?;
let _keywords = EnvGuard::set("TALK_RS_PROVIDERS_OPENAI_KEYWORDS", " , ")?;
let _languages = EnvGuard::set("TALK_RS_PROVIDERS_OPENAI_LANGUAGES", ",,")?;
let file = write_config(indoc! {"
output_dir: /tmp/test-output
providers: {}
"})?;
let config = Config::load(Some(file.path()))?;
let openai = config.providers.openai.as_ref().ok_or("openai missing")?;
assert!(openai.keywords.is_none());
assert!(openai.languages.is_none());
Ok(())
}
#[test]
fn openai_explicit_yaml_empty_languages_remains_configured() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let file = write_config(indoc! {"
output_dir: /tmp/test-output
providers:
openai:
api_key: sk-test-key
languages: []
"})?;
let config = Config::load(Some(file.path()))?;
let openai = config.providers.openai.as_ref().ok_or("openai missing")?;
assert_eq!(openai.languages.as_deref(), Some(&[][..]));
Ok(())
}
#[test]
fn test_config_openai_env_override() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let _set_key = EnvGuard::set("TALK_RS_PROVIDERS_OPENAI_API_KEY", "sk-env-key")?;
let yaml = r#"
output_dir: /tmp/test-output
providers: {}
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let o = config.providers.openai.as_ref().expect("openai from env");
assert_eq!(o.api_key, "sk-env-key");
assert_eq!(o.model, "gpt-transcribe");
assert_eq!(o.realtime_model, "gpt-live-transcribe");
Ok(())
}
#[test]
fn test_config_both_providers() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: mistral-key
openai:
api_key: openai-key
transcription:
default_provider: openai
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
assert!(config.providers.mistral.is_some());
assert!(config.providers.openai.is_some());
let t = config
.transcription
.as_ref()
.expect("transcription section");
assert_eq!(t.default_provider, Provider::OpenAI);
Ok(())
}
#[test]
fn test_config_no_providers() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers: {}
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
assert!(config.providers.mistral.is_none());
assert!(config.providers.openai.is_none());
Ok(())
}
#[test]
fn test_config_paste_section_absent_gives_none() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers: {}
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
assert!(config.paste.is_none());
Ok(())
}
#[test]
fn test_config_paste_chunk_chars_default() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers: {}
paste: {}
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let paste = config.paste.as_ref().expect("paste section present");
assert_eq!(paste.chunk_chars(), 150);
Ok(())
}
#[test]
fn test_config_paste_chunk_chars_custom() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers: {}
paste:
chunk_chars: 300
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let paste = config.paste.as_ref().expect("paste section present");
assert_eq!(paste.chunk_chars(), 300);
Ok(())
}
#[test]
fn test_config_paste_chunk_chars_zero_disables() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers: {}
paste:
chunk_chars: 0
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let paste = config.paste.as_ref().expect("paste section present");
assert_eq!(paste.chunk_chars(), 0);
Ok(())
}
#[test]
fn test_config_paste_restore_settle_ms_default() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers: {}
paste: {}
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let paste = config.paste.as_ref().expect("paste section present");
assert_eq!(paste.restore_settle_ms(), 200);
Ok(())
}
#[test]
fn test_config_paste_restore_settle_ms_custom() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers: {}
paste:
restore_settle_ms: 500
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let paste = config.paste.as_ref().expect("paste section present");
assert_eq!(paste.restore_settle_ms(), 500);
Ok(())
}
#[test]
fn test_config_paste_chunk_fetch_timeout_ms_default() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers: {}
paste: {}
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let paste = config.paste.as_ref().expect("paste section present");
assert_eq!(paste.chunk_fetch_timeout_ms(), 500);
Ok(())
}
#[test]
fn test_config_paste_target_fetch_retries_default_flat() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers: {}
paste: {}
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let paste = config.paste.as_ref().expect("paste section present");
match paste {
PasteConfig::Flat(f) => {
assert_eq!(f.target_fetch_retries, 2);
assert_eq!(f.chunk_fetch_timeout_ms, 500);
}
#[cfg(feature = "ui")]
PasteConfig::Tree(_) => panic!("expected flat variant for empty paste section"),
}
Ok(())
}
#[test]
fn test_config_paste_target_fetch_retries_custom_flat() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers: {}
paste:
target_fetch_retries: 5
chunk_fetch_timeout_ms: 500
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let paste = config.paste.as_ref().expect("paste section present");
match paste {
PasteConfig::Flat(f) => {
assert_eq!(f.target_fetch_retries, 5);
assert_eq!(f.chunk_fetch_timeout_ms, 500);
}
#[cfg(feature = "ui")]
PasteConfig::Tree(_) => panic!("expected flat variant"),
}
#[cfg(feature = "ui")]
match paste.to_tree() {
crate::paste::PasteNodeConfig::Chunk { child, .. } => match *child {
crate::paste::PasteNodeConfig::Clipboard {
target_fetch_retries,
chunk_fetch_timeout_ms,
..
} => {
assert_eq!(target_fetch_retries, 5);
assert_eq!(chunk_fetch_timeout_ms, 500);
}
other => panic!("expected Clipboard child, got {:?}", other),
},
other => panic!("expected Chunk root, got {:?}", other),
}
Ok(())
}
#[cfg(feature = "ui")]
#[test]
fn test_config_paste_target_fetch_retries_tree() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers: {}
paste:
node: chunk
chunk_chars: 120
child:
node: clipboard
chunk_fetch_timeout_ms: 500
target_fetch_retries: 4
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let paste = config.paste.as_ref().expect("paste section present");
match paste.to_tree() {
crate::paste::PasteNodeConfig::Chunk { child, .. } => match *child {
crate::paste::PasteNodeConfig::Clipboard {
target_fetch_retries,
chunk_fetch_timeout_ms,
..
} => {
assert_eq!(target_fetch_retries, 4);
assert_eq!(chunk_fetch_timeout_ms, 500);
}
other => panic!("expected Clipboard child, got {:?}", other),
},
other => panic!("expected Chunk root, got {:?}", other),
}
Ok(())
}
#[cfg(feature = "ui")]
#[test]
fn test_config_paste_target_fetch_retries_tree_default_when_absent(
) -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers: {}
paste:
node: clipboard
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let paste = config.paste.as_ref().expect("paste section present");
match paste.to_tree() {
crate::paste::PasteNodeConfig::Clipboard {
target_fetch_retries,
chunk_fetch_timeout_ms,
..
} => {
assert_eq!(target_fetch_retries, 2);
assert_eq!(chunk_fetch_timeout_ms, 500);
}
other => panic!("expected Clipboard root, got {:?}", other),
}
Ok(())
}
#[test]
fn test_config_paste_chunk_fetch_timeout_ms_custom() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers: {}
paste:
chunk_fetch_timeout_ms: 800
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let paste = config.paste.as_ref().expect("paste section present");
assert_eq!(paste.chunk_fetch_timeout_ms(), 800);
Ok(())
}
#[test]
fn test_recording_defaults_when_absent() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let resolved = config.recording.clone().unwrap_or_default().resolved();
assert_eq!(resolved.sample_rate, 48_000);
assert_eq!(resolved.channels, 1);
assert_eq!(resolved.bitrate, 128_000);
Ok(())
}
#[test]
fn test_recording_parsed_from_yaml() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
recording:
sample_rate: 24000
channels: 2
bitrate: 96000
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let rec = config
.recording
.as_ref()
.expect("recording section present");
let resolved = rec.resolved();
assert_eq!(resolved.sample_rate, 24_000);
assert_eq!(resolved.channels, 2);
assert_eq!(resolved.bitrate, 96_000);
Ok(())
}
#[test]
fn test_recording_partial_yaml_fills_defaults() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
recording:
channels: 2
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let resolved = config.recording.as_ref().expect("recording").resolved();
assert_eq!(resolved.sample_rate, 48_000); assert_eq!(resolved.channels, 2); assert_eq!(resolved.bitrate, 128_000); Ok(())
}
#[test]
fn test_recording_env_overrides() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let _r = EnvGuard::set("TALK_RS_RECORDING_SAMPLE_RATE", "32000")?;
let _c = EnvGuard::set("TALK_RS_RECORDING_CHANNELS", "2")?;
let _b = EnvGuard::set("TALK_RS_RECORDING_BITRATE", "64000")?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let resolved = config.recording.as_ref().expect("recording").resolved();
assert_eq!(resolved.sample_rate, 32_000);
assert_eq!(resolved.channels, 2);
assert_eq!(resolved.bitrate, 64_000);
Ok(())
}
#[test]
fn test_recording_env_invalid_channels_rejected() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let _c = EnvGuard::set("TALK_RS_RECORDING_CHANNELS", "not-a-number")?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
"#;
let file = write_config(yaml)?;
match Config::load(Some(file.path())) {
Ok(_) => Err("expected invalid channels env to fail".into()),
Err(err) => {
assert!(
err.to_string().contains("TALK_RS_RECORDING_CHANNELS"),
"error should name the variable, got: {}",
err
);
Ok(())
}
}
}
#[test]
fn test_config_url_defaults_to_none() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
openai:
api_key: sk-test-key
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let m = config.providers.mistral.as_ref().expect("mistral present");
assert!(m.url.is_none());
let o = config.providers.openai.as_ref().expect("openai present");
assert!(o.url.is_none());
Ok(())
}
#[test]
fn test_config_url_from_yaml() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
url: https://custom-mistral.example.com
openai:
api_key: sk-test-key
url: https://custom-openai.example.com
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let m = config.providers.mistral.as_ref().expect("mistral present");
assert_eq!(m.url.as_deref(), Some("https://custom-mistral.example.com"));
let o = config.providers.openai.as_ref().expect("openai present");
assert_eq!(o.url.as_deref(), Some("https://custom-openai.example.com"));
Ok(())
}
#[test]
fn test_config_url_env_override() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let _set_url = EnvGuard::set(
"TALK_RS_PROVIDERS_MISTRAL_URL",
"https://env-mistral.example.com",
)?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let m = config.providers.mistral.as_ref().expect("mistral present");
assert_eq!(m.url.as_deref(), Some("https://env-mistral.example.com"));
Ok(())
}
#[test]
fn test_parakeet_section_parsed_from_yaml() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
parakeet:
variant: fp32
num_threads: 4
model: my-parakeet-build
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let p = config
.providers
.parakeet
.as_ref()
.ok_or("parakeet section present")?;
assert_eq!(p.variant, ParakeetVariant::Fp32);
assert_eq!(p.num_threads, 4);
assert_eq!(p.model.as_deref(), Some("my-parakeet-build"));
assert!(p.model_dir.is_none());
Ok(())
}
#[test]
fn test_parakeet_section_absent_gives_none() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers: {}
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
assert!(config.providers.parakeet.is_none());
Ok(())
}
#[test]
fn test_parakeet_defaults_with_empty_block() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
parakeet: {}
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let p = config
.providers
.parakeet
.as_ref()
.ok_or("parakeet present")?;
assert_eq!(p.variant, ParakeetVariant::Int8);
assert_eq!(p.num_threads, 2);
assert!(p.model.is_none());
assert!(p.model_dir.is_none());
assert_eq!(p.resolved_variant(), ParakeetVariant::Int8);
assert_eq!(p.resolved_model_name(), "parakeet-tdt-0.6b-v3-int8");
let dir = p.resolved_model_dir()?;
let s = dir.to_string_lossy();
assert!(
s.ends_with("models/parakeet-tdt-0.6b-v3-int8"),
"expected default model_dir to end with models/parakeet-tdt-0.6b-v3-int8, got: {}",
s
);
Ok(())
}
#[test]
fn test_parakeet_resolved_model_dir_user_override() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
parakeet:
model_dir: /opt/models/my-parakeet
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let p = config
.providers
.parakeet
.as_ref()
.ok_or("parakeet present")?;
assert_eq!(
p.resolved_model_dir()?,
PathBuf::from("/opt/models/my-parakeet")
);
assert_eq!(p.resolved_model_name(), "parakeet-tdt-0.6b-v3-int8");
Ok(())
}
#[test]
fn test_parakeet_default_provider_resolves() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
parakeet: {}
transcription:
default_provider: parakeet
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let t = config
.transcription
.as_ref()
.ok_or("transcription section")?;
assert_eq!(t.default_provider, Provider::Parakeet);
Ok(())
}
#[test]
fn test_parakeet_env_overrides_existing_section() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let _v = EnvGuard::set("TALK_RS_PROVIDERS_PARAKEET_VARIANT", "fp32")?;
let _t = EnvGuard::set("TALK_RS_PROVIDERS_PARAKEET_NUM_THREADS", "4")?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
parakeet:
variant: int8
num_threads: 2
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let p = config
.providers
.parakeet
.as_ref()
.ok_or("parakeet present")?;
assert_eq!(p.variant, ParakeetVariant::Fp32);
assert_eq!(p.num_threads, 4);
assert_eq!(p.resolved_model_name(), "parakeet-tdt-0.6b-v3-fp32");
let dir = p.resolved_model_dir()?;
assert!(dir
.to_string_lossy()
.ends_with("models/parakeet-tdt-0.6b-v3-fp32"));
Ok(())
}
#[test]
fn test_parakeet_env_creates_section_from_scratch() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let _d = EnvGuard::set("TALK_RS_PROVIDERS_PARAKEET_MODEL_DIR", "/tmp/from-env")?;
let yaml = r#"
output_dir: /tmp/test-output
providers: {}
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let p = config
.providers
.parakeet
.as_ref()
.ok_or("parakeet section from env")?;
assert_eq!(p.variant, ParakeetVariant::Int8); assert_eq!(p.num_threads, 2); assert_eq!(p.model_dir.as_deref(), Some(Path::new("/tmp/from-env")));
Ok(())
}
#[test]
fn test_parakeet_env_invalid_variant_rejected() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let _v = EnvGuard::set("TALK_RS_PROVIDERS_PARAKEET_VARIANT", "bogus")?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
parakeet: {}
"#;
let file = write_config(yaml)?;
match Config::load(Some(file.path())) {
Ok(_) => Err("expected invalid variant env to fail".into()),
Err(err) => {
assert!(
err.to_string().contains("bogus"),
"error should name the bad value, got: {}",
err
);
Ok(())
}
}
}
#[test]
fn test_parakeet_env_invalid_num_threads_rejected() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let _n = EnvGuard::set("TALK_RS_PROVIDERS_PARAKEET_NUM_THREADS", "not-a-number")?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
parakeet: {}
"#;
let file = write_config(yaml)?;
match Config::load(Some(file.path())) {
Ok(_) => Err("expected invalid num_threads env to fail".into()),
Err(err) => {
assert!(
err.to_string()
.contains("TALK_RS_PROVIDERS_PARAKEET_NUM_THREADS"),
"error should name the variable, got: {}",
err
);
Ok(())
}
}
}
#[test]
fn test_provider_from_str_includes_parakeet() {
assert_eq!(
"parakeet".parse::<Provider>().unwrap_or(Provider::Mistral),
Provider::Parakeet
);
assert_eq!(
"PARAKEET".parse::<Provider>().unwrap_or(Provider::Mistral),
Provider::Parakeet
);
assert_eq!(
"mistral".parse::<Provider>().unwrap_or(Provider::OpenAI),
Provider::Mistral
);
assert_eq!(
"openai".parse::<Provider>().unwrap_or(Provider::Mistral),
Provider::OpenAI
);
let err = "bogus".parse::<Provider>().err().unwrap_or_default();
assert!(err.contains("mistral"), "err missing mistral: {}", err);
assert!(err.contains("openai"), "err missing openai: {}", err);
assert!(err.contains("parakeet"), "err missing parakeet: {}", err);
}
#[test]
fn test_provider_display_includes_parakeet() {
assert_eq!(Provider::Mistral.to_string(), "mistral");
assert_eq!(Provider::OpenAI.to_string(), "openai");
assert_eq!(Provider::Parakeet.to_string(), "parakeet");
}
#[test]
fn test_parakeet_variant_from_str_and_display() {
assert_eq!(
"int8"
.parse::<ParakeetVariant>()
.unwrap_or(ParakeetVariant::Fp32),
ParakeetVariant::Int8
);
assert_eq!(
"INT8"
.parse::<ParakeetVariant>()
.unwrap_or(ParakeetVariant::Fp32),
ParakeetVariant::Int8
);
assert_eq!(
"fp32"
.parse::<ParakeetVariant>()
.unwrap_or(ParakeetVariant::Int8),
ParakeetVariant::Fp32
);
assert_eq!(ParakeetVariant::Int8.to_string(), "int8");
assert_eq!(ParakeetVariant::Fp32.to_string(), "fp32");
for v in [ParakeetVariant::Int8, ParakeetVariant::Fp32] {
let s = v.to_string();
let parsed = s.parse::<ParakeetVariant>().unwrap_or_else(|_| {
if v == ParakeetVariant::Int8 {
ParakeetVariant::Fp32
} else {
ParakeetVariant::Int8
}
});
assert_eq!(parsed, v);
}
let err = "bogus".parse::<ParakeetVariant>().err().unwrap_or_default();
assert!(err.contains("bogus"), "err missing 'bogus': {}", err);
assert!(err.contains("int8"), "err missing 'int8': {}", err);
assert!(err.contains("fp32"), "err missing 'fp32': {}", err);
}
#[test]
fn test_config_url_env_override_openai() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let _set_key = EnvGuard::set("TALK_RS_PROVIDERS_OPENAI_API_KEY", "sk-env-key")?;
let _set_url = EnvGuard::set(
"TALK_RS_PROVIDERS_OPENAI_URL",
"https://env-openai.example.com",
)?;
let yaml = r#"
output_dir: /tmp/test-output
providers: {}
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let o = config.providers.openai.as_ref().expect("openai from env");
assert_eq!(o.api_key, "sk-env-key");
assert_eq!(o.url.as_deref(), Some("https://env-openai.example.com"));
Ok(())
}
#[test]
fn test_config_mistral_tts_defaults() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let m = config.providers.mistral.as_ref().expect("mistral present");
assert_eq!(m.tts_model, "voxtral-mini-tts-latest");
assert!(m.tts_voice.is_none());
Ok(())
}
#[test]
fn test_config_mistral_tts_fields_parse() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
tts_model: voxtral-mini-tts-2026
tts_voice: 5a271406-039d-46fe-835b-fbbb00eaf08d
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let m = config.providers.mistral.as_ref().expect("mistral present");
assert_eq!(m.tts_model, "voxtral-mini-tts-2026");
assert_eq!(
m.tts_voice.as_deref(),
Some("5a271406-039d-46fe-835b-fbbb00eaf08d")
);
Ok(())
}
#[test]
fn test_config_mistral_tts_voices_map_parses() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
tts_voices:
fr: 5a271406-039d-46fe-835b-fbbb00eaf08d
en: c69964a6-ab8b-4f8a-9465-ec0925096ec8
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let m = config.providers.mistral.as_ref().expect("mistral present");
assert!(m.tts_voice.is_none());
let map = m.tts_voices.as_ref().expect("tts_voices present");
assert_eq!(
map.get("fr").map(String::as_str),
Some("5a271406-039d-46fe-835b-fbbb00eaf08d")
);
assert_eq!(
map.get("en").map(String::as_str),
Some("c69964a6-ab8b-4f8a-9465-ec0925096ec8")
);
Ok(())
}
#[test]
fn test_config_kokoro_lang_auto_parses() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
kokoro:
lang: auto
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let k = config.providers.kokoro.as_ref().expect("kokoro present");
assert_eq!(k.lang.as_deref(), Some("auto"));
Ok(())
}
#[test]
fn test_config_mistral_tts_env_override() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let _m = EnvGuard::set("TALK_RS_PROVIDERS_MISTRAL_TTS_MODEL", "env-tts-model")?;
let _v = EnvGuard::set("TALK_RS_PROVIDERS_MISTRAL_TTS_VOICE", "env-voice-id")?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let m = config.providers.mistral.as_ref().expect("mistral present");
assert_eq!(m.tts_model, "env-tts-model");
assert_eq!(m.tts_voice.as_deref(), Some("env-voice-id"));
Ok(())
}
#[test]
fn test_config_kokoro_section_parses() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
kokoro:
voice: am_michael
num_threads: 8
lang: fr
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let k = config.providers.kokoro.as_ref().expect("kokoro present");
assert_eq!(k.voice.as_deref(), Some("am_michael"));
assert_eq!(k.resolved_num_threads(), 8);
assert_eq!(k.lang.as_deref(), Some("fr"));
let dir = k.resolved_model_dir()?;
assert!(
dir.ends_with("models/kokoro-multi-lang-v1_0"),
"unexpected default model dir: {}",
dir.display()
);
Ok(())
}
#[test]
fn test_config_kokoro_defaults() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
kokoro: {}
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let k = config.providers.kokoro.as_ref().expect("kokoro present");
assert_eq!(k.resolved_num_threads(), 4);
assert!(k.voice.is_none());
Ok(())
}
#[test]
fn test_config_kokoro_env_creates_section() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let _v = EnvGuard::set("TALK_RS_PROVIDERS_KOKORO_VOICE", "af_heart")?;
let _t = EnvGuard::set("TALK_RS_PROVIDERS_KOKORO_NUM_THREADS", "2")?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let k = config.providers.kokoro.as_ref().expect("kokoro from env");
assert_eq!(k.voice.as_deref(), Some("af_heart"));
assert_eq!(k.resolved_num_threads(), 2);
Ok(())
}
#[test]
fn test_config_speak_default_provider_parses() -> Result<(), Box<dyn Error>> {
let _lock = env_lock()?;
let _guards = clear_all_provider_env_vars()?;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test-api-key
speak:
default_provider: kokoro
"#;
let file = write_config(yaml)?;
let config = Config::load(Some(file.path()))?;
let speak = config.speak.as_ref().expect("speak section present");
assert_eq!(speak.default_provider, Some(SynthesisProvider::Kokoro));
Ok(())
}
#[test]
fn test_synthesis_provider_from_str_and_display() {
assert_eq!(
"kokoro"
.parse::<SynthesisProvider>()
.unwrap_or(SynthesisProvider::Mistral),
SynthesisProvider::Kokoro
);
assert_eq!(
"MISTRAL"
.parse::<SynthesisProvider>()
.unwrap_or(SynthesisProvider::Kokoro),
SynthesisProvider::Mistral
);
assert!("bogus".parse::<SynthesisProvider>().is_err());
assert_eq!(SynthesisProvider::Kokoro.to_string(), "kokoro");
assert_eq!(SynthesisProvider::Mistral.to_string(), "mistral");
}
}