1use std::collections::BTreeMap;
14use std::fs;
15use std::path::{Path, PathBuf};
16
17use serde::{Deserialize, Serialize};
18
19use crate::error::CoreError;
20
21#[derive(Debug, Clone, PartialEq, Serialize)]
30#[serde(untagged)]
31pub enum Value {
32 Null,
34 Bool(bool),
36 Int(i64),
38 Float(f64),
40 String(String),
42}
43
44impl<'de> Deserialize<'de> for Value {
45 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
46 where
47 D: serde::Deserializer<'de>,
48 {
49 use serde::de::{self, MapAccess, Visitor};
50
51 struct ValueVisitor;
52
53 impl<'de> Visitor<'de> for ValueVisitor {
54 type Value = Value;
55
56 fn expecting(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
57 f.write_str("a scalar (null, bool, number, or string)")
58 }
59
60 fn visit_unit<E>(self) -> Result<Value, E> {
61 Ok(Value::Null)
62 }
63
64 fn visit_bool<E>(self, b: bool) -> Result<Value, E> {
65 Ok(Value::Bool(b))
66 }
67
68 fn visit_i64<E>(self, i: i64) -> Result<Value, E> {
69 Ok(Value::Int(i))
70 }
71
72 fn visit_u64<E>(self, u: u64) -> Result<Value, E> {
73 #[allow(clippy::cast_precision_loss)]
75 Ok(i64::try_from(u).map_or(Value::Float(u as f64), Value::Int))
76 }
77
78 fn visit_f64<E>(self, f: f64) -> Result<Value, E> {
79 Ok(Value::Float(f))
80 }
81
82 fn visit_str<E>(self, s: &str) -> Result<Value, E> {
83 Ok(Value::String(s.to_owned()))
84 }
85
86 fn visit_string<E>(self, s: String) -> Result<Value, E> {
87 Ok(Value::String(s))
88 }
89
90 fn visit_map<A>(self, map: A) -> Result<Value, A::Error>
93 where
94 A: MapAccess<'de>,
95 {
96 let number =
97 serde_json::Number::deserialize(de::value::MapAccessDeserializer::new(map))?;
98 if let Some(i) = number.as_i64() {
99 Ok(Value::Int(i))
100 } else if let Some(f) = number.as_f64() {
101 Ok(Value::Float(f))
102 } else {
103 Err(de::Error::custom(format!(
104 "unrepresentable number {number}"
105 )))
106 }
107 }
108 }
109
110 deserializer.deserialize_any(ValueVisitor)
111 }
112}
113
114impl std::fmt::Display for Value {
115 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
116 match self {
117 Self::Null => f.write_str(""),
118 Self::Bool(b) => write!(f, "{b}"),
119 Self::Int(i) => write!(f, "{i}"),
120 Self::Float(x) => write!(f, "{x}"),
121 Self::String(s) => f.write_str(s),
122 }
123 }
124}
125
126#[derive(Debug, Clone, Default, PartialEq)]
131pub struct GlobalStore {
132 values: BTreeMap<String, Value>,
133}
134
135impl GlobalStore {
136 pub fn new() -> Self {
138 Self::default()
139 }
140
141 pub fn get(&self, name: &str) -> Option<&Value> {
143 self.values.get(name)
144 }
145
146 pub fn insert(&mut self, name: impl Into<String>, value: Value) {
148 self.values.insert(name.into(), value);
149 }
150
151 pub fn iter(&self) -> impl Iterator<Item = (&str, &Value)> {
153 self.values.iter().map(|(k, v)| (k.as_str(), v))
154 }
155
156 pub fn load(path: &Path) -> Result<Self, CoreError> {
159 let text = match fs::read_to_string(path) {
160 Ok(text) => text,
161 Err(err) if err.kind() == std::io::ErrorKind::NotFound => {
162 return Ok(Self::new());
163 }
164 Err(err) => {
165 return Err(CoreError::system_with(
166 format!("cannot read global state file {}", path.display()),
167 err,
168 ));
169 }
170 };
171 let values: BTreeMap<String, Value> = serde_json::from_str(&text).map_err(|err| {
172 CoreError::system_with(
173 format!("global state file {} is not valid JSON", path.display()),
174 err,
175 )
176 })?;
177 Ok(Self { values })
178 }
179
180 pub fn save(&self, path: &Path) -> Result<(), CoreError> {
184 let json = serde_json::to_string_pretty(&self.values)
185 .map_err(|err| CoreError::system_with("cannot serialize global state", err))?;
186 let tmp = sibling_tmp_path(path);
187 #[cfg(unix)]
188 {
189 use std::io::Write as _;
190 use std::os::unix::fs::{OpenOptionsExt, PermissionsExt};
191 let mut file = fs::OpenOptions::new()
192 .write(true)
193 .create(true)
194 .truncate(true)
195 .mode(0o600)
196 .open(&tmp)
197 .map_err(|err| {
198 CoreError::system_with(
199 format!("cannot write global state temp file {}", tmp.display()),
200 err,
201 )
202 })?;
203 file.write_all(json.as_bytes()).map_err(|err| {
204 CoreError::system_with(
205 format!("cannot write global state temp file {}", tmp.display()),
206 err,
207 )
208 })?;
209 fs::set_permissions(&tmp, fs::Permissions::from_mode(0o600)).map_err(|err| {
212 CoreError::system_with(format!("cannot set permissions on {}", tmp.display()), err)
213 })?;
214 }
215 #[cfg(not(unix))]
216 fs::write(&tmp, json).map_err(|err| {
217 CoreError::system_with(
218 format!("cannot write global state temp file {}", tmp.display()),
219 err,
220 )
221 })?;
222 fs::rename(&tmp, path).map_err(|err| {
223 CoreError::system_with(
224 format!("cannot move global state into place at {}", path.display()),
225 err,
226 )
227 })
228 }
229}
230
231fn sibling_tmp_path(path: &Path) -> PathBuf {
233 let mut os = path.as_os_str().to_owned();
234 os.push(format!(".{}.tmp", std::process::id()));
237 PathBuf::from(os)
238}
239
240#[derive(Debug, Default)]
245pub struct World {
246 scenario: BTreeMap<String, Value>,
247 global: GlobalStore,
248 promoted: std::collections::BTreeSet<String>,
252}
253
254impl World {
255 pub fn new(global: GlobalStore) -> Self {
257 Self {
258 scenario: BTreeMap::new(),
259 global,
260 promoted: std::collections::BTreeSet::new(),
261 }
262 }
263
264 pub fn get(&self, name: &str) -> Option<&Value> {
266 self.scenario.get(name).or_else(|| self.global.get(name))
267 }
268
269 pub fn set(&mut self, name: impl Into<String>, value: Value) {
271 self.scenario.insert(name.into(), value);
272 }
273
274 pub fn set_global(&mut self, name: impl Into<String>, value: Value) {
277 let name = name.into();
278 self.promoted.insert(name.clone());
279 self.global.insert(name, value);
280 }
281
282 pub fn global(&self) -> &GlobalStore {
284 &self.global
285 }
286
287 pub fn promotions(&self) -> impl Iterator<Item = (&str, &Value)> {
291 self.promoted
292 .iter()
293 .filter_map(|key| self.global.get(key).map(|value| (key.as_str(), value)))
294 }
295
296 pub fn merged(&self) -> BTreeMap<&str, &Value> {
299 let mut merged: BTreeMap<&str, &Value> = self.global.iter().collect();
300 for (k, v) in &self.scenario {
301 merged.insert(k.as_str(), v);
302 }
303 merged
304 }
305}
306
307#[cfg(test)]
308mod tests {
309 #![allow(clippy::unwrap_used)]
310
311 use super::*;
312
313 #[test]
314 fn scenario_scope_shadows_global() {
315 let mut store = GlobalStore::new();
316 store.insert("token", Value::String("global".into()));
317 let mut world = World::new(store);
318 assert_eq!(world.get("token"), Some(&Value::String("global".into())));
319
320 world.set("token", Value::String("scenario".into()));
321 assert_eq!(world.get("token"), Some(&Value::String("scenario".into())));
322 }
323
324 #[test]
325 fn merged_view_prefers_scenario_values() {
326 let mut store = GlobalStore::new();
327 store.insert("a", Value::Int(1));
328 store.insert("b", Value::Int(2));
329 let mut world = World::new(store);
330 world.set("b", Value::Int(20));
331 let merged = world.merged();
332 assert_eq!(merged["a"], &Value::Int(1));
333 assert_eq!(merged["b"], &Value::Int(20));
334 }
335
336 #[test]
337 fn promotions_are_the_write_set_only() {
338 let mut store = GlobalStore::new();
339 store.insert("seed", Value::Int(1));
340 let mut world = World::new(store);
341 world.set("scenario-only", Value::Bool(true));
342 world.set_global("promoted", Value::Int(2));
343
344 let promotions: Vec<_> = world.promotions().collect();
345 assert_eq!(promotions, vec![("promoted", &Value::Int(2))]);
346 }
347
348 #[test]
349 fn value_json_forms_round_trip() {
350 let cases = [
351 ("null", Value::Null),
352 ("true", Value::Bool(true)),
353 ("3", Value::Int(3)),
354 ("0.5", Value::Float(0.5)),
355 (r#""c-42""#, Value::String("c-42".into())),
356 ];
357 for (json, expected) in cases {
358 let parsed: Value = serde_json::from_str(json)
359 .unwrap_or_else(|err| panic!("cannot parse {json}: {err}"));
360 assert_eq!(parsed, expected, "for literal {json}");
361 }
362 }
363
364 #[test]
365 fn store_save_load_round_trips_atomically() {
366 let dir = tempfile::tempdir().unwrap();
367 let path = dir.path().join(".proef-state.json");
368
369 let mut store = GlobalStore::new();
370 store.insert("clientId", Value::String("c-42".into()));
371 store.insert("count", Value::Int(3));
372 store.insert("ratio", Value::Float(0.5));
373 store.save(&path).unwrap();
374
375 assert!(!sibling_tmp_path(&path).exists());
377
378 let loaded = GlobalStore::load(&path).unwrap();
379 assert_eq!(loaded, store);
380 }
381
382 #[test]
383 fn missing_state_file_loads_as_empty() {
384 let dir = tempfile::tempdir().unwrap();
385 let loaded = GlobalStore::load(&dir.path().join("absent.json")).unwrap();
386 assert_eq!(loaded, GlobalStore::new());
387 }
388
389 #[cfg(unix)]
390 #[test]
391 fn state_file_is_created_private() {
392 use std::os::unix::fs::PermissionsExt;
393 let dir = tempfile::tempdir().unwrap();
394 let path = dir.path().join(".proef-state.json");
395 GlobalStore::new().save(&path).unwrap();
396 let mode = fs::metadata(&path).unwrap().permissions().mode();
397 assert_eq!(mode & 0o777, 0o600);
398 }
399
400 #[test]
401 fn corrupt_state_file_is_a_system_fault() {
402 let dir = tempfile::tempdir().unwrap();
403 let path = dir.path().join(".proef-state.json");
404 fs::write(&path, "not json").unwrap();
405 let err = GlobalStore::load(&path).unwrap_err();
406 assert_eq!(err.exit_code().code(), 3);
407 }
408}