1use crate::context::Context;
9use crate::error::{Error, Result};
10use crate::{atomic, home};
11use serde::{Deserialize, Serialize};
12use serde_json::Value;
13use std::path::PathBuf;
14
15const SCHEMA: u32 = 3;
16
17#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
21pub struct Park {
22 pub service: String,
23 pub parked_at: i64,
24 pub refresh_fingerprint: String,
25 pub access_expires_at: Option<i64>,
27 pub refresh_expires_at: Option<i64>,
29}
30
31impl Park {
32 pub fn restorable_at(&self, now: i64) -> bool {
33 self.refresh_expires_at.is_none_or(|at| at > now)
34 }
35
36 pub fn askable_at(&self, now: i64) -> bool {
37 self.access_expires_at.is_none_or(|at| at > now)
38 }
39}
40
41#[derive(Serialize, Deserialize, Debug, Clone)]
42pub struct Account {
43 pub label: String,
44 pub account_uuid: String,
45 pub email: String,
46 pub organization_uuid: String,
47 pub oauth_account: Value,
50 pub parked: Option<Park>,
51}
52
53#[derive(Serialize, Deserialize, Debug)]
54pub struct State {
55 pub schema: u32,
56 pub machine: String,
57 pub accounts: Vec<Account>,
58 pub active: Option<String>,
59 #[serde(default)]
62 pub discarded: Vec<String>,
63}
64
65impl Default for State {
66 fn default() -> Self {
67 State {
68 schema: SCHEMA,
69 machine: machine_id(),
70 accounts: Vec::new(),
71 active: None,
72 discarded: Vec::new(),
73 }
74 }
75}
76
77impl State {
78 pub fn get(&self, label: &str) -> Option<&Account> {
79 self.accounts.iter().find(|a| a.label == label)
80 }
81
82 pub fn by_uuid(&self, uuid: &str) -> Option<&Account> {
83 self.accounts.iter().find(|a| a.account_uuid == uuid)
84 }
85
86 fn get_mut(&mut self, label: &str) -> Option<&mut Account> {
87 self.accounts.iter_mut().find(|a| a.label == label)
88 }
89
90 pub fn park(&mut self, label: &str, park: Park) {
92 let service = park.service.clone();
93 if let Some(previous) = self
94 .get_mut(label)
95 .and_then(|account| account.parked.replace(park))
96 && previous.service != service
97 {
98 self.discard(&previous.service);
99 }
100 }
101
102 pub fn discard(&mut self, service: &str) {
105 for account in &mut self.accounts {
106 account.parked.take_if(|p| p.service == service);
107 }
108 if !self.discarded.iter().any(|listed| listed == service) {
109 self.discarded.push(service.to_string());
110 }
111 }
112
113 pub fn references(&self, service: &str) -> bool {
114 self.accounts
115 .iter()
116 .any(|a| a.parked.as_ref().is_some_and(|p| p.service == service))
117 }
118
119 pub fn upsert(&mut self, account: Account) {
120 match self.get_mut(&account.label) {
121 Some(existing) => *existing = account,
122 None => self.accounts.push(account),
123 }
124 }
125
126 pub fn relabel(&mut self, from: &str, to: &str) -> Result<&Account> {
129 if from != to
130 && let Some(taken) = self.get(to)
131 {
132 return Err(Error::LabelTaken {
133 label: to.to_string(),
134 email: taken.email.clone(),
135 });
136 }
137 if self.active.as_deref() == Some(from) {
138 self.active = Some(to.to_string());
139 }
140 let account = self.get_mut(from).ok_or_else(|| Error::AccountUnknown {
141 label: from.to_string(),
142 })?;
143 account.label = to.to_string();
144 Ok(account)
145 }
146
147 pub fn remove(&mut self, label: &str) -> Option<Account> {
149 let index = self.accounts.iter().position(|a| a.label == label)?;
150 let account = self.accounts.remove(index);
151 if let Some(park) = &account.parked {
152 self.discard(&park.service);
153 }
154 Some(account)
155 }
156}
157
158pub fn machine_id() -> String {
160 use sha2::{Digest, Sha256};
161 match machine_uid::get() {
162 Ok(raw) => hex::encode(Sha256::digest(raw.as_bytes())),
163 Err(_) => String::from("unknown"),
164 }
165}
166
167fn file(ctx: &Context) -> PathBuf {
168 home::dir(ctx).join("state.json")
169}
170
171pub fn load(ctx: &Context) -> Result<State> {
172 let path = file(ctx);
173 home::check_location(&home::dir(ctx))?;
174 let raw = match std::fs::read_to_string(&path) {
175 Ok(s) => s,
176 Err(e) if e.kind() == std::io::ErrorKind::NotFound => return Ok(State::default()),
177 Err(source) => return Err(Error::StateUnreadable { path, source }),
178 };
179 let state: State = serde_json::from_str(&raw).map_err(|source| Error::StateCorrupt {
180 path: path.clone(),
181 source,
182 })?;
183 if state.schema != SCHEMA {
184 return Err(Error::StateVersionMismatch {
185 path,
186 found: state.schema,
187 expected: SCHEMA,
188 });
189 }
190 if state.machine != machine_id() {
191 return Err(Error::StateWrongMachine { path });
192 }
193 Ok(state)
194}
195
196pub fn save(ctx: &Context, state: &State) -> Result<()> {
197 home::check_location(&home::dir(ctx))?;
198 let path = file(ctx);
199 let write = |source| Error::StateWriteFailed {
200 path: path.clone(),
201 source,
202 };
203 home::ensure(ctx).map_err(write)?;
204 let body = serde_json::to_string_pretty(state).expect("State is always serialisable");
205 atomic::write(&path, body.as_bytes(), atomic::Perms::Secret).map_err(write)
206}
207
208#[cfg(test)]
209mod tests {
210 use super::*;
211
212 #[test]
213 fn machine_id_is_stable_and_real() {
214 let a = machine_id();
215 assert_eq!(a, machine_id());
216 assert_eq!(
217 a.len(),
218 64,
219 "expected a sha256 of the platform id, got {a:?}"
220 );
221 assert_ne!(
222 a, "unknown",
223 "this platform should report a stable machine id"
224 );
225 }
226
227 fn park(service: &str) -> Park {
228 Park {
229 service: service.into(),
230 parked_at: 100,
231 refresh_fingerprint: "f".into(),
232 access_expires_at: Some(200),
233 refresh_expires_at: Some(300),
234 }
235 }
236
237 fn account(label: &str, parked: Option<Park>) -> Account {
238 Account {
239 label: label.into(),
240 account_uuid: format!("{label}-uuid"),
241 email: format!("{label}@example.com"),
242 organization_uuid: "o".into(),
243 oauth_account: serde_json::json!({}),
244 parked,
245 }
246 }
247
248 #[test]
249 fn a_new_park_discards_the_one_it_replaces() {
250 let mut s = State::default();
251 s.upsert(account("work", Some(park("old"))));
252 s.park("work", park("new"));
253 assert_eq!(
254 s.get("work").unwrap().parked.as_ref().unwrap().service,
255 "new"
256 );
257 assert_eq!(s.discarded, ["old"]);
258 assert!(!s.references("old"));
259 }
260
261 #[test]
262 fn discarding_releases_whichever_account_held_it_and_lists_it_once() {
263 let mut s = State::default();
264 s.upsert(account("work", Some(park("current"))));
265 s.discard("something-else");
266 assert!(s.get("work").unwrap().parked.is_some());
267 s.discard("current");
268 s.discard("current");
269 assert!(s.get("work").unwrap().parked.is_none());
270 assert_eq!(s.discarded, ["something-else", "current"]);
271 }
272
273 #[test]
274 fn relabelling_keeps_the_account_its_park_and_whether_it_is_active() {
275 let mut s = State::default();
276 s.upsert(account("wrong", Some(park("p"))));
277 s.upsert(account("other", None));
278 s.active = Some("wrong".into());
279
280 assert_eq!(
281 s.relabel("wrong", "right").unwrap().email,
282 "wrong@example.com"
283 );
284 assert!(s.get("wrong").is_none());
285 let renamed = s.get("right").unwrap();
286 assert_eq!(renamed.account_uuid, "wrong-uuid");
287 assert_eq!(renamed.parked.as_ref().unwrap().service, "p");
288 assert_eq!(s.active.as_deref(), Some("right"));
289 assert!(s.discarded.is_empty(), "nothing is deleted by a rename");
290 }
291
292 #[test]
293 fn relabelling_refuses_a_taken_label_and_an_unknown_one() {
294 let mut s = State::default();
295 s.upsert(account("a", None));
296 s.upsert(account("b", None));
297 s.active = Some("a".into());
298 assert!(matches!(s.relabel("a", "b"), Err(Error::LabelTaken { .. })));
299 assert!(matches!(
300 s.relabel("nobody", "c"),
301 Err(Error::AccountUnknown { .. })
302 ));
303 assert_eq!(
304 s.active.as_deref(),
305 Some("a"),
306 "a refused rename changes nothing"
307 );
308 assert!(s.relabel("a", "a").is_ok());
309 }
310
311 #[test]
312 fn removing_an_account_lists_its_park_for_deletion() {
313 let mut s = State::default();
314 s.upsert(account("work", Some(park("p"))));
315 assert_eq!(s.remove("work").unwrap().label, "work");
316 assert!(s.accounts.is_empty());
317 assert_eq!(s.discarded, ["p"]);
318 }
319
320 #[test]
321 fn a_park_is_restorable_until_its_login_expires() {
322 let p = park("p");
323 assert!(p.askable_at(199) && !p.askable_at(200));
324 assert!(p.restorable_at(299) && !p.restorable_at(300));
325 let unknown = Park {
326 access_expires_at: None,
327 refresh_expires_at: None,
328 ..park("p")
329 };
330 assert!(
331 unknown.restorable_at(i64::MAX),
332 "no expiry recorded is not expired"
333 );
334 }
335
336 #[test]
337 fn accounts_are_replaced_by_label_not_duplicated() {
338 let mut s = State::default();
339 s.upsert(account("work", None));
340 s.upsert(Account {
341 email: "d@e.f".into(),
342 ..account("work", None)
343 });
344 assert_eq!(s.accounts.len(), 1);
345 assert_eq!(s.get("work").unwrap().email, "d@e.f");
346 }
347}