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>,
}
impl RoleConfig {
pub(crate) fn upsert(
configs: &mut Vec<RoleConfig>,
role: impl Into<String>,
set_field: impl FnOnce(&mut RoleConfig),
) {
let role = role.into();
if let Some(existing) = configs.iter_mut().find(|rc| rc.role == role) {
set_field(existing);
} else {
let mut new = RoleConfig {
role,
model: None,
reasoning_effort: None,
};
set_field(&mut new);
configs.push(new);
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct ModelRouting {
pub model: String,
pub provider_order: Option<String>,
pub allow_fallbacks: Option<bool>,
}
impl ModelRouting {
pub(crate) fn upsert(
routings: &mut Vec<ModelRouting>,
model: impl Into<String>,
set_field: impl FnOnce(&mut ModelRouting),
) {
let model = model.into();
if let Some(existing) = routings.iter_mut().find(|mr| mr.model == model) {
set_field(existing);
} else {
let mut new = ModelRouting {
model,
provider_order: None,
allow_fallbacks: None,
};
set_field(&mut new);
routings.push(new);
}
}
}
#[derive(Debug, Clone, PartialEq, 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 firecrawl_key: Option<String>,
pub exa_key: Option<String>,
pub web_search_provider: 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)],
firecrawl_key [non_empty],
exa_key [non_empty],
web_search_provider [non_empty],
telegram_bot_token [non_empty],
}
impl ConfigData {
pub(crate) fn normalize_entries(&mut self) {
for rc in &mut self.per_role_configs {
rc.model = non_empty(rc.model.take());
rc.reasoning_effort = non_empty(rc.reasoning_effort.take());
}
for mr in &mut self.model_routings {
mr.provider_order = non_empty(mr.provider_order.take());
}
}
pub(crate) fn finalize(&mut self) {
self.normalize_string_fields();
self.normalize_entries();
self.per_role_configs.sort_by(|a, b| a.role.cmp(&b.role));
self.model_routings.sort_by(|a, b| a.model.cmp(&b.model));
}
}
#[must_use]
pub(crate) fn trimmed_or_none(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| trimmed_or_none(&s))
}
#[must_use]
fn resolve_or(val: Option<String>, fallback: &str) -> String {
non_empty(val).unwrap_or(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::STRUCT_FIELDS_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::STRUCT_FIELDS_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.finalize();
CONFIG.swap(config);
tracing::info!("Config reloaded from DB");
Ok(())
}
pub async fn save_and_reload(mut 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();
let tx = store.begin_tx().await?;
for (key, value) in config.string_fields() {
if let Some(v) = value.filter(|v| !v.is_empty()) {
store.set_kv_tx(&tx, key, v).await?;
} else {
store.delete_kv_tx(&tx, key).await?;
}
}
store
.save_role_and_routing_configs_tx(&tx, &config.per_role_configs, &config.model_routings)
.await?;
tx.commit().await?;
config.finalize();
let new_token = config.telegram_bot_token.clone();
CONFIG.swap(config);
tracing::info!("Config saved and swapped into runtime");
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::*;
use std::collections::HashSet;
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"),
("firecrawl_key", "fc-test-key"),
("exa_key", "exa-test-key"),
("web_search_provider", "exa"),
("telegram_bot_token", "123:abc"),
];
#[test]
fn string_fields_roundtrip() {
let mut config = ConfigData::default();
let field_keys: HashSet<&str> = config.string_fields().iter().map(|(k, _)| *k).collect();
let test_keys: HashSet<&str> = STRING_FIELD_TEST_VALUES.iter().map(|(k, _)| *k).collect();
assert_eq!(
field_keys, test_keys,
"STRING_FIELD_TEST_VALUES is out of sync with string_config_fields! \
macro invocation — add (or remove) an entry for each field listed \
above/below"
);
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 trimmed_or_none_trims_whitespace() {
assert_eq!(trimmed_or_none(" value "), Some("value".to_string()));
assert_eq!(trimmed_or_none(" "), None);
assert_eq!(trimmed_or_none(""), 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"
);
}
#[test]
fn struct_fields_default_matches_derive_default() {
assert_eq!(ConfigData::default(), ConfigData::STRUCT_FIELDS_DEFAULT);
}
#[test]
fn normalize_entries_works() {
let mut config = ConfigData {
per_role_configs: vec![
RoleConfig {
role: "engineer".into(),
model: Some(String::new()),
reasoning_effort: Some(" high ".into()),
},
RoleConfig {
role: "manager".into(),
model: Some(" gpt-4 ".into()),
reasoning_effort: None,
},
],
model_routings: vec![ModelRouting {
model: "gpt-4".into(),
provider_order: Some(" ".into()),
allow_fallbacks: None,
}],
..ConfigData::default()
};
config.normalize_entries();
assert_eq!(config.per_role_configs[0].model, None);
assert_eq!(
config.per_role_configs[0].reasoning_effort,
Some("high".into())
);
assert_eq!(config.per_role_configs[1].model, Some("gpt-4".into()));
assert_eq!(config.per_role_configs[1].reasoning_effort, None);
assert_eq!(config.model_routings[0].provider_order, None);
}
fn role_config(role: &str, model: Option<&str>, reasoning_effort: Option<&str>) -> RoleConfig {
RoleConfig {
role: role.into(),
model: model.map(String::from),
reasoning_effort: reasoning_effort.map(String::from),
}
}
fn model_routing(
model: &str,
provider_order: Option<&str>,
allow_fallbacks: Option<bool>,
) -> ModelRouting {
ModelRouting {
model: model.into(),
provider_order: provider_order.map(String::from),
allow_fallbacks,
}
}
macro_rules! test_upsert_three_scenarios {
(
$helper:ident, // role_config or model_routing
$upsert:path, // RoleConfig::upsert or ModelRouting::upsert
$key:expr, // "engineer" or "gpt-4"
$key_accessor:ident, // .role (RoleConfig) or .model (ModelRouting)
$field_a:ident, $field_b:ident, $start_a:expr, $start_b:expr, $update_a:expr, $update_b:expr, $new_a:expr, $new_b:expr, $none_start:expr, ) => {
{
let (a, b) = $start_a;
let mut items = vec![$helper($key, a, b)];
$upsert(&mut items, $key, |item| item.$field_a = $update_a);
assert_eq!(items.len(), 1);
assert_eq!(
items[0].$field_a, $update_a,
concat!("[", stringify!($field_a), "] target field updated")
);
assert_eq!(
items[0].$field_b,
b.map(Into::into),
concat!("[", stringify!($field_a), "] other field preserved")
);
}
{
let (a, b) = $start_b;
let mut items = vec![$helper($key, a, b)];
$upsert(&mut items, $key, |item| item.$field_b = $update_b);
assert_eq!(items.len(), 1);
assert_eq!(
items[0].$field_b, $update_b,
concat!("[", stringify!($field_b), "] target field updated")
);
assert_eq!(
items[0].$field_a,
a.map(Into::into),
concat!("[", stringify!($field_b), "] other field preserved")
);
}
{
let mut items = vec![];
$upsert(&mut items, $key, |item| item.$field_a = $new_a);
assert_eq!(items.len(), 1);
assert_eq!(items[0].$key_accessor, $key);
assert_eq!(
items[0].$field_a, $new_a,
concat!("[", stringify!($field_a), "] set on new entry")
);
assert_eq!(items[0].$field_b, None);
}
{
let mut items = vec![];
$upsert(&mut items, $key, |item| item.$field_b = $new_b);
assert_eq!(items.len(), 1);
assert_eq!(items[0].$key_accessor, $key);
assert_eq!(
items[0].$field_b, $new_b,
concat!("[", stringify!($field_b), "] set on new entry")
);
assert_eq!(items[0].$field_a, None);
}
{
let (a, b) = $none_start;
let mut items = vec![$helper($key, a, b)];
$upsert(&mut items, $key, |item| item.$field_a = None);
assert_eq!(
items[0].$field_a, None,
concat!("[", stringify!($field_a), "] cleared to None")
);
assert_eq!(
items[0].$field_b,
b.map(Into::into),
concat!(
"[",
stringify!($field_a),
"] other field preserved when clearing"
)
);
}
{
let (a, b) = $none_start;
let mut items = vec![$helper($key, a, b)];
$upsert(&mut items, $key, |item| item.$field_b = None);
assert_eq!(
items[0].$field_b, None,
concat!("[", stringify!($field_b), "] cleared to None")
);
assert_eq!(
items[0].$field_a,
a.map(Into::into),
concat!(
"[",
stringify!($field_b),
"] other field preserved when clearing"
)
);
}
};
}
#[test]
fn upsert_role_config_fields() {
test_upsert_three_scenarios! {
role_config,
RoleConfig::upsert,
"engineer",
role,
model,
reasoning_effort,
(Some("old"), Some("high")),
(Some("gpt-4"), Some("low")),
Some("new".into()),
Some("high".into()),
Some("gpt-4".into()),
Some("high".into()),
(Some("gpt-4"), Some("high")),
}
}
#[test]
fn upsert_model_routing_fields() {
test_upsert_three_scenarios! {
model_routing,
ModelRouting::upsert,
"gpt-4",
model,
provider_order,
allow_fallbacks,
(Some("OpenAi"), Some(true)),
(Some("OpenAi"), Some(true)),
Some("Anthropic".into()),
Some(false),
Some("OpenAi".into()),
Some(false),
(Some("OpenAi"), Some(true)),
}
}
#[test]
fn upsert_multiple_entries_independent_keys() {
let mut configs = vec![
role_config("engineer", Some("model-a"), None),
role_config("manager", Some("model-b"), Some("high")),
];
let mut routings = vec![
model_routing("gpt-4", Some("OpenAi"), None),
model_routing("claude-3", Some("Anthropic"), Some(true)),
];
RoleConfig::upsert(&mut configs, "engineer", |c| {
c.model = Some("model-c".into());
});
assert_eq!(configs[0].model, Some("model-c".into()));
assert_eq!(configs[1].model, Some("model-b".into()));
RoleConfig::upsert(&mut configs, "manager", |c| {
c.reasoning_effort = Some("low".into());
});
assert_eq!(configs[1].reasoning_effort, Some("low".into()));
assert_eq!(configs[0].reasoning_effort, None);
ModelRouting::upsert(&mut routings, "gpt-4", |mr| {
mr.provider_order = Some("Google".into());
});
assert_eq!(routings[0].provider_order, Some("Google".into()));
assert_eq!(routings[1].provider_order, Some("Anthropic".into()));
ModelRouting::upsert(&mut routings, "claude-3", |mr| {
mr.allow_fallbacks = Some(false);
});
assert_eq!(routings[1].allow_fallbacks, Some(false));
assert_eq!(routings[0].allow_fallbacks, None);
assert_eq!(configs.len(), 2);
assert_eq!(routings.len(), 2);
}
}