use std::collections::{BTreeMap, HashMap};
use serde::{Deserialize, Serialize};
pub trait AccountState: Sized {
fn new_for_account(account_id: &str, region: &str, endpoint: &str) -> Self;
fn inherit_from(&mut self, _sibling: &Self) {}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MultiAccountState<T> {
default_account_id: String,
region: String,
endpoint: String,
accounts: HashMap<String, T>,
}
impl<T: AccountState> MultiAccountState<T> {
pub fn new(default_account_id: &str, region: &str, endpoint: &str) -> Self {
let mut accounts = HashMap::new();
accounts.insert(
default_account_id.to_string(),
T::new_for_account(default_account_id, region, endpoint),
);
Self {
default_account_id: default_account_id.to_string(),
region: region.to_string(),
endpoint: endpoint.to_string(),
accounts,
}
}
pub fn map<U>(&self, mut f: impl FnMut(&T) -> U) -> MultiAccountState<U> {
MultiAccountState {
default_account_id: self.default_account_id.clone(),
region: self.region.clone(),
endpoint: self.endpoint.clone(),
accounts: self
.accounts
.iter()
.map(|(k, v)| (k.clone(), f(v)))
.collect(),
}
}
pub fn map_into<U>(self, mut f: impl FnMut(&str, T) -> U) -> MultiAccountState<U> {
MultiAccountState {
default_account_id: self.default_account_id,
region: self.region,
endpoint: self.endpoint,
accounts: self
.accounts
.into_iter()
.map(|(k, v)| {
let mapped = f(&k, v);
(k, mapped)
})
.collect(),
}
}
pub fn get_or_create(&mut self, account_id: &str) -> &mut T {
if !self.accounts.contains_key(account_id) {
let mut state = T::new_for_account(account_id, &self.region, &self.endpoint);
if let Some(sibling) = self.accounts.get(&self.default_account_id) {
state.inherit_from(sibling);
}
self.accounts.insert(account_id.to_string(), state);
}
self.accounts.get_mut(account_id).unwrap()
}
pub fn get_or_create_with<F>(&mut self, account_id: &str, init: F) -> &mut T
where
F: FnOnce(&mut T),
{
if !self.accounts.contains_key(account_id) {
let mut state = T::new_for_account(account_id, &self.region, &self.endpoint);
init(&mut state);
self.accounts.insert(account_id.to_string(), state);
}
self.accounts.get_mut(account_id).unwrap()
}
pub fn get(&self, account_id: &str) -> Option<&T> {
self.accounts.get(account_id)
}
pub fn get_mut(&mut self, account_id: &str) -> Option<&mut T> {
self.accounts.get_mut(account_id)
}
pub fn iter(&self) -> impl Iterator<Item = (&str, &T)> {
self.accounts.iter().map(|(k, v)| (k.as_str(), v))
}
pub fn iter_mut(&mut self) -> impl Iterator<Item = (&str, &mut T)> {
self.accounts.iter_mut().map(|(k, v)| (k.as_str(), v))
}
pub fn default_account_id(&self) -> &str {
&self.default_account_id
}
pub fn default_mut(&mut self) -> &mut T {
self.accounts.get_mut(&self.default_account_id).unwrap()
}
pub fn default_ref(&self) -> &T {
self.accounts.get(&self.default_account_id).unwrap()
}
pub fn reset(&mut self) {
self.accounts.clear();
self.accounts.insert(
self.default_account_id.clone(),
T::new_for_account(&self.default_account_id, &self.region, &self.endpoint),
);
}
pub fn find_account<F>(&self, predicate: F) -> Option<&str>
where
F: Fn(&T) -> bool,
{
self.accounts
.iter()
.find(|(_, v)| predicate(v))
.map(|(k, _)| k.as_str())
}
pub fn account_count(&self) -> usize {
self.accounts.len()
}
pub fn region(&self) -> &str {
&self.region
}
pub fn endpoint(&self) -> &str {
&self.endpoint
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RegionalState<T> {
account_id: String,
default_region: String,
endpoint: String,
#[serde(default = "BTreeMap::new")]
regions: BTreeMap<String, T>,
}
impl<T> RegionalState<T> {
pub fn new(account_id: &str, default_region: &str, endpoint: &str) -> Self {
Self {
account_id: account_id.to_string(),
default_region: default_region.to_string(),
endpoint: endpoint.to_string(),
regions: BTreeMap::new(),
}
}
pub fn account_id(&self) -> &str {
&self.account_id
}
pub fn default_region(&self) -> &str {
&self.default_region
}
pub fn endpoint(&self) -> &str {
&self.endpoint
}
pub fn region(&self, region: &str) -> Option<&T> {
self.regions.get(region)
}
pub fn get_region_mut(&mut self, region: &str) -> Option<&mut T> {
self.regions.get_mut(region)
}
pub fn regions(&self) -> impl Iterator<Item = (&str, &T)> {
self.regions.iter().map(|(k, v)| (k.as_str(), v))
}
pub fn regions_mut(&mut self) -> impl Iterator<Item = (&str, &mut T)> {
self.regions.iter_mut().map(|(k, v)| (k.as_str(), v))
}
pub fn insert_region(&mut self, region: &str, state: T) -> Option<T> {
self.regions.insert(region.to_string(), state)
}
pub fn clear(&mut self) {
self.regions.clear();
}
pub fn map<U>(&self, mut f: impl FnMut(&T) -> U) -> RegionalState<U> {
RegionalState {
account_id: self.account_id.clone(),
default_region: self.default_region.clone(),
endpoint: self.endpoint.clone(),
regions: self
.regions
.iter()
.map(|(k, v)| (k.clone(), f(v)))
.collect(),
}
}
}
impl<T: AccountState> RegionalState<T> {
pub fn region_mut(&mut self, region: &str) -> &mut T {
if !self.regions.contains_key(region) {
let mut state = T::new_for_account(&self.account_id, region, &self.endpoint);
if let Some(sibling) = self.regions.values().next() {
state.inherit_from(sibling);
}
self.regions.insert(region.to_string(), state);
}
self.regions.get_mut(region).expect("inserted above")
}
pub fn region_or_default_mut(&mut self, region: Option<&str>) -> &mut T {
let region = region
.filter(|r| !r.is_empty())
.map(str::to_string)
.unwrap_or_else(|| self.default_region.clone());
self.region_mut(®ion)
}
}
impl<T: AccountState> AccountState for RegionalState<T> {
fn new_for_account(account_id: &str, region: &str, endpoint: &str) -> Self {
Self::new(account_id, region, endpoint)
}
fn inherit_from(&mut self, sibling: &Self) {
if let Some(shared) = sibling.regions.values().next() {
for state in self.regions.values_mut() {
state.inherit_from(shared);
}
}
}
}
pub type MultiRegionState<T> = MultiAccountState<RegionalState<T>>;
impl<T: AccountState> MultiAccountState<RegionalState<T>> {
pub fn regional(&self, account_id: &str, region: &str) -> Option<&T> {
self.get(account_id).and_then(|a| a.region(region))
}
pub fn regional_mut(&mut self, account_id: &str, region: &str) -> &mut T {
let exists = self
.get(account_id)
.is_some_and(|a| a.region(region).is_some());
if !exists {
let mut state = T::new_for_account(account_id, region, &self.endpoint);
let sibling = self
.get(account_id)
.and_then(|a| a.regions.values().next())
.or_else(|| {
let default = self.get(&self.default_account_id)?;
default
.region(region)
.or_else(|| default.regions.values().next())
});
if let Some(sibling) = sibling {
state.inherit_from(sibling);
}
self.get_or_create(account_id)
.regions
.insert(region.to_string(), state);
}
self.get_mut(account_id)
.and_then(|a| a.regions.get_mut(region))
.expect("created above")
}
pub fn regional_get_mut(&mut self, account_id: &str, region: &str) -> Option<&mut T> {
self.get_mut(account_id)
.and_then(|a| a.get_region_mut(region))
}
pub fn by_arn(&self, arn: &str) -> Option<&T> {
let account = fakecloud_aws::arn::account_of(arn)?;
let region = fakecloud_aws::arn::region_of(arn)?;
self.regional(account, region)
}
pub fn by_arn_mut(&mut self, arn: &str) -> Option<&mut T> {
let account = fakecloud_aws::arn::account_of(arn)?.to_string();
let region = fakecloud_aws::arn::region_of(arn)?.to_string();
self.regional_get_mut(&account, ®ion)
}
pub fn default_regional_mut(&mut self) -> &mut T {
let region = self.region().to_string();
self.default_mut().region_mut(®ion)
}
pub fn default_regional(&self) -> Option<&T> {
self.default_ref().region(self.region())
}
pub fn iter_regional(&self) -> impl Iterator<Item = (&str, &str, &T)> {
self.iter()
.flat_map(|(account, a)| a.regions().map(move |(region, s)| (account, region, s)))
}
pub fn iter_regional_mut(&mut self) -> impl Iterator<Item = (&str, &str, &mut T)> {
self.iter_mut()
.flat_map(|(account, a)| a.regions_mut().map(move |(region, s)| (account, region, s)))
}
}
pub trait SplitByRegion: AccountState {
fn split_by_region(self, into: &mut RegionalState<Self>);
}
impl<T: SplitByRegion> RegionalState<T> {
pub fn from_legacy(account_id: &str, default_region: &str, endpoint: &str, legacy: T) -> Self {
let mut regional = Self::new(account_id, default_region, endpoint);
legacy.split_by_region(&mut regional);
regional
}
}
impl<T: SplitByRegion> MultiAccountState<T> {
pub fn into_regional(self) -> MultiRegionState<T> {
let region = self.region.clone();
let endpoint = self.endpoint.clone();
self.map_into(|account, state| {
RegionalState::from_legacy(account, ®ion, &endpoint, state)
})
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RegionalSnapshot<T> {
pub schema_version: u32,
#[serde(default = "none")]
pub accounts: Option<MultiRegionState<T>>,
#[serde(default = "none", skip_serializing_if = "Option::is_none")]
pub state: Option<RegionalState<T>>,
}
fn none<T>() -> Option<T> {
None
}
impl<T> RegionalSnapshot<T> {
pub fn of(schema_version: u32, accounts: MultiRegionState<T>) -> Self {
Self {
schema_version,
accounts: Some(accounts),
state: None,
}
}
}
#[derive(Deserialize)]
struct SnapshotVersionProbe {
schema_version: u32,
}
#[derive(Deserialize)]
#[serde(bound = "T: serde::de::DeserializeOwned")]
struct LegacySnapshot<T> {
#[serde(default = "none")]
accounts: Option<MultiAccountState<T>>,
#[serde(default = "none")]
state: Option<T>,
}
pub fn parse_regional_snapshot<T>(
bytes: &[u8],
current: u32,
legacy_single: impl FnOnce(T) -> RegionalState<T>,
) -> Result<RegionalSnapshot<T>, serde_json::Error>
where
T: SplitByRegion + serde::de::DeserializeOwned,
{
let SnapshotVersionProbe { schema_version } = serde_json::from_slice(bytes)?;
if schema_version > current {
return Ok(RegionalSnapshot {
schema_version,
accounts: None,
state: None,
});
}
if schema_version == current {
return serde_json::from_slice(bytes);
}
let legacy: LegacySnapshot<T> = serde_json::from_slice(bytes)?;
Ok(RegionalSnapshot {
schema_version: current,
accounts: legacy.accounts.map(MultiAccountState::into_regional),
state: legacy.state.map(legacy_single),
})
}
pub fn arn_region_or<'a>(arn: &'a str, default: &'a str) -> &'a str {
fakecloud_aws::arn::region_of(arn).unwrap_or(default)
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Debug, Clone, Serialize, Deserialize)]
struct TestState {
account_id: String,
items: Vec<String>,
}
impl AccountState for TestState {
fn new_for_account(account_id: &str, _region: &str, _endpoint: &str) -> Self {
Self {
account_id: account_id.to_string(),
items: Vec::new(),
}
}
}
#[test]
fn default_account_exists_on_creation() {
let mas: MultiAccountState<TestState> =
MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
assert_eq!(mas.account_count(), 1);
assert!(mas.get("111111111111").is_some());
}
#[test]
fn get_or_create_makes_new_account() {
let mut mas: MultiAccountState<TestState> =
MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
let state = mas.get_or_create("222222222222");
assert_eq!(state.account_id, "222222222222");
assert_eq!(mas.account_count(), 2);
}
#[test]
fn get_returns_none_for_unknown() {
let mas: MultiAccountState<TestState> =
MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
assert!(mas.get("999999999999").is_none());
}
#[test]
fn reset_clears_all_but_default() {
let mut mas: MultiAccountState<TestState> =
MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
mas.get_or_create("222222222222");
mas.get_or_create("333333333333");
assert_eq!(mas.account_count(), 3);
mas.reset();
assert_eq!(mas.account_count(), 1);
assert!(mas.get("111111111111").is_some());
assert!(mas.get("222222222222").is_none());
}
#[test]
fn iter_visits_all_accounts() {
let mut mas: MultiAccountState<TestState> =
MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
mas.get_or_create("222222222222");
let ids: Vec<&str> = mas.iter().map(|(id, _)| id).collect();
assert_eq!(ids.len(), 2);
assert!(ids.contains(&"111111111111"));
assert!(ids.contains(&"222222222222"));
}
impl SplitByRegion for TestState {
fn split_by_region(self, into: &mut RegionalState<Self>) {
for item in self.items {
let region = fakecloud_aws::arn::region_of(&item).map(str::to_string);
into.region_or_default_mut(region.as_deref())
.items
.push(item);
}
}
}
#[test]
fn regional_state_isolates_regions() {
let mut mrs: MultiRegionState<TestState> =
MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
mrs.regional_mut("111111111111", "us-east-1")
.items
.push("q".into());
mrs.regional_mut("111111111111", "eu-west-1")
.items
.push("q".into());
assert_eq!(
mrs.regional("111111111111", "us-east-1").unwrap().items,
["q"]
);
assert_eq!(
mrs.regional("111111111111", "eu-west-1").unwrap().items,
["q"]
);
assert!(mrs.regional("111111111111", "ap-south-1").is_none());
assert!(mrs.regional("222222222222", "us-east-1").is_none());
assert_eq!(mrs.iter_regional().count(), 2);
}
#[test]
fn by_arn_resolves_account_and_region() {
let mut mrs: MultiRegionState<TestState> =
MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
mrs.regional_mut("222222222222", "eu-west-1")
.items
.push("x".into());
let s = mrs
.by_arn("arn:aws:sqs:eu-west-1:222222222222:x")
.expect("state");
assert_eq!(s.account_id, "222222222222");
assert!(mrs.by_arn("arn:aws:sqs:us-east-1:222222222222:x").is_none());
assert!(mrs.by_arn("arn:aws:iam::222222222222:role/x").is_none());
}
#[test]
fn regional_state_round_trips_through_json() {
let mut mrs: MultiRegionState<TestState> =
MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
mrs.regional_mut("111111111111", "eu-west-1")
.items
.push("q".into());
let json = serde_json::to_string(&mrs).unwrap();
let back: MultiRegionState<TestState> = serde_json::from_str(&json).unwrap();
assert_eq!(
back.regional("111111111111", "eu-west-1").unwrap().items,
["q"]
);
}
#[test]
fn legacy_state_splits_by_arn_region() {
let mut legacy: MultiAccountState<TestState> =
MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
let acct = legacy.default_mut();
acct.items
.push("arn:aws:sqs:eu-west-1:111111111111:a".into());
acct.items
.push("arn:aws:sqs:us-east-1:111111111111:b".into());
acct.items.push("plain".into());
let regional = legacy.into_regional();
assert_eq!(
regional
.regional("111111111111", "eu-west-1")
.unwrap()
.items,
["arn:aws:sqs:eu-west-1:111111111111:a"]
);
assert_eq!(
regional
.regional("111111111111", "us-east-1")
.unwrap()
.items,
["arn:aws:sqs:us-east-1:111111111111:b", "plain"]
);
assert_eq!(regional.region(), "us-east-1");
assert_eq!(regional.default_account_id(), "111111111111");
}
#[test]
fn default_regional_mut_uses_server_region() {
let mut mrs: MultiRegionState<TestState> =
MultiAccountState::new("111111111111", "eu-central-1", "http://localhost:4566");
mrs.default_regional_mut().items.push("z".into());
assert!(mrs.regional("111111111111", "eu-central-1").is_some());
}
#[test]
fn parse_regional_snapshot_migrates_legacy_and_reports_newer() {
let mut legacy: MultiAccountState<TestState> =
MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
legacy
.default_mut()
.items
.push("arn:aws:sqs:eu-west-1:111111111111:a".into());
let bytes =
serde_json::to_vec(&serde_json::json!({"schema_version": 1, "accounts": legacy}))
.unwrap();
let snap = parse_regional_snapshot::<TestState>(&bytes, 2, |_| unreachable!()).unwrap();
assert_eq!(snap.schema_version, 2);
let accounts = snap.accounts.unwrap();
assert_eq!(
accounts
.regional("111111111111", "eu-west-1")
.unwrap()
.items
.len(),
1
);
let current = serde_json::to_vec(&RegionalSnapshot::of(2, accounts)).unwrap();
let again = parse_regional_snapshot::<TestState>(¤t, 2, |_| unreachable!()).unwrap();
assert!(again
.accounts
.unwrap()
.regional("111111111111", "eu-west-1")
.is_some());
let single = serde_json::to_vec(&serde_json::json!({
"schema_version": 1,
"state": {"account_id": "111111111111", "items": ["x"]}
}))
.unwrap();
let snap = parse_regional_snapshot::<TestState>(&single, 2, |s| {
RegionalState::from_legacy("111111111111", "us-east-1", "", s)
})
.unwrap();
assert_eq!(
snap.state.unwrap().region("us-east-1").unwrap().items,
["x"]
);
let newer = parse_regional_snapshot::<TestState>(
br#"{"schema_version": 9}"#,
2,
|_| unreachable!(),
)
.unwrap();
assert_eq!(newer.schema_version, 9);
assert!(newer.accounts.is_none());
}
#[derive(Debug, Clone, Serialize, Deserialize)]
struct SharedCacheState {
cache: Option<String>,
}
impl AccountState for SharedCacheState {
fn new_for_account(_account_id: &str, _region: &str, _endpoint: &str) -> Self {
Self { cache: None }
}
fn inherit_from(&mut self, sibling: &Self) {
self.cache = sibling.cache.clone();
}
}
#[test]
fn a_new_accounts_first_region_inherits_from_the_default_account() {
let mut mrs: MultiRegionState<SharedCacheState> =
MultiAccountState::new("111111111111", "us-east-1", "http://localhost:4566");
mrs.default_regional_mut().cache = Some("shared".into());
assert_eq!(
mrs.regional_mut("222222222222", "us-east-1")
.cache
.as_deref(),
Some("shared")
);
assert_eq!(
mrs.regional_mut("333333333333", "eu-west-1")
.cache
.as_deref(),
Some("shared")
);
mrs.regional_mut("222222222222", "us-east-1").cache = Some("own".into());
assert_eq!(
mrs.regional_mut("222222222222", "ap-south-1")
.cache
.as_deref(),
Some("own")
);
assert_eq!(
mrs.regional_mut("333333333333", "eu-west-1")
.cache
.as_deref(),
Some("shared")
);
}
}