fakecloud_sagemaker/
state.rs1use std::collections::BTreeMap;
17use std::sync::Arc;
18
19use parking_lot::RwLock;
20use serde::{Deserialize, Serialize};
21use serde_json::Value;
22
23use fakecloud_core::multi_account::{AccountState, MultiAccountState};
24
25pub const SAGEMAKER_SNAPSHOT_SCHEMA_VERSION: u32 = 1;
26
27#[derive(Debug, Clone, Default, Serialize, Deserialize)]
29pub struct SageMakerData {
30 #[serde(default)]
33 pub resources: BTreeMap<String, BTreeMap<String, Value>>,
34 #[serde(default)]
36 pub tags: BTreeMap<String, BTreeMap<String, String>>,
37 #[serde(default)]
39 pub singletons: BTreeMap<String, Value>,
40 #[serde(default)]
42 pub seq: u64,
43}
44
45impl SageMakerData {
46 pub fn get_resource(&self, family: &str, id: &str) -> Option<&Value> {
48 self.resources.get(family).and_then(|m| m.get(id))
49 }
50
51 pub fn get_resource_mut(&mut self, family: &str, id: &str) -> Option<&mut Value> {
53 self.resources.get_mut(family).and_then(|m| m.get_mut(id))
54 }
55
56 pub fn put_resource(&mut self, family: &str, id: &str, record: Value) {
58 self.resources
59 .entry(family.to_string())
60 .or_default()
61 .insert(id.to_string(), record);
62 }
63
64 pub fn remove_resource(&mut self, family: &str, id: &str) -> Option<Value> {
66 let removed = self.resources.get_mut(family).and_then(|m| m.remove(id));
67 if let Some(m) = self.resources.get(family) {
68 if m.is_empty() {
69 self.resources.remove(family);
70 }
71 }
72 removed
73 }
74
75 pub fn resolve_key(&self, family: &str, value: &str) -> Option<String> {
87 let m = self.resources.get(family)?;
88 if m.contains_key(value) {
89 return Some(value.to_string());
90 }
91 let canonical = [
92 format!("{family}Name"),
93 format!("{family}Id"),
94 format!("{family}Arn"),
95 ];
96 for (k, rec) in m {
97 if let Some(obj) = rec.as_object() {
98 if canonical
99 .iter()
100 .any(|cand| obj.get(cand).and_then(Value::as_str) == Some(value))
101 {
102 return Some(k.clone());
103 }
104 }
105 }
106 None
107 }
108
109 pub fn list_resource_entries(&self, family: &str) -> Vec<(String, Value)> {
111 self.resources
112 .get(family)
113 .map(|m| m.iter().map(|(k, v)| (k.clone(), v.clone())).collect())
114 .unwrap_or_default()
115 }
116
117 pub fn next_seq(&mut self) -> u64 {
119 self.seq += 1;
120 self.seq
121 }
122}
123
124impl AccountState for SageMakerData {
125 fn new_for_account(_account_id: &str, _region: &str, _endpoint: &str) -> Self {
126 Self::default()
127 }
128}
129
130pub type SharedSageMakerState = Arc<RwLock<MultiAccountState<SageMakerData>>>;
131
132#[derive(Debug, Serialize, Deserialize)]
133pub struct SageMakerSnapshot {
134 pub schema_version: u32,
135 pub accounts: MultiAccountState<SageMakerData>,
136}
137
138#[cfg(test)]
139mod tests {
140 use super::*;
141 use serde_json::json;
142
143 #[test]
144 fn new_account_is_empty() {
145 let data = SageMakerData::new_for_account("000000000000", "us-east-1", "");
146 assert!(data.resources.is_empty());
147 assert!(data.tags.is_empty());
148 }
149
150 #[test]
151 fn resource_round_trips() {
152 let mut d = SageMakerData::default();
153 d.put_resource(
154 "Model",
155 "m1",
156 json!({"ModelName": "m1", "ModelArn": "arn:aws:sagemaker:us-east-1:0:model/m1"}),
157 );
158 assert_eq!(d.get_resource("Model", "m1").unwrap()["ModelName"], "m1");
159 assert_eq!(d.resolve_key("Model", "m1").as_deref(), Some("m1"));
160 assert_eq!(
162 d.resolve_key("Model", "arn:aws:sagemaker:us-east-1:0:model/m1")
163 .as_deref(),
164 Some("m1")
165 );
166 assert!(d.remove_resource("Model", "m1").is_some());
167 assert!(d.get_resource("Model", "m1").is_none());
168 }
169
170 #[test]
171 fn resolve_key_ignores_noncanonical_sibling_member() {
172 let mut d = SageMakerData::default();
173 let shared = "arn:aws:iam::0:role/shared";
174 d.put_resource("Model", "m1", json!({"ModelName": "m1", "RoleArn": shared}));
176 assert_eq!(d.resolve_key("Model", shared), None);
179 assert_eq!(d.resolve_key("Model", "m1").as_deref(), Some("m1"));
181 }
182}