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, default_value = "")]
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");
}
}