use std::{net::SocketAddr, path::PathBuf};
use clap::{Parser, ValueEnum};
pub fn env_bool(var:&str) -> bool {
match std::env::var(var) {
Ok(v) => matches!(v.to_lowercase().as_str(), "1" | "true"),
Err(_) => false,
}
}
pub fn env_parse_warn<T:std::str::FromStr>(var:&str) -> Option<T> {
match std::env::var(var) {
Ok(v) => {
match v.parse::<T>() {
Ok(parsed) => Some(parsed),
Err(_) => {
tracing::warn!("{}={:?} could not be parsed; ignoring override", var, v);
None
},
}
},
Err(_) => None,
}
}
#[derive(Debug, Clone, Copy, ValueEnum)]
pub enum ProxyMode {
Cache,
Token,
}
#[derive(clap::Subcommand, Debug, Clone)]
pub enum Command {
Run,
Setup {
#[arg(long, env = "APHRODITE_API_KEY", hide_env_values = true)]
api_key:Option<String>,
#[arg(long, env = "APHRODITE_API_URL", default_value = "https://api.deepseek.com")]
api_url:String,
#[arg(long, env = "APHRODITE_MODEL", default_value = "deepseek-v4-pro")]
model:String,
#[arg(long, env = "APHRODITE_CACHE_PORT", default_value = "9797")]
cache_port:u16,
#[arg(long, env = "APHRODITE_TOKEN_PORT", default_value = "9798")]
token_port:u16,
#[arg(long)]
no_launch:bool,
#[arg(long)]
force:bool,
},
}
#[derive(Debug, Clone)]
pub struct SetupArgs {
pub api_key:Option<String>,
pub api_url:String,
pub model:String,
pub cache_port:u16,
pub token_port:u16,
pub no_launch:bool,
pub force:bool,
}
impl From<Command> for SetupArgs {
fn from(cmd:Command) -> Self {
match cmd {
Command::Setup { api_key, api_url, model, cache_port, token_port, no_launch, force } => {
Self { api_key, api_url, model, cache_port, token_port, no_launch, force }
},
_ => {
Self {
api_key:None,
api_url:"https://api.deepseek.com".into(),
model:"deepseek-v4-pro".into(),
cache_port:9797,
token_port:9798,
no_launch:false,
force:false,
}
},
}
}
}
#[derive(Parser, Debug, Clone)]
#[command(name = "aphrodite", version, about)]
pub struct Cli {
#[command(subcommand)]
pub command:Option<Command>,
#[arg(long, default_value = "token", env = "APHRODITE_MODE")]
pub mode:ProxyMode,
#[arg(long, default_value = "127.0.0.1:9797", env = "APHRODITE_LISTEN")]
pub listen:SocketAddr,
#[arg(long, default_value = "https://api.openai.com", env = "APHRODITE_API_URL")]
pub api_url:String,
#[arg(long, env = "APHRODITE_API_KEY", hide_env_values = true)]
pub api_key:String,
#[arg(long, default_value = "default-model", env = "APHRODITE_MODEL")]
pub model:String,
#[arg(long, default_value = "1000000")]
pub max_context:usize,
#[arg(long, default_value = "384000")]
pub max_output:usize,
#[arg(long, env = "APHRODITE_DB")]
pub ccr_db_path:Option<PathBuf>,
#[arg(long, default_value = "3600", env = "APHRODITE_CCR_TTL")]
pub ccr_ttl_seconds:u64,
#[arg(long)]
pub no_ccr_marker:bool,
#[arg(long)]
pub tool_relay:bool,
#[arg(long, env = "APHRODITE_NOTIFY_URL")]
pub notify_url:Option<String>,
#[arg(long, env = "APHRODITE_NOTIFY_KEY", hide_env_values = true)]
pub notify_key:Option<String>,
#[arg(long)]
pub dev:bool,
#[arg(long, env = "APHRODITE_LOG_COMPACT")]
pub log_compact:bool,
#[arg(long, default_value = "300")]
pub timeout:u64,
}
#[derive(Debug, Clone, serde::Deserialize)]
pub struct MultiConfig {
pub defaults:Option<Defaults>,
pub proxies:Vec<ProxyConfig>,
pub compression:Option<CompressionConfig>,
pub previews:Option<PreviewsConfig>,
pub prompts:Option<PromptsConfig>,
}
#[derive(Debug, Clone, serde::Deserialize)]
pub struct Defaults {
pub api_url:Option<String>,
pub model:Option<String>,
pub ccr_ttl_seconds:Option<u64>,
pub api_key:Option<String>,
}
#[derive(Debug, Clone, serde::Deserialize)]
pub struct CompressionConfig {
pub engine_threshold_pct:Option<u32>,
pub engine_protect_first:Option<u32>,
pub engine_protect_last:Option<u32>,
pub engine_min_msgs:Option<u32>,
pub tool_threshold_token:Option<u32>,
pub tool_threshold_cache:Option<u32>,
pub terminal_threshold:Option<u32>,
pub inline_threshold:Option<u32>,
pub auto_expand:Option<bool>,
pub auto_expand_limit:Option<u32>,
pub catalog_mode:Option<String>,
pub classifier_poll:Option<bool>,
pub code_multiplier:Option<f64>,
}
#[derive(Debug, Clone, serde::Deserialize)]
pub struct PreviewsConfig {
pub model_family:Option<String>,
pub code_structure_map:Option<bool>,
pub preview_max_chars:Option<u32>,
pub rust_preview_lines:Option<u32>,
}
#[derive(Debug, Clone, serde::Deserialize)]
pub struct PromptsConfig {
pub retrieve_guidance:Option<String>,
pub ccr_marker_hint:Option<bool>,
pub catalog_intent_hints:Option<bool>,
}
#[derive(Debug, Clone, serde::Deserialize)]
pub struct ProxyConfig {
pub name:Option<String>,
#[serde(default)]
pub listen:Option<String>,
pub mode:Option<String>,
pub api_key:Option<String>,
pub api_url:Option<String>,
pub model:Option<String>,
pub tool_relay:Option<bool>,
pub dev:Option<bool>,
pub ccr_ttl_seconds:Option<u64>,
pub ccr_db_path:Option<String>,
pub notify_url:Option<String>,
pub notify_key:Option<String>,
pub timeout:Option<u64>,
pub max_context:Option<usize>,
pub max_output:Option<usize>,
}
impl MultiConfig {
pub fn load(path:&str) -> anyhow::Result<Self> {
let content = std::fs::read_to_string(path)?;
Ok(toml::from_str(&content)?)
}
pub fn resolve(&self, cfg:&ProxyConfig) -> anyhow::Result<Cli> {
let d = self.defaults.as_ref();
let api_key:String = cfg
.api_key
.clone()
.or_else(|| d.and_then(|d| d.api_key.clone()))
.or_else(|| std::env::var("APHRODITE_API_KEY").ok())
.or_else(|| std::env::var("DEEPSEEK_API_KEY").ok())
.or_else(|| std::env::var("HEADROOM_DEEPSEEK_KEY").ok())
.unwrap_or_default();
if api_key.is_empty() {
anyhow::bail!("no API key configured - set APHRODITE_API_KEY env var or api_key in aphrodite.toml");
}
let listen:SocketAddr = match cfg.listen.as_deref() {
Some(s) => s.parse().map_err(|_| anyhow::anyhow!("invalid listen address: {s}"))?,
None => "127.0.0.1:9797".parse().unwrap(),
};
let listen = match (cfg.mode.as_deref(), cfg.name.as_deref()) {
(_, Some("cache")) | (Some("cache"), _) => Self::apply_port_override(listen, "APHRODITE_CACHE_PORT"),
(_, Some("token")) | (Some("token"), _) => Self::apply_port_override(listen, "APHRODITE_TOKEN_PORT"),
_ => listen,
};
let max_context = cfg.max_context.unwrap_or(1_000_000);
let max_output = cfg.max_output.unwrap_or(384_000);
if max_output >= max_context {
anyhow::bail!("max_output ({max_output}) must be less than max_context ({max_context})");
}
Ok(Cli {
command:None,
mode:match cfg.mode.as_deref() {
Some("token") => ProxyMode::Token,
Some("cache") => ProxyMode::Cache,
None => {
tracing::info!("no mode specified, defaulting to token");
ProxyMode::Token
},
Some(other) => {
tracing::warn!("unknown mode {:?}, defaulting to token", other);
ProxyMode::Token
},
},
listen,
api_url:std::env::var("APHRODITE_API_URL")
.ok()
.or_else(|| cfg.api_url.clone())
.or_else(|| d.and_then(|d| d.api_url.clone()))
.unwrap_or_else(|| "https://api.openai.com".into()),
api_key,
model:std::env::var("APHRODITE_MODEL")
.ok()
.or_else(|| cfg.model.clone())
.or_else(|| d.and_then(|d| d.model.clone()))
.unwrap_or_else(|| "default-model".into()),
max_context,
max_output,
ccr_db_path:std::env::var("APHRODITE_DB")
.ok()
.or_else(|| cfg.ccr_db_path.clone())
.filter(|s| !s.is_empty())
.map(Into::into),
ccr_ttl_seconds:env_parse_warn::<u64>("APHRODITE_CCR_TTL")
.or(cfg.ccr_ttl_seconds)
.or_else(|| d.and_then(|d| d.ccr_ttl_seconds))
.unwrap_or(3600),
no_ccr_marker:false,
tool_relay:cfg.tool_relay.unwrap_or(false),
notify_url:std::env::var("APHRODITE_NOTIFY_URL").ok().or_else(|| cfg.notify_url.clone()),
notify_key:std::env::var("APHRODITE_NOTIFY_KEY").ok().or_else(|| cfg.notify_key.clone()),
dev:cfg.dev.unwrap_or(false),
log_compact:false,
timeout:{
let t = cfg.timeout.unwrap_or(300);
if t > 600 {
tracing::warn!("timeout {}s exceeds maximum 600s, clamping", t);
600
} else {
t
}
},
})
}
fn apply_port_override(listen:SocketAddr, env_var:&str) -> SocketAddr {
match std::env::var(env_var) {
Ok(p) => {
match p.parse::<u16>() {
Ok(port) => {
let mut addr = listen;
addr.set_port(port);
tracing::info!("{}={} overriding listen to {}", env_var, port, addr);
addr
},
Err(_) => {
tracing::warn!(
"{}={:?} is not a valid port (1-65535); ignoring override, using {}",
env_var,
p,
listen,
);
listen
},
}
},
Err(_) => listen,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn env_guard() -> std::sync::MutexGuard<'static, ()> {
static G:std::sync::OnceLock<std::sync::Mutex<()>> = std::sync::OnceLock::new();
G.get_or_init(|| std::sync::Mutex::new(()))
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
#[test]
fn test_env_bool_true_values_case_insensitive() {
let _g = env_guard();
for v in ["1", "true", "TRUE", "True"] {
std::env::set_var("APHRODITE_TEST_BOOL", v);
assert!(env_bool("APHRODITE_TEST_BOOL"), "{v:?} should be true");
}
std::env::remove_var("APHRODITE_TEST_BOOL");
}
#[test]
fn test_env_bool_false_values() {
let _g = env_guard();
for v in ["0", "false", "yes", ""] {
std::env::set_var("APHRODITE_TEST_BOOL", v);
assert!(!env_bool("APHRODITE_TEST_BOOL"), "{v:?} should be false");
}
std::env::remove_var("APHRODITE_TEST_BOOL");
assert!(!env_bool("APHRODITE_TEST_BOOL"), "absent should be false");
}
fn multi_config_from_toml(toml_str:&str) -> MultiConfig { toml::from_str(toml_str).expect("valid test TOML") }
#[test]
fn test_resolve_default_ports_per_mode() {
let _g = env_guard();
std::env::remove_var("APHRODITE_CACHE_PORT");
std::env::remove_var("APHRODITE_TOKEN_PORT");
let mc = multi_config_from_toml(
r#"
[[proxies]]
name = "cache"
mode = "cache"
api_key = "test-key"
"#,
);
let cli = mc.resolve(&mc.proxies[0]).unwrap();
assert_eq!(cli.listen.port(), 9797);
let mc = multi_config_from_toml(
r#"
[[proxies]]
name = "token"
mode = "token"
api_key = "test-key"
"#,
);
let cli = mc.resolve(&mc.proxies[0]).unwrap();
assert_eq!(cli.listen.port(), 9797); }
#[test]
fn test_resolve_explicit_port_override_via_env() {
let _g = env_guard();
std::env::set_var("APHRODITE_CACHE_PORT", "19797");
let mc = multi_config_from_toml(
r#"
[[proxies]]
name = "cache"
mode = "cache"
api_key = "test-key"
listen = "127.0.0.1:9797"
"#,
);
let cli = mc.resolve(&mc.proxies[0]).unwrap();
assert_eq!(cli.listen.port(), 19797);
std::env::remove_var("APHRODITE_CACHE_PORT");
}
#[test]
fn test_resolve_env_overrides_toml_for_api_url_model_ttl_db_notify() {
let _g = env_guard();
for (k, v) in [
("APHRODITE_API_URL", "https://env-api.example.com"),
("APHRODITE_MODEL", "env-model"),
("APHRODITE_CCR_TTL", "42"),
("APHRODITE_DB", "/tmp/env-ccr.db"),
("APHRODITE_NOTIFY_URL", "https://env-notify.example.com"),
("APHRODITE_NOTIFY_KEY", "env-notify-key"),
] {
std::env::set_var(k, v);
}
let mc = multi_config_from_toml(
r#"
[[proxies]]
name = "token"
mode = "token"
api_key = "test-key"
api_url = "https://toml-api.example.com"
model = "toml-model"
ccr_ttl_seconds = 111
ccr_db_path = "/tmp/toml-ccr.db"
notify_url = "https://toml-notify.example.com"
notify_key = "toml-notify-key"
"#,
);
let cli = mc.resolve(&mc.proxies[0]).unwrap();
for k in [
"APHRODITE_API_URL",
"APHRODITE_MODEL",
"APHRODITE_CCR_TTL",
"APHRODITE_DB",
"APHRODITE_NOTIFY_URL",
"APHRODITE_NOTIFY_KEY",
] {
std::env::remove_var(k);
}
assert_eq!(cli.api_url, "https://env-api.example.com");
assert_eq!(cli.model, "env-model");
assert_eq!(cli.ccr_ttl_seconds, 42);
assert_eq!(cli.ccr_db_path.unwrap().to_string_lossy(), "/tmp/env-ccr.db");
assert_eq!(cli.notify_url.as_deref(), Some("https://env-notify.example.com"));
assert_eq!(cli.notify_key.as_deref(), Some("env-notify-key"));
}
#[test]
fn test_resolve_falls_back_to_toml_when_env_unset() {
let _g = env_guard();
for k in ["APHRODITE_API_URL", "APHRODITE_MODEL", "APHRODITE_CCR_TTL", "APHRODITE_DB"] {
std::env::remove_var(k);
}
let mc = multi_config_from_toml(
r#"
[[proxies]]
name = "token"
mode = "token"
api_key = "test-key"
api_url = "https://toml-api.example.com"
model = "toml-model"
ccr_ttl_seconds = 111
"#,
);
let cli = mc.resolve(&mc.proxies[0]).unwrap();
assert_eq!(cli.api_url, "https://toml-api.example.com");
assert_eq!(cli.model, "toml-model");
assert_eq!(cli.ccr_ttl_seconds, 111);
}
#[test]
fn test_resolve_invalid_mode_falls_back_to_token() {
let mc = multi_config_from_toml(
r#"
[[proxies]]
name = "weird"
mode = "not_a_real_mode"
api_key = "test-key"
"#,
);
let cli = mc.resolve(&mc.proxies[0]).unwrap();
assert!(matches!(cli.mode, ProxyMode::Token));
}
#[test]
fn test_resolve_missing_mode_falls_back_to_token() {
let mc = multi_config_from_toml(
r#"
[[proxies]]
name = "no_mode"
api_key = "test-key"
"#,
);
let cli = mc.resolve(&mc.proxies[0]).unwrap();
assert!(matches!(cli.mode, ProxyMode::Token));
}
#[test]
fn test_resolve_timeout_clamped_to_600() {
let mc = multi_config_from_toml(
r#"
[[proxies]]
name = "slow"
api_key = "test-key"
timeout = 9999
"#,
);
let cli = mc.resolve(&mc.proxies[0]).unwrap();
assert_eq!(cli.timeout, 600);
}
#[test]
fn test_resolve_timeout_under_max_is_unchanged() {
let mc = multi_config_from_toml(
r#"
[[proxies]]
name = "normal"
api_key = "test-key"
timeout = 120
"#,
);
let cli = mc.resolve(&mc.proxies[0]).unwrap();
assert_eq!(cli.timeout, 120);
}
#[test]
fn test_resolve_missing_api_key_errors() {
let _g = env_guard();
std::env::remove_var("APHRODITE_API_KEY");
std::env::remove_var("DEEPSEEK_API_KEY");
std::env::remove_var("HEADROOM_DEEPSEEK_KEY");
let mc = multi_config_from_toml(
r#"
[[proxies]]
name = "no_key"
"#,
);
let result = mc.resolve(&mc.proxies[0]);
assert!(result.is_err());
}
#[test]
fn test_resolve_invalid_listen_address_errors() {
let mc = multi_config_from_toml(
r#"
[[proxies]]
name = "bad_listen"
api_key = "test-key"
listen = "not-an-address"
"#,
);
let result = mc.resolve(&mc.proxies[0]);
assert!(result.is_err());
}
#[test]
fn test_resolve_max_output_must_be_less_than_max_context() {
let mc = multi_config_from_toml(
r#"
[[proxies]]
name = "bad_budget"
api_key = "test-key"
max_context = 100
max_output = 200
"#,
);
let result = mc.resolve(&mc.proxies[0]);
assert!(result.is_err());
}
#[test]
fn test_resolve_defaults_fill_in_missing_proxy_fields() {
let _g = env_guard();
std::env::remove_var("APHRODITE_API_URL");
std::env::remove_var("APHRODITE_MODEL");
let mc = multi_config_from_toml(
r#"
[defaults]
api_key = "default-key"
api_url = "https://default.example.com"
model = "default-model-name"
[[proxies]]
name = "uses_defaults"
"#,
);
let cli = mc.resolve(&mc.proxies[0]).unwrap();
assert_eq!(cli.api_key, "default-key");
assert_eq!(cli.api_url, "https://default.example.com");
assert_eq!(cli.model, "default-model-name");
}
}