1use std::collections::{BTreeMap, HashMap};
9
10use serde::{Deserialize, Serialize};
11
12pub trait AccountState: Sized {
15 fn new_for_account(account_id: &str, region: &str, endpoint: &str) -> Self;
17
18 fn inherit_from(&mut self, _sibling: &Self) {}
22}
23
24#[derive(Debug, Clone, Serialize, Deserialize)]
30pub struct MultiAccountState<T> {
31 default_account_id: String,
32 region: String,
33 endpoint: String,
34 accounts: HashMap<String, T>,
35}
36
37impl<T: AccountState> MultiAccountState<T> {
38 pub fn new(default_account_id: &str, region: &str, endpoint: &str) -> Self {
40 let mut accounts = HashMap::new();
41 accounts.insert(
42 default_account_id.to_string(),
43 T::new_for_account(default_account_id, region, endpoint),
44 );
45 Self {
46 default_account_id: default_account_id.to_string(),
47 region: region.to_string(),
48 endpoint: endpoint.to_string(),
49 accounts,
50 }
51 }
52
53 pub fn map<U>(&self, mut f: impl FnMut(&T) -> U) -> MultiAccountState<U> {
55 MultiAccountState {
56 default_account_id: self.default_account_id.clone(),
57 region: self.region.clone(),
58 endpoint: self.endpoint.clone(),
59 accounts: self
60 .accounts
61 .iter()
62 .map(|(k, v)| (k.clone(), f(v)))
63 .collect(),
64 }
65 }
66
67 pub fn map_into<U>(self, mut f: impl FnMut(&str, T) -> U) -> MultiAccountState<U> {
71 MultiAccountState {
72 default_account_id: self.default_account_id,
73 region: self.region,
74 endpoint: self.endpoint,
75 accounts: self
76 .accounts
77 .into_iter()
78 .map(|(k, v)| {
79 let mapped = f(&k, v);
80 (k, mapped)
81 })
82 .collect(),
83 }
84 }
85
86 pub fn get_or_create(&mut self, account_id: &str) -> &mut T {
92 if !self.accounts.contains_key(account_id) {
93 let mut state = T::new_for_account(account_id, &self.region, &self.endpoint);
94 if let Some(sibling) = self.accounts.get(&self.default_account_id) {
96 state.inherit_from(sibling);
97 }
98 self.accounts.insert(account_id.to_string(), state);
99 }
100 self.accounts.get_mut(account_id).unwrap()
101 }
102
103 pub fn get_or_create_with<F>(&mut self, account_id: &str, init: F) -> &mut T
107 where
108 F: FnOnce(&mut T),
109 {
110 if !self.accounts.contains_key(account_id) {
111 let mut state = T::new_for_account(account_id, &self.region, &self.endpoint);
112 init(&mut state);
113 self.accounts.insert(account_id.to_string(), state);
114 }
115 self.accounts.get_mut(account_id).unwrap()
116 }
117
118 pub fn get(&self, account_id: &str) -> Option<&T> {
120 self.accounts.get(account_id)
121 }
122
123 pub fn get_mut(&mut self, account_id: &str) -> Option<&mut T> {
125 self.accounts.get_mut(account_id)
126 }
127
128 pub fn iter(&self) -> impl Iterator<Item = (&str, &T)> {
130 self.accounts.iter().map(|(k, v)| (k.as_str(), v))
131 }
132
133 pub fn iter_mut(&mut self) -> impl Iterator<Item = (&str, &mut T)> {
135 self.accounts.iter_mut().map(|(k, v)| (k.as_str(), v))
136 }
137
138 pub fn default_account_id(&self) -> &str {
140 &self.default_account_id
141 }
142
143 pub fn default_mut(&mut self) -> &mut T {
145 self.accounts.get_mut(&self.default_account_id).unwrap()
146 }
147
148 pub fn default_ref(&self) -> &T {
150 self.accounts.get(&self.default_account_id).unwrap()
151 }
152
153 pub fn reset(&mut self) {
156 self.accounts.clear();
157 self.accounts.insert(
158 self.default_account_id.clone(),
159 T::new_for_account(&self.default_account_id, &self.region, &self.endpoint),
160 );
161 }
162
163 pub fn find_account<F>(&self, predicate: F) -> Option<&str>
167 where
168 F: Fn(&T) -> bool,
169 {
170 self.accounts
171 .iter()
172 .find(|(_, v)| predicate(v))
173 .map(|(k, _)| k.as_str())
174 }
175
176 pub fn account_count(&self) -> usize {
178 self.accounts.len()
179 }
180
181 pub fn region(&self) -> &str {
183 &self.region
184 }
185
186 pub fn endpoint(&self) -> &str {
188 &self.endpoint
189 }
190}
191
192#[derive(Debug, Clone, Serialize, Deserialize)]
205pub struct RegionalState<T> {
206 account_id: String,
207 default_region: String,
209 endpoint: String,
210 #[serde(default = "BTreeMap::new")]
212 regions: BTreeMap<String, T>,
213}
214
215impl<T> RegionalState<T> {
216 pub fn new(account_id: &str, default_region: &str, endpoint: &str) -> Self {
218 Self {
219 account_id: account_id.to_string(),
220 default_region: default_region.to_string(),
221 endpoint: endpoint.to_string(),
222 regions: BTreeMap::new(),
223 }
224 }
225
226 pub fn account_id(&self) -> &str {
228 &self.account_id
229 }
230
231 pub fn default_region(&self) -> &str {
233 &self.default_region
234 }
235
236 pub fn endpoint(&self) -> &str {
238 &self.endpoint
239 }
240
241 pub fn region(&self, region: &str) -> Option<&T> {
243 self.regions.get(region)
244 }
245
246 pub fn get_region_mut(&mut self, region: &str) -> Option<&mut T> {
248 self.regions.get_mut(region)
249 }
250
251 pub fn regions(&self) -> impl Iterator<Item = (&str, &T)> {
253 self.regions.iter().map(|(k, v)| (k.as_str(), v))
254 }
255
256 pub fn regions_mut(&mut self) -> impl Iterator<Item = (&str, &mut T)> {
258 self.regions.iter_mut().map(|(k, v)| (k.as_str(), v))
259 }
260
261 pub fn insert_region(&mut self, region: &str, state: T) -> Option<T> {
263 self.regions.insert(region.to_string(), state)
264 }
265
266 pub fn clear(&mut self) {
268 self.regions.clear();
269 }
270
271 pub fn map<U>(&self, mut f: impl FnMut(&T) -> U) -> RegionalState<U> {
273 RegionalState {
274 account_id: self.account_id.clone(),
275 default_region: self.default_region.clone(),
276 endpoint: self.endpoint.clone(),
277 regions: self
278 .regions
279 .iter()
280 .map(|(k, v)| (k.clone(), f(v)))
281 .collect(),
282 }
283 }
284}
285
286impl<T: AccountState> RegionalState<T> {
287 pub fn region_mut(&mut self, region: &str) -> &mut T {
291 if !self.regions.contains_key(region) {
292 let mut state = T::new_for_account(&self.account_id, region, &self.endpoint);
293 if let Some(sibling) = self.regions.values().next() {
294 state.inherit_from(sibling);
295 }
296 self.regions.insert(region.to_string(), state);
297 }
298 self.regions.get_mut(region).expect("inserted above")
299 }
300
301 pub fn region_or_default_mut(&mut self, region: Option<&str>) -> &mut T {
305 let region = region
306 .filter(|r| !r.is_empty())
307 .map(str::to_string)
308 .unwrap_or_else(|| self.default_region.clone());
309 self.region_mut(®ion)
310 }
311}
312
313impl<T: AccountState> AccountState for RegionalState<T> {
314 fn new_for_account(account_id: &str, region: &str, endpoint: &str) -> Self {
315 Self::new(account_id, region, endpoint)
316 }
317
318 fn inherit_from(&mut self, sibling: &Self) {
319 if let Some(shared) = sibling.regions.values().next() {
320 for state in self.regions.values_mut() {
321 state.inherit_from(shared);
322 }
323 }
324 }
325}
326
327pub type MultiRegionState<T> = MultiAccountState<RegionalState<T>>;
329
330impl<T: AccountState> MultiAccountState<RegionalState<T>> {
331 pub fn regional(&self, account_id: &str, region: &str) -> Option<&T> {
335 self.get(account_id).and_then(|a| a.region(region))
336 }
337
338 pub fn regional_mut(&mut self, account_id: &str, region: &str) -> &mut T {
346 let exists = self
347 .get(account_id)
348 .is_some_and(|a| a.region(region).is_some());
349 if !exists {
350 let mut state = T::new_for_account(account_id, region, &self.endpoint);
351 let sibling = self
352 .get(account_id)
353 .and_then(|a| a.regions.values().next())
354 .or_else(|| {
355 let default = self.get(&self.default_account_id)?;
356 default
357 .region(region)
358 .or_else(|| default.regions.values().next())
359 });
360 if let Some(sibling) = sibling {
361 state.inherit_from(sibling);
362 }
363 self.get_or_create(account_id)
364 .regions
365 .insert(region.to_string(), state);
366 }
367 self.get_mut(account_id)
368 .and_then(|a| a.regions.get_mut(region))
369 .expect("created above")
370 }
371
372 pub fn regional_get_mut(&mut self, account_id: &str, region: &str) -> Option<&mut T> {
374 self.get_mut(account_id)
375 .and_then(|a| a.get_region_mut(region))
376 }
377
378 pub fn by_arn(&self, arn: &str) -> Option<&T> {
381 let account = fakecloud_aws::arn::account_of(arn)?;
382 let region = fakecloud_aws::arn::region_of(arn)?;
383 self.regional(account, region)
384 }
385
386 pub fn by_arn_mut(&mut self, arn: &str) -> Option<&mut T> {
388 let account = fakecloud_aws::arn::account_of(arn)?.to_string();
389 let region = fakecloud_aws::arn::region_of(arn)?.to_string();
390 self.regional_get_mut(&account, ®ion)
391 }
392
393 pub fn default_regional_mut(&mut self) -> &mut T {
396 let region = self.region().to_string();
397 self.default_mut().region_mut(®ion)
398 }
399
400 pub fn default_regional(&self) -> Option<&T> {
403 self.default_ref().region(self.region())
404 }
405
406 pub fn iter_regional(&self) -> impl Iterator<Item = (&str, &str, &T)> {
408 self.iter()
409 .flat_map(|(account, a)| a.regions().map(move |(region, s)| (account, region, s)))
410 }
411
412 pub fn iter_regional_mut(&mut self) -> impl Iterator<Item = (&str, &str, &mut T)> {
414 self.iter_mut()
415 .flat_map(|(account, a)| a.regions_mut().map(move |(region, s)| (account, region, s)))
416 }
417}
418
419pub trait SplitByRegion: AccountState {
423 fn split_by_region(self, into: &mut RegionalState<Self>);
428}
429
430impl<T: SplitByRegion> RegionalState<T> {
431 pub fn from_legacy(account_id: &str, default_region: &str, endpoint: &str, legacy: T) -> Self {
433 let mut regional = Self::new(account_id, default_region, endpoint);
434 legacy.split_by_region(&mut regional);
435 regional
436 }
437}
438
439impl<T: SplitByRegion> MultiAccountState<T> {
440 pub fn into_regional(self) -> MultiRegionState<T> {
443 let region = self.region.clone();
444 let endpoint = self.endpoint.clone();
445 self.map_into(|account, state| {
446 RegionalState::from_legacy(account, ®ion, &endpoint, state)
447 })
448 }
449}
450
451#[derive(Debug, Clone, Serialize, Deserialize)]
457pub struct RegionalSnapshot<T> {
458 pub schema_version: u32,
459 #[serde(default = "none")]
460 pub accounts: Option<MultiRegionState<T>>,
461 #[serde(default = "none", skip_serializing_if = "Option::is_none")]
462 pub state: Option<RegionalState<T>>,
463}
464
465fn none<T>() -> Option<T> {
466 None
467}
468
469impl<T> RegionalSnapshot<T> {
470 pub fn of(schema_version: u32, accounts: MultiRegionState<T>) -> Self {
472 Self {
473 schema_version,
474 accounts: Some(accounts),
475 state: None,
476 }
477 }
478}
479
480#[derive(Deserialize)]
481struct SnapshotVersionProbe {
482 schema_version: u32,
483}
484
485#[derive(Deserialize)]
486#[serde(bound = "T: serde::de::DeserializeOwned")]
487struct LegacySnapshot<T> {
488 #[serde(default = "none")]
489 accounts: Option<MultiAccountState<T>>,
490 #[serde(default = "none")]
491 state: Option<T>,
492}
493
494pub fn parse_regional_snapshot<T>(
502 bytes: &[u8],
503 current: u32,
504 legacy_single: impl FnOnce(T) -> RegionalState<T>,
505) -> Result<RegionalSnapshot<T>, serde_json::Error>
506where
507 T: SplitByRegion + serde::de::DeserializeOwned,
508{
509 let SnapshotVersionProbe { schema_version } = serde_json::from_slice(bytes)?;
510 if schema_version > current {
511 return Ok(RegionalSnapshot {
512 schema_version,
513 accounts: None,
514 state: None,
515 });
516 }
517 if schema_version == current {
518 return serde_json::from_slice(bytes);
519 }
520 let legacy: LegacySnapshot<T> = serde_json::from_slice(bytes)?;
521 Ok(RegionalSnapshot {
522 schema_version: current,
523 accounts: legacy.accounts.map(MultiAccountState::into_regional),
524 state: legacy.state.map(legacy_single),
525 })
526}
527
528pub fn arn_region_or<'a>(arn: &'a str, default: &'a str) -> &'a str {
531 fakecloud_aws::arn::region_of(arn).unwrap_or(default)
532}
533
534#[cfg(test)]
535mod tests {
536 use super::*;
537
538 #[derive(Debug, Clone, Serialize, Deserialize)]
539 struct TestState {
540 account_id: String,
541 items: Vec<String>,
542 }
543
544 impl AccountState for TestState {
545 fn new_for_account(account_id: &str, _region: &str, _endpoint: &str) -> Self {
546 Self {
547 account_id: account_id.to_string(),
548 items: Vec::new(),
549 }
550 }
551 }
552
553 #[test]
554 fn default_account_exists_on_creation() {
555 let mas: MultiAccountState<TestState> =
556 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
557 assert_eq!(mas.account_count(), 1);
558 assert!(mas.get("111111111111").is_some());
559 }
560
561 #[test]
562 fn get_or_create_makes_new_account() {
563 let mut mas: MultiAccountState<TestState> =
564 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
565 let state = mas.get_or_create("222222222222");
566 assert_eq!(state.account_id, "222222222222");
567 assert_eq!(mas.account_count(), 2);
568 }
569
570 #[test]
571 fn get_returns_none_for_unknown() {
572 let mas: MultiAccountState<TestState> =
573 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
574 assert!(mas.get("999999999999").is_none());
575 }
576
577 #[test]
578 fn reset_clears_all_but_default() {
579 let mut mas: MultiAccountState<TestState> =
580 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
581 mas.get_or_create("222222222222");
582 mas.get_or_create("333333333333");
583 assert_eq!(mas.account_count(), 3);
584 mas.reset();
585 assert_eq!(mas.account_count(), 1);
586 assert!(mas.get("111111111111").is_some());
587 assert!(mas.get("222222222222").is_none());
588 }
589
590 #[test]
591 fn iter_visits_all_accounts() {
592 let mut mas: MultiAccountState<TestState> =
593 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
594 mas.get_or_create("222222222222");
595 let ids: Vec<&str> = mas.iter().map(|(id, _)| id).collect();
596 assert_eq!(ids.len(), 2);
597 assert!(ids.contains(&"111111111111"));
598 assert!(ids.contains(&"222222222222"));
599 }
600
601 impl SplitByRegion for TestState {
602 fn split_by_region(self, into: &mut RegionalState<Self>) {
603 for item in self.items {
604 let region = fakecloud_aws::arn::region_of(&item).map(str::to_string);
606 into.region_or_default_mut(region.as_deref())
607 .items
608 .push(item);
609 }
610 }
611 }
612
613 #[test]
614 fn regional_state_isolates_regions() {
615 let mut mrs: MultiRegionState<TestState> =
616 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
617 mrs.regional_mut("111111111111", "us-east-1")
618 .items
619 .push("q".into());
620 mrs.regional_mut("111111111111", "eu-west-1")
621 .items
622 .push("q".into());
623 assert_eq!(
624 mrs.regional("111111111111", "us-east-1").unwrap().items,
625 ["q"]
626 );
627 assert_eq!(
628 mrs.regional("111111111111", "eu-west-1").unwrap().items,
629 ["q"]
630 );
631 assert!(mrs.regional("111111111111", "ap-south-1").is_none());
632 assert!(mrs.regional("222222222222", "us-east-1").is_none());
633 assert_eq!(mrs.iter_regional().count(), 2);
635 }
636
637 #[test]
638 fn by_arn_resolves_account_and_region() {
639 let mut mrs: MultiRegionState<TestState> =
640 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
641 mrs.regional_mut("222222222222", "eu-west-1")
642 .items
643 .push("x".into());
644 let s = mrs
645 .by_arn("arn:aws:sqs:eu-west-1:222222222222:x")
646 .expect("state");
647 assert_eq!(s.account_id, "222222222222");
648 assert!(mrs.by_arn("arn:aws:sqs:us-east-1:222222222222:x").is_none());
649 assert!(mrs.by_arn("arn:aws:iam::222222222222:role/x").is_none());
650 }
651
652 #[test]
653 fn regional_state_round_trips_through_json() {
654 let mut mrs: MultiRegionState<TestState> =
655 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
656 mrs.regional_mut("111111111111", "eu-west-1")
657 .items
658 .push("q".into());
659 let json = serde_json::to_string(&mrs).unwrap();
660 let back: MultiRegionState<TestState> = serde_json::from_str(&json).unwrap();
661 assert_eq!(
662 back.regional("111111111111", "eu-west-1").unwrap().items,
663 ["q"]
664 );
665 }
666
667 #[test]
668 fn legacy_state_splits_by_arn_region() {
669 let mut legacy: MultiAccountState<TestState> =
670 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
671 let acct = legacy.default_mut();
672 acct.items
673 .push("arn:aws:sqs:eu-west-1:111111111111:a".into());
674 acct.items
675 .push("arn:aws:sqs:us-east-1:111111111111:b".into());
676 acct.items.push("plain".into());
677 let regional = legacy.into_regional();
678 assert_eq!(
679 regional
680 .regional("111111111111", "eu-west-1")
681 .unwrap()
682 .items,
683 ["arn:aws:sqs:eu-west-1:111111111111:a"]
684 );
685 assert_eq!(
686 regional
687 .regional("111111111111", "us-east-1")
688 .unwrap()
689 .items,
690 ["arn:aws:sqs:us-east-1:111111111111:b", "plain"]
691 );
692 assert_eq!(regional.region(), "us-east-1");
693 assert_eq!(regional.default_account_id(), "111111111111");
694 }
695
696 #[test]
697 fn default_regional_mut_uses_server_region() {
698 let mut mrs: MultiRegionState<TestState> =
699 MultiAccountState::new("111111111111", "eu-central-1", "http://localhost:4566");
700 mrs.default_regional_mut().items.push("z".into());
701 assert!(mrs.regional("111111111111", "eu-central-1").is_some());
702 }
703
704 #[test]
705 fn parse_regional_snapshot_migrates_legacy_and_reports_newer() {
706 let mut legacy: MultiAccountState<TestState> =
707 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
708 legacy
709 .default_mut()
710 .items
711 .push("arn:aws:sqs:eu-west-1:111111111111:a".into());
712 let bytes =
713 serde_json::to_vec(&serde_json::json!({"schema_version": 1, "accounts": legacy}))
714 .unwrap();
715 let snap = parse_regional_snapshot::<TestState>(&bytes, 2, |_| unreachable!()).unwrap();
716 assert_eq!(snap.schema_version, 2);
717 let accounts = snap.accounts.unwrap();
718 assert_eq!(
719 accounts
720 .regional("111111111111", "eu-west-1")
721 .unwrap()
722 .items
723 .len(),
724 1
725 );
726
727 let current = serde_json::to_vec(&RegionalSnapshot::of(2, accounts)).unwrap();
728 let again = parse_regional_snapshot::<TestState>(¤t, 2, |_| unreachable!()).unwrap();
729 assert!(again
730 .accounts
731 .unwrap()
732 .regional("111111111111", "eu-west-1")
733 .is_some());
734
735 let single = serde_json::to_vec(&serde_json::json!({
736 "schema_version": 1,
737 "state": {"account_id": "111111111111", "items": ["x"]}
738 }))
739 .unwrap();
740 let snap = parse_regional_snapshot::<TestState>(&single, 2, |s| {
741 RegionalState::from_legacy("111111111111", "us-east-1", "", s)
742 })
743 .unwrap();
744 assert_eq!(
745 snap.state.unwrap().region("us-east-1").unwrap().items,
746 ["x"]
747 );
748
749 let newer = parse_regional_snapshot::<TestState>(
750 br#"{"schema_version": 9}"#,
751 2,
752 |_| unreachable!(),
753 )
754 .unwrap();
755 assert_eq!(newer.schema_version, 9);
756 assert!(newer.accounts.is_none());
757 }
758
759 #[derive(Debug, Clone, Serialize, Deserialize)]
760 struct SharedCacheState {
761 cache: Option<String>,
762 }
763
764 impl AccountState for SharedCacheState {
765 fn new_for_account(_account_id: &str, _region: &str, _endpoint: &str) -> Self {
766 Self { cache: None }
767 }
768
769 fn inherit_from(&mut self, sibling: &Self) {
770 self.cache = sibling.cache.clone();
771 }
772 }
773
774 #[test]
775 fn a_new_accounts_first_region_inherits_from_the_default_account() {
776 let mut mrs: MultiRegionState<SharedCacheState> =
777 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
778 mrs.default_regional_mut().cache = Some("shared".into());
779 assert_eq!(
781 mrs.regional_mut("222222222222", "us-east-1")
782 .cache
783 .as_deref(),
784 Some("shared")
785 );
786 assert_eq!(
787 mrs.regional_mut("333333333333", "eu-west-1")
788 .cache
789 .as_deref(),
790 Some("shared")
791 );
792 mrs.regional_mut("222222222222", "us-east-1").cache = Some("own".into());
794 assert_eq!(
795 mrs.regional_mut("222222222222", "ap-south-1")
796 .cache
797 .as_deref(),
798 Some("own")
799 );
800 assert_eq!(
802 mrs.regional_mut("333333333333", "eu-west-1")
803 .cache
804 .as_deref(),
805 Some("shared")
806 );
807 }
808}