use std::collections::BTreeMap;
use std::sync::Arc;
use parking_lot::RwLock;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use fakecloud_core::multi_account::{AccountState, MultiAccountState};
pub const SAGEMAKER_SNAPSHOT_SCHEMA_VERSION: u32 = 1;
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct SageMakerData {
#[serde(default)]
pub resources: BTreeMap<String, BTreeMap<String, Value>>,
#[serde(default)]
pub tags: BTreeMap<String, BTreeMap<String, String>>,
#[serde(default)]
pub singletons: BTreeMap<String, Value>,
#[serde(default)]
pub seq: u64,
}
impl SageMakerData {
pub fn get_resource(&self, family: &str, id: &str) -> Option<&Value> {
self.resources.get(family).and_then(|m| m.get(id))
}
pub fn get_resource_mut(&mut self, family: &str, id: &str) -> Option<&mut Value> {
self.resources.get_mut(family).and_then(|m| m.get_mut(id))
}
pub fn put_resource(&mut self, family: &str, id: &str, record: Value) {
self.resources
.entry(family.to_string())
.or_default()
.insert(id.to_string(), record);
}
pub fn remove_resource(&mut self, family: &str, id: &str) -> Option<Value> {
let removed = self.resources.get_mut(family).and_then(|m| m.remove(id));
if let Some(m) = self.resources.get(family) {
if m.is_empty() {
self.resources.remove(family);
}
}
removed
}
pub fn resolve_key(&self, family: &str, value: &str) -> Option<String> {
let m = self.resources.get(family)?;
if m.contains_key(value) {
return Some(value.to_string());
}
let canonical = [
format!("{family}Name"),
format!("{family}Id"),
format!("{family}Arn"),
];
for (k, rec) in m {
if let Some(obj) = rec.as_object() {
if canonical
.iter()
.any(|cand| obj.get(cand).and_then(Value::as_str) == Some(value))
{
return Some(k.clone());
}
}
}
None
}
pub fn list_resource_entries(&self, family: &str) -> Vec<(String, Value)> {
self.resources
.get(family)
.map(|m| m.iter().map(|(k, v)| (k.clone(), v.clone())).collect())
.unwrap_or_default()
}
pub fn next_seq(&mut self) -> u64 {
self.seq += 1;
self.seq
}
}
impl AccountState for SageMakerData {
fn new_for_account(_account_id: &str, _region: &str, _endpoint: &str) -> Self {
Self::default()
}
}
pub type SharedSageMakerState = Arc<RwLock<MultiAccountState<SageMakerData>>>;
#[derive(Debug, Serialize, Deserialize)]
pub struct SageMakerSnapshot {
pub schema_version: u32,
pub accounts: MultiAccountState<SageMakerData>,
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn new_account_is_empty() {
let data = SageMakerData::new_for_account("000000000000", "us-east-1", "");
assert!(data.resources.is_empty());
assert!(data.tags.is_empty());
}
#[test]
fn resource_round_trips() {
let mut d = SageMakerData::default();
d.put_resource(
"Model",
"m1",
json!({"ModelName": "m1", "ModelArn": "arn:aws:sagemaker:us-east-1:0:model/m1"}),
);
assert_eq!(d.get_resource("Model", "m1").unwrap()["ModelName"], "m1");
assert_eq!(d.resolve_key("Model", "m1").as_deref(), Some("m1"));
assert_eq!(
d.resolve_key("Model", "arn:aws:sagemaker:us-east-1:0:model/m1")
.as_deref(),
Some("m1")
);
assert!(d.remove_resource("Model", "m1").is_some());
assert!(d.get_resource("Model", "m1").is_none());
}
#[test]
fn resolve_key_ignores_noncanonical_sibling_member() {
let mut d = SageMakerData::default();
let shared = "arn:aws:iam::0:role/shared";
d.put_resource("Model", "m1", json!({"ModelName": "m1", "RoleArn": shared}));
assert_eq!(d.resolve_key("Model", shared), None);
assert_eq!(d.resolve_key("Model", "m1").as_deref(), Some("m1"));
}
}