use std::collections::HashMap;
use std::path::{Path, PathBuf};
use std::sync::OnceLock;
use serde::{Deserialize, Serialize};
use crate::types::ReasoningEffort;
static DOTENV_VARS: OnceLock<HashMap<String, String>> = OnceLock::new();
fn load_dotenv_once(path: &Path) -> &'static HashMap<String, String> {
DOTENV_VARS.get_or_init(|| {
let mut map = HashMap::new();
let Ok(content) = std::fs::read_to_string(path) else {
return map;
};
for line in content.lines() {
let line = line.trim();
if line.is_empty() || line.starts_with('#') {
continue;
}
if let Some((k, v)) = line.split_once('=') {
let k = k.trim().to_string();
let v = v.trim().trim_matches('"').trim_matches('\'').to_string();
map.insert(k, v);
}
}
map
})
}
fn env_or_dotenv(key: &str, dotenv: &HashMap<String, String>) -> Option<String> {
std::env::var(key)
.ok()
.filter(|v| !v.is_empty())
.or_else(|| dotenv.get(key).filter(|v| !v.is_empty()).cloned())
}
pub fn get_secret(key: &str) -> Option<String> {
std::env::var(key)
.ok()
.filter(|v| !v.is_empty())
.or_else(|| {
DOTENV_VARS
.get()?
.get(key)
.filter(|v| !v.is_empty())
.cloned()
})
}
#[derive(Debug, Clone, Copy)]
pub struct BuiltinProvider {
pub name: &'static str,
pub base_url: &'static str,
pub api_key_env: &'static str,
pub tokens_param: &'static str,
}
pub const BUILTIN_PROVIDERS: &[BuiltinProvider] = &[
BuiltinProvider {
name: "openai",
base_url: "https://api.openai.com/v1",
api_key_env: "OPENAI_API_KEY",
tokens_param: "max_completion_tokens",
},
BuiltinProvider {
name: "gemini",
base_url: "https://generativelanguage.googleapis.com/v1beta/openai",
api_key_env: "GEMINI_API_KEY",
tokens_param: "max_tokens",
},
BuiltinProvider {
name: "groq",
base_url: "https://api.groq.com/openai/v1",
api_key_env: "GROQ_API_KEY",
tokens_param: "max_tokens",
},
BuiltinProvider {
name: "mistral",
base_url: "https://api.mistral.ai/v1",
api_key_env: "MISTRAL_API_KEY",
tokens_param: "max_tokens",
},
BuiltinProvider {
name: "deepseek",
base_url: "https://api.deepseek.com/v1",
api_key_env: "DEEPSEEK_API_KEY",
tokens_param: "max_tokens",
},
BuiltinProvider {
name: "xai",
base_url: "https://api.x.ai/v1",
api_key_env: "XAI_API_KEY",
tokens_param: "max_tokens",
},
BuiltinProvider {
name: "together",
base_url: "https://api.together.xyz/v1",
api_key_env: "TOGETHER_API_KEY",
tokens_param: "max_tokens",
},
BuiltinProvider {
name: "fireworks",
base_url: "https://api.fireworks.ai/inference/v1",
api_key_env: "FIREWORKS_API_KEY",
tokens_param: "max_tokens",
},
BuiltinProvider {
name: "cerebras",
base_url: "https://api.cerebras.ai/v1",
api_key_env: "CEREBRAS_API_KEY",
tokens_param: "max_tokens",
},
BuiltinProvider {
name: "perplexity",
base_url: "https://api.perplexity.ai",
api_key_env: "PERPLEXITY_API_KEY",
tokens_param: "max_tokens",
},
BuiltinProvider {
name: "cohere",
base_url: "https://api.cohere.com/compatibility/v1",
api_key_env: "COHERE_API_KEY",
tokens_param: "max_tokens",
},
BuiltinProvider {
name: "nvidia",
base_url: "https://integrate.api.nvidia.com/v1",
api_key_env: "NVIDIA_API_KEY",
tokens_param: "max_tokens",
},
BuiltinProvider {
name: "alibaba",
base_url: "https://dashscope.aliyuncs.com/compatible-mode/v1",
api_key_env: "DASHSCOPE_API_KEY",
tokens_param: "max_tokens",
},
BuiltinProvider {
name: "doubao",
base_url: "https://ark.cn-beijing.volces.com/api/v3",
api_key_env: "ARK_API_KEY",
tokens_param: "max_tokens",
},
BuiltinProvider {
name: "zhipu",
base_url: "https://open.bigmodel.cn/api/paas/v4",
api_key_env: "ZHIPU_API_KEY",
tokens_param: "max_tokens",
},
BuiltinProvider {
name: "moonshot",
base_url: "https://api.moonshot.cn/v1",
api_key_env: "MOONSHOT_API_KEY",
tokens_param: "max_tokens",
},
BuiltinProvider {
name: "baidu",
base_url: "https://qianfan.baidubce.com/v2",
api_key_env: "QIANFAN_API_KEY",
tokens_param: "max_tokens",
},
BuiltinProvider {
name: "thaillm",
base_url: "http://thaillm.or.th/api/v1",
api_key_env: "THAILLM_API_KEY",
tokens_param: "max_completion_tokens",
},
BuiltinProvider {
name: "vllm",
base_url: "http://localhost:8000/v1",
api_key_env: "VLLM_API_KEY",
tokens_param: "max_completion_tokens",
},
BuiltinProvider {
name: "openrouter",
base_url: "https://openrouter.ai/api/v1",
api_key_env: "OPENROUTER_API_KEY",
tokens_param: "max_completion_tokens",
},
];
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct ProviderProfile {
#[serde(default)]
pub name: Option<String>,
#[serde(default)]
pub url: Option<String>,
#[serde(default)]
pub key: Option<String>,
#[serde(default)]
pub model: Option<String>,
}
impl ProviderProfile {
pub fn resolved_key(&self) -> Option<String> {
let k = self.key.as_deref()?;
if let Some(var) = k.strip_prefix("${").and_then(|s| s.strip_suffix('}')) {
get_secret(var)
} else {
Some(k.to_string())
}
}
pub fn resolved_base_url(&self) -> Option<String> {
if let Some(url) = &self.url {
return Some(url.clone());
}
let name = self.name.as_deref()?;
BUILTIN_PROVIDERS
.iter()
.find(|p| p.name == name)
.map(|p| p.base_url.to_string())
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AgentConfig {
#[serde(skip)]
pub home_dir: PathBuf,
#[serde(default = "default_model")]
pub model: String,
#[serde(default = "default_max_iterations")]
pub max_iterations: u32,
#[serde(default)]
pub sub_agent_max_iterations: Option<u32>,
#[serde(default = "default_max_delegation_depth")]
pub max_delegation_depth: u32,
#[serde(default)]
pub tool_delay_ms: u64,
#[serde(default = "default_provider")]
pub provider: String,
pub base_url: Option<String>,
#[serde(default)]
pub providers: std::collections::HashMap<String, ProviderProfile>,
#[serde(default)]
pub routing: std::collections::HashMap<String, String>,
#[serde(default)]
pub tools: std::collections::HashMap<String, std::collections::HashMap<String, String>>,
#[serde(default)]
pub skills: std::collections::HashMap<String, std::collections::HashMap<String, String>>,
#[serde(skip)]
pub api_key: Option<String>,
#[serde(skip)]
pub fallback_api_keys: Vec<String>,
#[serde(default)]
pub compression: CompressionConfig,
#[serde(default)]
pub mcp_servers: Vec<McpServerConfig>,
#[serde(default)]
pub max_concurrent_requests: Option<usize>,
#[serde(default)]
pub security: SecurityConfig,
#[serde(default)]
pub memory_expiry: MemoryExpiryConfig,
#[serde(default = "default_nudge_interval")]
pub nudge_interval: u32,
#[serde(default = "default_llm_max_retries")]
pub llm_max_retries: u32,
#[serde(default = "default_llm_retry_base_ms")]
pub llm_retry_base_ms: u64,
#[serde(default)]
pub platform: PlatformConfig,
#[serde(default = "default_auto_skill_threshold")]
pub auto_skill_threshold: u32,
#[serde(default = "default_llm_timeout_secs")]
pub llm_timeout_secs: u64,
#[serde(default = "default_tool_timeout_secs")]
pub tool_timeout_secs: u64,
#[serde(default = "default_shutdown_timeout_secs")]
pub shutdown_timeout_secs: u64,
#[serde(default)]
pub max_tokens_per_task: Option<u32>,
#[serde(default)]
pub max_output_tokens: Option<u32>,
#[serde(default)]
pub reasoning_effort: Option<ReasoningEffort>,
#[serde(default)]
pub context_window: Option<usize>,
#[serde(default = "default_disabled_toolsets")]
pub disabled_toolsets: Vec<String>,
#[serde(default)]
pub disabled_tools: Vec<String>,
#[serde(default)]
pub show_usage_footer: bool,
#[serde(default)]
pub max_memory_tokens: Option<u32>,
#[serde(default)]
pub platforms: WebhookPlatformsConfig,
#[serde(default)]
pub server: ServerConfig,
#[serde(default)]
pub cron: CronConfig,
#[serde(default)]
pub roles: RolesConfig,
}
pub const DEFAULT_MODEL: &str = "anthropic/claude-sonnet-4-6";
pub const DEFAULT_PROVIDER: &str = "openrouter";
fn default_model() -> String {
DEFAULT_MODEL.into()
}
fn default_provider() -> String {
DEFAULT_PROVIDER.into()
}
fn default_max_iterations() -> u32 {
90
}
fn default_max_delegation_depth() -> u32 {
1
}
fn default_nudge_interval() -> u32 {
5
}
fn default_auto_skill_threshold() -> u32 {
5
}
fn default_llm_max_retries() -> u32 {
3
}
fn default_llm_retry_base_ms() -> u64 {
1000
}
fn default_llm_timeout_secs() -> u64 {
120
}
fn default_tool_timeout_secs() -> u64 {
60
}
fn default_shutdown_timeout_secs() -> u64 {
30
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MemoryExpiryConfig {
#[serde(default = "default_fact_days")]
pub fact_days: Option<u32>,
#[serde(default = "default_project_days")]
pub project_days: Option<u32>,
#[serde(default = "default_other_days")]
pub other_days: Option<u32>,
#[serde(default)]
pub preference_days: Option<u32>,
#[serde(default)]
pub skill_days: Option<u32>,
}
#[allow(clippy::unnecessary_wraps)]
fn default_fact_days() -> Option<u32> {
Some(90)
}
#[allow(clippy::unnecessary_wraps)]
fn default_project_days() -> Option<u32> {
Some(30)
}
#[allow(clippy::unnecessary_wraps)]
fn default_other_days() -> Option<u32> {
Some(60)
}
impl Default for MemoryExpiryConfig {
fn default() -> Self {
Self {
fact_days: default_fact_days(),
project_days: default_project_days(),
other_days: default_other_days(),
preference_days: None,
skill_days: None,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum TerminalSandbox {
#[default]
None,
Docker,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SecurityConfig {
#[serde(skip)]
pub gateway_api_key: Option<String>,
#[serde(default)]
pub allowed_read_paths: Vec<PathBuf>,
#[serde(default)]
pub allowed_write_paths: Vec<PathBuf>,
#[serde(default = "default_approval_mode")]
pub approval_mode: String,
#[serde(default)]
pub rate_limit_rpm: Option<u32>,
#[serde(default)]
pub terminal_sandbox: TerminalSandbox,
#[serde(default = "default_sandbox_image")]
pub terminal_sandbox_image: String,
#[serde(default)]
pub terminal_sandbox_opts: Vec<String>,
}
fn default_approval_mode() -> String {
"smart".to_string()
}
fn default_sandbox_image() -> String {
"ubuntu:24.04".to_string()
}
impl Default for SecurityConfig {
fn default() -> Self {
Self {
gateway_api_key: None,
allowed_read_paths: Vec::new(),
allowed_write_paths: Vec::new(),
approval_mode: default_approval_mode(),
rate_limit_rpm: None,
terminal_sandbox: TerminalSandbox::None,
terminal_sandbox_image: default_sandbox_image(),
terminal_sandbox_opts: Vec::new(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PlatformConfig {
#[serde(default)]
pub require_mention: bool,
#[serde(default)]
pub bot_username: String,
#[serde(default = "default_true")]
pub session_per_user: bool,
}
fn default_true() -> bool {
true
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct RoleDefinition {
#[serde(default)]
pub approval_mode: Option<String>,
#[serde(default)]
pub allowed_toolsets: Vec<String>,
#[serde(default)]
pub allowed_tools: Vec<String>,
#[serde(default)]
pub denied_tools: Vec<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct InviteCode {
pub role: String,
#[serde(default = "default_invite_max_uses")]
pub max_uses: u32,
#[serde(default)]
pub uses: u32,
#[serde(default)]
pub expires_at: Option<u64>,
}
fn default_invite_max_uses() -> u32 {
1
}
impl InviteCode {
pub fn is_valid(&self) -> bool {
if self.max_uses > 0 && self.uses >= self.max_uses {
return false;
}
if let Some(exp) = self.expires_at {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
if now > exp {
return false;
}
}
true
}
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct RolesConfig {
#[serde(default)]
pub definitions: std::collections::HashMap<String, RoleDefinition>,
#[serde(default)]
pub users: std::collections::HashMap<String, std::collections::HashMap<String, String>>,
#[serde(default)]
pub default_role: Option<String>,
#[serde(default)]
pub invites: std::collections::HashMap<String, InviteCode>,
}
impl RolesConfig {
pub fn lookup_role(
&self,
platform: &str,
user_id: &str,
username: Option<&str>,
) -> Option<String> {
let map = self.users.get(platform)?;
if let Some(role) = map.get(user_id) {
return Some(role.clone());
}
if platform == "telegram" {
if let Some(uname) = username {
let with_at = if uname.starts_with('@') {
uname.to_string()
} else {
format!("@{uname}")
};
if let Some(role) = map.get(&with_at) {
return Some(role.clone());
}
}
}
None
}
pub fn set_user_role(&mut self, platform: &str, user_id: &str, role: &str) {
self.users
.entry(platform.to_string())
.or_default()
.insert(user_id.to_string(), role.to_string());
}
pub fn remove_user(&mut self, platform: &str, user_id: &str) -> bool {
if let Some(map) = self.users.get_mut(platform) {
return map.remove(user_id).is_some();
}
false
}
pub fn redeem_invite(&mut self, code: &str, platform: &str, user_id: &str) -> Option<String> {
let invite = self.invites.get_mut(code)?;
if !invite.is_valid() {
return None;
}
let role = invite.role.clone();
let max_uses = invite.max_uses;
invite.uses += 1;
let exhausted = max_uses > 0 && invite.uses >= max_uses;
if exhausted {
self.invites.remove(code);
}
self.set_user_role(platform, user_id, &role);
Some(role)
}
}
impl Default for PlatformConfig {
fn default() -> Self {
Self {
require_mention: false,
bot_username: String::new(),
session_per_user: true,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct McpServerConfig {
pub name: String,
pub command: String,
#[serde(default)]
pub args: Vec<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WebhookPlatformConfig {
#[serde(default)]
pub enabled: bool,
pub port: u16,
pub webhook_path: String,
#[serde(default)]
pub hmac_secret: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct WebhookPlatformsConfig {
#[serde(default)]
pub line: Option<WebhookPlatformConfig>,
#[serde(default)]
pub whatsapp: Option<WebhookPlatformConfig>,
#[serde(default)]
pub webhook: Option<WebhookPlatformConfig>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ServerConfig {
#[serde(default = "default_server_port")]
pub port: u16,
}
fn default_server_port() -> u16 {
3000
}
fn default_disabled_toolsets() -> Vec<String> {
vec![]
}
fn parse_cron_jobs_str(s: &str) -> Vec<CronJob> {
s.split(',')
.filter_map(|entry| {
let (expr, task) = entry.trim().split_once('=')?;
Some(CronJob {
schedule: expr.trim().to_string(),
task: task.trim().to_string(),
})
})
.collect()
}
impl Default for ServerConfig {
fn default() -> Self {
Self {
port: default_server_port(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CronJob {
pub schedule: String,
pub task: String,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct CronConfig {
#[serde(default)]
pub jobs: Vec<CronJob>,
#[serde(default)]
pub timezone: Option<String>,
#[serde(default)]
pub memory_consolidation: Option<String>,
#[serde(default)]
pub memory_expiry: Option<String>,
}
impl WebhookPlatformConfig {
pub fn default_webhook() -> Self {
Self {
enabled: true,
port: 3001,
webhook_path: "/webhook".to_string(),
hmac_secret: None,
}
}
pub fn default_line() -> Self {
Self {
enabled: true,
port: 3002,
webhook_path: "/line".to_string(),
hmac_secret: None,
}
}
pub fn default_whatsapp() -> Self {
Self {
enabled: true,
port: 3003,
webhook_path: "/whatsapp".to_string(),
hmac_secret: None,
}
}
}
impl Default for AgentConfig {
fn default() -> Self {
let cwd = std::env::current_dir().unwrap_or_default();
let home = dirs::home_dir().unwrap_or_default();
Self {
home_dir: Self::garudust_dir(),
model: DEFAULT_MODEL.into(),
max_iterations: 90,
sub_agent_max_iterations: None,
max_delegation_depth: 1,
tool_delay_ms: 0,
provider: DEFAULT_PROVIDER.into(),
base_url: None,
providers: std::collections::HashMap::new(),
routing: std::collections::HashMap::new(),
tools: std::collections::HashMap::new(),
skills: std::collections::HashMap::new(),
api_key: None,
fallback_api_keys: Vec::new(),
compression: CompressionConfig::default(),
mcp_servers: Vec::new(),
max_concurrent_requests: None,
security: SecurityConfig {
gateway_api_key: None,
allowed_read_paths: vec![cwd.clone(), home],
allowed_write_paths: vec![cwd],
approval_mode: default_approval_mode(),
rate_limit_rpm: None,
terminal_sandbox: TerminalSandbox::None,
terminal_sandbox_image: default_sandbox_image(),
terminal_sandbox_opts: Vec::new(),
},
memory_expiry: MemoryExpiryConfig::default(),
nudge_interval: default_nudge_interval(),
llm_max_retries: default_llm_max_retries(),
llm_retry_base_ms: default_llm_retry_base_ms(),
platform: PlatformConfig::default(),
auto_skill_threshold: default_auto_skill_threshold(),
llm_timeout_secs: default_llm_timeout_secs(),
tool_timeout_secs: default_tool_timeout_secs(),
shutdown_timeout_secs: default_shutdown_timeout_secs(),
max_tokens_per_task: None,
max_output_tokens: None,
reasoning_effort: None,
context_window: None,
disabled_toolsets: default_disabled_toolsets(),
disabled_tools: Vec::new(),
show_usage_footer: false,
max_memory_tokens: None,
platforms: WebhookPlatformsConfig {
webhook: Some(WebhookPlatformConfig::default_webhook()),
line: None,
whatsapp: None,
},
server: ServerConfig::default(),
cron: CronConfig::default(),
roles: RolesConfig::default(),
}
}
}
pub(crate) fn resolve_key_for_provider(
provider: &str,
dotenv: &HashMap<String, String>,
) -> Option<String> {
if matches!(provider, "ollama" | "bedrock" | "codex") {
return None;
}
if provider == "anthropic" {
return env_or_dotenv("ANTHROPIC_API_KEY", dotenv);
}
if let Some(p) = BUILTIN_PROVIDERS.iter().find(|p| p.name == provider) {
return env_or_dotenv(p.api_key_env, dotenv);
}
env_or_dotenv("OPENROUTER_API_KEY", dotenv)
}
pub(crate) fn detect_provider_from_env(config: &mut AgentConfig, dotenv: &HashMap<String, String>) {
if let Some(k) = env_or_dotenv("ANTHROPIC_API_KEY", dotenv) {
config.api_key = Some(k);
config.provider = "anthropic".into();
return;
}
for p in BUILTIN_PROVIDERS {
if matches!(p.name, "thaillm" | "vllm" | "openrouter") {
continue;
}
if let Some(k) = env_or_dotenv(p.api_key_env, dotenv) {
config.api_key = Some(k);
config.provider = p.name.into();
return;
}
}
if let Some(url) = env_or_dotenv("OLLAMA_BASE_URL", dotenv) {
config.provider = "ollama".into();
config.base_url = Some(url);
return;
}
if let Some(url) = env_or_dotenv("VLLM_BASE_URL", dotenv) {
config.provider = "vllm".into();
config.base_url = Some(url);
config.api_key = env_or_dotenv("VLLM_API_KEY", dotenv);
return;
}
if let Some(k) = env_or_dotenv("THAILLM_API_KEY", dotenv) {
config.api_key = Some(k);
config.provider = "thaillm".into();
return;
}
if let Some(k) = env_or_dotenv("OPENROUTER_API_KEY", dotenv) {
config.api_key = Some(k);
config.provider = "openrouter".into();
}
}
impl AgentConfig {
pub fn effective_base_url(&self) -> Option<String> {
if let Some(p) = self.providers.get("default") {
if let Some(url) = p.resolved_base_url() {
return Some(url);
}
}
self.base_url.clone()
}
pub fn effective_api_key(&self) -> Option<String> {
if let Some(p) = self.providers.get("default") {
if let Some(k) = p.resolved_key() {
return Some(k);
}
}
self.api_key.clone()
}
pub fn garudust_dir() -> PathBuf {
dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("/tmp"))
.join(".garudust")
}
pub fn load() -> Self {
let home_dir = Self::garudust_dir();
let env_file = home_dir.join(".env");
let dotenv = load_dotenv_once(&env_file);
let yaml_path = home_dir.join("config.yaml");
let mut config: AgentConfig = if yaml_path.exists() {
let src = std::fs::read_to_string(&yaml_path).unwrap_or_default();
serde_yaml::from_str(&src).unwrap_or_default()
} else {
AgentConfig::default()
};
config.home_dir = home_dir;
if let Some(default_profile) = config.providers.get("default") {
if let Some(name) = &default_profile.name {
if !name.is_empty() {
config.provider = name.clone();
}
}
if let Some(model) = &default_profile.model {
if !model.is_empty() {
config.model = model.clone();
}
}
}
if config.security.allowed_read_paths.is_empty() {
let cwd = std::env::current_dir().unwrap_or_default();
let home = dirs::home_dir().unwrap_or_default();
config.security.allowed_read_paths = vec![cwd.clone(), home];
config.security.allowed_write_paths = vec![cwd];
}
let yaml_authoritative = yaml_path.exists();
if yaml_authoritative {
if config.api_key.is_none() {
config.api_key = resolve_key_for_provider(&config.provider, dotenv);
}
} else {
detect_provider_from_env(&mut config, dotenv);
}
if let Some(m) = env_or_dotenv("GARUDUST_MODEL", dotenv) {
config.model = m;
}
if let Some(u) = env_or_dotenv("GARUDUST_BASE_URL", dotenv) {
config.base_url = Some(u);
}
if let Some(v) = env_or_dotenv("LLM_FALLBACK_API_KEYS", dotenv) {
config.fallback_api_keys = v
.split(',')
.map(str::trim)
.filter(|s| !s.is_empty())
.map(str::to_string)
.collect();
}
if let Some(k) = env_or_dotenv("GARUDUST_API_KEY", dotenv) {
config.security.gateway_api_key = Some(k);
}
if let Some(v) = env_or_dotenv("GARUDUST_RATE_LIMIT", dotenv) {
if let Ok(n) = v.parse::<u32>() {
config.security.rate_limit_rpm = Some(n);
}
}
if let Some(mode) = env_or_dotenv("GARUDUST_APPROVAL_MODE", dotenv) {
config.security.approval_mode = mode;
}
if let Some(sandbox) = env_or_dotenv("GARUDUST_TERMINAL_SANDBOX", dotenv) {
config.security.terminal_sandbox = match sandbox.to_lowercase().as_str() {
"docker" => TerminalSandbox::Docker,
_ => TerminalSandbox::None,
};
}
if let Some(image) = env_or_dotenv("GARUDUST_SANDBOX_IMAGE", dotenv) {
config.security.terminal_sandbox_image = image;
}
if let Some(v) = env_or_dotenv("GARUDUST_PORT", dotenv) {
if let Ok(n) = v.parse::<u16>() {
config.server.port = n;
}
}
if let Some(v) = env_or_dotenv("GARUDUST_MEMORY_CRON", dotenv) {
config.cron.memory_consolidation = Some(v);
}
if let Some(v) = env_or_dotenv("GARUDUST_MEMORY_EXPIRY_CRON", dotenv) {
config.cron.memory_expiry = Some(v);
}
if let Some(v) = env_or_dotenv("GARUDUST_CRON_JOBS", dotenv) {
config.cron.jobs = parse_cron_jobs_str(&v);
}
config
}
pub fn save_yaml(&self) -> std::io::Result<()> {
std::fs::create_dir_all(&self.home_dir)?;
let yaml = serde_yaml::to_string(self).map_err(std::io::Error::other)?;
let tmp = self.home_dir.join("config.yaml.tmp");
std::fs::write(&tmp, yaml)?;
std::fs::rename(tmp, self.home_dir.join("config.yaml"))
}
pub fn set_env_var(home_dir: &Path, key: &str, value: &str) -> std::io::Result<()> {
std::fs::create_dir_all(home_dir)?;
let env_path = home_dir.join(".env");
let existing = if env_path.exists() {
std::fs::read_to_string(&env_path)?
} else {
String::new()
};
let prefix = format!("{key}=");
let mut lines: Vec<String> = existing
.lines()
.filter(|l| !l.starts_with(&prefix))
.map(String::from)
.collect();
lines.push(format!("{key}={value}"));
std::fs::write(&env_path, lines.join("\n") + "\n")
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CompressionConfig {
pub enabled: bool,
pub threshold_fraction: f32,
pub model: Option<String>,
}
impl Default for CompressionConfig {
fn default() -> Self {
Self {
enabled: true,
threshold_fraction: 0.8,
model: None,
}
}
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use super::{detect_provider_from_env, resolve_key_for_provider, AgentConfig, RolesConfig};
fn dotenv(pairs: &[(&str, &str)]) -> HashMap<String, String> {
pairs
.iter()
.map(|(k, v)| ((*k).to_string(), (*v).to_string()))
.collect()
}
fn roles_with_admin() -> RolesConfig {
let mut r = RolesConfig::default();
r.set_user_role("telegram", "111", "admin");
r
}
#[test]
fn roles_lookup_by_id() {
let r = roles_with_admin();
assert_eq!(r.lookup_role("telegram", "111", None), Some("admin".into()));
}
#[test]
fn roles_lookup_no_match_returns_none() {
let r = roles_with_admin();
assert_eq!(r.lookup_role("telegram", "999", None), None);
}
#[test]
fn roles_lookup_wrong_platform_returns_none() {
let r = roles_with_admin();
assert_eq!(r.lookup_role("discord", "111", None), None);
}
#[test]
fn roles_lookup_telegram_username_with_at() {
let mut r = RolesConfig::default();
r.set_user_role("telegram", "@somchai", "member");
assert_eq!(
r.lookup_role("telegram", "0", Some("@somchai")),
Some("member".into())
);
}
#[test]
fn roles_lookup_telegram_username_without_at() {
let mut r = RolesConfig::default();
r.set_user_role("telegram", "@somchai", "member");
assert_eq!(
r.lookup_role("telegram", "0", Some("somchai")),
Some("member".into())
);
}
#[test]
fn roles_lookup_username_only_works_on_telegram() {
let mut r = RolesConfig::default();
r.set_user_role("discord", "@somchai", "member");
assert_eq!(r.lookup_role("discord", "0", Some("somchai")), None);
}
#[test]
fn roles_set_creates_new_entry() {
let mut r = RolesConfig::default();
r.set_user_role("line", "Uabc", "member");
assert_eq!(r.lookup_role("line", "Uabc", None), Some("member".into()));
}
#[test]
fn roles_set_updates_existing_entry() {
let mut r = roles_with_admin();
r.set_user_role("telegram", "111", "member");
assert_eq!(
r.lookup_role("telegram", "111", None),
Some("member".into())
);
}
#[test]
fn roles_remove_user_returns_true_when_found() {
let mut r = roles_with_admin();
assert!(r.remove_user("telegram", "111"));
assert!(r.lookup_role("telegram", "111", None).is_none());
}
#[test]
fn roles_remove_user_returns_false_when_missing() {
let mut r = roles_with_admin();
assert!(!r.remove_user("telegram", "999"));
}
#[test]
fn roles_remove_user_wrong_platform_returns_false() {
let mut r = roles_with_admin();
assert!(!r.remove_user("discord", "111"));
}
#[test]
fn resolve_openai_key() {
let map = dotenv(&[("OPENAI_API_KEY", "sk-test-openai")]);
assert_eq!(
resolve_key_for_provider("openai", &map),
Some("sk-test-openai".into())
);
}
#[test]
fn resolve_gemini_key() {
let map = dotenv(&[("GEMINI_API_KEY", "AIza-test")]);
assert_eq!(
resolve_key_for_provider("gemini", &map),
Some("AIza-test".into())
);
}
#[test]
fn resolve_groq_key() {
let map = dotenv(&[("GROQ_API_KEY", "gsk-test")]);
assert_eq!(
resolve_key_for_provider("groq", &map),
Some("gsk-test".into())
);
}
#[test]
fn resolve_mistral_key() {
let map = dotenv(&[("MISTRAL_API_KEY", "ms-test")]);
assert_eq!(
resolve_key_for_provider("mistral", &map),
Some("ms-test".into())
);
}
#[test]
fn resolve_deepseek_key() {
let map = dotenv(&[("DEEPSEEK_API_KEY", "ds-test")]);
assert_eq!(
resolve_key_for_provider("deepseek", &map),
Some("ds-test".into())
);
}
#[test]
fn resolve_xai_key() {
let map = dotenv(&[("XAI_API_KEY", "xai-test")]);
assert_eq!(
resolve_key_for_provider("xai", &map),
Some("xai-test".into())
);
}
#[test]
fn resolve_ollama_returns_none() {
let map = dotenv(&[("OPENROUTER_API_KEY", "or-test")]);
assert_eq!(resolve_key_for_provider("ollama", &map), None);
}
#[test]
fn resolve_unknown_provider_falls_back_to_openrouter() {
let map = dotenv(&[("OPENROUTER_API_KEY", "or-test")]);
assert_eq!(
resolve_key_for_provider("custom-provider", &map),
Some("or-test".into())
);
}
fn detect(pairs: &[(&str, &str)]) -> AgentConfig {
let mut cfg = AgentConfig::default();
detect_provider_from_env(&mut cfg, &dotenv(pairs));
cfg
}
#[test]
fn detect_openai_only() {
let cfg = detect(&[("OPENAI_API_KEY", "sk-test-openai")]);
assert_eq!(cfg.provider, "openai");
assert_eq!(cfg.api_key.as_deref(), Some("sk-test-openai"));
}
#[test]
fn detect_gemini_only() {
let cfg = detect(&[("GEMINI_API_KEY", "AIza-test")]);
assert_eq!(cfg.provider, "gemini");
assert_eq!(cfg.api_key.as_deref(), Some("AIza-test"));
}
#[test]
fn detect_groq_only() {
let cfg = detect(&[("GROQ_API_KEY", "gsk-test")]);
assert_eq!(cfg.provider, "groq");
assert_eq!(cfg.api_key.as_deref(), Some("gsk-test"));
}
#[test]
fn detect_mistral_only() {
let cfg = detect(&[("MISTRAL_API_KEY", "ms-test")]);
assert_eq!(cfg.provider, "mistral");
assert_eq!(cfg.api_key.as_deref(), Some("ms-test"));
}
#[test]
fn detect_deepseek_only() {
let cfg = detect(&[("DEEPSEEK_API_KEY", "ds-test")]);
assert_eq!(cfg.provider, "deepseek");
assert_eq!(cfg.api_key.as_deref(), Some("ds-test"));
}
#[test]
fn detect_xai_only() {
let cfg = detect(&[("XAI_API_KEY", "xai-test")]);
assert_eq!(cfg.provider, "xai");
assert_eq!(cfg.api_key.as_deref(), Some("xai-test"));
}
#[test]
fn detect_openrouter_only() {
let cfg = detect(&[("OPENROUTER_API_KEY", "or-test")]);
assert_eq!(cfg.provider, "openrouter");
assert_eq!(cfg.api_key.as_deref(), Some("or-test"));
}
#[test]
fn detect_ollama_sets_base_url_not_key() {
let cfg = detect(&[("OLLAMA_BASE_URL", "http://localhost:11434")]);
assert_eq!(cfg.provider, "ollama");
assert_eq!(cfg.base_url.as_deref(), Some("http://localhost:11434"));
assert!(cfg.api_key.is_none());
}
#[test]
fn detect_vllm_sets_base_url_and_key() {
let cfg = detect(&[
("VLLM_BASE_URL", "http://localhost:8000/v1"),
("VLLM_API_KEY", "vllm-test"),
]);
assert_eq!(cfg.provider, "vllm");
assert_eq!(cfg.base_url.as_deref(), Some("http://localhost:8000/v1"));
assert_eq!(cfg.api_key.as_deref(), Some("vllm-test"));
}
#[test]
fn detect_empty_env_leaves_defaults() {
let cfg = detect(&[]);
assert_eq!(cfg.provider, "openrouter");
assert!(cfg.api_key.is_none());
}
#[test]
fn detect_anthropic_wins_over_openai_in_dotenv() {
let cfg = detect(&[
("ANTHROPIC_API_KEY", "sk-ant-test"),
("OPENAI_API_KEY", "sk-oai-test"),
]);
assert_eq!(cfg.provider, "anthropic");
assert_eq!(cfg.api_key.as_deref(), Some("sk-ant-test"));
}
#[test]
fn resolve_together_key() {
let map = dotenv(&[("TOGETHER_API_KEY", "tog-test")]);
assert_eq!(
resolve_key_for_provider("together", &map),
Some("tog-test".into())
);
}
#[test]
fn resolve_fireworks_key() {
let map = dotenv(&[("FIREWORKS_API_KEY", "fw-test")]);
assert_eq!(
resolve_key_for_provider("fireworks", &map),
Some("fw-test".into())
);
}
#[test]
fn resolve_cerebras_key() {
let map = dotenv(&[("CEREBRAS_API_KEY", "cb-test")]);
assert_eq!(
resolve_key_for_provider("cerebras", &map),
Some("cb-test".into())
);
}
#[test]
fn resolve_perplexity_key() {
let map = dotenv(&[("PERPLEXITY_API_KEY", "pplx-test")]);
assert_eq!(
resolve_key_for_provider("perplexity", &map),
Some("pplx-test".into())
);
}
#[test]
fn resolve_cohere_key() {
let map = dotenv(&[("COHERE_API_KEY", "co-test")]);
assert_eq!(
resolve_key_for_provider("cohere", &map),
Some("co-test".into())
);
}
#[test]
fn resolve_nvidia_key() {
let map = dotenv(&[("NVIDIA_API_KEY", "nvapi-test")]);
assert_eq!(
resolve_key_for_provider("nvidia", &map),
Some("nvapi-test".into())
);
}
#[test]
fn resolve_alibaba_key() {
let map = dotenv(&[("DASHSCOPE_API_KEY", "sk-ds-test")]);
assert_eq!(
resolve_key_for_provider("alibaba", &map),
Some("sk-ds-test".into())
);
}
#[test]
fn resolve_doubao_key() {
let map = dotenv(&[("ARK_API_KEY", "ark-test")]);
assert_eq!(
resolve_key_for_provider("doubao", &map),
Some("ark-test".into())
);
}
#[test]
fn resolve_zhipu_key() {
let map = dotenv(&[("ZHIPU_API_KEY", "zp-test")]);
assert_eq!(
resolve_key_for_provider("zhipu", &map),
Some("zp-test".into())
);
}
#[test]
fn resolve_moonshot_key() {
let map = dotenv(&[("MOONSHOT_API_KEY", "ms-kimi-test")]);
assert_eq!(
resolve_key_for_provider("moonshot", &map),
Some("ms-kimi-test".into())
);
}
#[test]
fn resolve_baidu_key() {
let map = dotenv(&[("QIANFAN_API_KEY", "qf-test")]);
assert_eq!(
resolve_key_for_provider("baidu", &map),
Some("qf-test".into())
);
}
#[test]
fn detect_together_only() {
let cfg = detect(&[("TOGETHER_API_KEY", "tog-test")]);
assert_eq!(cfg.provider, "together");
assert_eq!(cfg.api_key.as_deref(), Some("tog-test"));
}
#[test]
fn detect_fireworks_only() {
let cfg = detect(&[("FIREWORKS_API_KEY", "fw-test")]);
assert_eq!(cfg.provider, "fireworks");
assert_eq!(cfg.api_key.as_deref(), Some("fw-test"));
}
#[test]
fn detect_cerebras_only() {
let cfg = detect(&[("CEREBRAS_API_KEY", "cb-test")]);
assert_eq!(cfg.provider, "cerebras");
assert_eq!(cfg.api_key.as_deref(), Some("cb-test"));
}
#[test]
fn detect_perplexity_only() {
let cfg = detect(&[("PERPLEXITY_API_KEY", "pplx-test")]);
assert_eq!(cfg.provider, "perplexity");
assert_eq!(cfg.api_key.as_deref(), Some("pplx-test"));
}
#[test]
fn detect_cohere_only() {
let cfg = detect(&[("COHERE_API_KEY", "co-test")]);
assert_eq!(cfg.provider, "cohere");
assert_eq!(cfg.api_key.as_deref(), Some("co-test"));
}
#[test]
fn detect_nvidia_only() {
let cfg = detect(&[("NVIDIA_API_KEY", "nvapi-test")]);
assert_eq!(cfg.provider, "nvidia");
assert_eq!(cfg.api_key.as_deref(), Some("nvapi-test"));
}
#[test]
fn detect_alibaba_only() {
let cfg = detect(&[("DASHSCOPE_API_KEY", "sk-ds-test")]);
assert_eq!(cfg.provider, "alibaba");
assert_eq!(cfg.api_key.as_deref(), Some("sk-ds-test"));
}
#[test]
fn detect_doubao_only() {
let cfg = detect(&[("ARK_API_KEY", "ark-test")]);
assert_eq!(cfg.provider, "doubao");
assert_eq!(cfg.api_key.as_deref(), Some("ark-test"));
}
#[test]
fn detect_zhipu_only() {
let cfg = detect(&[("ZHIPU_API_KEY", "zp-test")]);
assert_eq!(cfg.provider, "zhipu");
assert_eq!(cfg.api_key.as_deref(), Some("zp-test"));
}
#[test]
fn detect_moonshot_only() {
let cfg = detect(&[("MOONSHOT_API_KEY", "ms-kimi-test")]);
assert_eq!(cfg.provider, "moonshot");
assert_eq!(cfg.api_key.as_deref(), Some("ms-kimi-test"));
}
#[test]
fn detect_baidu_only() {
let cfg = detect(&[("QIANFAN_API_KEY", "qf-test")]);
assert_eq!(cfg.provider, "baidu");
assert_eq!(cfg.api_key.as_deref(), Some("qf-test"));
}
#[test]
fn profile_resolved_key_literal() {
let p = super::ProviderProfile {
key: Some("sk-literal".into()),
..Default::default()
};
assert_eq!(p.resolved_key(), Some("sk-literal".into()));
}
#[test]
fn profile_resolved_key_none_when_absent() {
let p = super::ProviderProfile::default();
assert!(p.resolved_key().is_none());
}
#[test]
fn profile_resolved_key_env_var_interpolation() {
std::env::set_var("GARUDUST_TEST_KEY_INTERP", "env-value-123");
let p = super::ProviderProfile {
key: Some("${GARUDUST_TEST_KEY_INTERP}".into()),
..Default::default()
};
assert_eq!(p.resolved_key(), Some("env-value-123".into()));
std::env::remove_var("GARUDUST_TEST_KEY_INTERP");
}
#[test]
fn profile_resolved_key_missing_env_var_returns_none() {
std::env::remove_var("GARUDUST_TEST_KEY_MISSING");
let p = super::ProviderProfile {
key: Some("${GARUDUST_TEST_KEY_MISSING}".into()),
..Default::default()
};
assert!(p.resolved_key().is_none());
}
#[test]
fn providers_default_overrides_provider_and_model() {
let yaml = "
providers:
default:
name: groq
model: llama-3.3-70b-versatile
";
let mut cfg: AgentConfig = serde_yaml::from_str(yaml).unwrap();
if let Some(default_profile) = cfg.providers.get("default") {
if let Some(name) = &default_profile.name.clone() {
if !name.is_empty() {
cfg.provider = name.clone();
}
}
if let Some(model) = &default_profile.model.clone() {
if !model.is_empty() {
cfg.model = model.clone();
}
}
}
assert_eq!(cfg.provider, "groq");
assert_eq!(cfg.model, "llama-3.3-70b-versatile");
}
#[test]
fn providers_map_deserializes_correctly() {
let yaml = r#"
providers:
groq-backup:
name: groq
key: "${GROQ_API_KEY_2}"
local:
url: "http://192.168.1.10:8000/v1"
"#;
let cfg: AgentConfig = serde_yaml::from_str(yaml).unwrap();
assert!(cfg.providers.contains_key("groq-backup"));
assert!(cfg.providers.contains_key("local"));
let backup = &cfg.providers["groq-backup"];
assert_eq!(backup.name.as_deref(), Some("groq"));
assert_eq!(backup.key.as_deref(), Some("${GROQ_API_KEY_2}"));
let local = &cfg.providers["local"];
assert_eq!(local.url.as_deref(), Some("http://192.168.1.10:8000/v1"));
}
}