Skip to main content

ares_llm/
config.rs

1use serde::{Deserialize, Serialize};
2
3// ============= Provider Configuration =============
4
5/// LLM provider configuration.
6///
7/// `openai` covers the OpenAI API and OpenAI-compatible endpoints
8/// (NVIDIA NIM, Groq, Azure OpenAI, local vLLM, etc.).
9///
10/// `azure` covers Azure AI Foundry's OpenAI-compatible `/openai/v1`
11/// endpoints while reading both the API key and base URL from environment
12/// variables.
13///
14/// `bedrock` requires the `bedrock` feature on `ares-llm`. `anthropic`
15/// requires the `anthropic` feature on `ares-llm`. `ollama`
16/// requires the `ollama` feature. These are compiled out of the default
17/// production build (`--no-default-features --features openai,postgres,mcp`),
18/// so they are not present in `ares-server` binaries that omit those
19/// features. The `ares.toml` schema is the same regardless: the runtime
20/// `Provider::from_config_with_params` returns a clear configuration error
21/// when a variant is selected without the corresponding feature enabled.
22#[derive(Debug, Clone, Serialize, Deserialize)]
23#[serde(tag = "type", rename_all = "lowercase")]
24#[non_exhaustive]
25pub enum ProviderConfig {
26    /// OpenAI API (or compatible endpoints, including NVIDIA NIM).
27    OpenAI {
28        /// Environment variable containing API key.
29        api_key_env: String,
30        /// API base URL (default: `https://api.openai.com/v1`).
31        #[serde(default = "default_openai_base")]
32        api_base: String,
33        /// Default model to use with this provider.
34        default_model: String,
35    },
36    /// Azure AI Foundry OpenAI-compatible chat completions.
37    Azure {
38        /// Environment variable containing the Foundry API key.
39        #[serde(default = "default_azure_api_key_env")]
40        api_key_env: String,
41        /// Environment variable containing the Foundry base URL.
42        #[serde(default = "default_azure_base_url_env")]
43        base_url_env: String,
44        /// Default Foundry model id.
45        #[serde(default = "default_azure_default_model")]
46        default_model: String,
47    },
48    /// Anthropic Claude API.
49    Anthropic {
50        /// Environment variable containing API key (default: `ANTHROPIC_API_KEY`).
51        #[serde(default = "default_anthropic_api_key_env")]
52        api_key_env: String,
53        /// Default model (default: `claude-3-5-sonnet-20241022`).
54        #[serde(default = "default_anthropic_default_model")]
55        default_model: String,
56    },
57    /// AWS Bedrock Claude API.
58    Bedrock {
59        /// Environment variable containing the Bedrock bearer token.
60        #[serde(default = "default_bedrock_api_key_env")]
61        api_key_env: String,
62        /// Environment variable containing the AWS region.
63        #[serde(default = "default_bedrock_region_env")]
64        region_env: String,
65        /// Default Bedrock model id.
66        #[serde(default = "default_bedrock_default_model")]
67        default_model: String,
68    },
69    /// Local Ollama server. The `api_key_env` is a dummy field (Ollama does
70    /// not require authentication) so the same fleet-secrets storage layer can
71    /// be used uniformly for all providers.
72    Ollama {
73        /// Dummy env var; unused at runtime.
74        #[serde(default = "default_ollama_api_key_env")]
75        api_key_env: String,
76        /// Base URL of the Ollama server.
77        #[serde(default = "default_ollama_base_url")]
78        base_url: String,
79        /// Default model id (e.g. `ministral-3:3b`).
80        #[serde(default = "default_ollama_default_model")]
81        default_model: String,
82    },
83}
84
85impl ProviderConfig {
86    /// Returns the provider type discriminator.
87    pub fn type_name(&self) -> &'static str {
88        match self {
89            ProviderConfig::OpenAI { .. } => "openai",
90            ProviderConfig::Azure { .. } => "azure",
91            ProviderConfig::Anthropic { .. } => "anthropic",
92            ProviderConfig::Bedrock { .. } => "bedrock",
93            ProviderConfig::Ollama { .. } => "ollama",
94        }
95    }
96}
97
98impl std::str::FromStr for ProviderConfig {
99    type Err = String;
100
101    /// Parse a provider type name into a default `ProviderConfig` variant.
102    /// Accepts `openai` (and the `nvidia` alias), `azure`, `anthropic`, `bedrock`, and `ollama`.
103    fn from_str(s: &str) -> Result<Self, Self::Err> {
104        match s.trim().to_lowercase().as_str() {
105            "openai" | "nvidia" => Ok(ProviderConfig::OpenAI {
106                api_key_env: default_nvidia_api_key_env(),
107                api_base: default_nvidia_api_base(),
108                default_model: default_nvidia_default_model(),
109            }),
110            "azure" => Ok(ProviderConfig::Azure {
111                api_key_env: default_azure_api_key_env(),
112                base_url_env: default_azure_base_url_env(),
113                default_model: default_azure_default_model(),
114            }),
115            "anthropic" => Ok(ProviderConfig::Anthropic {
116                api_key_env: default_anthropic_api_key_env(),
117                default_model: default_anthropic_default_model(),
118            }),
119            "bedrock" => Ok(ProviderConfig::Bedrock {
120                api_key_env: default_bedrock_api_key_env(),
121                region_env: default_bedrock_region_env(),
122                default_model: default_bedrock_default_model(),
123            }),
124            "ollama" => Ok(ProviderConfig::Ollama {
125                api_key_env: default_ollama_api_key_env(),
126                base_url: default_ollama_base_url(),
127                default_model: default_ollama_default_model(),
128            }),
129            other => Err(format!(
130                "Unknown provider type: {other}. Use: openai (or nvidia), azure, anthropic, bedrock, ollama"
131            )),
132        }
133    }
134}
135
136fn default_openai_base() -> String {
137    "https://api.openai.com/v1".to_string()
138}
139
140fn default_nvidia_api_key_env() -> String {
141    "NVIDIA_API_KEY".to_string()
142}
143
144fn default_nvidia_api_base() -> String {
145    "https://integrate.api.nvidia.com/v1".to_string()
146}
147
148fn default_nvidia_default_model() -> String {
149    "nvidia/nemotron-3-ultra-550b-a55b".to_string()
150}
151
152fn default_azure_api_key_env() -> String {
153    "AZURE_FOUNDRY_API_KEY".to_string()
154}
155
156fn default_azure_base_url_env() -> String {
157    "AZURE_FOUNDRY_BASE_URL".to_string()
158}
159
160fn default_azure_default_model() -> String {
161    "DeepSeek-V4-Flash".to_string()
162}
163
164fn default_anthropic_api_key_env() -> String {
165    "ANTHROPIC_API_KEY".to_string()
166}
167
168fn default_anthropic_default_model() -> String {
169    "claude-3-5-sonnet-20241022".to_string()
170}
171
172fn default_bedrock_api_key_env() -> String {
173    "AWS_BEARER_TOKEN_BEDROCK".to_string()
174}
175
176fn default_bedrock_region_env() -> String {
177    "AWS_REGION".to_string()
178}
179
180fn default_bedrock_default_model() -> String {
181    "us.anthropic.claude-haiku-4-5-20251001-v1:0".to_string()
182}
183
184fn default_ollama_api_key_env() -> String {
185    "OLLAMA_API_KEY".to_string()
186}
187
188fn default_ollama_base_url() -> String {
189    "http://localhost:11434".to_string()
190}
191
192fn default_ollama_default_model() -> String {
193    "ministral-3:3b".to_string()
194}
195
196// ============= Model Configuration =============
197
198/// Model configuration referencing a provider.
199#[derive(Debug, Clone, Serialize, Deserialize)]
200pub struct ModelConfig {
201    /// Reference to a provider name defined in \[providers\].
202    pub provider: String,
203
204    /// Model name/identifier to use with the provider.
205    pub model: String,
206
207    /// Sampling temperature (0.0 = deterministic, 1.0+ = creative). Default: 0.7.
208    #[serde(default = "default_temperature")]
209    pub temperature: f32,
210
211    /// Maximum tokens to generate (default: 512).
212    #[serde(default = "default_model_max_tokens")]
213    pub max_tokens: u32,
214}
215
216fn default_temperature() -> f32 {
217    0.7
218}
219
220fn default_model_max_tokens() -> u32 {
221    512
222}