use std::collections::BTreeMap;
use std::path::{Path, PathBuf};
use serde::Deserialize;
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(default)]
pub struct Config {
pub default: DefaultSelection,
pub provider: BTreeMap<String, ProviderConfig>,
pub mcp: BTreeMap<String, McpConfig>,
pub proxy: ProxyConfig,
pub agent: AgentConfig,
pub retry: Option<RetryConfig>,
pub compaction: Option<CompactionConfig>,
pub permission: PermissionConfig,
pub tools: Option<ToolsConfig>,
pub failover: Option<FailoverConfig>,
pub quiet_startup: Option<bool>,
}
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(default)]
pub struct FailoverConfig {
pub enabled: bool,
pub providers: Vec<String>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(default)]
pub struct ToolsConfig {
pub bash_filter: Option<String>,
pub shell: Option<String>,
pub timeout_secs: u64,
}
impl Default for ToolsConfig {
fn default() -> Self {
Self {
bash_filter: None,
shell: None,
timeout_secs: 120,
}
}
}
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(default)]
pub struct DefaultSelection {
pub provider: Option<String>,
pub model: Option<String>,
}
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(default)]
pub struct McpConfig {
pub command: String,
pub args: Vec<String>,
pub prefix: Option<String>,
pub url: Option<String>,
pub respawn: bool,
}
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(default)]
pub struct ProviderConfig {
pub protocol: Option<String>,
pub base_url: String,
pub api_key: Option<String>,
pub models: Vec<String>,
pub effort: Option<crate::llm::ir::EffortLevel>,
pub legacy_thinking: bool,
pub context_window: Option<u64>,
pub max_tokens: Option<u64>,
pub headers: std::collections::BTreeMap<String, String>,
}
impl ProviderConfig {
pub fn default_model(&self) -> String {
self.models
.first()
.cloned()
.unwrap_or_else(|| "model".to_string())
}
pub fn dummy() -> Self {
Self {
base_url: "http://localhost:9999/v1".into(),
api_key: None,
models: vec!["test-model".into()],
..Default::default()
}
}
}
#[derive(Debug, Clone, Deserialize)]
#[serde(default)]
pub struct ProxyConfig {
pub url: Option<String>,
pub connect_timeout_secs: u64,
pub read_timeout_secs: u64,
}
impl Default for ProxyConfig {
fn default() -> Self {
Self {
url: None,
connect_timeout_secs: 30,
read_timeout_secs: 0,
}
}
}
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(default)]
pub struct AgentConfig {
pub max_turns: Option<u32>,
pub budget_tokens: Option<u64>,
pub parallel_tools: Option<bool>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(default)]
pub struct RetryConfig {
pub enabled: bool,
pub max_retries: u32,
pub base_delay_ms: u64,
pub max_delay_ms: u64,
pub jitter: bool,
pub respect_retry_after: bool,
}
impl Default for RetryConfig {
fn default() -> Self {
Self {
enabled: true,
max_retries: 5,
base_delay_ms: 2_000,
max_delay_ms: 60_000,
jitter: true,
respect_retry_after: true,
}
}
}
#[derive(Debug, Clone, Deserialize)]
#[serde(default)]
pub struct CompactionConfig {
pub enabled: bool,
pub dedup: bool,
pub purge_errors_after: u32,
pub nudge_threshold: f64,
pub compress_tool: bool,
}
impl Default for CompactionConfig {
fn default() -> Self {
Self {
enabled: true,
dedup: true,
purge_errors_after: 4,
nudge_threshold: 0.7,
compress_tool: true,
}
}
}
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(default)]
pub struct PermissionConfig {
pub allow: Vec<String>,
}
impl Config {
pub fn from_files(paths: &[PathBuf]) -> anyhow::Result<Self> {
let mut cfg = Self::default();
for p in paths {
if p.is_file() {
let raw = std::fs::read_to_string(p)?;
let parsed: Config = toml::from_str(&raw)
.map_err(|e| anyhow::anyhow!("bad config {}: {e}", p.display()))?;
cfg.merge(parsed);
}
}
cfg.expand_env_refs();
Ok(cfg)
}
fn merge(&mut self, other: Self) {
if other.default.provider.is_some() {
self.default.provider = other.default.provider;
}
if other.default.model.is_some() {
self.default.model = other.default.model;
}
for (k, v) in other.provider {
self.provider.insert(k, v);
}
for (k, v) in other.mcp {
self.mcp.insert(k, v);
}
if other.proxy.url.is_some() {
self.proxy.url = other.proxy.url;
}
if other.agent.max_turns.is_some() {
self.agent.max_turns = other.agent.max_turns;
}
if other.agent.budget_tokens.is_some() {
self.agent.budget_tokens = other.agent.budget_tokens;
}
if other.agent.parallel_tools.is_some() {
self.agent.parallel_tools = other.agent.parallel_tools;
}
if other.retry.is_some() {
self.retry = other.retry;
}
if other.compaction.is_some() {
self.compaction = other.compaction;
}
if !other.permission.allow.is_empty() {
self.permission.allow = other.permission.allow;
}
if other.tools.is_some() {
self.tools = other.tools;
}
if other.failover.is_some() {
self.failover = other.failover;
}
if other.quiet_startup.is_some() {
self.quiet_startup = other.quiet_startup;
}
}
fn expand_env_refs(&mut self) {
for p in self.provider.values_mut() {
if let Some(key) = &p.api_key {
p.api_key = expand_env_ref(key);
}
}
}
pub fn config_paths(cli_config: Option<&Path>) -> Vec<PathBuf> {
let mut paths = Vec::new();
if let Some(home) = dirs::home_dir() {
paths.push(home.join(".config").join("hey").join("config.toml"));
}
if let Some(proj) = find_project_root(std::env::current_dir().unwrap_or_default()) {
paths.push(proj.join(".hey").join("config.toml"));
}
if let Some(p) = cli_config {
paths.push(p.to_path_buf());
}
paths
}
pub fn resolve(
&self,
base_url: Option<&str>,
protocol: Option<&str>,
api_key: Option<&str>,
cli_model: Option<&str>,
) -> Result<(String, String, ProviderConfig), String> {
let env_provider = std::env::var("HEY_PROVIDER").ok().filter(|s| !s.is_empty());
let env_model = std::env::var("HEY_MODEL").ok().filter(|s| !s.is_empty());
if let Some(base_url) = base_url {
let m = cli_model
.or(env_model.as_deref())
.ok_or_else(|| "--base-url requires --model <name>".to_string())?;
let (pname, model) = m.split_once('/').unwrap_or(("cli", m));
let protocol = protocol.unwrap_or("openai").to_string();
let pcfg = ProviderConfig {
protocol: Some(protocol),
base_url: base_url.to_string(),
api_key: api_key.map(str::to_string),
models: vec![model.to_string()],
..Default::default()
};
Ok((pname.to_string(), model.to_string(), pcfg))
} else {
let m = cli_model.or(env_model.as_deref());
let (pname, model) = if let Some(m) = m {
let (p, rest) = m.split_once('/').ok_or_else(|| {
format!(
"model \"{m}\" must be provider/model (e.g. zen/deepseek-v4-flash) or use --base-url"
)
})?;
(p.to_string(), rest.to_string())
} else {
let p = env_provider
.or_else(|| self.default.provider.clone())
.or_else(|| self.provider.keys().next().cloned())
.ok_or_else(|| {
"no provider configured; add [provider.*] or use --base-url".to_string()
})?;
let base = self
.provider
.get(&p)
.ok_or_else(|| format!("provider '{p}' not found in config"))?;
let mm = env_model
.or_else(|| self.default.model.clone())
.unwrap_or_else(|| base.default_model());
(p, mm)
};
let pcfg = self
.provider
.get(&pname)
.cloned()
.ok_or_else(|| format!("provider '{pname}' not found in config"))?;
Ok((pname, model, pcfg))
}
}
}
pub fn find_project_root(dir: PathBuf) -> Option<PathBuf> {
let mut cur = Some(dir.as_path());
while let Some(d) = cur {
if d.join(".git").exists() {
return Some(d.to_path_buf());
}
cur = d.parent();
}
Some(dir)
}
pub fn expand_env_ref(s: &str) -> Option<String> {
if let Some(rest) = s.strip_prefix("${").and_then(|r| r.strip_suffix('}')) {
if let Some((var, default)) = rest.split_once(":-") {
return std::env::var(var)
.ok()
.or_else(|| Some(default.to_string()));
}
std::env::var(rest).ok()
} else {
Some(s.to_string())
}
}
#[cfg(test)]
mod tests {
use super::*;
static ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
#[test]
fn expand_env_ref_handles_var_and_default() {
let _env = ENV_LOCK
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner); unsafe {
std::env::set_var("HEY_TEST_KEY", "abc");
}
assert_eq!(expand_env_ref("${HEY_TEST_KEY}"), Some("abc".into()));
assert_eq!(expand_env_ref("${NO_SUCH_VAR_XYZ}"), None);
assert_eq!(
expand_env_ref("${NO_SUCH_VAR_XYZ:-fallback}"),
Some("fallback".into())
);
assert_eq!(expand_env_ref("plain"), Some("plain".into()));
unsafe {
std::env::remove_var("HEY_TEST_KEY");
}
}
#[test]
fn parse_mcp_http_url() {
let raw = r#"
[mcp.remote]
url = "https://mcp.example.com/api"
[mcp.local]
command = "python3"
args = ["s.py"]
"#;
let cfg: Config = toml::from_str(raw).unwrap();
assert_eq!(
cfg.mcp["remote"].url.as_deref(),
Some("https://mcp.example.com/api")
);
assert_eq!(cfg.mcp["remote"].command, ""); assert!(cfg.mcp["local"].url.is_none());
assert_eq!(cfg.mcp["local"].command, "python3");
}
#[test]
fn provider_headers_deserialize() {
let raw = r#"
[provider.test]
base_url = "http://localhost"
api_key = "x"
models = ["m"]
headers = { "X-API-Key" = "abc", "OpenAI-Organization" = "org-123" }
"#;
let cfg: Config = toml::from_str(raw).unwrap();
let p = &cfg.provider["test"];
assert_eq!(p.headers.get("X-API-Key").unwrap(), "abc");
assert_eq!(p.headers.get("OpenAI-Organization").unwrap(), "org-123");
assert_eq!(p.headers.len(), 2);
}
#[test]
fn provider_headers_default_empty() {
let raw = r#"
[provider.test]
base_url = "http://localhost"
api_key = "x"
models = ["m"]
"#;
let cfg: Config = toml::from_str(raw).unwrap();
assert!(cfg.provider["test"].headers.is_empty());
}
#[test]
fn merge_keeps_mcp_and_tools() {
let raw = r#"
[provider.zen]
base_url = "http://x/v1"
models = ["m"]
[mcp.fake]
command = "python3"
args = ["s.py"]
[tools]
bash_filter = "grep warn"
"#;
let parsed: Config = toml::from_str(raw).unwrap();
let mut cfg = Config::default();
cfg.merge(parsed);
assert!(cfg.mcp.contains_key("fake"), "mcp 段必须被合并");
assert_eq!(cfg.mcp["fake"].command, "python3");
assert_eq!(
cfg.tools.as_ref().and_then(|t| t.bash_filter.as_deref()),
Some("grep warn")
);
}
#[test]
fn merge_keeps_parallel_tools() {
let raw = r"
[agent]
parallel_tools = false
";
let parsed: Config = toml::from_str(raw).unwrap();
let mut cfg = Config::default();
cfg.merge(parsed);
assert!(
!cfg.agent.parallel_tools.unwrap_or(true),
"显式 false 必须保留"
);
}
#[test]
fn merge_does_not_clobber_with_defaults() {
let user_global: Config = toml::from_str(
"quiet_startup = true
",
)
.unwrap();
let project_default: Config = toml::from_str("[agent]\nmax_turns = 30\n").unwrap();
let mut cfg = Config::default();
cfg.merge(user_global); assert!(
cfg.quiet_startup.unwrap_or(false),
"user 显式 true 必须保留"
);
cfg.merge(project_default); assert!(
cfg.quiet_startup.unwrap_or(false),
"后加载的默认值不应覆盖用户显式设置(actual={:?})",
cfg.quiet_startup
);
assert_eq!(cfg.agent.max_turns, Some(30));
}
#[test]
fn tools_timeout_secs_parse_and_default() {
let cfg: Config = toml::from_str("[tools]\ntimeout_secs = 300\n").unwrap();
assert_eq!(cfg.tools.as_ref().unwrap().timeout_secs, 300);
assert_eq!(
Config::default().tools.unwrap_or_default().timeout_secs,
120
);
let mut merged = Config::default();
merged.merge(toml::from_str("[tools]\ntimeout_secs = 0\n").unwrap());
assert_eq!(
merged.tools.unwrap().timeout_secs,
0,
"0 = 禁用(不包 timeout)"
);
}
#[test]
fn config_from_toml() {
let raw = r#"
[default]
provider = "local"
model = "qwen3:14b"
[provider.local]
base_url = "http://localhost:11434/v1"
models = ["qwen3:14b"]
[agent]
max_turns = 5
"#;
let cfg: Config = toml::from_str(raw).unwrap();
assert_eq!(cfg.agent.max_turns, Some(5));
assert_eq!(cfg.default.provider.as_deref(), Some("local"));
assert_eq!(cfg.default.model.as_deref(), Some("qwen3:14b"));
}
#[test]
fn resolve_base_url_creates_ad_hoc_provider() {
let cfg = Config::default();
let (p, m, pcfg) = cfg
.resolve(
Some("http://localhost:4000/v1"),
None,
Some("k"),
Some("gpt-4o"),
)
.unwrap();
assert_eq!((p.as_str(), m.as_str()), ("cli", "gpt-4o"));
assert_eq!(pcfg.base_url, "http://localhost:4000/v1");
assert_eq!(pcfg.api_key.as_deref(), Some("k"));
assert_eq!(pcfg.protocol.as_deref(), Some("openai"));
let (p, m, _) = cfg
.resolve(
Some("http://x/v1"),
Some("anthropic"),
None,
Some("mybox/haiku"),
)
.unwrap();
assert_eq!((p.as_str(), m.as_str()), ("mybox", "haiku"));
assert!(cfg.resolve(Some("http://x/v1"), None, None, None).is_err());
}
#[test]
fn resolve_config_requires_provider_prefix() {
let mut cfg = Config::default();
cfg.provider.insert(
"zen".into(),
ProviderConfig {
base_url: "http://z/v1".into(),
models: vec!["m1".into()],
..Default::default()
},
);
let (p, m, pcfg) = cfg.resolve(None, None, None, Some("zen/m2")).unwrap();
assert_eq!((p.as_str(), m.as_str()), ("zen", "m2"));
assert_eq!(pcfg.base_url, "http://z/v1");
assert!(cfg.resolve(None, None, None, Some("m1")).is_err());
let (p, m, _) = cfg.resolve(None, None, None, None).unwrap();
assert_eq!((p.as_str(), m.as_str()), ("zen", "m1"));
}
#[test]
fn parse_effort_and_legacy_thinking() {
let raw = r#"
[provider.a]
base_url = "http://a/v1"
effort = "high"
[provider.b]
base_url = "http://b/v1"
protocol = "anthropic"
legacy_thinking = true
"#;
let cfg: Config = toml::from_str(raw).unwrap();
assert_eq!(
cfg.provider["a"].effort,
Some(crate::llm::ir::EffortLevel::High)
);
assert!(!cfg.provider["a"].legacy_thinking); assert_eq!(cfg.provider["b"].effort, None); assert!(cfg.provider["b"].legacy_thinking);
}
}