fix(store): verify restored story snapshots
This commit is contained in:
@@ -339,6 +339,9 @@ fn map_store_error(error: StoreError) -> TurnFailure {
|
||||
StoreError::NodeAlreadyExists(_)
|
||||
| StoreError::ParentNotFound(_)
|
||||
| StoreError::StateMismatch(_)
|
||||
| StoreError::StateHashMismatch { .. }
|
||||
| StoreError::Sqlite(_)
|
||||
| StoreError::Serialization(_)
|
||||
| StoreError::Poisoned => internal_failure("turn could not be committed"),
|
||||
}
|
||||
}
|
||||
@@ -816,7 +819,8 @@ mod persistent_turn_tests {
|
||||
user_input: String::new(),
|
||||
beats: Vec::new(),
|
||||
delta: StateDelta { ops: Vec::new() },
|
||||
state_hash: format!("hash_{id}"),
|
||||
state_hash: hash_runtime_state(&state(id, branch))
|
||||
.expect("serializable test state"),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ use std::{
|
||||
sync::{Mutex, MutexGuard},
|
||||
};
|
||||
|
||||
use nana_domain::{RuntimeState, StoryNode};
|
||||
use nana_domain::{RuntimeState, StoryNode, stable_json_hash};
|
||||
use rusqlite::{params, Connection, OptionalExtension, TransactionBehavior};
|
||||
use thiserror::Error;
|
||||
|
||||
@@ -22,6 +22,8 @@ pub enum StoreError {
|
||||
StaleBranchHead { expected: String, actual: String },
|
||||
#[error("node and materialized state disagree: {0}")]
|
||||
StateMismatch(&'static str),
|
||||
#[error("materialized state hash does not match node: {node_id}")]
|
||||
StateHashMismatch { node_id: String },
|
||||
#[error("sqlite storage error: {0}")]
|
||||
Sqlite(String),
|
||||
#[error("JSON serialization error: {0}")]
|
||||
@@ -46,6 +48,8 @@ pub trait StoryStore: Send + Sync {
|
||||
fn append_node(&self, node: &StoryNode, state: &RuntimeState) -> Result<(), StoreError>;
|
||||
|
||||
fn load_state(&self, story_id: &str, branch_id: &str) -> Result<RuntimeState, StoreError>;
|
||||
|
||||
fn load_node(&self, story_id: &str, node_id: &str) -> Result<StoryNode, StoreError>;
|
||||
}
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
@@ -169,6 +173,14 @@ impl StoryStore for InMemoryStoryStore {
|
||||
.cloned()
|
||||
.ok_or(StoreError::StateMismatch("branch head has no state"))
|
||||
}
|
||||
|
||||
fn load_node(&self, story_id: &str, node_id: &str) -> Result<StoryNode, StoreError> {
|
||||
let data = self.lock()?;
|
||||
data.nodes
|
||||
.get(&(story_id.to_owned(), node_id.to_owned()))
|
||||
.cloned()
|
||||
.ok_or_else(|| StoreError::ParentNotFound(node_id.to_owned()))
|
||||
}
|
||||
}
|
||||
|
||||
/// Durable SQLite implementation of the append-only story store.
|
||||
@@ -222,7 +234,8 @@ impl SqliteStoryStore {
|
||||
let connection = self.lock()?;
|
||||
let stored = connection
|
||||
.query_row(
|
||||
"SELECT nodes.branch_id, materialized_states.state_json
|
||||
"SELECT nodes.branch_id, nodes.parent_id, nodes.node_json,
|
||||
materialized_states.state_json
|
||||
FROM nodes
|
||||
LEFT JOIN materialized_states
|
||||
ON materialized_states.story_id = nodes.story_id
|
||||
@@ -233,17 +246,22 @@ impl SqliteStoryStore {
|
||||
Ok((
|
||||
row.get::<_, String>(0)?,
|
||||
row.get::<_, Option<String>>(1)?,
|
||||
row.get::<_, String>(2)?,
|
||||
row.get::<_, Option<String>>(3)?,
|
||||
))
|
||||
},
|
||||
)
|
||||
.optional()?
|
||||
.ok_or_else(|| StoreError::ParentNotFound(node_id.to_owned()))?;
|
||||
|
||||
let node = deserialize_node(&stored.2)?;
|
||||
validate_loaded_node(&node, story_id, node_id, &stored.0, stored.1.as_deref())?;
|
||||
let state_json = stored
|
||||
.1
|
||||
.3
|
||||
.ok_or(StoreError::StateMismatch("node has no materialized state"))?;
|
||||
let state = deserialize_state(&state_json)?;
|
||||
validate_loaded_state(&state, story_id, node_id, &stored.0)?;
|
||||
validate_state_hash(&node, &state)?;
|
||||
Ok(state)
|
||||
}
|
||||
|
||||
@@ -364,8 +382,13 @@ impl StoryStore for SqliteStoryStore {
|
||||
|
||||
let stored = connection
|
||||
.query_row(
|
||||
"SELECT branch_heads.head_node_id, materialized_states.state_json
|
||||
"SELECT branch_heads.head_node_id, nodes.branch_id, nodes.parent_id,
|
||||
nodes.node_json, materialized_states.state_json
|
||||
FROM branch_heads
|
||||
JOIN nodes
|
||||
ON nodes.story_id = branch_heads.story_id
|
||||
AND nodes.node_id = branch_heads.head_node_id
|
||||
AND nodes.branch_id = branch_heads.branch_id
|
||||
LEFT JOIN materialized_states
|
||||
ON materialized_states.story_id = branch_heads.story_id
|
||||
AND materialized_states.node_id = branch_heads.head_node_id
|
||||
@@ -375,7 +398,10 @@ impl StoryStore for SqliteStoryStore {
|
||||
|row| {
|
||||
Ok((
|
||||
row.get::<_, String>(0)?,
|
||||
row.get::<_, Option<String>>(1)?,
|
||||
row.get::<_, String>(1)?,
|
||||
row.get::<_, Option<String>>(2)?,
|
||||
row.get::<_, String>(3)?,
|
||||
row.get::<_, Option<String>>(4)?,
|
||||
))
|
||||
},
|
||||
)
|
||||
@@ -385,13 +411,45 @@ impl StoryStore for SqliteStoryStore {
|
||||
branch_id: branch_id.to_owned(),
|
||||
})?;
|
||||
|
||||
let node = deserialize_node(&stored.3)?;
|
||||
validate_loaded_node(
|
||||
&node,
|
||||
story_id,
|
||||
&stored.0,
|
||||
&stored.1,
|
||||
stored.2.as_deref(),
|
||||
)?;
|
||||
let state_json = stored
|
||||
.1
|
||||
.4
|
||||
.ok_or(StoreError::StateMismatch("branch head has no state"))?;
|
||||
let state = deserialize_state(&state_json)?;
|
||||
validate_loaded_state(&state, story_id, &stored.0, branch_id)?;
|
||||
validate_state_hash(&node, &state)?;
|
||||
Ok(state)
|
||||
}
|
||||
|
||||
fn load_node(&self, story_id: &str, node_id: &str) -> Result<StoryNode, StoreError> {
|
||||
let connection = self.lock()?;
|
||||
let stored = connection
|
||||
.query_row(
|
||||
"SELECT branch_id, parent_id, node_json
|
||||
FROM nodes
|
||||
WHERE story_id = ?1 AND node_id = ?2",
|
||||
params![story_id, node_id],
|
||||
|row| {
|
||||
Ok((
|
||||
row.get::<_, String>(0)?,
|
||||
row.get::<_, Option<String>>(1)?,
|
||||
row.get::<_, String>(2)?,
|
||||
))
|
||||
},
|
||||
)
|
||||
.optional()?
|
||||
.ok_or_else(|| StoreError::ParentNotFound(node_id.to_owned()))?;
|
||||
let node = deserialize_node(&stored.2)?;
|
||||
validate_loaded_node(&node, story_id, node_id, &stored.0, stored.1.as_deref())?;
|
||||
Ok(node)
|
||||
}
|
||||
}
|
||||
|
||||
fn initialize_schema(connection: &Connection) -> Result<(), StoreError> {
|
||||
@@ -453,13 +511,40 @@ fn validate_materialized_state(
|
||||
if node.branch_id != state.current_branch {
|
||||
return Err(StoreError::StateMismatch("current_branch"));
|
||||
}
|
||||
validate_state_hash(node, state)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn deserialize_node(node_json: &str) -> Result<StoryNode, StoreError> {
|
||||
serde_json::from_str(node_json).map_err(StoreError::from)
|
||||
}
|
||||
|
||||
fn deserialize_state(state_json: &str) -> Result<RuntimeState, StoreError> {
|
||||
serde_json::from_str(state_json).map_err(StoreError::from)
|
||||
}
|
||||
|
||||
fn validate_loaded_node(
|
||||
node: &StoryNode,
|
||||
story_id: &str,
|
||||
node_id: &str,
|
||||
branch_id: &str,
|
||||
parent_id: Option<&str>,
|
||||
) -> Result<(), StoreError> {
|
||||
if node.story_id != story_id {
|
||||
return Err(StoreError::StateMismatch("story_id"));
|
||||
}
|
||||
if node.id != node_id {
|
||||
return Err(StoreError::StateMismatch("current_node"));
|
||||
}
|
||||
if node.branch_id != branch_id {
|
||||
return Err(StoreError::StateMismatch("current_branch"));
|
||||
}
|
||||
if node.parent_id.as_deref() != parent_id {
|
||||
return Err(StoreError::StateMismatch("parent_id"));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn validate_loaded_state(
|
||||
state: &RuntimeState,
|
||||
story_id: &str,
|
||||
@@ -478,6 +563,16 @@ fn validate_loaded_state(
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn validate_state_hash(node: &StoryNode, state: &RuntimeState) -> Result<(), StoreError> {
|
||||
let bytes = serde_json::to_vec(state)?;
|
||||
if stable_json_hash(&bytes) != node.state_hash {
|
||||
return Err(StoreError::StateHashMismatch {
|
||||
node_id: node.id.clone(),
|
||||
});
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::{
|
||||
@@ -487,7 +582,7 @@ mod tests {
|
||||
time::{SystemTime, UNIX_EPOCH},
|
||||
};
|
||||
|
||||
use nana_domain::{RuntimeState, StateDelta, StoryNode};
|
||||
use nana_domain::{RuntimeState, StateDelta, StoryNode, stable_json_hash};
|
||||
use rusqlite::params;
|
||||
|
||||
use super::{InMemoryStoryStore, SqliteStoryStore, StoreError, StoryStore};
|
||||
@@ -509,6 +604,7 @@ mod tests {
|
||||
}
|
||||
|
||||
fn node(id: &str, parent_id: Option<&str>, branch: &str) -> StoryNode {
|
||||
let materialized = state(id, branch);
|
||||
StoryNode {
|
||||
id: id.to_owned(),
|
||||
story_id: "story_demo".to_owned(),
|
||||
@@ -518,7 +614,9 @@ mod tests {
|
||||
user_input: String::new(),
|
||||
beats: Vec::new(),
|
||||
delta: StateDelta { ops: Vec::new() },
|
||||
state_hash: format!("hash_{id}"),
|
||||
state_hash: stable_json_hash(
|
||||
&serde_json::to_vec(&materialized).expect("serializable test state"),
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -600,6 +698,14 @@ mod tests {
|
||||
.current_node,
|
||||
"node_001"
|
||||
);
|
||||
assert_eq!(
|
||||
store
|
||||
.load_node("story_demo", "node_002")
|
||||
.expect("stored node")
|
||||
.parent_id
|
||||
.as_deref(),
|
||||
Some("node_001")
|
||||
);
|
||||
}
|
||||
|
||||
fn assert_forks_from_an_old_node(store: &impl InspectableStoryStore) {
|
||||
@@ -814,6 +920,102 @@ mod tests {
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sqlite_rejects_valid_state_json_with_a_mismatched_hash() {
|
||||
let store = SqliteStoryStore::open_in_memory().expect("in-memory SQLite store");
|
||||
store
|
||||
.append_node(
|
||||
&node("node_001", None, "branch_main"),
|
||||
&state("node_001", "branch_main"),
|
||||
)
|
||||
.expect("root append");
|
||||
let mut corrupted = state("node_001", "branch_main");
|
||||
corrupted
|
||||
.world_flags
|
||||
.insert("silently_changed".to_owned(), true);
|
||||
let corrupted_json = serde_json::to_string(&corrupted).expect("corrupt test JSON");
|
||||
store
|
||||
.connection
|
||||
.lock()
|
||||
.expect("SQLite connection lock")
|
||||
.execute(
|
||||
"UPDATE materialized_states SET state_json = ?1",
|
||||
params![corrupted_json],
|
||||
)
|
||||
.expect("corrupt test state");
|
||||
|
||||
assert_eq!(
|
||||
store.load_state("story_demo", "branch_main"),
|
||||
Err(StoreError::StateHashMismatch {
|
||||
node_id: "node_001".to_owned()
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sqlite_rejects_node_json_that_disagrees_with_structure_columns() {
|
||||
let store = SqliteStoryStore::open_in_memory().expect("in-memory SQLite store");
|
||||
store
|
||||
.append_node(
|
||||
&node("node_001", None, "branch_main"),
|
||||
&state("node_001", "branch_main"),
|
||||
)
|
||||
.expect("root append");
|
||||
let mut corrupted = node("node_001", None, "branch_main");
|
||||
corrupted.branch_id = "hidden_branch".to_owned();
|
||||
let corrupted_json = serde_json::to_string(&corrupted).expect("corrupt node JSON");
|
||||
store
|
||||
.connection
|
||||
.lock()
|
||||
.expect("SQLite connection lock")
|
||||
.execute(
|
||||
"UPDATE nodes SET node_json = ?1",
|
||||
params![corrupted_json],
|
||||
)
|
||||
.expect("corrupt test node");
|
||||
|
||||
assert_eq!(
|
||||
store.load_node("story_demo", "node_001"),
|
||||
Err(StoreError::StateMismatch("current_branch"))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sqlite_rejects_a_branch_head_pointing_at_another_branch() {
|
||||
let store = SqliteStoryStore::open_in_memory().expect("in-memory SQLite store");
|
||||
store
|
||||
.append_node(
|
||||
&node("node_001", None, "branch_main"),
|
||||
&state("node_001", "branch_main"),
|
||||
)
|
||||
.expect("root append");
|
||||
store
|
||||
.append_node(
|
||||
&node("node_other", Some("node_001"), "branch_other"),
|
||||
&state("node_other", "branch_other"),
|
||||
)
|
||||
.expect("other branch append");
|
||||
store
|
||||
.connection
|
||||
.lock()
|
||||
.expect("SQLite connection lock")
|
||||
.execute(
|
||||
"UPDATE branch_heads
|
||||
SET head_node_id = 'node_other'
|
||||
WHERE story_id = 'story_demo' AND branch_id = 'branch_main'",
|
||||
[],
|
||||
)
|
||||
.expect("corrupt branch head");
|
||||
|
||||
assert_eq!(
|
||||
store.load_state("story_demo", "branch_main"),
|
||||
Err(StoreError::BranchNotFound {
|
||||
story_id: "story_demo".to_owned(),
|
||||
branch_id: "branch_main".to_owned()
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
struct TemporaryDatabase {
|
||||
path: PathBuf,
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user