1use std::collections::HashMap;
2use std::path::{Path, PathBuf};
3
4use serde::{Deserialize, Serialize};
5use sha2::{Digest, Sha256};
6use tracing::info;
7
8use crate::error::{CoreError, Result};
9use ahu::Bundle;
10
11pub const MAX_EPOCH_JUMP: u64 = 10_000;
16
17#[derive(Debug, Clone, Serialize, Deserialize, Default)]
18struct PersistedState {
19 high_water_marks: HashMap<String, u64>,
20 manifest_digests: HashMap<String, String>,
21 #[serde(default)]
23 active_bundles: HashMap<String, PathBuf>,
24}
25
26impl PersistedState {
27 fn make_key(producer_id: &str, issuer_key_hash_hex: &str) -> String {
28 format!("{producer_id}:{issuer_key_hash_hex}")
29 }
30}
31
32#[derive(Clone)]
33pub struct StateStore {
34 path: PathBuf,
35 state: PersistedState,
36 staged: bool,
37}
38
39impl StateStore {
40 pub fn open(path: &Path) -> Result<Self> {
41 if let Some(parent) = path.parent() {
42 if !parent.exists() {
43 std::fs::create_dir_all(parent).map_err(|e| {
44 CoreError::StateStore(format!(
45 "failed to create state directory {}: {e}",
46 parent.display()
47 ))
48 })?;
49 }
50 }
51
52 let state = if path.exists() {
53 let contents = std::fs::read_to_string(path).map_err(|e| {
54 CoreError::StateStore(format!("failed to read state file {}: {e}", path.display()))
55 })?;
56 serde_json::from_str(&contents).map_err(|e| {
57 CoreError::StateStore(format!(
58 "failed to parse state file {}: {e}",
59 path.display()
60 ))
61 })?
62 } else {
63 info!(path = %path.display(), "no existing state file — initializing fresh state");
64 PersistedState::default()
65 };
66
67 Ok(StateStore {
68 path: path.to_path_buf(),
69 state,
70 staged: false,
71 })
72 }
73
74 pub fn get_high_water(&self, producer_id: &str, issuer_key_hash_hex: &str) -> Option<u64> {
75 let key = PersistedState::make_key(producer_id, issuer_key_hash_hex);
76 self.state.high_water_marks.get(&key).copied()
77 }
78
79 pub fn get_manifest_digest(
80 &self,
81 producer_id: &str,
82 issuer_key_hash_hex: &str,
83 ) -> Option<[u8; 32]> {
84 let key = PersistedState::make_key(producer_id, issuer_key_hash_hex);
85 self.state.manifest_digests.get(&key).and_then(|hex_str| {
86 let bytes = hex::decode(hex_str).ok()?;
87 <[u8; 32]>::try_from(bytes.as_slice()).ok()
88 })
89 }
90
91 pub fn advance(
92 &mut self,
93 producer_id: &str,
94 issuer_key_hash_hex: &str,
95 epoch: u64,
96 manifest_digest: [u8; 32],
97 ) -> Result<()> {
98 let mut next = self.clone();
99 let key = PersistedState::make_key(producer_id, issuer_key_hash_hex);
100 let current = next.state.high_water_marks.get(&key).copied();
101 if current.is_none_or(|current| epoch > current) {
102 next.state.high_water_marks.insert(key.clone(), epoch);
103 next.state
104 .manifest_digests
105 .insert(key, hex::encode(manifest_digest));
106 if !self.staged {
107 next.persist()?;
108 }
109 self.state = next.state;
110 }
111 Ok(())
112 }
113
114 pub(crate) fn transaction(&self) -> Self {
115 let mut candidate = self.clone();
116 candidate.staged = true;
117 candidate.state.active_bundles.clear();
118 candidate
119 }
120
121 pub(crate) fn commit(&mut self, candidate: Self) -> Result<()> {
122 candidate.persist()?;
123 self.state = candidate.state;
124 let dir = self
127 .path
128 .parent()
129 .unwrap_or(Path::new("."))
130 .join("generations");
131 if let Ok(files) = std::fs::read_dir(&dir) {
132 for file in files.flatten() {
133 let path = file.path();
134 let generated = path.extension().is_some_and(|ext| ext == "ahu")
135 && path
136 .file_stem()
137 .and_then(|s| s.to_str())
138 .is_some_and(|s| s.len() == 64 && s.bytes().all(|b| b.is_ascii_hexdigit()));
139 if generated
140 && !self
141 .state
142 .active_bundles
143 .values()
144 .any(|active| active == &path)
145 {
146 if let Err(error) = std::fs::remove_file(&path) {
147 tracing::warn!(%error, "could not remove obsolete generation snapshot");
148 }
149 }
150 }
151 }
152 Ok(())
153 }
154
155 pub(crate) fn active_bundles(&self) -> &HashMap<String, PathBuf> {
156 &self.state.active_bundles
157 }
158
159 pub(crate) fn snapshot(&mut self, label: &str, bundle: &Bundle) -> Result<()> {
161 let bytes = bundle.to_bytes()?;
162 let digest = hex::encode(Sha256::digest(&bytes));
163 let dir = self
164 .path
165 .parent()
166 .unwrap_or(Path::new("."))
167 .join("generations");
168 std::fs::create_dir_all(&dir)?;
169 let path = dir.join(format!("{digest}.ahu"));
170 if !path.exists() || std::fs::read(&path)? != bytes {
172 Self::write_atomic(&path, &bytes)?;
173 }
174 self.state.active_bundles.insert(label.to_owned(), path);
175 Ok(())
176 }
177
178 pub fn check_rollback(&self, bundle: &Bundle) -> Result<()> {
179 let producer_id = &bundle.manifest.producer_id;
180 let manifest_digest: [u8; 32] = Sha256::digest(&bundle.manifest_bytes).into();
181 for scope in &bundle.manifest.ca_scopes {
182 let ikh = hex::encode(&scope.issuer_key_hash);
183 if let Some(hw) = self.get_high_water(producer_id, &ikh) {
184 if scope.epoch < hw {
186 return Err(CoreError::EpochRollback {
187 scope: format!("{}:{}", producer_id, &ikh[..16.min(ikh.len())]),
188 epoch: scope.epoch,
189 high_water: hw,
190 });
191 }
192 if scope.epoch == hw
197 && self.get_manifest_digest(producer_id, &ikh) != Some(manifest_digest)
198 {
199 return Err(CoreError::EpochRollback {
200 scope: format!("{}:{}", producer_id, &ikh[..16.min(ikh.len())]),
201 epoch: scope.epoch,
202 high_water: hw,
203 });
204 }
205 let jump = scope.epoch - hw;
206 if jump > MAX_EPOCH_JUMP {
207 return Err(CoreError::EpochJumpTooLarge {
208 scope: format!("{}:{}", producer_id, &ikh[..16.min(ikh.len())]),
209 epoch: scope.epoch,
210 high_water: hw,
211 jump,
212 max_jump: MAX_EPOCH_JUMP,
213 });
214 }
215 }
216 }
217 Ok(())
218 }
219
220 pub fn check_continuity(&self, bundle: &Bundle) -> Result<()> {
221 if let Some(prev_digest) = &bundle.manifest.continuity.prev_manifest_digest {
222 let producer_id = &bundle.manifest.producer_id;
223 for scope in &bundle.manifest.ca_scopes {
224 let ikh = hex::encode(&scope.issuer_key_hash);
225 if let Some(recorded) = self.get_manifest_digest(producer_id, &ikh) {
226 let digest: [u8; 32] = Sha256::digest(&bundle.manifest_bytes).into();
227 let identical = self.get_high_water(producer_id, &ikh) == Some(scope.epoch)
228 && digest == recorded;
229 if !identical && *prev_digest != recorded {
230 return Err(CoreError::ForkDetected {
231 scope: format!("{}:{}", producer_id, &ikh[..16.min(ikh.len())]),
232 });
233 }
234 }
235 }
236 }
237 Ok(())
238 }
239
240 pub fn advance_from_bundle(&mut self, bundle: &Bundle) -> Result<()> {
241 let mut candidate = self.clone();
242 candidate.staged = true;
243 let digest: [u8; 32] = Sha256::digest(&bundle.manifest_bytes).into();
244 for scope in &bundle.manifest.ca_scopes {
245 candidate.advance(
246 &bundle.manifest.producer_id,
247 &hex::encode(&scope.issuer_key_hash),
248 scope.epoch,
249 digest,
250 )?;
251 }
252 if !self.staged {
253 candidate.persist()?;
254 }
255 self.state = candidate.state;
256 Ok(())
257 }
258
259 fn persist(&self) -> Result<()> {
260 let json = serde_json::to_vec_pretty(&self.state)
261 .map_err(|e| CoreError::StateStore(format!("failed to serialize state: {e}")))?;
262 Self::write_atomic(&self.path, &json)
263 }
264
265 fn write_atomic(path: &Path, bytes: &[u8]) -> Result<()> {
266 use std::io::Write;
267 use std::sync::atomic::{AtomicU64, Ordering};
268 static COUNTER: AtomicU64 = AtomicU64::new(0);
269 let tmp = path.with_extension(format!(
270 "tmp-{}-{}",
271 std::process::id(),
272 COUNTER.fetch_add(1, Ordering::Relaxed)
273 ));
274 let result = (|| -> Result<()> {
275 let mut file = std::fs::OpenOptions::new()
276 .write(true)
277 .create_new(true)
278 .open(&tmp)?;
279 file.write_all(bytes)?;
280 file.sync_all()?;
281 std::fs::rename(&tmp, path)?;
282 std::fs::File::open(path.parent().unwrap_or(Path::new(".")))?.sync_all()?;
283 Ok(())
284 })();
285 if result.is_err() {
286 let _ = std::fs::remove_file(&tmp);
287 }
288 result
289 }
290}