Skip to main content

cooklang_import/
config.rs

1use config::{Config, ConfigError, Environment, File};
2use serde::Deserialize;
3use std::collections::HashMap;
4
5/// Main AI configuration structure
6#[derive(Debug, Deserialize, Clone)]
7pub struct AiConfig {
8    /// Default provider to use when not specified
9    #[serde(default = "default_provider")]
10    pub default_provider: String,
11    /// Map of provider name to provider configuration
12    #[serde(default)]
13    pub providers: HashMap<String, ProviderConfig>,
14    /// Fallback configuration for automatic provider switching
15    #[serde(default)]
16    pub fallback: FallbackConfig,
17    /// Extractors configuration
18    #[serde(default)]
19    pub extractors: ExtractorsConfig,
20    /// Converters configuration
21    #[serde(default)]
22    pub converters: ConvertersConfig,
23    /// Page scriber configuration for browser-based fetching
24    #[serde(default)]
25    pub page_scriber: PageScriberConfig,
26    /// Request timeout in seconds
27    #[serde(default = "default_timeout")]
28    pub timeout: u64,
29}
30
31/// Configuration for a specific AI provider
32#[derive(Debug, Deserialize, Clone)]
33pub struct ProviderConfig {
34    /// Whether this provider is enabled
35    pub enabled: bool,
36    /// Model identifier (e.g., "gpt-4", "claude-3-5-sonnet-20250929")
37    pub model: String,
38    /// Temperature for generation (0.0-1.0)
39    #[serde(default = "default_temperature")]
40    pub temperature: f32,
41    /// Maximum tokens to generate
42    #[serde(default = "default_max_tokens")]
43    pub max_tokens: u32,
44
45    // Optional provider-specific fields
46    /// API key for authentication (can also be set via environment variable)
47    pub api_key: Option<String>,
48    /// Base URL for API endpoint (for custom or proxy endpoints)
49    pub base_url: Option<String>,
50    /// Specific endpoint path (for Azure or custom deployments)
51    pub endpoint: Option<String>,
52    /// Deployment name (Azure OpenAI specific)
53    pub deployment_name: Option<String>,
54    /// API version (Azure OpenAI specific)
55    pub api_version: Option<String>,
56    /// Project ID (Google Cloud specific)
57    pub project_id: Option<String>,
58}
59
60/// Configuration for provider fallback and retry behavior
61#[derive(Debug, Deserialize, Clone)]
62pub struct FallbackConfig {
63    /// Whether fallback is enabled
64    #[serde(default)]
65    pub enabled: bool,
66    /// Order of providers to try (first to last)
67    #[serde(default)]
68    pub order: Vec<String>,
69    /// Number of retry attempts per provider before fallback
70    #[serde(default = "default_retry_attempts")]
71    pub retry_attempts: u32,
72    /// Initial delay between retries in milliseconds (uses exponential backoff)
73    #[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/// Configuration for recipe extractors
89#[derive(Debug, Clone, Deserialize, Default)]
90pub struct ExtractorsConfig {
91    /// List of enabled extractors
92    #[serde(default = "default_extractors")]
93    pub enabled: Vec<String>,
94    /// Order in which extractors should be tried
95    #[serde(default = "default_extractors")]
96    pub order: Vec<String>,
97}
98
99/// Configuration for recipe converters
100#[derive(Debug, Clone, Deserialize, Default)]
101pub struct ConvertersConfig {
102    /// List of enabled converters
103    #[serde(default)]
104    pub enabled: Vec<String>,
105    /// Order in which converters should be tried
106    #[serde(default)]
107    pub order: Vec<String>,
108    /// Default converter to use
109    #[serde(default)]
110    pub default: String,
111}
112
113/// Configuration for the page scriber service (browser-based fetching)
114#[derive(Debug, Deserialize, Clone, Default)]
115pub struct PageScriberConfig {
116    /// Base URL of the page scriber service (e.g., "http://localhost:4000")
117    pub url: Option<String>,
118    /// Domains that should use page scriber directly (suffix-matched)
119    /// e.g., ["seriouseats.com", "allrecipes.com"]
120    #[serde(default)]
121    pub domains: Vec<String>,
122}
123
124// Default value functions
125fn 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    /// Load configuration from file and environment variables
159    ///
160    /// Configuration is loaded with the following priority (highest to lowest):
161    /// 1. Environment variables with COOKLANG__ prefix
162    /// 2. config.toml file in current directory
163    /// 3. Default values
164    ///
165    /// Environment variable format: COOKLANG__PROVIDERS__OPENAI__API_KEY
166    pub fn load() -> Result<Self, ConfigError> {
167        load_config()
168    }
169}
170
171/// Load configuration from file and environment variables
172///
173/// Configuration is loaded with the following priority (highest to lowest):
174/// 1. Environment variables with COOKLANG__ prefix
175/// 2. config.toml file in current directory
176/// 3. Default values
177///
178/// Environment variable format: COOKLANG__PROVIDERS__OPENAI__API_KEY
179pub fn load_config() -> Result<AiConfig, ConfigError> {
180    let settings = Config::builder()
181        // Optional config file (can be missing)
182        .add_source(File::with_name("config").required(false))
183        // Environment variables with COOKLANG_ prefix
184        // Use double underscore for nested: COOKLANG__PROVIDERS__OPENAI__API_KEY
185        .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        // Test that ProviderConfig can be created with None for optional fields
221        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        // Clear any environment variables that might interfere
241        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        // Loading config without a file should use defaults (will fail because no providers configured)
251        // This is expected behavior - we need at least one provider configured
252        let result = load_config();
253
254        // We expect this to fail because no providers are configured
255        // The important thing is it doesn't panic
256        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        // Test that we can construct AiConfig with proper structure
269        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}