use crate::Role;
use crate::role::role_info;
use anyhow::{Context, Result};
use directories::UserDirs;
use std::path::PathBuf;
use std::sync::{OnceLock, RwLock, RwLockReadGuard};
use tokio::fs;
pub(crate) const DEFAULT_PROVIDER_ENDPOINT: &str = "https://openrouter.ai/api/v1";
const DEFAULT_IMAGE_GEN_MODEL: &str = "google/gemini-3.1-flash-image-preview";
const DEFAULT_VIDEO_GEN_MODEL: &str = "google/veo-3.1-lite";
pub(crate) const DEFAULT_IMAGE_TRANSCRIPTION_MODEL: &str = "qwen/qwen3.6-plus";
pub(crate) const DEFAULT_AUDIO_TRANSCRIPTION_MODEL: &str = "xiaomi/mimo-v2.5";
#[derive(Debug, Clone, PartialEq)]
pub struct RoleConfig {
pub role: String,
pub model: Option<String>,
pub reasoning_effort: Option<String>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct ModelRouting {
pub model: String,
pub provider_order: Option<String>,
pub allow_fallbacks: Option<bool>,
}
#[derive(Debug, Clone, Default)]
pub struct ConfigData {
pub provider_key: Option<String>,
pub provider_endpoint: Option<String>,
pub image_transcription_model: Option<String>,
pub audio_transcription_model: Option<String>,
pub transcription_provider: Option<String>,
pub audio_transcription_provider: Option<String>,
pub image_gen_model: Option<String>,
pub image_gen_models: Option<String>,
pub video_gen_model: Option<String>,
pub video_gen_models: Option<String>,
pub exa_key: Option<String>,
pub telegram_bot_token: Option<String>,
pub per_role_configs: Vec<RoleConfig>,
pub model_routings: Vec<ModelRouting>,
}
macro_rules! string_config_fields {
(
$(
$field:ident [ $($annotation:tt)* ]
),* $(,)?
) => {
impl ConfigData {
pub(crate) const STRUCT_FIELDS_DEFAULT: Self = Self {
$($field: None,)*
per_role_configs: Vec::new(),
model_routings: Vec::new(),
};
#[must_use]
pub fn string_fields(&self) -> Vec<(&'static str, Option<&str>)> {
vec![$((stringify!($field), self.$field.as_deref())),*]
}
#[must_use]
pub fn set_string_field(&mut self, key: &str, value: &str) -> bool {
match key {
$(stringify!($field) => self.$field = non_empty(Some(value.to_owned())),)*
_ => return false,
}
true
}
pub(crate) fn normalize_string_fields(&mut self) {
$(self.$field = non_empty(self.$field.take());)*
}
}
impl ConfigReload {
$(
string_config_fields!(@accessor $field $($annotation)*);
)*
}
};
(@accessor $field:ident non_empty) => {
#[doc = concat!(
"Returns the configured `", stringify!($field),
"`, with empty/whitespace values collapsed to `None`."
)]
#[must_use]
pub fn $field(&self) -> Option<String> {
non_empty(self.read().$field.clone())
}
};
(@accessor $field:ident or($default:expr)) => {
#[doc = concat!(
"Returns the configured `", stringify!($field),
"`, falling back to the default if unset."
)]
#[must_use]
pub fn $field(&self) -> String {
resolve_or(self.read().$field.clone(), $default)
}
};
(@accessor $field:ident list_or(fallback = $fallback:ident, default = $default:expr)) => {
#[doc = concat!(
"Returns the list of available `", stringify!($field), "`."
)]
#[must_use]
pub fn $field(&self) -> Vec<String> {
let guard = self.read();
resolve_list_or(
guard.$field.as_ref(),
guard.$fallback.clone(),
$default,
)
}
};
}
string_config_fields! {
provider_key [non_empty],
provider_endpoint [or(DEFAULT_PROVIDER_ENDPOINT)],
image_transcription_model [or(DEFAULT_IMAGE_TRANSCRIPTION_MODEL)],
audio_transcription_model [or(DEFAULT_AUDIO_TRANSCRIPTION_MODEL)],
transcription_provider [non_empty],
audio_transcription_provider [non_empty],
image_gen_model [or(DEFAULT_IMAGE_GEN_MODEL)],
image_gen_models [list_or(fallback = image_gen_model, default = DEFAULT_IMAGE_GEN_MODEL)],
video_gen_model [or(DEFAULT_VIDEO_GEN_MODEL)],
video_gen_models [list_or(fallback = video_gen_model, default = DEFAULT_VIDEO_GEN_MODEL)],
exa_key [non_empty],
telegram_bot_token [non_empty],
}
#[must_use]
pub(crate) fn trim_non_empty(s: &str) -> Option<String> {
let trimmed = s.trim();
if trimmed.is_empty() {
None
} else {
Some(trimmed.to_owned())
}
}
#[must_use]
pub(crate) fn parse_newline_list(s: &str) -> Vec<String> {
s.split('\n')
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.collect()
}
#[must_use]
fn non_empty(val: Option<String>) -> Option<String> {
val.and_then(|s| trim_non_empty(&s))
}
#[must_use]
fn resolve_or(val: Option<String>, fallback: &str) -> String {
non_empty(val).unwrap_or_else(|| fallback.to_string())
}
#[must_use]
fn resolve_list_or(
list_field: Option<&String>,
fallback_field: Option<String>,
default_value: &str,
) -> Vec<String> {
if let Some(raw) = list_field {
let parsed = parse_newline_list(raw);
if !parsed.is_empty() {
return parsed;
}
}
vec![resolve_or(fallback_field, default_value)]
}
pub(crate) fn expand_tilde(path: &str) -> PathBuf {
if let Some(stripped) = path.strip_prefix('~') {
let home = std::env::var("HOME").or_else(|_| std::env::var("USERPROFILE"));
if let Ok(home) = home {
return PathBuf::from(home).join(stripped.trim_start_matches('/'));
}
}
PathBuf::from(path)
}
pub static CONFIG: ConfigReload = ConfigReload::const_new();
pub struct ConfigReload {
storage_root: OnceLock<PathBuf>,
inner: RwLock<ConfigData>,
}
impl ConfigReload {
#[must_use]
pub const fn const_new() -> Self {
Self {
storage_root: OnceLock::new(),
inner: RwLock::new(ConfigData::STRUCT_FIELDS_DEFAULT),
}
}
pub fn global_storage_root(&self) -> PathBuf {
self.storage_root
.get()
.expect("CONFIG storage_root not initialized")
.clone()
}
#[must_use]
pub fn try_storage_root(&self) -> Option<PathBuf> {
self.storage_root.get().cloned()
}
pub(crate) fn set_storage_root(&self, root: PathBuf) {
self.storage_root
.set(root)
.expect("CONFIG storage_root already set");
}
#[cfg(test)]
pub(crate) fn try_set_storage_root(&self, root: PathBuf) -> std::result::Result<(), PathBuf> {
self.storage_root.set(root)
}
fn read(&self) -> RwLockReadGuard<'_, ConfigData> {
self.inner.read().expect("CONFIG inner poisoned")
}
pub(crate) fn swap(&self, new_config: ConfigData) {
*self.inner.write().expect("CONFIG inner poisoned") = new_config;
}
#[must_use]
pub fn snapshot(&self) -> ConfigData {
self.read().clone()
}
#[must_use]
pub fn set_string_field_and_apply(&self, key: &str, value: &str) -> bool {
let mut config = self.snapshot();
let recognized = config.set_string_field(key, value);
if recognized {
*self.inner.write().expect("CONFIG inner poisoned") = config;
}
recognized
}
#[must_use]
pub fn model_routing(&self, model: &str) -> ModelRouting {
let guard = self.read();
if let Some(mr) = guard.model_routings.iter().find(|mr| mr.model == model) {
mr.clone()
} else {
ModelRouting {
model: model.to_string(),
provider_order: None,
allow_fallbacks: None,
}
}
}
fn find_role_config(&self, role: Role) -> Option<RoleConfig> {
let role_key: &str = role.into();
let guard = self.read();
guard
.per_role_configs
.iter()
.find(|rc| rc.role == role_key)
.cloned()
}
#[must_use]
pub fn role_model(&self, role: Role) -> String {
if let Some(rc) = self.find_role_config(role)
&& let Some(ref m) = rc.model
&& !m.is_empty()
{
return m.clone();
}
role_info(&role).default_model.to_string()
}
#[must_use]
pub fn role_reasoning_effort(&self, role: Role) -> Option<String> {
if let Some(rc) = self.find_role_config(role)
&& let Some(ref r) = rc.reasoning_effort
&& !r.is_empty()
{
return Some(r.clone());
}
Some(role_info(&role).default_reasoning_effort.to_string())
}
}
pub fn default_config_dir() -> Result<PathBuf> {
if let Ok(home) = std::env::var("HOME")
&& !home.is_empty()
{
return Ok(PathBuf::from(home).join(".mahbot"));
}
let home = UserDirs::new()
.map(|u| u.home_dir().to_path_buf())
.context("Could not find home directory")?;
Ok(home.join(".mahbot"))
}
pub async fn load_or_init() -> Result<()> {
let mahbot_dir = default_config_dir()?;
fs::create_dir_all(&mahbot_dir)
.await
.context("Failed to create config directory")?;
CONFIG.set_storage_root(mahbot_dir.clone());
CONFIG.swap(ConfigData::default());
tracing::info!(
"Config system initialised (storage root: {}).",
mahbot_dir.display()
);
Ok(())
}
pub async fn reload_from_db() -> Result<()> {
let store = crate::config_db::store();
let mut config = ConfigData::default();
let kvs = store.get_all_kv().await?;
for (key, value) in &kvs {
if !config.set_string_field(key, value) {
tracing::debug!(key, "Unknown config key, ignoring");
}
}
let roles = store.get_all_role_configs().await?;
config.per_role_configs = roles;
let routings = store.get_all_model_routings().await?;
config.model_routings = routings;
CONFIG.swap(config);
tracing::info!("Config reloaded from DB");
Ok(())
}
pub async fn save_and_reload(config: &ConfigData) -> Result<()> {
validate_config(config)?;
let old_token = CONFIG.telegram_bot_token();
if config.telegram_bot_token != old_token
&& let Some(ref new_token) = config.telegram_bot_token
{
crate::channels::telegram::TelegramChannel::validate_token(new_token).await?;
}
crate::providers::warmup_provider_from_config(config).await?;
let store = crate::config_db::store();
for (key, value) in config.string_fields() {
if let Some(v) = value.filter(|v| !v.is_empty()) {
store.set_kv(key, v).await?;
} else {
store.delete_kv(key).await?;
}
}
store
.save_routing_configs(&config.per_role_configs, &config.model_routings)
.await?;
let mut config = config.clone();
config.normalize_string_fields();
config.per_role_configs.sort_by(|a, b| a.role.cmp(&b.role));
config.model_routings.sort_by(|a, b| a.model.cmp(&b.model));
let new_token = config.telegram_bot_token.clone();
CONFIG.swap(config);
tracing::info!("Config reloaded from DB");
crate::providers::recreate_all().await?;
if new_token != old_token {
tracing::info!(
old = ?old_token,
new = ?new_token,
"Telegram bot token changed — restarting listener",
);
crate::channels::telegram::restart_telegram_listener(new_token.as_deref()).await?;
}
tracing::info!("Config saved and reloaded successfully");
Ok(())
}
fn validate_config(config: &ConfigData) -> Result<()> {
if let Some(ref ep) = config.provider_endpoint
&& !ep.trim().is_empty()
&& !ep.starts_with("https://")
&& !ep.starts_with("http://")
{
anyhow::bail!("Provider endpoint must be a valid URL starting with https:// or http://");
}
if let Some(ref key) = config.provider_key
&& (key.trim() == "sk-..." || key.trim().starts_with("sk-.."))
{
anyhow::bail!("Provider key is still the placeholder value — please set a real key");
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
const STRING_FIELD_TEST_VALUES: &[(&str, &str)] = &[
("provider_key", "sk-test-key"),
("provider_endpoint", "https://example.com/api"),
("image_transcription_model", "gpt-4-vision"),
("audio_transcription_model", "whisper-1"),
("transcription_provider", "OpenAI"),
("audio_transcription_provider", "Deepgram"),
("image_gen_model", "dall-e-3"),
("image_gen_models", "dall-e-3\nstable-diffusion\nmidjourney"),
("video_gen_model", "sora"),
("video_gen_models", "sora\npika\nrunway"),
("exa_key", "exa-test-key"),
("telegram_bot_token", "123:abc"),
];
#[test]
fn string_fields_roundtrip() {
let mut config = ConfigData::default();
for (_key, value) in config.string_fields() {
assert!(value.is_none(), "field should start as None");
}
for &(key, value) in STRING_FIELD_TEST_VALUES {
let recognized = config.set_string_field(key, value);
assert!(recognized, "key '{key}' should be recognized");
let found = config
.string_fields()
.iter()
.find(|(k, _)| *k == key)
.and_then(|(_, v)| *v);
assert_eq!(
found,
Some(value),
"value for '{key}' should match after set"
);
}
let _ = config.set_string_field("provider_key", "");
let pk = config
.string_fields()
.iter()
.find(|(k, _)| *k == "provider_key")
.and_then(|(_, v)| *v);
assert!(pk.is_none(), "empty string should be stored as None");
let _ = config.set_string_field("provider_key", " ");
let pk = config
.string_fields()
.iter()
.find(|(k, _)| *k == "provider_key")
.and_then(|(_, v)| *v);
assert!(
pk.is_none(),
"whitespace-only string should be stored as None"
);
assert!(!config.set_string_field("nonexistent_key", "value"));
}
#[test]
fn config_reload_accessors_roundtrip() {
let reload = ConfigReload::const_new();
assert_eq!(reload.provider_key(), None, "unset provider_key is None");
let mut config = ConfigData::default();
assert!(config.set_string_field("provider_key", "sk-test"));
reload.swap(config);
assert_eq!(reload.provider_key(), Some("sk-test".to_string()));
reload.swap(ConfigData::default());
assert_eq!(
reload.provider_endpoint(),
DEFAULT_PROVIDER_ENDPOINT,
"unset provider_endpoint falls back to default"
);
let mut empty = ConfigData::default();
assert!(empty.set_string_field("provider_key", ""));
reload.swap(empty);
assert_eq!(
reload.provider_key(),
None,
"empty string is collapsed to None"
);
reload.swap(ConfigData::default());
assert_eq!(
reload.image_gen_models(),
vec![DEFAULT_IMAGE_GEN_MODEL.to_string()],
"unset image_gen_models falls back to active model"
);
let mut list_config = ConfigData::default();
assert!(list_config.set_string_field("image_gen_models", "model-a\nmodel-b\nmodel-c"));
reload.swap(list_config);
assert_eq!(
reload.image_gen_models(),
vec!["model-a", "model-b", "model-c"]
);
}
#[test]
fn trim_non_empty_trims_whitespace() {
assert_eq!(trim_non_empty(" value "), Some("value".to_string()));
assert_eq!(trim_non_empty(" "), None);
assert_eq!(trim_non_empty(""), None);
assert_eq!(non_empty(None), None);
assert_eq!(resolve_or(None, "fallback"), "fallback");
}
#[test]
fn set_string_field_and_apply_updates_in_memory() {
let reload = ConfigReload::const_new();
assert!(!reload.set_string_field_and_apply("nonexistent", "value"));
assert!(reload.set_string_field_and_apply("image_gen_model", "test-model"));
assert_eq!(reload.image_gen_model(), "test-model");
assert!(reload.set_string_field_and_apply("image_gen_model", ""));
assert_eq!(
reload.image_gen_model(),
"google/gemini-3.1-flash-image-preview"
);
}
}