Skip to main content

squigit/
settings.rs

1// Copyright 2026 a7mddra
2// SPDX-License-Identifier: Apache-2.0
3
4use crate::auth::{
5    check_reveal_authorization, delete_api_key as delete_stored_api_key, encrypt_and_save_api_key,
6    get_api_key_status, reveal_api_key as reveal_stored_api_key, session_api_key_width,
7    session_api_keys_active, validate_api_key, ApiKeyProvider, RevealAuthResult,
8};
9use crate::brain::provider::gemini::models::{
10    DEFAULT_MODEL_EFFORT, MODEL_EFFORTS, PRIMARY_FAST_MODEL, SELECTABLE_MODELS,
11};
12use crate::storage::{self, ProfileStore};
13use serde::{Deserialize, Serialize};
14use squigit_ocr::models::{DEFAULT_OCR_MODEL_ID, OCR_MODELS};
15use std::fs;
16use std::path::{Path, PathBuf};
17use std::str::FromStr;
18
19use crate::services::ocr_models;
20
21pub type SettingsResult<T> = std::result::Result<T, String>;
22
23const CONFIG_FILE_NAME: &str = "config.toml";
24
25const VALID_THEMES: &[&str] = &["system", "dark", "light"];
26const VALID_CAPTURE_TYPES: &[&str] = &["traditional", "squiggle"];
27
28#[derive(Clone, Debug, Deserialize, Serialize)]
29#[serde(rename_all = "camelCase")]
30pub struct SquigitConfig {
31    pub model: String,
32    pub effort: String,
33    pub ocr_enabled: bool,
34    pub ocr_language: String,
35}
36
37#[derive(Default, Deserialize)]
38#[serde(rename_all = "camelCase")]
39pub struct ConfigUpdate {
40    pub model: Option<String>,
41    pub effort: Option<String>,
42    pub ocr_enabled: Option<bool>,
43    pub ocr_language: Option<String>,
44}
45
46impl Default for SquigitConfig {
47    fn default() -> Self {
48        Self {
49            model: PRIMARY_FAST_MODEL.to_string(),
50            effort: DEFAULT_MODEL_EFFORT.to_string(),
51            ocr_enabled: true,
52            ocr_language: DEFAULT_OCR_MODEL_ID.to_string(),
53        }
54    }
55}
56
57#[derive(Clone, Debug, Deserialize, Serialize)]
58#[serde(rename_all = "camelCase")]
59pub struct DesktopConfig {
60    pub theme: String,
61    pub capture_type: String,
62    pub show_tray_icon: bool,
63}
64
65#[derive(Default, Deserialize)]
66#[serde(rename_all = "camelCase")]
67pub struct DesktopConfigUpdate {
68    pub theme: Option<String>,
69    pub capture_type: Option<String>,
70    pub show_tray_icon: Option<bool>,
71}
72
73impl Default for DesktopConfig {
74    fn default() -> Self {
75        Self {
76            theme: "system".to_string(),
77            capture_type: "traditional".to_string(),
78            show_tray_icon: true,
79        }
80    }
81}
82
83fn config_path() -> SettingsResult<PathBuf> {
84    storage::config_path(CONFIG_FILE_NAME)
85        .ok_or_else(|| "Could not locate Squigit's config directory".to_string())
86}
87
88fn default_config_table() -> toml::Table {
89    let mut table = toml::Table::new();
90    let root_defaults = SquigitConfig::default();
91    table.insert(
92        "model".to_string(),
93        toml::Value::String(root_defaults.model),
94    );
95    table.insert(
96        "effort".to_string(),
97        toml::Value::String(root_defaults.effort),
98    );
99    table.insert(
100        "ocr_enabled".to_string(),
101        toml::Value::Boolean(root_defaults.ocr_enabled),
102    );
103    table.insert(
104        "ocr_language".to_string(),
105        toml::Value::String(root_defaults.ocr_language),
106    );
107
108    let mut desktop = toml::Table::new();
109    let desktop_defaults = DesktopConfig::default();
110    desktop.insert(
111        "theme".to_string(),
112        toml::Value::String(desktop_defaults.theme),
113    );
114    desktop.insert(
115        "capture_type".to_string(),
116        toml::Value::String(desktop_defaults.capture_type),
117    );
118    desktop.insert(
119        "show_tray_icon".to_string(),
120        toml::Value::Boolean(desktop_defaults.show_tray_icon),
121    );
122
123    table.insert("desktop".to_string(), toml::Value::Table(desktop));
124    table
125}
126
127fn write_raw_config(table: &toml::Table) -> SettingsResult<()> {
128    let path = config_path()?;
129    if let Some(parent) = path.parent() {
130        fs::create_dir_all(parent).map_err(|error| error.to_string())?;
131    }
132    let serialized = toml::to_string_pretty(table).map_err(|error| error.to_string())?;
133    fs::write(path, serialized).map_err(|error| error.to_string())
134}
135
136fn read_raw_config() -> SettingsResult<toml::Table> {
137    let path = config_path()?;
138    if !path.exists() {
139        let defaults = default_config_table();
140        write_raw_config(&defaults)?;
141        return Ok(defaults);
142    }
143
144    let content = match fs::read_to_string(&path) {
145        Ok(c) => c,
146        Err(err) => return Err(err.to_string()),
147    };
148
149    match toml::from_str::<toml::Table>(&content) {
150        Ok(table) => Ok(table),
151        Err(_) => {
152            // Auto repair corrupt or invalid TOML with default values
153            let defaults = default_config_table();
154            write_raw_config(&defaults)?;
155            Ok(defaults)
156        }
157    }
158}
159
160fn normalize_root_config(table: &toml::Table) -> (SquigitConfig, bool) {
161    let defaults = SquigitConfig::default();
162    let mut modified = false;
163
164    let model = table
165        .get("model")
166        .and_then(|v| v.as_str())
167        .filter(|id| SELECTABLE_MODELS.iter().any(|m| m.id == *id))
168        .map(|s| s.to_string())
169        .unwrap_or_else(|| {
170            modified = true;
171            defaults.model
172        });
173
174    let effort = table
175        .get("effort")
176        .and_then(|v| v.as_str())
177        .filter(|e| MODEL_EFFORTS.contains(e))
178        .map(|s| s.to_string())
179        .unwrap_or_else(|| {
180            modified = true;
181            defaults.effort
182        });
183
184    let ocr_enabled = table
185        .get("ocr_enabled")
186        .and_then(|v| v.as_bool())
187        .unwrap_or_else(|| {
188            modified = true;
189            defaults.ocr_enabled
190        });
191
192    let ocr_language = table
193        .get("ocr_language")
194        .and_then(|v| v.as_str())
195        .filter(|id| OCR_MODELS.iter().any(|m| m.id == *id))
196        .map(|s| s.to_string())
197        .unwrap_or_else(|| {
198            modified = true;
199            defaults.ocr_language
200        });
201
202    (
203        SquigitConfig {
204            model,
205            effort,
206            ocr_enabled,
207            ocr_language,
208        },
209        modified,
210    )
211}
212
213pub fn load_config() -> SettingsResult<SquigitConfig> {
214    let mut table = read_raw_config()?;
215    let (config, modified) = normalize_root_config(&table);
216
217    if modified {
218        table.insert(
219            "model".to_string(),
220            toml::Value::String(config.model.clone()),
221        );
222        table.insert(
223            "effort".to_string(),
224            toml::Value::String(config.effort.clone()),
225        );
226        table.insert(
227            "ocr_enabled".to_string(),
228            toml::Value::Boolean(config.ocr_enabled),
229        );
230        table.insert(
231            "ocr_language".to_string(),
232            toml::Value::String(config.ocr_language.clone()),
233        );
234        write_raw_config(&table)?;
235    }
236
237    Ok(config)
238}
239
240pub fn update_config(updates: ConfigUpdate) -> SettingsResult<SquigitConfig> {
241    let mut table = read_raw_config()?;
242    let (current, _) = normalize_root_config(&table);
243
244    let next_model = updates
245        .model
246        .filter(|id| SELECTABLE_MODELS.iter().any(|m| m.id == id))
247        .unwrap_or(current.model);
248    let next_effort = updates
249        .effort
250        .filter(|e| MODEL_EFFORTS.contains(&e.as_str()))
251        .unwrap_or(current.effort);
252    let next_ocr_enabled = updates.ocr_enabled.unwrap_or(current.ocr_enabled);
253    let next_ocr_language = updates
254        .ocr_language
255        .filter(|id| OCR_MODELS.iter().any(|m| m.id == id))
256        .unwrap_or(current.ocr_language);
257
258    let next = SquigitConfig {
259        model: next_model,
260        effort: next_effort,
261        ocr_enabled: next_ocr_enabled,
262        ocr_language: next_ocr_language,
263    };
264
265    table.insert("model".to_string(), toml::Value::String(next.model.clone()));
266    table.insert(
267        "effort".to_string(),
268        toml::Value::String(next.effort.clone()),
269    );
270    table.insert(
271        "ocr_enabled".to_string(),
272        toml::Value::Boolean(next.ocr_enabled),
273    );
274    table.insert(
275        "ocr_language".to_string(),
276        toml::Value::String(next.ocr_language.clone()),
277    );
278    write_raw_config(&table)?;
279
280    Ok(next)
281}
282
283fn normalize_desktop_config(table: &toml::Table) -> (DesktopConfig, bool) {
284    let defaults = DesktopConfig::default();
285    let desktop_table = match table.get("desktop").and_then(|v| v.as_table()) {
286        Some(t) => t,
287        None => return (defaults, true),
288    };
289
290    let mut modified = false;
291
292    let theme = desktop_table
293        .get("theme")
294        .and_then(|v| v.as_str())
295        .filter(|t| VALID_THEMES.contains(t))
296        .map(|s| s.to_string())
297        .unwrap_or_else(|| {
298            modified = true;
299            defaults.theme
300        });
301
302    let capture_type = desktop_table
303        .get("capture_type")
304        .and_then(|v| v.as_str())
305        .filter(|c| VALID_CAPTURE_TYPES.contains(c))
306        .map(|s| s.to_string())
307        .unwrap_or_else(|| {
308            modified = true;
309            defaults.capture_type
310        });
311
312    let show_tray_icon = desktop_table
313        .get("show_tray_icon")
314        .and_then(|v| v.as_bool())
315        .unwrap_or_else(|| {
316            modified = true;
317            defaults.show_tray_icon
318        });
319
320    (
321        DesktopConfig {
322            theme,
323            capture_type,
324            show_tray_icon,
325        },
326        modified,
327    )
328}
329
330fn set_desktop_table(table: &mut toml::Table, config: &DesktopConfig) {
331    let mut desktop = match table.remove("desktop") {
332        Some(toml::Value::Table(t)) => t,
333        _ => toml::Table::new(),
334    };
335    desktop.insert(
336        "theme".to_string(),
337        toml::Value::String(config.theme.clone()),
338    );
339    desktop.insert(
340        "capture_type".to_string(),
341        toml::Value::String(config.capture_type.clone()),
342    );
343    desktop.insert(
344        "show_tray_icon".to_string(),
345        toml::Value::Boolean(config.show_tray_icon),
346    );
347    table.insert("desktop".to_string(), toml::Value::Table(desktop));
348}
349
350pub fn load_desktop_config() -> SettingsResult<DesktopConfig> {
351    let mut table = read_raw_config()?;
352    let (config, modified) = normalize_desktop_config(&table);
353
354    if modified {
355        set_desktop_table(&mut table, &config);
356        write_raw_config(&table)?;
357    }
358
359    Ok(config)
360}
361
362pub fn update_desktop_config(updates: DesktopConfigUpdate) -> SettingsResult<DesktopConfig> {
363    let mut table = read_raw_config()?;
364    let (current, _) = normalize_desktop_config(&table);
365
366    let next_theme = updates
367        .theme
368        .filter(|t| VALID_THEMES.contains(&t.as_str()))
369        .unwrap_or(current.theme);
370    let next_capture_type = updates
371        .capture_type
372        .filter(|c| VALID_CAPTURE_TYPES.contains(&c.as_str()))
373        .unwrap_or(current.capture_type);
374    let next_show_tray_icon = updates.show_tray_icon.unwrap_or(current.show_tray_icon);
375    let next = DesktopConfig {
376        theme: next_theme,
377        capture_type: next_capture_type,
378        show_tray_icon: next_show_tray_icon,
379    };
380
381    set_desktop_table(&mut table, &next);
382    write_raw_config(&table)?;
383
384    Ok(next)
385}
386
387#[derive(Serialize)]
388#[serde(rename_all = "camelCase")]
389pub struct CredentialState {
390    pub configured: bool,
391    pub width: Option<u32>,
392}
393
394#[derive(Serialize)]
395#[serde(rename_all = "camelCase")]
396pub struct SettingsSnapshot {
397    pub active_profile_id: Option<String>,
398    pub config: SquigitConfig,
399    pub persona: String,
400    pub google_ai_studio: CredentialState,
401    pub imgbb: CredentialState,
402}
403
404fn credential_state(
405    store: &ProfileStore,
406    profile_id: Option<&str>,
407    provider: ApiKeyProvider,
408) -> SettingsResult<CredentialState> {
409    let Some(profile_id) = profile_id else {
410        return Ok(CredentialState {
411            configured: false,
412            width: None,
413        });
414    };
415    let session_width = session_api_key_width(provider);
416    Ok(CredentialState {
417        configured: get_api_key_status(store, provider, profile_id)
418            .map_err(|error| error.to_string())?,
419        width: if session_api_keys_active() {
420            session_width
421        } else {
422            store
423                .get_key_width(profile_id, provider.storage_key_name())
424                .map_err(|error| error.to_string())?
425        },
426    })
427}
428
429pub fn load_settings() -> SettingsResult<SettingsSnapshot> {
430    let store = storage::profile_store().map_err(|error| error.to_string())?;
431    let active_profile_id = store
432        .get_active_profile_id()
433        .map_err(|error| error.to_string())?;
434    Ok(SettingsSnapshot {
435        google_ai_studio: credential_state(
436            &store,
437            active_profile_id.as_deref(),
438            ApiKeyProvider::GoogleAiStudio,
439        )?,
440        imgbb: credential_state(&store, active_profile_id.as_deref(), ApiKeyProvider::ImgBb)?,
441        active_profile_id,
442        config: load_config()?,
443        persona: load_persona()?,
444    })
445}
446
447pub fn load_persona() -> SettingsResult<String> {
448    storage::load_rules()
449}
450
451pub fn save_persona(content: &str) -> SettingsResult<()> {
452    storage::save_rules(content)
453}
454
455pub fn import_persona(source_path: &str) -> SettingsResult<String> {
456    let path = Path::new(source_path);
457    let is_markdown = path
458        .extension()
459        .and_then(|extension| extension.to_str())
460        .is_some_and(|extension| extension.eq_ignore_ascii_case("md"));
461    if !is_markdown {
462        return Err("RULES.md import requires a Markdown file".to_string());
463    }
464    let imported = fs::read_to_string(path).map_err(|error| error.to_string())?;
465    save_persona(&imported)?;
466    Ok(imported)
467}
468
469fn provider(value: &str) -> SettingsResult<ApiKeyProvider> {
470    ApiKeyProvider::from_str(value).map_err(|error| error.to_string())
471}
472
473pub fn validate_api_key_format(provider_name: &str, plaintext: &str) -> SettingsResult<bool> {
474    if plaintext.trim().is_empty() {
475        return Ok(true);
476    }
477    Ok(validate_api_key(provider(provider_name)?, plaintext.trim()).is_ok())
478}
479
480pub fn set_api_key(profile_id: &str, provider_name: &str, plaintext: &str) -> SettingsResult<()> {
481    let store = storage::profile_store().map_err(|error| error.to_string())?;
482    encrypt_and_save_api_key(
483        &store,
484        profile_id,
485        provider(provider_name)?,
486        plaintext.trim(),
487    )
488    .map_err(|error| error.to_string())
489}
490
491pub fn delete_api_key(profile_id: &str, provider_name: &str) -> SettingsResult<bool> {
492    let store = storage::profile_store().map_err(|error| error.to_string())?;
493    delete_stored_api_key(&store, profile_id, provider(provider_name)?)
494        .map_err(|error| error.to_string())
495}
496
497pub fn reveal_api_key(
498    profile_id: &str,
499    provider_name: &str,
500    captcha_passed: bool,
501) -> SettingsResult<Option<String>> {
502    let store = storage::profile_store().map_err(|error| error.to_string())?;
503    let authorization = check_reveal_authorization(&store).map_err(|error| error.to_string())?;
504    if authorization != RevealAuthResult::Authorized && !captcha_passed {
505        return Err("captcha-required".to_string());
506    }
507    if captcha_passed {
508        store
509            .update_last_trusted_reveal()
510            .map_err(|error| error.to_string())?;
511    }
512    reveal_stored_api_key(&store, provider(provider_name)?, profile_id)
513        .map(|secret| secret.map(|value| value.into_inner()))
514        .map_err(|error| error.to_string())
515}
516
517#[derive(Serialize)]
518#[serde(rename_all = "camelCase")]
519pub struct OcrModelStatus {
520    pub id: String,
521    pub name: String,
522    pub lang: String,
523    pub size: String,
524    pub state: String,
525    pub progress: u8,
526    pub loaded: u64,
527    pub total: u64,
528    pub error: Option<String>,
529}
530
531#[derive(Serialize)]
532#[serde(rename_all = "camelCase")]
533pub struct OcrModelsSnapshot {
534    pub default_ocr_model_id: String,
535    pub models: Vec<OcrModelStatus>,
536}
537
538pub fn load_ocr_models() -> SettingsResult<OcrModelsSnapshot> {
539    let manager = ocr_models()?;
540    let downloaded = manager
541        .list_downloaded_models()
542        .map_err(|error| error.to_string())?;
543    let jobs = manager.download_jobs_snapshot();
544    let models = OCR_MODELS
545        .iter()
546        .map(|model| {
547            let installed =
548                model.id == DEFAULT_OCR_MODEL_ID || downloaded.iter().any(|id| id == model.id);
549            let job = jobs.iter().find(|job| job.id == model.id);
550            let state = match job.map(|job| job.status.as_str()) {
551                Some("downloaded" | "already_installed") => "downloaded".to_string(),
552                Some(status) => status.to_string(),
553                None if installed => "downloaded".to_string(),
554                None => "idle".to_string(),
555            };
556            OcrModelStatus {
557                id: model.id.to_string(),
558                name: model.name.to_string(),
559                lang: model.lang.to_string(),
560                size: model.size.to_string(),
561                state,
562                progress: job.map_or(if installed { 100 } else { 0 }, |job| job.progress),
563                loaded: job.map_or(0, |job| job.loaded),
564                total: job.map_or(0, |job| job.total),
565                error: job.and_then(|job| job.error.clone()),
566            }
567        })
568        .collect();
569    Ok(OcrModelsSnapshot {
570        default_ocr_model_id: DEFAULT_OCR_MODEL_ID.to_string(),
571        models,
572    })
573}
574
575pub fn download_ocr_model(model_id: &str) -> SettingsResult<OcrModelsSnapshot> {
576    ocr_models()?
577        .start_download(model_id)
578        .map_err(|error| error.to_string())?;
579    load_ocr_models()
580}
581
582pub fn cancel_ocr_model_download(model_id: &str) -> SettingsResult<OcrModelsSnapshot> {
583    ocr_models()?.cancel_download(model_id);
584    load_ocr_models()
585}
586
587pub fn delete_ocr_model(model_id: &str) -> SettingsResult<OcrModelsSnapshot> {
588    if model_id == DEFAULT_OCR_MODEL_ID {
589        return Err("The bundled OCR model cannot be deleted".to_string());
590    }
591    ocr_models()?
592        .trash_downloaded_model(model_id)
593        .map_err(|error| error.to_string())?;
594    let config = load_config()?;
595    if config.ocr_language == model_id {
596        update_config(ConfigUpdate {
597            ocr_language: Some(DEFAULT_OCR_MODEL_ID.to_string()),
598            ..Default::default()
599        })?;
600    }
601    load_ocr_models()
602}
603
604#[derive(Serialize)]
605#[serde(rename_all = "camelCase")]
606pub struct HelpDiagnostics {
607    pub squigit_version: Option<String>,
608    pub ocr_version: Option<String>,
609    pub os: String,
610}
611
612fn read_ocr_version() -> SettingsResult<Option<String>> {
613    let sidecar_path = squigit_ocr::sidecar::resolve_sidecar_path();
614    Ok(squigit_ocr::sidecar::read_sidecar_version(&sidecar_path).ok())
615}
616
617pub fn ocr_available() -> SettingsResult<bool> {
618    Ok(read_ocr_version()?.is_some())
619}
620
621fn current_squigit_version() -> SettingsResult<Option<String>> {
622    Ok(storage::version_store()
623        .map_err(|error| error.to_string())?
624        .load()
625        .map_err(|error| error.to_string())?
626        .and_then(|version| version.app.current_version)
627        .map(|version| version.trim().to_string())
628        .filter(|version| !version.is_empty()))
629}
630
631#[cfg(any(target_os = "macos", target_os = "windows"))]
632fn command_output(program: &str, arguments: &[&str]) -> Option<String> {
633    std::process::Command::new(program)
634        .args(arguments)
635        .output()
636        .ok()
637        .filter(|output| output.status.success())
638        .map(|output| String::from_utf8_lossy(&output.stdout).trim().to_string())
639        .filter(|output| !output.is_empty())
640}
641
642fn os_release() -> Option<String> {
643    #[cfg(target_os = "linux")]
644    {
645        return fs::read_to_string("/proc/sys/kernel/osrelease")
646            .ok()
647            .map(|release| release.trim().to_string())
648            .filter(|release| !release.is_empty());
649    }
650    #[cfg(target_os = "macos")]
651    {
652        return command_output("sw_vers", &["-productVersion"]);
653    }
654    #[cfg(target_os = "windows")]
655    {
656        return command_output("cmd", &["/C", "ver"]);
657    }
658    #[allow(unreachable_code)]
659    None
660}
661
662fn os_diagnostics() -> String {
663    let name = match std::env::consts::OS {
664        "linux" => "Linux",
665        "macos" => "macOS",
666        "windows" => "Windows",
667        other => other,
668    };
669    let architecture = match std::env::consts::ARCH {
670        "x86_64" => "x64",
671        "aarch64" => "arm64",
672        other => other,
673    };
674    match os_release() {
675        Some(release) => format!("{name} {architecture} {release}"),
676        None => format!("{name} {architecture}"),
677    }
678}
679
680pub fn load_help_diagnostics() -> SettingsResult<HelpDiagnostics> {
681    Ok(HelpDiagnostics {
682        squigit_version: current_squigit_version()?,
683        ocr_version: read_ocr_version()?,
684        os: os_diagnostics(),
685    })
686}