Skip to main content

hoike_core/
state.rs

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
11/// Maximum allowed epoch jump from the current high-water mark.
12/// Prevents a poisoned bundle with epoch = u64::MAX from permanently
13/// locking out a CA. Kept as defense-in-depth even with CMS seal
14/// verification, since seal trust-anchor enforcement is optional.
15pub 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    /// Immutable bundle snapshots committed with the rollback marks.
22    #[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        // Only collect after the new descriptor is durable. In-flight requests
125        // hold heap bundles; an unlinked prior snapshot cannot change their data.
126        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    /// Persist immutable content before the descriptor that references it.
160    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        // Existing blobs are verified, never trusted solely by their filename.
171        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                // A strictly older epoch is always a rollback.
185                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                // Re-loading at the current high-water epoch is legitimate only
193                // when it is the *same* bundle (identical manifest digest) — e.g.
194                // a process restart reloading its own state. A *different* bundle
195                // at the same epoch is a fork/rollback attack and is rejected.
196                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}