use std::{collections::HashMap, path::PathBuf};
pub struct Config {
raw:toml::Table,
overrides:HashMap<String, String>,
}
impl Default for Config {
fn default() -> Self { Self { raw:toml::Table::new(), overrides:HashMap::new() } }
}
impl Config {
pub fn load() -> Self {
let search_paths = vec![
PathBuf::from("aphrodite.toml"),
dirs::home_dir()
.unwrap_or_default()
.join(".hermes")
.join("aphrodite")
.join("aphrodite.toml"),
];
for path in &search_paths {
if let Ok(content) = std::fs::read_to_string(path) {
match content.parse::<toml::Table>() {
Ok(table) => return Self { raw:table, overrides:HashMap::new() },
Err(err) => {
tracing::warn!(
path = %path.display(),
error = %err,
"aphrodite.toml found but failed to parse; skipping"
);
},
}
}
}
Self::default()
}
pub fn reload(&mut self) { *self = Self::load(); }
pub fn load_from(path:&str) -> Self {
if let Ok(content) = std::fs::read_to_string(path) {
match content.parse::<toml::Table>() {
Ok(table) => return Self { raw:table, overrides:HashMap::new() },
Err(err) => {
tracing::warn!(
path = %path,
error = %err,
"aphrodite.toml found but failed to parse; using defaults"
);
},
}
}
Self::default()
}
pub fn set_override(&mut self, key:&str, value:&str) { self.overrides.insert(key.to_string(), value.to_string()); }
fn section(&self, name:&str) -> Option<&toml::Table> { self.raw.get(name).and_then(|v| v.as_table()) }
pub fn get_bool(&self, env_key:&str, section:&str, key:&str, default:bool) -> bool {
if let Some(v) = self.overrides.get(env_key) {
let v = v.to_ascii_lowercase();
return v == "true" || v == "1";
}
if let Ok(v) = std::env::var(env_key) {
let v = v.to_ascii_lowercase();
return v == "true" || v == "1";
}
self.section(section)
.and_then(|s| s.get(key))
.and_then(|v| v.as_bool())
.unwrap_or(default)
}
pub fn get_u64(&self, env_key:&str, section:&str, key:&str, default:u64) -> u64 {
if let Some(v) = self.overrides.get(env_key) {
return v.parse().unwrap_or(default);
}
if let Ok(v) = std::env::var(env_key) {
return v.parse().unwrap_or(default);
}
self.section(section)
.and_then(|s| s.get(key))
.and_then(|v| v.as_integer())
.map(|v| v as u64)
.unwrap_or(default)
}
pub fn get_usize(&self, env_key:&str, section:&str, key:&str, default:usize) -> usize {
self.get_u64(env_key, section, key, default as u64) as usize
}
pub fn get_string(&self, env_key:&str, section:&str, key:&str, default:&str) -> String {
if let Some(v) = self.overrides.get(env_key) {
return v.clone();
}
if let Ok(v) = std::env::var(env_key) {
return v;
}
self.section(section)
.and_then(|s| s.get(key))
.and_then(|v| v.as_str())
.map(|v| v.to_string())
.unwrap_or_else(|| default.to_string())
}
pub fn get_string_list(&self, section:&str, key:&str) -> Vec<String> {
self.section(section)
.and_then(|s| s.get(key))
.and_then(|v| v.as_array())
.map(|a| a.iter().filter_map(|v| v.as_str().map(|s| s.to_string())).collect())
.unwrap_or_default()
}
pub fn apply_compression(&self, state:&mut crate::state::AphroditeState) {
state.context_engine_enabled = self.get_bool("APHRODITE_CONTEXT_ENGINE", "compression", "context_engine", true);
state.engine_threshold_pct =
self.get_u64("APHRODITE_ENGINE_THRESHOLD_PCT", "compression", "engine_threshold_pct", 45);
state.engine_min_msgs = self.get_usize("APHRODITE_ENGINE_MIN_MSGS", "compression", "engine_min_msgs", 8);
state.engine_protect_first =
self.get_usize("APHRODITE_ENGINE_PROTECT_FIRST", "compression", "engine_protect_first", 2);
state.engine_protect_last =
self.get_usize("APHRODITE_ENGINE_PROTECT_LAST", "compression", "engine_protect_last", 5);
state.tool_threshold =
self.get_usize("APHRODITE_TOOL_THRESHOLD_TOKEN", "compression", "tool_threshold_token", 4096);
state.terminal_threshold =
self.get_usize("APHRODITE_TERMINAL_THRESHOLD", "compression", "terminal_threshold", 1024);
state.model = self.get_string("APHRODITE_MODEL", "defaults", "model", "gpt-4o");
state.api_url = self.get_string("APHRODITE_API_URL", "defaults", "api_url", "");
state.flow_budget_chars = self.get_usize("APHRODITE_FLOW_BUDGET_CHARS", "flow", "budget_chars", 4000);
state.poll_worker_enabled = self.get_bool("APHRODITE_POLL_WORKER", "compression", "poll_worker", true);
state.chain_split_enabled = self.get_bool("APHRODITE_CHAIN_SPLIT", "compression", "chain_split", false);
state.chain_split_min_segments = self
.get_usize(
"APHRODITE_CHAIN_SPLIT_MIN_SEGMENTS",
"compression",
"chain_split_min_segments",
2,
)
.max(2);
state.chain_split_floor = state.chain_split_min_segments;
state.chain_split_max_segments = self
.get_usize(
"APHRODITE_CHAIN_SPLIT_MAX_SEGMENTS",
"compression",
"chain_split_max_segments",
6,
)
.max(state.chain_split_min_segments);
let home_aphrodite = dirs::home_dir().unwrap_or_default().join(".hermes").join("aphrodite");
let bin_relative = std::env::current_exe()
.ok()
.and_then(|p| p.parent().map(|d| d.join("directives")));
let mut dirs:Vec<std::path::PathBuf> = Vec::new();
if let Ok(env_dir) = std::env::var("APHRODITE_DIRECTIVES_DIR")
&& !env_dir.trim().is_empty()
{
dirs.push(std::path::PathBuf::from(env_dir));
}
dirs.push(std::path::PathBuf::from("directives"));
dirs.push(home_aphrodite.join("directives"));
if let Some(bin_dir) = bin_relative {
dirs.push(bin_dir);
}
let mut probed_unusable:Vec<String> = Vec::new();
let mut selected:Option<std::path::PathBuf> = None;
for dir in dirs {
let exists = dir.is_dir();
tracing::info!(
directive_source = "none",
path = %dir.display(),
exists = exists,
"probing directives candidate"
);
if exists {
let loaded = crate::directives::load_directives(&dir);
tracing::info!(
directive_source = "disk",
path = %dir.display(),
count = loaded.len(),
"selected directives source: {} ({} directive(s) loaded)",
dir.display(),
loaded.len()
);
state.directives = loaded;
selected = Some(dir);
break;
}
probed_unusable.push(dir.display().to_string());
}
if selected.is_none() {
tracing::warn!(
directive_source = "builtins",
probed = ?probed_unusable,
"no usable directives directory found (probed: {}); falling back to built-in directives",
probed_unusable.join(", ")
);
state.directives = crate::directives::loaded_builtins();
}
let active = self.get_string_list("directives", "active");
state.active_directives = active.into_iter().filter(|name| state.directives.contains_key(name)).collect();
if state.active_directives.is_empty() && !state.directives.is_empty() {
for name in ["focus", "foresight", "lazy"] {
if state.directives.contains_key(name) {
state.active_directives.push(name.to_string());
}
}
}
state.session_inject = self.get_string(
"APHRODITE_SESSION_INJECT",
"prompts",
"session_inject",
crate::flow::SHIPPED_SESSION_INJECT,
);
}
pub fn apply_previews(&self) {
let max = self.get_u64("APHRODITE_PREVIEW_MAX_CHARS", "previews", "preview_max_chars", 0);
crate::preview::set_preview_max_chars(if max == 0 { None } else { Some(max.min(u32::MAX as u64) as u32) });
}
}
#[cfg(test)]
mod tests {
use super::*;
static CWD_GUARD:std::sync::OnceLock<std::sync::Mutex<()>> = std::sync::OnceLock::new();
#[test]
fn test_defaults() {
let cfg = Config::default();
assert_eq!(cfg.get_u64("NONEXISTENT", "compression", "threshold", 42), 42);
assert!(cfg.get_bool("NONEXISTENT", "compression", "enabled", true));
assert_eq!(cfg.get_string("NONEXISTENT", "defaults", "model", "gpt-4o"), "gpt-4o");
}
#[test]
fn test_override() {
let mut cfg = Config::default();
cfg.set_override("APHRODITE_ENGINE_THRESHOLD_PCT", "90");
assert_eq!(
cfg.get_u64("APHRODITE_ENGINE_THRESHOLD_PCT", "compression", "engine_threshold_pct", 45),
90
);
}
#[test]
fn test_apply_compression_reads_tool_threshold_token_key() {
let mut cfg = Config::default();
cfg.set_override("APHRODITE_TOOL_THRESHOLD_TOKEN", "777");
let mut state = crate::state::AphroditeState::default();
cfg.apply_compression(&mut state);
assert_eq!(state.tool_threshold, 777);
}
#[test]
fn test_apply_compression_from_toml_table() {
let cfg = Config {
raw:"[compression]\ntool_threshold_token = 321\n".parse().unwrap(),
overrides:HashMap::new(),
};
let mut state = crate::state::AphroditeState::default();
cfg.apply_compression(&mut state);
assert_eq!(state.tool_threshold, 321);
}
#[test]
fn test_flow_budget_from_toml() {
let cfg = Config { raw:"[flow]\nbudget_chars = 1234\n".parse().unwrap(), overrides:HashMap::new() };
let mut state = crate::state::AphroditeState::default();
cfg.apply_compression(&mut state);
assert_eq!(state.flow_budget_chars, 1234);
let mut cfg2 = Config::default();
cfg2.set_override("APHRODITE_FLOW_BUDGET_CHARS", "555");
let mut state2 = crate::state::AphroditeState::default();
cfg2.apply_compression(&mut state2);
assert_eq!(state2.flow_budget_chars, 555);
let mut state3 = crate::state::AphroditeState::default();
Config::default().apply_compression(&mut state3);
assert_eq!(state3.flow_budget_chars, 4000);
}
#[test]
fn test_directives_loaded_even_when_active_empty() {
let _g = CWD_GUARD.get_or_init(|| std::sync::Mutex::new(())).lock().unwrap();
let tmp = std::env::temp_dir().join(format!(
"aphrodite-cfg-directives-{}",
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
));
std::fs::create_dir_all(tmp.join("directives")).unwrap();
std::fs::write(tmp.join("directives").join("focus.md"), "# focus\nstay targeted").unwrap();
let original = std::env::current_dir().unwrap();
std::env::set_current_dir(&tmp).unwrap();
let cfg = Config { raw:"[directives]\nactive = []\n".parse().unwrap(), overrides:HashMap::new() };
let mut state = crate::state::AphroditeState::default();
cfg.apply_compression(&mut state);
std::env::set_current_dir(&original).unwrap();
let _ = std::fs::remove_dir_all(&tmp);
assert!(
state.directives.contains_key("focus"),
"directives must load even when [directives] active is empty"
);
assert!(
!state.active_directives.is_empty(),
"empty active list should seed focus + foresight defaults"
);
assert!(state.active_directives.contains(&"focus".to_string()));
}
#[test]
fn test_directives_env_override_wins_over_cwd() {
let _g = CWD_GUARD.get_or_init(|| std::sync::Mutex::new(())).lock().unwrap();
let stamp = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let env_dir = std::env::temp_dir().join(format!("aphrodite-cfg-envdir-{stamp}"));
let cwd_dir = std::env::temp_dir().join(format!("aphrodite-cfg-env-cwd-{stamp}"));
std::fs::create_dir_all(&env_dir).unwrap();
std::fs::create_dir_all(cwd_dir.join("directives")).unwrap();
std::fs::write(env_dir.join("envwin.md"), "# envwin\nfrom env override").unwrap();
std::fs::write(cwd_dir.join("directives").join("cwdwin.md"), "# cwdwin\nfrom cwd").unwrap();
let original = std::env::current_dir().unwrap();
std::env::set_current_dir(&cwd_dir).unwrap();
unsafe { std::env::set_var("APHRODITE_DIRECTIVES_DIR", &env_dir) };
let mut state = crate::state::AphroditeState::default();
Config::default().apply_compression(&mut state);
unsafe { std::env::remove_var("APHRODITE_DIRECTIVES_DIR") };
std::env::set_current_dir(&original).unwrap();
let _ = std::fs::remove_dir_all(&env_dir);
let _ = std::fs::remove_dir_all(&cwd_dir);
assert!(
state.directives.contains_key("envwin"),
"$APHRODITE_DIRECTIVES_DIR must win as the first candidate"
);
assert!(
!state.directives.contains_key("cwdwin"),
"the cwd directives/ candidate must not beat the env override"
);
}
#[test]
fn test_poll_worker_enabled_default_true() {
let cfg = Config::default();
let mut state = crate::state::AphroditeState::default();
cfg.apply_compression(&mut state);
assert!(state.poll_worker_enabled, "poll_worker must default to true");
}
#[test]
fn test_poll_worker_disabled_via_env_override() {
let mut cfg = Config::default();
cfg.set_override("APHRODITE_POLL_WORKER", "false");
let mut state = crate::state::AphroditeState::default();
cfg.apply_compression(&mut state);
assert!(!state.poll_worker_enabled);
}
#[test]
fn test_poll_worker_disabled_via_toml() {
let cfg = Config {
raw:"[compression]\npoll_worker = false\n".parse().unwrap(),
overrides:HashMap::new(),
};
let mut state = crate::state::AphroditeState::default();
cfg.apply_compression(&mut state);
assert!(!state.poll_worker_enabled);
}
#[test]
fn test_apply_previews_wires_preview_max_chars() {
let _g = crate::preview::preview_cap_test_guard();
let env_backup = std::env::var("APHRODITE_PREVIEW_MAX_CHARS").ok();
unsafe { std::env::remove_var("APHRODITE_PREVIEW_MAX_CHARS") };
let cfg = Config {
raw:"[previews]\npreview_max_chars = 77\n".parse().unwrap(),
overrides:HashMap::new(),
};
cfg.apply_previews();
assert_eq!(crate::preview::preview_max_chars(), 77);
let mut cfg2 = Config::default();
cfg2.set_override("APHRODITE_PREVIEW_MAX_CHARS", "123");
cfg2.apply_previews();
assert_eq!(crate::preview::preview_max_chars(), 123);
Config::default().apply_previews();
assert_eq!(crate::preview::preview_max_chars(), 0);
if let Some(v) = env_backup {
unsafe { std::env::set_var("APHRODITE_PREVIEW_MAX_CHARS", v) };
}
}
#[test]
fn test_get_bool_override_case_insensitive() {
let mut cfg = Config::default();
for (value, expected) in [
("TRUE", true),
("True", true),
("tRuE", true),
("1", true),
("YES", false),
("on", false),
("0", false),
("FALSE", false),
] {
cfg.set_override("APHRODITE_TEST_BOOL_CASE", value);
assert_eq!(
cfg.get_bool("APHRODITE_TEST_BOOL_CASE", "compression", "enabled", true),
expected,
"override value {value:?} must resolve to {expected}"
);
}
}
#[test]
fn test_get_bool_env_case_insensitive() {
let _g = CWD_GUARD.get_or_init(|| std::sync::Mutex::new(())).lock().unwrap();
const KEY:&str = "APHRODITE_TEST_BOOL_CASE_ENV";
let backup = std::env::var(KEY).ok();
unsafe { std::env::remove_var(KEY) };
let cfg = Config::default();
for (value, expected) in [
("TRUE", true),
("True", true),
("tRuE", true),
("1", true),
("YES", false),
("on", false),
("0", false),
("FALSE", false),
] {
unsafe { std::env::set_var(KEY, value) };
assert_eq!(
cfg.get_bool(KEY, "compression", "enabled", true),
expected,
"env value {value:?} must resolve to {expected}"
);
}
unsafe { std::env::remove_var(KEY) };
if let Some(v) = backup {
unsafe { std::env::set_var(KEY, v) };
}
}
struct WarnCapture {
tx:std::sync::mpsc::Sender<String>,
}
impl tracing::Subscriber for WarnCapture {
fn enabled(&self, _m:&tracing::Metadata<'_>) -> bool { true }
fn new_span(&self, _s:&tracing::span::Attributes<'_>) -> tracing::span::Id { tracing::span::Id::from_u64(1) }
fn record(&self, _span:&tracing::span::Id, _values:&tracing::span::Record<'_>) {}
fn record_follows_from(&self, _span:&tracing::span::Id, _follows:&tracing::span::Id) {}
fn enter(&self, _span:&tracing::span::Id) {}
fn exit(&self, _span:&tracing::span::Id) {}
fn event(&self, e:&tracing::Event<'_>) {
if *e.metadata().level() != tracing::Level::WARN {
return;
}
let mut s = String::new();
struct Rec<'a>(&'a mut String);
impl tracing::field::Visit for Rec<'_> {
fn record_debug(&mut self, _f:&tracing::field::Field, v:&dyn std::fmt::Debug) {
self.0.push_str(&format!("{v:?} "));
}
}
e.record(&mut Rec(&mut s));
let _ = self.tx.send(s);
}
}
#[test]
fn test_load_from_broken_toml_warns_and_returns_defaults() {
let stamp = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let broken = std::env::temp_dir().join(format!("aphrodite-cfg-broken-{stamp}.toml"));
std::fs::write(&broken, "[compression\nenabled = false\n").unwrap();
let (tx, rx) = std::sync::mpsc::channel();
let cfg = tracing::subscriber::with_default(WarnCapture { tx }, || Config::load_from(broken.to_str().unwrap()));
let _ = std::fs::remove_file(&broken);
assert!(cfg.get_bool("NONEXISTENT", "compression", "enabled", true));
let msg = rx
.recv_timeout(std::time::Duration::from_secs(5))
.expect("warn must be emitted for a broken TOML file");
assert!(msg.contains("failed to parse"), "warn must mention the parse failure: {msg}");
assert!(
msg.contains(broken.to_str().unwrap()),
"warn must include the offending path: {msg}"
);
}
#[test]
fn test_load_from_valid_and_missing() {
let stamp = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let valid = std::env::temp_dir().join(format!("aphrodite-cfg-valid-{stamp}.toml"));
std::fs::write(&valid, "[compression]\nenabled = false\n").unwrap();
let cfg = Config::load_from(valid.to_str().unwrap());
let _ = std::fs::remove_file(&valid);
assert!(
!cfg.get_bool("NONEXISTENT", "compression", "enabled", true),
"a valid file must parse and win over the default"
);
let cfg = Config::load_from(&format!("/nonexistent/aphrodite-{stamp}.toml"));
assert!(cfg.get_bool("NONEXISTENT", "compression", "enabled", true));
}
}