1use cerno_host::HostKind;
8use cerno_types::Calibration;
9use serde::Deserialize;
10use std::collections::BTreeMap;
11use std::net::SocketAddr;
12use std::time::Duration;
13
14const DEFAULT_BIND: &str = "127.0.0.1:3000";
18const DEFAULT_MODEL: &str = "gemma4:e2b-it-qat";
19const DEFAULT_KEEP_ALIVE: &str = "5m";
20const DEFAULT_CONCURRENCY: usize = 4;
21const DEFAULT_TIMEOUT_SECS: u64 = 30;
22const DEFAULT_REQUEST_TIMEOUT_SECS: u64 = 50;
25
26#[derive(Debug, thiserror::Error)]
27pub enum ConfigError {
28 #[error("{var} is not valid: {source}")]
29 Var {
30 var: &'static str,
31 #[source]
32 source: Box<dyn std::error::Error + Send + Sync>,
33 },
34
35 #[error("could not read config file {path}: {source}")]
36 Read {
37 path: String,
38 #[source]
39 source: std::io::Error,
40 },
41
42 #[error("could not parse config file {path}: {source}")]
43 Parse {
44 path: String,
45 #[source]
46 source: toml::de::Error,
47 },
48
49 #[error("{setting} {reason}")]
51 Unusable {
52 setting: &'static str,
53 reason: &'static str,
54 },
55
56 #[error(
57 "default_model {model:?} is not in the model table, and strict_models is on; \
58 configure it or pick one of: {known:?}"
59 )]
60 UnknownDefault { model: String, known: Vec<String> },
61
62 #[error(
65 "model {alias:?} has calibration temperature {temperature}; it must be finite and above zero"
66 )]
67 InvalidCalibration { alias: String, temperature: f64 },
68}
69
70#[derive(Debug, Clone, Deserialize)]
72pub struct ModelEntry {
73 pub model: String,
75 #[serde(default = "default_calibration")]
76 pub calibration: CalibrationEntry,
77}
78
79#[derive(Debug, Clone, Copy, Deserialize)]
80pub struct CalibrationEntry {
81 pub temperature: f64,
82}
83
84fn default_calibration() -> CalibrationEntry {
85 CalibrationEntry { temperature: 1.0 }
86}
87
88impl From<CalibrationEntry> for Calibration {
89 fn from(entry: CalibrationEntry) -> Self {
90 Calibration {
91 temperature: entry.temperature,
92 }
93 }
94}
95
96#[derive(Debug, Default, Deserialize)]
97struct FileConfig {
98 #[serde(default)]
99 default_model: Option<String>,
100 #[serde(default)]
101 strict_models: Option<bool>,
102 #[serde(default)]
103 models: BTreeMap<String, ModelEntry>,
104}
105
106#[derive(Debug, Clone)]
107pub struct Config {
108 pub bind: SocketAddr,
109 pub host: HostKind,
111 pub host_url: String,
113 pub host_api_key: Option<String>,
115 pub default_model: String,
117 pub models: BTreeMap<String, ModelEntry>,
119 pub strict_models: bool,
122 pub max_concurrent_questions: usize,
123 pub keep_alive: Option<String>,
124 pub host_timeout: Duration,
125 pub request_timeout: Duration,
128}
129
130fn var<T>(name: &'static str, fallback: T) -> Result<T, ConfigError>
131where
132 T: std::str::FromStr,
133 T::Err: std::error::Error + Send + Sync + 'static,
134{
135 match std::env::var(name) {
136 Err(_) => Ok(fallback),
137 Ok(raw) => raw.trim().parse().map_err(|e: T::Err| ConfigError::Var {
138 var: name,
139 source: Box::new(e),
140 }),
141 }
142}
143
144impl Config {
145 pub fn from_env() -> Result<Self, ConfigError> {
147 let file = match std::env::var("CERNO_CONFIG") {
148 Err(_) => FileConfig::default(),
149 Ok(path) => {
150 let text = std::fs::read_to_string(&path).map_err(|source| ConfigError::Read {
151 path: path.clone(),
152 source,
153 })?;
154 toml::from_str(&text).map_err(|source| ConfigError::Parse { path, source })?
155 }
156 };
157
158 let default_model = std::env::var("CERNO_DEFAULT_MODEL")
159 .ok()
160 .map(|model| model.trim().to_string())
161 .or(file.default_model)
162 .unwrap_or_else(|| DEFAULT_MODEL.to_string());
163
164 let strict_models = match std::env::var("CERNO_STRICT_MODELS") {
165 Ok(raw) => {
166 raw.trim()
167 .parse()
168 .map_err(|e: std::str::ParseBoolError| ConfigError::Var {
169 var: "CERNO_STRICT_MODELS",
170 source: Box::new(e),
171 })?
172 }
173 Err(_) => file.strict_models.unwrap_or(false),
174 };
175
176 let host: HostKind = var("CERNO_HOST", HostKind::Ollama)?;
177
178 let config = Self {
179 bind: var(
180 "CERNO_BIND",
181 DEFAULT_BIND.parse().expect("valid default bind"),
182 )?,
183 host,
184 host_url: std::env::var("CERNO_HOST_URL")
185 .unwrap_or_else(|_| host.default_url().to_string())
186 .trim()
187 .trim_end_matches('/')
188 .to_string(),
189 host_api_key: std::env::var("CERNO_HOST_API_KEY")
190 .ok()
191 .filter(|key| !key.trim().is_empty()),
192 default_model,
193 models: file.models,
194 strict_models,
195 max_concurrent_questions: var("CERNO_MAX_CONCURRENT_QUESTIONS", DEFAULT_CONCURRENCY)?,
196 keep_alive: match std::env::var("CERNO_KEEP_ALIVE") {
197 Ok(v) if v.trim().is_empty() => None,
200 Ok(v) => Some(v),
201 Err(_) => Some(DEFAULT_KEEP_ALIVE.to_string()),
202 },
203 host_timeout: Duration::from_secs(var(
204 "CERNO_HOST_TIMEOUT_SECS",
205 DEFAULT_TIMEOUT_SECS,
206 )?),
207 request_timeout: Duration::from_secs(var(
208 "CERNO_REQUEST_TIMEOUT_SECS",
209 DEFAULT_REQUEST_TIMEOUT_SECS,
210 )?),
211 };
212
213 config.validate()?;
214
215 if config.host_timeout >= config.request_timeout {
218 tracing::warn!(
219 host_timeout_secs = config.host_timeout.as_secs(),
220 request_timeout_secs = config.request_timeout.as_secs(),
221 "CERNO_HOST_TIMEOUT_SECS is not below CERNO_REQUEST_TIMEOUT_SECS, so it never \
222 takes effect"
223 );
224 }
225
226 for (model, aliases) in config.conflicting_aliases() {
228 tracing::warn!(
229 model,
230 ?aliases,
231 "aliases point at the same model with different calibrations; a request naming \
232 the model itself gets the first alias's"
233 );
234 }
235
236 Ok(config)
237 }
238
239 fn validate(&self) -> Result<(), ConfigError> {
241 if self.default_model.trim().is_empty() {
243 return Err(ConfigError::Unusable {
244 setting: "default_model",
245 reason: "must not be empty",
246 });
247 }
248 if self.host_url.trim().is_empty() {
249 return Err(ConfigError::Unusable {
250 setting: "CERNO_HOST_URL",
251 reason: "must not be empty; unset it to use the host's default",
252 });
253 }
254 if self.host_timeout.is_zero() {
255 return Err(ConfigError::Unusable {
256 setting: "CERNO_HOST_TIMEOUT_SECS",
257 reason: "must be above zero",
258 });
259 }
260 if self.request_timeout.is_zero() {
261 return Err(ConfigError::Unusable {
262 setting: "CERNO_REQUEST_TIMEOUT_SECS",
263 reason: "must be above zero",
264 });
265 }
266 if self.max_concurrent_questions == 0 {
269 return Err(ConfigError::Unusable {
270 setting: "CERNO_MAX_CONCURRENT_QUESTIONS",
271 reason: "must be above zero",
272 });
273 }
274
275 for (alias, entry) in &self.models {
276 let temperature = entry.calibration.temperature;
277 if !temperature.is_finite() || temperature <= 0.0 {
278 return Err(ConfigError::InvalidCalibration {
279 alias: alias.clone(),
280 temperature,
281 });
282 }
283 }
284
285 if self.strict_models && self.resolve(&self.default_model).is_none() {
287 return Err(ConfigError::UnknownDefault {
288 model: self.default_model.clone(),
289 known: self.models.keys().cloned().collect(),
290 });
291 }
292
293 Ok(())
294 }
295
296 pub fn resolve(&self, name: &str) -> Option<(String, Calibration)> {
302 if let Some(entry) = self.models.get(name) {
303 return Some((entry.model.clone(), entry.calibration.into()));
304 }
305 if let Some(entry) = self.models.values().find(|e| e.model == name) {
306 return Some((entry.model.clone(), entry.calibration.into()));
307 }
308 if self.strict_models {
309 return None;
310 }
311 Some((name.to_string(), Calibration::default()))
312 }
313
314 pub fn conflicting_aliases(&self) -> Vec<(String, Vec<String>)> {
319 let mut by_model: BTreeMap<&str, Vec<(&str, f64)>> = BTreeMap::new();
320 for (alias, entry) in &self.models {
321 by_model
322 .entry(&entry.model)
323 .or_default()
324 .push((alias, entry.calibration.temperature));
325 }
326
327 by_model
328 .into_iter()
329 .filter(|(_, aliases)| aliases.iter().any(|(_, t)| *t != aliases[0].1))
330 .map(|(model, aliases)| {
331 (
332 model.to_string(),
333 aliases.iter().map(|(a, _)| a.to_string()).collect(),
334 )
335 })
336 .collect()
337 }
338
339 pub fn known_models(&self) -> Vec<String> {
340 self.models.keys().cloned().collect()
341 }
342}
343
344#[cfg(test)]
345mod tests {
346 use super::*;
347
348 fn entry(model: &str, temperature: f64) -> ModelEntry {
349 ModelEntry {
350 model: model.into(),
351 calibration: CalibrationEntry { temperature },
352 }
353 }
354
355 fn config(strict: bool) -> Config {
356 Config {
357 bind: DEFAULT_BIND.parse().unwrap(),
358 host: HostKind::Ollama,
359 host_url: HostKind::Ollama.default_url().into(),
360 host_api_key: None,
361 default_model: "small".into(),
362 models: BTreeMap::from([("small".to_string(), entry("gemma4:e2b-it-qat", 2.5))]),
363 strict_models: strict,
364 max_concurrent_questions: 4,
365 keep_alive: Some("5m".into()),
366 host_timeout: Duration::from_secs(30),
367 request_timeout: Duration::from_secs(50),
368 }
369 }
370
371 #[test]
372 fn an_alias_resolves_to_its_model_and_calibration() {
373 let (model, calibration) = config(false).resolve("small").unwrap();
374
375 assert_eq!(model, "gemma4:e2b-it-qat");
376 assert_eq!(calibration.temperature, 2.5);
377 }
378
379 #[test]
382 fn a_models_own_name_inherits_the_alias_calibration() {
383 let (model, calibration) = config(false).resolve("gemma4:e2b-it-qat").unwrap();
384
385 assert_eq!(model, "gemma4:e2b-it-qat");
386 assert_eq!(calibration.temperature, 2.5);
387 }
388
389 #[test]
390 fn an_unconfigured_model_passes_through_uncalibrated() {
391 let (model, calibration) = config(false).resolve("granite4:3b").unwrap();
392
393 assert_eq!(model, "granite4:3b");
394 assert_eq!(calibration.temperature, 1.0);
395 }
396
397 #[test]
398 fn strict_mode_rejects_anything_not_configured() {
399 assert!(config(true).resolve("granite4:3b").is_none());
400 assert!(config(true).resolve("small").is_some());
401 assert!(config(true).resolve("gemma4:e2b-it-qat").is_some());
402 }
403
404 #[test]
407 fn a_non_positive_or_non_finite_temperature_is_refused_at_startup() {
408 for temperature in [0.0, -1.0, f64::NAN, f64::INFINITY] {
409 let mut config = config(false);
410 config
411 .models
412 .insert("bad".into(), entry("gemma4:e2b-it-qat", temperature));
413
414 assert!(
415 matches!(
416 config.validate(),
417 Err(ConfigError::InvalidCalibration { ref alias, .. }) if alias == "bad"
418 ),
419 "temperature {temperature} was accepted"
420 );
421 }
422
423 assert!(config(false).validate().is_ok());
424 }
425
426 #[test]
429 fn an_empty_model_or_url_and_a_zero_timeout_are_refused_at_startup() {
430 let mut empty_model = config(false);
431 empty_model.default_model = " ".into();
432 let mut empty_url = config(false);
433 empty_url.host_url = String::new();
434 let mut no_time = config(false);
435 no_time.host_timeout = Duration::ZERO;
436 let mut no_request_time = config(false);
437 no_request_time.request_timeout = Duration::ZERO;
438 let mut no_concurrency = config(false);
439 no_concurrency.max_concurrent_questions = 0;
440
441 for (config, setting) in [
442 (empty_model, "default_model"),
443 (empty_url, "CERNO_HOST_URL"),
444 (no_time, "CERNO_HOST_TIMEOUT_SECS"),
445 (no_request_time, "CERNO_REQUEST_TIMEOUT_SECS"),
446 (no_concurrency, "CERNO_MAX_CONCURRENT_QUESTIONS"),
447 ] {
448 assert!(
449 matches!(
450 config.validate(),
451 Err(ConfigError::Unusable { setting: s, .. }) if s == setting
452 ),
453 "{setting} was accepted"
454 );
455 }
456 }
457
458 #[test]
459 fn aliases_disagreeing_about_one_model_are_found() {
460 let mut config = config(false);
461 config
462 .models
463 .insert("same".into(), entry("gemma4:e2b-it-qat", 2.5));
464 assert!(config.conflicting_aliases().is_empty(), "same temperature");
465
466 config
467 .models
468 .insert("warm".into(), entry("gemma4:e2b-it-qat", 4.0));
469 assert_eq!(
470 config.conflicting_aliases(),
471 vec![(
472 "gemma4:e2b-it-qat".to_string(),
473 vec!["same".to_string(), "small".to_string(), "warm".to_string()]
474 )]
475 );
476 }
477
478 #[test]
479 fn strict_mode_refuses_an_unreachable_default() {
480 let mut config = config(true);
481 config.default_model = "missing".into();
482
483 assert!(matches!(
484 config.validate(),
485 Err(ConfigError::UnknownDefault { .. })
486 ));
487 }
488
489 #[test]
490 fn a_toml_table_parses_into_the_model_map() {
491 let file: FileConfig = toml::from_str(
492 r#"
493 default_model = "small"
494 strict_models = true
495
496 [models.small]
497 model = "gemma4:e2b-it-qat"
498 calibration = { temperature = 2.5 }
499
500 [models.big]
501 model = "gemma4:26b-a4b-it-q4_K_M"
502 "#,
503 )
504 .unwrap();
505
506 assert_eq!(file.default_model.as_deref(), Some("small"));
507 assert_eq!(file.strict_models, Some(true));
508 assert_eq!(file.models["small"].calibration.temperature, 2.5);
509 assert_eq!(file.models["big"].calibration.temperature, 1.0);
511 }
512}