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 .validate_cloud_endpoint(&self.api_key, &self.cloud_model)
95 .await
96 .map_err(|_| ValidationError::CloudEndpointUnreachable)?;
97
98 Ok(())
99 }
100
101 pub async fn validate_local_only(&self) -> Result<(), ValidationError> {
102 if self.local_model == LOCAL_MODEL_PLACEHOLDER {
103 return Err(ValidationError::LocalModel);
104 }
105
106 let validator = ModelValidator::new();
107 validator
108 .validate_local_endpoint(&self.endpoint, &self.local_model)
109 .await
110 .map_err(|_| ValidationError::LocalEndpointUnreachable)?;
111
112 Ok(())
113 }
114
115 pub async fn validate_cloud_only(&self) -> Result<(), ValidationError> {
116 if self.cloud_model == CLOUD_MODEL_PLACEHOLDER {
117 return Err(ValidationError::CloudModel);
118 }
119 if self.api_key == API_KEY_PLACEHOLDER {
120 return Err(ValidationError::ApiKey);
121 }
122
123 let validator = ModelValidator::new();
124 validator
125 .validate_cloud_endpoint(&self.api_key, &self.cloud_model)
126 .await
127 .map_err(|_| ValidationError::CloudEndpointUnreachable)?;
128
129 Ok(())
130 }
131}