Skip to main content

agentic_core/
settings.rs

1use crate::models::ModelValidator;
2use crate::theme::ThemeVariant;
3use figment::{
4    providers::{Format, Toml},
5    Figment,
6};
7use serde::{Deserialize, Serialize};
8use std::fs;
9
10const LOCAL_MODEL_PLACEHOLDER: &str = "[SELECT]";
11const CLOUD_MODEL_PLACEHOLDER: &str = "[SELECT]";
12const API_KEY_PLACEHOLDER: &str = "sk-or-v1-982...b52";
13
14#[derive(Debug, PartialEq, Eq)]
15pub enum ValidationError {
16    LocalModel,
17    CloudModel,
18    ApiKey,
19    LocalEndpointUnreachable,
20    LocalModelNotFound,
21    CloudEndpointUnreachable,
22    CloudModelNotFound,
23}
24
25#[derive(Debug, Clone, Deserialize, Serialize)]
26pub struct Settings {
27    pub theme: ThemeVariant,
28    pub endpoint: String,
29    pub local_model: String,
30    pub api_key: String,
31    pub cloud_model: String,
32}
33
34impl Default for Settings {
35    fn default() -> Self {
36        Self {
37            theme: ThemeVariant::default(),
38            endpoint: "localhost:11434".to_string(),
39            local_model: LOCAL_MODEL_PLACEHOLDER.to_string(),
40            api_key: API_KEY_PLACEHOLDER.to_string(),
41            cloud_model: CLOUD_MODEL_PLACEHOLDER.to_string(),
42        }
43    }
44}
45
46impl Settings {
47    pub fn new() -> Result<Self, Box<dyn std::error::Error>> {
48        // This will create a default config if it doesn't exist
49        let config_path = "config.toml";
50        let figment = Figment::new().merge(Toml::file(config_path));
51
52        match figment.extract() {
53            Ok(settings) => Ok(settings),
54            Err(_) => {
55                let default_settings = Settings::default();
56                default_settings.save().unwrap_or_default();
57                Ok(default_settings)
58            }
59        }
60    }
61
62    pub fn save(&self) -> Result<(), std::io::Error> {
63        let toml_string =
64            toml::to_string_pretty(self).expect("Failed to serialize settings to TOML");
65        fs::write("config.toml", toml_string)
66    }
67
68    pub fn is_valid(&self) -> Result<(), ValidationError> {
69        if self.local_model == LOCAL_MODEL_PLACEHOLDER {
70            return Err(ValidationError::LocalModel);
71        }
72        if self.cloud_model == CLOUD_MODEL_PLACEHOLDER {
73            return Err(ValidationError::CloudModel);
74        }
75        if self.api_key == API_KEY_PLACEHOLDER {
76            return Err(ValidationError::ApiKey);
77        }
78        Ok(())
79    }
80
81    pub async fn validate_endpoints(&self) -> Result<(), ValidationError> {
82        let validator = ModelValidator::new();
83
84        // First do basic validation
85        self.is_valid()?;
86
87        // Then validate actual endpoints and generation capabilities
88        validator
89            .validate_local_endpoint(&self.endpoint, &self.local_model)
90            .await
91            .map_err(|_| ValidationError::LocalEndpointUnreachable)?;
92
93        validator
94            .test_local_generation(&self.endpoint, &self.local_model)
95            .await
96            .map_err(|_| ValidationError::LocalEndpointUnreachable)?;
97
98        validator
99            .validate_cloud_endpoint(&self.api_key, &self.cloud_model)
100            .await
101            .map_err(|_| ValidationError::CloudEndpointUnreachable)?;
102
103        validator
104            .test_cloud_generation(&self.api_key, &self.cloud_model)
105            .await
106            .map_err(|_| ValidationError::CloudEndpointUnreachable)?;
107
108        Ok(())
109    }
110
111    pub async fn validate_local_only(&self) -> Result<(), ValidationError> {
112        if self.local_model == LOCAL_MODEL_PLACEHOLDER {
113            return Err(ValidationError::LocalModel);
114        }
115
116        let validator = ModelValidator::new();
117
118        // First validate the model exists
119        validator
120            .validate_local_endpoint(&self.endpoint, &self.local_model)
121            .await
122            .map_err(|_| ValidationError::LocalEndpointUnreachable)?;
123
124        // Then test actual generation capability
125        validator
126            .test_local_generation(&self.endpoint, &self.local_model)
127            .await
128            .map_err(|_| ValidationError::LocalEndpointUnreachable)?;
129
130        Ok(())
131    }
132
133    pub async fn validate_cloud_only(&self) -> Result<(), ValidationError> {
134        if self.cloud_model == CLOUD_MODEL_PLACEHOLDER {
135            return Err(ValidationError::CloudModel);
136        }
137        if self.api_key == API_KEY_PLACEHOLDER {
138            return Err(ValidationError::ApiKey);
139        }
140
141        let validator = ModelValidator::new();
142
143        // First validate the model exists
144        validator
145            .validate_cloud_endpoint(&self.api_key, &self.cloud_model)
146            .await
147            .map_err(|_| ValidationError::CloudEndpointUnreachable)?;
148
149        // Then test actual generation capability
150        validator
151            .test_cloud_generation(&self.api_key, &self.cloud_model)
152            .await
153            .map_err(|_| ValidationError::CloudEndpointUnreachable)?;
154
155        Ok(())
156    }
157}