1use std::fs;
9use std::path::{Path, PathBuf};
10use std::time::Duration;
11
12use crate::cache::{acquire_lock, atomic_write};
13use crate::error::{AppError, Result};
14use crate::vendor::VendorId;
15
16const LOCK_TIMEOUT: Duration = Duration::from_secs(5);
21
22fn state_dir() -> Result<PathBuf> {
23 let base = directories::BaseDirs::new()
24 .ok_or_else(|| AppError::Other("could not resolve XDG cache dir".into()))?;
25 Ok(base.cache_dir().join("ai-usagebar"))
26}
27
28fn state_path() -> Result<PathBuf> {
29 Ok(state_dir()?.join("active_vendor"))
30}
31
32pub fn read() -> Option<VendorId> {
35 read_from(&state_path().ok()?)
36}
37
38pub fn read_from(path: &Path) -> Option<VendorId> {
42 let raw = fs::read_to_string(path).ok()?;
43 parse_slug(raw.trim())
44}
45
46pub fn write(vendor: VendorId) -> Result<()> {
48 write_to(&state_path()?, vendor)
49}
50
51pub fn write_to(path: &Path, vendor: VendorId) -> Result<()> {
54 atomic_write(path, vendor.slug().as_bytes())
55}
56
57pub fn cycle(enabled: &[VendorId], start: VendorId, delta: i32) -> Result<VendorId> {
61 cycle_at(&state_path()?, enabled, start, delta)
62}
63
64fn lock_path_for(state: &Path) -> PathBuf {
67 let mut p = state.as_os_str().to_os_string();
68 p.push(".lock");
69 PathBuf::from(p)
70}
71
72pub fn cycle_at(
76 path: &Path,
77 enabled: &[VendorId],
78 start: VendorId,
79 delta: i32,
80) -> Result<VendorId> {
81 if enabled.is_empty() {
82 return Err(AppError::Other("no enabled vendors to cycle".into()));
83 }
84 let _lock = acquire_lock(&lock_path_for(path), LOCK_TIMEOUT)?;
89 let current = read_from(path)
90 .filter(|v| enabled.contains(v))
91 .unwrap_or(start);
92 let cur_idx = enabled.iter().position(|v| *v == current).unwrap_or(0);
93 let n = enabled.len() as i32;
94 let next_idx = ((cur_idx as i32 + delta).rem_euclid(n)) as usize;
95 let next = enabled[next_idx];
96 write_to(path, next)?;
97 Ok(next)
98}
99
100fn parse_slug(s: &str) -> Option<VendorId> {
101 match s {
102 "anthropic" => Some(VendorId::Anthropic),
103 "anthropic_api" => Some(VendorId::AnthropicApi),
104 "openai" => Some(VendorId::Openai),
105 "copilot" => Some(VendorId::Copilot),
106 "zai" => Some(VendorId::Zai),
107 "openrouter" => Some(VendorId::Openrouter),
108 "deepseek" => Some(VendorId::Deepseek),
109 "kimi" => Some(VendorId::Kimi),
110 "kilo" => Some(VendorId::Kilo),
111 "novita" => Some(VendorId::Novita),
112 "moonshot" => Some(VendorId::Moonshot),
113 "grok" => Some(VendorId::Grok),
114 "supergrok" => Some(VendorId::Supergrok),
115 "grokbot" => Some(VendorId::Grokbot),
116 "antigravity" => Some(VendorId::Antigravity),
117 "cursor" => Some(VendorId::Cursor),
118 "minimax" => Some(VendorId::Minimax),
119 "kiro" => Some(VendorId::Kiro),
120 "nous" => Some(VendorId::NousResearch),
121 "opencode-go" => Some(VendorId::OpenCodeGo),
122 "commandcode" => Some(VendorId::CommandCode),
123 "ollama" => Some(VendorId::Ollama),
124 "orcarouter" => Some(VendorId::OrcaRouter),
125 "modelstudio" => Some(VendorId::ModelStudio),
126 _ => None,
127 }
128}
129
130#[cfg(test)]
131mod tests {
132 use super::*;
133 use tempfile::TempDir;
134
135 const CYCLE_SET: [VendorId; 4] = [
138 VendorId::Anthropic,
139 VendorId::Openai,
140 VendorId::Zai,
141 VendorId::Openrouter,
142 ];
143
144 #[test]
150 fn concurrent_cycles_do_not_lose_a_step() {
151 let td = TempDir::new().unwrap();
152 let path = td.path().join("active_vendor");
153 let start = VendorId::Anthropic;
154
155 const THREADS: usize = 4;
158 std::thread::scope(|s| {
159 for _ in 0..THREADS {
160 s.spawn(|| {
161 let _ = cycle_at(&path, &CYCLE_SET, start, 1);
162 });
163 }
164 });
165
166 let landed = read_from(&path).expect("a vendor must have been persisted");
167 assert_eq!(
168 landed,
169 start,
170 "{THREADS} single steps over {} vendors must return to the start; \
171 landing on {landed:?} means a step was lost to a race",
172 CYCLE_SET.len()
173 );
174 }
175
176 #[test]
177 fn parse_slug_round_trip() {
178 for id in VendorId::all() {
179 assert_eq!(parse_slug(id.slug()), Some(*id));
180 }
181 }
182
183 #[test]
184 fn parse_slug_unknown_returns_none() {
185 assert!(parse_slug("not-a-vendor").is_none());
186 assert!(parse_slug("").is_none());
187 }
188
189 #[test]
190 fn read_from_missing_or_garbage_returns_none() {
191 let td = TempDir::new().unwrap();
192 assert!(read_from(&td.path().join("active_vendor")).is_none());
194 let path = td.path().join("active_vendor");
196 write_to(&path, VendorId::Zai).unwrap();
197 assert_eq!(read_from(&path), Some(VendorId::Zai));
198 fs::write(&path, "not-a-vendor").unwrap();
200 assert!(read_from(&path).is_none());
201 }
202
203 #[test]
204 fn cycle_at_persists_state_across_calls() {
205 let td = TempDir::new().unwrap();
206 let path = td.path().join("active_vendor");
207
208 let v = cycle_at(&path, &CYCLE_SET, VendorId::Anthropic, 1).unwrap();
210 assert_eq!(v, VendorId::Openai);
211 assert_eq!(read_from(&path), Some(VendorId::Openai));
212
213 let v = cycle_at(&path, &CYCLE_SET, VendorId::Anthropic, 1).unwrap();
215 assert_eq!(v, VendorId::Zai);
216 assert_eq!(read_from(&path), Some(VendorId::Zai));
217 }
218
219 #[test]
220 fn cycle_at_wraps_forward_and_backward() {
221 let td = TempDir::new().unwrap();
222 let path = td.path().join("active_vendor");
223 write_to(&path, VendorId::Anthropic).unwrap();
224
225 assert_eq!(
227 cycle_at(&path, &CYCLE_SET, VendorId::Anthropic, -1).unwrap(),
228 VendorId::Openrouter
229 );
230 assert_eq!(
232 cycle_at(&path, &CYCLE_SET, VendorId::Anthropic, 1).unwrap(),
233 VendorId::Anthropic
234 );
235 }
236
237 #[test]
238 fn cycle_at_ignores_persisted_vendor_not_in_enabled_set() {
239 let td = TempDir::new().unwrap();
240 let path = td.path().join("active_vendor");
241 write_to(&path, VendorId::Deepseek).unwrap();
243 let enabled = [VendorId::Anthropic, VendorId::Openai];
244 let v = cycle_at(&path, &enabled, VendorId::Openai, 1).unwrap();
246 assert_eq!(v, VendorId::Anthropic);
247 }
248
249 #[test]
250 fn cycle_at_errors_on_empty_enabled() {
251 let td = TempDir::new().unwrap();
252 let path = td.path().join("active_vendor");
253 let res = cycle_at(&path, &[], VendorId::Anthropic, 1);
254 assert!(matches!(res, Err(AppError::Other(_))));
255 }
256}