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 reset_account(&mut self, account_id: &str) {
166 if let Some(state) = self.accounts.get_mut(account_id) {
167 *state = T::new_for_account(account_id, &self.region, &self.endpoint);
168 }
169 }
170
171 pub fn find_account<F>(&self, predicate: F) -> Option<&str>
175 where
176 F: Fn(&T) -> bool,
177 {
178 self.accounts
179 .iter()
180 .find(|(_, v)| predicate(v))
181 .map(|(k, _)| k.as_str())
182 }
183
184 pub fn account_count(&self) -> usize {
186 self.accounts.len()
187 }
188
189 pub fn region(&self) -> &str {
191 &self.region
192 }
193
194 pub fn endpoint(&self) -> &str {
196 &self.endpoint
197 }
198}
199
200#[derive(Debug, Clone, Serialize, Deserialize)]
213pub struct RegionalState<T> {
214 account_id: String,
215 default_region: String,
217 endpoint: String,
218 #[serde(default = "BTreeMap::new")]
220 regions: BTreeMap<String, T>,
221}
222
223impl<T> RegionalState<T> {
224 pub fn new(account_id: &str, default_region: &str, endpoint: &str) -> Self {
226 Self {
227 account_id: account_id.to_string(),
228 default_region: default_region.to_string(),
229 endpoint: endpoint.to_string(),
230 regions: BTreeMap::new(),
231 }
232 }
233
234 pub fn account_id(&self) -> &str {
236 &self.account_id
237 }
238
239 pub fn default_region(&self) -> &str {
241 &self.default_region
242 }
243
244 pub fn endpoint(&self) -> &str {
246 &self.endpoint
247 }
248
249 pub fn region(&self, region: &str) -> Option<&T> {
251 self.regions.get(region)
252 }
253
254 pub fn get_region_mut(&mut self, region: &str) -> Option<&mut T> {
256 self.regions.get_mut(region)
257 }
258
259 pub fn regions(&self) -> impl Iterator<Item = (&str, &T)> {
261 self.regions.iter().map(|(k, v)| (k.as_str(), v))
262 }
263
264 pub fn regions_mut(&mut self) -> impl Iterator<Item = (&str, &mut T)> {
266 self.regions.iter_mut().map(|(k, v)| (k.as_str(), v))
267 }
268
269 pub fn insert_region(&mut self, region: &str, state: T) -> Option<T> {
271 self.regions.insert(region.to_string(), state)
272 }
273
274 pub fn clear(&mut self) {
276 self.regions.clear();
277 }
278
279 pub fn map<U>(&self, mut f: impl FnMut(&T) -> U) -> RegionalState<U> {
281 RegionalState {
282 account_id: self.account_id.clone(),
283 default_region: self.default_region.clone(),
284 endpoint: self.endpoint.clone(),
285 regions: self
286 .regions
287 .iter()
288 .map(|(k, v)| (k.clone(), f(v)))
289 .collect(),
290 }
291 }
292}
293
294impl<T: AccountState> RegionalState<T> {
295 pub fn region_mut(&mut self, region: &str) -> &mut T {
299 if !self.regions.contains_key(region) {
300 let mut state = T::new_for_account(&self.account_id, region, &self.endpoint);
301 if let Some(sibling) = self.regions.values().next() {
302 state.inherit_from(sibling);
303 }
304 self.regions.insert(region.to_string(), state);
305 }
306 self.regions.get_mut(region).expect("inserted above")
307 }
308
309 pub fn region_or_default_mut(&mut self, region: Option<&str>) -> &mut T {
313 let region = region
314 .filter(|r| !r.is_empty())
315 .map(str::to_string)
316 .unwrap_or_else(|| self.default_region.clone());
317 self.region_mut(®ion)
318 }
319}
320
321impl<T: AccountState> AccountState for RegionalState<T> {
322 fn new_for_account(account_id: &str, region: &str, endpoint: &str) -> Self {
323 Self::new(account_id, region, endpoint)
324 }
325
326 fn inherit_from(&mut self, sibling: &Self) {
327 if let Some(shared) = sibling.regions.values().next() {
328 for state in self.regions.values_mut() {
329 state.inherit_from(shared);
330 }
331 }
332 }
333}
334
335pub type MultiRegionState<T> = MultiAccountState<RegionalState<T>>;
337
338impl<T: AccountState> MultiAccountState<RegionalState<T>> {
339 pub fn regional(&self, account_id: &str, region: &str) -> Option<&T> {
343 self.get(account_id).and_then(|a| a.region(region))
344 }
345
346 pub fn regional_mut(&mut self, account_id: &str, region: &str) -> &mut T {
354 let exists = self
355 .get(account_id)
356 .is_some_and(|a| a.region(region).is_some());
357 if !exists {
358 let mut state = T::new_for_account(account_id, region, &self.endpoint);
359 let sibling = self
360 .get(account_id)
361 .and_then(|a| a.regions.values().next())
362 .or_else(|| {
363 let default = self.get(&self.default_account_id)?;
364 default
365 .region(region)
366 .or_else(|| default.regions.values().next())
367 });
368 if let Some(sibling) = sibling {
369 state.inherit_from(sibling);
370 }
371 self.get_or_create(account_id)
372 .regions
373 .insert(region.to_string(), state);
374 }
375 self.get_mut(account_id)
376 .and_then(|a| a.regions.get_mut(region))
377 .expect("created above")
378 }
379
380 pub fn regional_get_mut(&mut self, account_id: &str, region: &str) -> Option<&mut T> {
382 self.get_mut(account_id)
383 .and_then(|a| a.get_region_mut(region))
384 }
385
386 pub fn by_arn(&self, arn: &str) -> Option<&T> {
389 let account = fakecloud_aws::arn::account_of(arn)?;
390 let region = fakecloud_aws::arn::region_of(arn)?;
391 self.regional(account, region)
392 }
393
394 pub fn by_arn_mut(&mut self, arn: &str) -> Option<&mut T> {
396 let account = fakecloud_aws::arn::account_of(arn)?.to_string();
397 let region = fakecloud_aws::arn::region_of(arn)?.to_string();
398 self.regional_get_mut(&account, ®ion)
399 }
400
401 pub fn default_regional_mut(&mut self) -> &mut T {
404 let region = self.region().to_string();
405 self.default_mut().region_mut(®ion)
406 }
407
408 pub fn default_regional(&self) -> Option<&T> {
411 self.default_ref().region(self.region())
412 }
413
414 pub fn iter_regional(&self) -> impl Iterator<Item = (&str, &str, &T)> {
416 self.iter()
417 .flat_map(|(account, a)| a.regions().map(move |(region, s)| (account, region, s)))
418 }
419
420 pub fn iter_regional_mut(&mut self) -> impl Iterator<Item = (&str, &str, &mut T)> {
422 self.iter_mut()
423 .flat_map(|(account, a)| a.regions_mut().map(move |(region, s)| (account, region, s)))
424 }
425}
426
427pub trait SplitByRegion: AccountState {
431 fn split_by_region(self, into: &mut RegionalState<Self>);
436}
437
438impl<T: SplitByRegion> RegionalState<T> {
439 pub fn from_legacy(account_id: &str, default_region: &str, endpoint: &str, legacy: T) -> Self {
441 let mut regional = Self::new(account_id, default_region, endpoint);
442 legacy.split_by_region(&mut regional);
443 regional
444 }
445}
446
447impl<T: SplitByRegion> MultiAccountState<T> {
448 pub fn into_regional(self) -> MultiRegionState<T> {
451 let region = self.region.clone();
452 let endpoint = self.endpoint.clone();
453 self.map_into(|account, state| {
454 RegionalState::from_legacy(account, ®ion, &endpoint, state)
455 })
456 }
457}
458
459#[derive(Debug, Clone, Serialize, Deserialize)]
465pub struct RegionalSnapshot<T> {
466 pub schema_version: u32,
467 #[serde(default = "none")]
468 pub accounts: Option<MultiRegionState<T>>,
469 #[serde(default = "none", skip_serializing_if = "Option::is_none")]
470 pub state: Option<RegionalState<T>>,
471}
472
473fn none<T>() -> Option<T> {
474 None
475}
476
477impl<T> RegionalSnapshot<T> {
478 pub fn of(schema_version: u32, accounts: MultiRegionState<T>) -> Self {
480 Self {
481 schema_version,
482 accounts: Some(accounts),
483 state: None,
484 }
485 }
486}
487
488#[derive(Deserialize)]
489struct SnapshotVersionProbe {
490 schema_version: u32,
491}
492
493#[derive(Deserialize)]
494#[serde(bound = "T: serde::de::DeserializeOwned")]
495struct LegacySnapshot<T> {
496 #[serde(default = "none")]
497 accounts: Option<MultiAccountState<T>>,
498 #[serde(default = "none")]
499 state: Option<T>,
500}
501
502pub fn parse_regional_snapshot<T>(
510 bytes: &[u8],
511 current: u32,
512 legacy_single: impl FnOnce(T) -> RegionalState<T>,
513) -> Result<RegionalSnapshot<T>, serde_json::Error>
514where
515 T: SplitByRegion + serde::de::DeserializeOwned,
516{
517 let SnapshotVersionProbe { schema_version } = serde_json::from_slice(bytes)?;
518 if schema_version > current {
519 return Ok(RegionalSnapshot {
520 schema_version,
521 accounts: None,
522 state: None,
523 });
524 }
525 if schema_version == current {
526 return serde_json::from_slice(bytes);
527 }
528 let legacy: LegacySnapshot<T> = serde_json::from_slice(bytes)?;
529 Ok(RegionalSnapshot {
530 schema_version: current,
531 accounts: legacy.accounts.map(MultiAccountState::into_regional),
532 state: legacy.state.map(legacy_single),
533 })
534}
535
536pub fn arn_region_or<'a>(arn: &'a str, default: &'a str) -> &'a str {
539 fakecloud_aws::arn::region_of(arn).unwrap_or(default)
540}
541
542#[cfg(test)]
543mod tests {
544 use super::*;
545
546 #[derive(Debug, Clone, Serialize, Deserialize)]
547 struct TestState {
548 account_id: String,
549 items: Vec<String>,
550 }
551
552 impl AccountState for TestState {
553 fn new_for_account(account_id: &str, _region: &str, _endpoint: &str) -> Self {
554 Self {
555 account_id: account_id.to_string(),
556 items: Vec::new(),
557 }
558 }
559 }
560
561 #[test]
562 fn default_account_exists_on_creation() {
563 let mas: MultiAccountState<TestState> =
564 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
565 assert_eq!(mas.account_count(), 1);
566 assert!(mas.get("111111111111").is_some());
567 }
568
569 #[test]
570 fn get_or_create_makes_new_account() {
571 let mut mas: MultiAccountState<TestState> =
572 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
573 let state = mas.get_or_create("222222222222");
574 assert_eq!(state.account_id, "222222222222");
575 assert_eq!(mas.account_count(), 2);
576 }
577
578 #[test]
579 fn get_returns_none_for_unknown() {
580 let mas: MultiAccountState<TestState> =
581 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
582 assert!(mas.get("999999999999").is_none());
583 }
584
585 #[test]
586 fn reset_clears_all_but_default() {
587 let mut mas: MultiAccountState<TestState> =
588 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
589 mas.get_or_create("222222222222");
590 mas.get_or_create("333333333333");
591 assert_eq!(mas.account_count(), 3);
592 mas.reset();
593 assert_eq!(mas.account_count(), 1);
594 assert!(mas.get("111111111111").is_some());
595 assert!(mas.get("222222222222").is_none());
596 }
597
598 #[test]
599 fn reset_account_clears_only_that_account() {
600 let mut mas: MultiAccountState<TestState> =
601 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
602 mas.default_mut().items.push("a".into());
603 mas.get_or_create("222222222222").items.push("b".into());
604 mas.reset_account("222222222222");
605 assert_eq!(mas.account_count(), 2);
606 assert!(mas.get("222222222222").unwrap().items.is_empty());
607 assert_eq!(
608 mas.get("111111111111").unwrap().items,
609 vec!["a".to_string()]
610 );
611 mas.reset_account("333333333333");
613 assert!(mas.get("333333333333").is_none());
614 }
615
616 #[test]
617 fn iter_visits_all_accounts() {
618 let mut mas: MultiAccountState<TestState> =
619 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
620 mas.get_or_create("222222222222");
621 let ids: Vec<&str> = mas.iter().map(|(id, _)| id).collect();
622 assert_eq!(ids.len(), 2);
623 assert!(ids.contains(&"111111111111"));
624 assert!(ids.contains(&"222222222222"));
625 }
626
627 impl SplitByRegion for TestState {
628 fn split_by_region(self, into: &mut RegionalState<Self>) {
629 for item in self.items {
630 let region = fakecloud_aws::arn::region_of(&item).map(str::to_string);
632 into.region_or_default_mut(region.as_deref())
633 .items
634 .push(item);
635 }
636 }
637 }
638
639 #[test]
640 fn regional_state_isolates_regions() {
641 let mut mrs: MultiRegionState<TestState> =
642 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
643 mrs.regional_mut("111111111111", "us-east-1")
644 .items
645 .push("q".into());
646 mrs.regional_mut("111111111111", "eu-west-1")
647 .items
648 .push("q".into());
649 assert_eq!(
650 mrs.regional("111111111111", "us-east-1").unwrap().items,
651 ["q"]
652 );
653 assert_eq!(
654 mrs.regional("111111111111", "eu-west-1").unwrap().items,
655 ["q"]
656 );
657 assert!(mrs.regional("111111111111", "ap-south-1").is_none());
658 assert!(mrs.regional("222222222222", "us-east-1").is_none());
659 assert_eq!(mrs.iter_regional().count(), 2);
661 }
662
663 #[test]
664 fn by_arn_resolves_account_and_region() {
665 let mut mrs: MultiRegionState<TestState> =
666 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
667 mrs.regional_mut("222222222222", "eu-west-1")
668 .items
669 .push("x".into());
670 let s = mrs
671 .by_arn("arn:aws:sqs:eu-west-1:222222222222:x")
672 .expect("state");
673 assert_eq!(s.account_id, "222222222222");
674 assert!(mrs.by_arn("arn:aws:sqs:us-east-1:222222222222:x").is_none());
675 assert!(mrs.by_arn("arn:aws:iam::222222222222:role/x").is_none());
676 }
677
678 #[test]
679 fn regional_state_round_trips_through_json() {
680 let mut mrs: MultiRegionState<TestState> =
681 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
682 mrs.regional_mut("111111111111", "eu-west-1")
683 .items
684 .push("q".into());
685 let json = serde_json::to_string(&mrs).unwrap();
686 let back: MultiRegionState<TestState> = serde_json::from_str(&json).unwrap();
687 assert_eq!(
688 back.regional("111111111111", "eu-west-1").unwrap().items,
689 ["q"]
690 );
691 }
692
693 #[test]
694 fn legacy_state_splits_by_arn_region() {
695 let mut legacy: MultiAccountState<TestState> =
696 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
697 let acct = legacy.default_mut();
698 acct.items
699 .push("arn:aws:sqs:eu-west-1:111111111111:a".into());
700 acct.items
701 .push("arn:aws:sqs:us-east-1:111111111111:b".into());
702 acct.items.push("plain".into());
703 let regional = legacy.into_regional();
704 assert_eq!(
705 regional
706 .regional("111111111111", "eu-west-1")
707 .unwrap()
708 .items,
709 ["arn:aws:sqs:eu-west-1:111111111111:a"]
710 );
711 assert_eq!(
712 regional
713 .regional("111111111111", "us-east-1")
714 .unwrap()
715 .items,
716 ["arn:aws:sqs:us-east-1:111111111111:b", "plain"]
717 );
718 assert_eq!(regional.region(), "us-east-1");
719 assert_eq!(regional.default_account_id(), "111111111111");
720 }
721
722 #[test]
723 fn default_regional_mut_uses_server_region() {
724 let mut mrs: MultiRegionState<TestState> =
725 MultiAccountState::new("111111111111", "eu-central-1", "http://localhost:4566");
726 mrs.default_regional_mut().items.push("z".into());
727 assert!(mrs.regional("111111111111", "eu-central-1").is_some());
728 }
729
730 #[test]
731 fn parse_regional_snapshot_migrates_legacy_and_reports_newer() {
732 let mut legacy: MultiAccountState<TestState> =
733 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
734 legacy
735 .default_mut()
736 .items
737 .push("arn:aws:sqs:eu-west-1:111111111111:a".into());
738 let bytes =
739 serde_json::to_vec(&serde_json::json!({"schema_version": 1, "accounts": legacy}))
740 .unwrap();
741 let snap = parse_regional_snapshot::<TestState>(&bytes, 2, |_| unreachable!()).unwrap();
742 assert_eq!(snap.schema_version, 2);
743 let accounts = snap.accounts.unwrap();
744 assert_eq!(
745 accounts
746 .regional("111111111111", "eu-west-1")
747 .unwrap()
748 .items
749 .len(),
750 1
751 );
752
753 let current = serde_json::to_vec(&RegionalSnapshot::of(2, accounts)).unwrap();
754 let again = parse_regional_snapshot::<TestState>(¤t, 2, |_| unreachable!()).unwrap();
755 assert!(again
756 .accounts
757 .unwrap()
758 .regional("111111111111", "eu-west-1")
759 .is_some());
760
761 let single = serde_json::to_vec(&serde_json::json!({
762 "schema_version": 1,
763 "state": {"account_id": "111111111111", "items": ["x"]}
764 }))
765 .unwrap();
766 let snap = parse_regional_snapshot::<TestState>(&single, 2, |s| {
767 RegionalState::from_legacy("111111111111", "us-east-1", "", s)
768 })
769 .unwrap();
770 assert_eq!(
771 snap.state.unwrap().region("us-east-1").unwrap().items,
772 ["x"]
773 );
774
775 let newer = parse_regional_snapshot::<TestState>(
776 br#"{"schema_version": 9}"#,
777 2,
778 |_| unreachable!(),
779 )
780 .unwrap();
781 assert_eq!(newer.schema_version, 9);
782 assert!(newer.accounts.is_none());
783 }
784
785 #[derive(Debug, Clone, Serialize, Deserialize)]
786 struct SharedCacheState {
787 cache: Option<String>,
788 }
789
790 impl AccountState for SharedCacheState {
791 fn new_for_account(_account_id: &str, _region: &str, _endpoint: &str) -> Self {
792 Self { cache: None }
793 }
794
795 fn inherit_from(&mut self, sibling: &Self) {
796 self.cache = sibling.cache.clone();
797 }
798 }
799
800 #[test]
801 fn a_new_accounts_first_region_inherits_from_the_default_account() {
802 let mut mrs: MultiRegionState<SharedCacheState> =
803 MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
804 mrs.default_regional_mut().cache = Some("shared".into());
805 assert_eq!(
807 mrs.regional_mut("222222222222", "us-east-1")
808 .cache
809 .as_deref(),
810 Some("shared")
811 );
812 assert_eq!(
813 mrs.regional_mut("333333333333", "eu-west-1")
814 .cache
815 .as_deref(),
816 Some("shared")
817 );
818 mrs.regional_mut("222222222222", "us-east-1").cache = Some("own".into());
820 assert_eq!(
821 mrs.regional_mut("222222222222", "ap-south-1")
822 .cache
823 .as_deref(),
824 Some("own")
825 );
826 assert_eq!(
828 mrs.regional_mut("333333333333", "eu-west-1")
829 .cache
830 .as_deref(),
831 Some("shared")
832 );
833 }
834}