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 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 self.is_valid()?;
86
87 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 validator
120 .validate_local_endpoint(&self.endpoint, &self.local_model)
121 .await
122 .map_err(|_| ValidationError::LocalEndpointUnreachable)?;
123
124 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 validator
145 .validate_cloud_endpoint(&self.api_key, &self.cloud_model)
146 .await
147 .map_err(|_| ValidationError::CloudEndpointUnreachable)?;
148
149 validator
151 .test_cloud_generation(&self.api_key, &self.cloud_model)
152 .await
153 .map_err(|_| ValidationError::CloudEndpointUnreachable)?;
154
155 Ok(())
156 }
157}