1use config::{Config, ConfigError, Environment, File};
2use serde::Deserialize;
3use std::collections::HashMap;
4
5#[derive(Debug, Deserialize, Clone)]
7pub struct AiConfig {
8 #[serde(default = "default_provider")]
10 pub default_provider: String,
11 #[serde(default)]
13 pub providers: HashMap<String, ProviderConfig>,
14 #[serde(default)]
16 pub fallback: FallbackConfig,
17 #[serde(default)]
19 pub extractors: ExtractorsConfig,
20 #[serde(default)]
22 pub converters: ConvertersConfig,
23 #[serde(default)]
25 pub page_scriber: PageScriberConfig,
26 #[serde(default = "default_timeout")]
28 pub timeout: u64,
29}
30
31#[derive(Debug, Deserialize, Clone)]
33pub struct ProviderConfig {
34 pub enabled: bool,
36 pub model: String,
38 #[serde(default = "default_temperature")]
40 pub temperature: f32,
41 #[serde(default = "default_max_tokens")]
43 pub max_tokens: u32,
44
45 pub api_key: Option<String>,
48 pub base_url: Option<String>,
50 pub endpoint: Option<String>,
52 pub deployment_name: Option<String>,
54 pub api_version: Option<String>,
56 pub project_id: Option<String>,
58}
59
60#[derive(Debug, Deserialize, Clone)]
62pub struct FallbackConfig {
63 #[serde(default)]
65 pub enabled: bool,
66 #[serde(default)]
68 pub order: Vec<String>,
69 #[serde(default = "default_retry_attempts")]
71 pub retry_attempts: u32,
72 #[serde(default = "default_retry_delay_ms")]
74 pub retry_delay_ms: u64,
75}
76
77impl Default for FallbackConfig {
78 fn default() -> Self {
79 Self {
80 enabled: false,
81 order: Vec::new(),
82 retry_attempts: default_retry_attempts(),
83 retry_delay_ms: default_retry_delay_ms(),
84 }
85 }
86}
87
88#[derive(Debug, Clone, Deserialize, Default)]
90pub struct ExtractorsConfig {
91 #[serde(default = "default_extractors")]
93 pub enabled: Vec<String>,
94 #[serde(default = "default_extractors")]
96 pub order: Vec<String>,
97}
98
99#[derive(Debug, Clone, Deserialize, Default)]
101pub struct ConvertersConfig {
102 #[serde(default)]
104 pub enabled: Vec<String>,
105 #[serde(default)]
107 pub order: Vec<String>,
108 #[serde(default)]
110 pub default: String,
111}
112
113#[derive(Debug, Deserialize, Clone, Default)]
115pub struct PageScriberConfig {
116 pub url: Option<String>,
118 #[serde(default)]
121 pub domains: Vec<String>,
122}
123
124fn default_provider() -> String {
126 "open_ai".to_string()
127}
128
129fn default_temperature() -> f32 {
130 0.7
131}
132
133fn default_max_tokens() -> u32 {
134 2000
135}
136
137fn default_retry_attempts() -> u32 {
138 3
139}
140
141fn default_retry_delay_ms() -> u64 {
142 1000
143}
144
145fn default_extractors() -> Vec<String> {
146 vec![
147 "json_ld".to_string(),
148 "microdata".to_string(),
149 "html_class".to_string(),
150 ]
151}
152
153fn default_timeout() -> u64 {
154 30
155}
156
157impl AiConfig {
158 pub fn load() -> Result<Self, ConfigError> {
167 load_config()
168 }
169}
170
171pub fn load_config() -> Result<AiConfig, ConfigError> {
180 let settings = Config::builder()
181 .add_source(File::with_name("config").required(false))
183 .add_source(
186 Environment::with_prefix("COOKLANG")
187 .separator("__")
188 .try_parsing(true),
189 )
190 .build()?;
191
192 settings.try_deserialize()
193}
194
195#[cfg(test)]
196mod tests {
197 use super::*;
198 use std::env;
199
200 #[test]
201 fn test_default_values() {
202 assert_eq!(default_provider(), "open_ai");
203 assert_eq!(default_temperature(), 0.7);
204 assert_eq!(default_max_tokens(), 2000);
205 assert_eq!(default_retry_attempts(), 3);
206 assert_eq!(default_retry_delay_ms(), 1000);
207 }
208
209 #[test]
210 fn test_fallback_config_default() {
211 let fallback = FallbackConfig::default();
212 assert!(!fallback.enabled);
213 assert!(fallback.order.is_empty());
214 assert_eq!(fallback.retry_attempts, 3);
215 assert_eq!(fallback.retry_delay_ms, 1000);
216 }
217
218 #[test]
219 fn test_provider_config_has_optional_fields() {
220 let config = ProviderConfig {
222 enabled: true,
223 model: "gpt-4.1-mini".to_string(),
224 temperature: 0.7,
225 max_tokens: 2000,
226 api_key: None,
227 base_url: None,
228 endpoint: None,
229 deployment_name: None,
230 api_version: None,
231 project_id: None,
232 };
233
234 assert!(config.api_key.is_none());
235 assert!(config.base_url.is_none());
236 }
237
238 #[test]
239 fn test_load_config_without_file() {
240 let keys_to_clear: Vec<String> = env::vars()
242 .filter(|(k, _)| k.starts_with("COOKLANG__"))
243 .map(|(k, _)| k)
244 .collect();
245
246 for key in keys_to_clear {
247 env::remove_var(&key);
248 }
249
250 let result = load_config();
253
254 assert!(result.is_ok() || result.is_err());
257 }
258
259 #[test]
260 fn test_page_scriber_config_default() {
261 let config = PageScriberConfig::default();
262 assert!(config.url.is_none());
263 assert!(config.domains.is_empty());
264 }
265
266 #[test]
267 fn test_ai_config_structure() {
268 let mut providers = HashMap::new();
270 providers.insert(
271 "openai".to_string(),
272 ProviderConfig {
273 enabled: true,
274 model: "gpt-4.1-mini".to_string(),
275 temperature: 0.7,
276 max_tokens: 2000,
277 api_key: Some("test-key".to_string()),
278 base_url: None,
279 endpoint: None,
280 deployment_name: None,
281 api_version: None,
282 project_id: None,
283 },
284 );
285
286 let config = AiConfig {
287 default_provider: "openai".to_string(),
288 providers,
289 fallback: FallbackConfig::default(),
290 extractors: ExtractorsConfig::default(),
291 converters: ConvertersConfig::default(),
292 page_scriber: PageScriberConfig::default(),
293 timeout: default_timeout(),
294 };
295
296 assert_eq!(config.default_provider, "openai");
297 assert_eq!(config.providers.len(), 1);
298 assert!(config.providers.contains_key("openai"));
299 }
300}