Skip to main content

coven_database/
changeset_identity.rs

1use fallible_streaming_iterator::FallibleStreamingIterator;
2use rusqlite::hooks::Action;
3use rusqlite::session::{ChangesetItem, ChangesetIter};
4use rusqlite::types::ValueRef;
5
6use coven_protocol::synced_schema::{RowIdentityError, SyncedTable};
7
8#[derive(Debug, thiserror::Error)]
9pub enum ChangesetIdentityError {
10    #[error("changeset row identity validation failed: {0}")]
11    Parse(#[from] rusqlite::Error),
12    #[error("changeset contains undeclared table {0:?}")]
13    UndeclaredTable(String),
14    #[error(transparent)]
15    Row(#[from] RowIdentityError),
16}
17
18/// Captured writes also retain Coven's derived routing rows. Validate host
19/// identities against the host declaration without treating those internal rows
20/// as host tables. Incoming changesets still use the strict validator below.
21pub(crate) fn validate_captured_row_identities(
22    bytes: &[u8],
23    tables: &[SyncedTable],
24) -> Result<(), crate::DbError> {
25    let host = crate::gate::recorded_host_changeset(bytes)?;
26    validate_changeset_row_identities(&host, tables)?;
27    Ok(())
28}
29
30pub(crate) fn validate_changeset_row_identities(
31    bytes: &[u8],
32    tables: &[SyncedTable],
33) -> Result<(), ChangesetIdentityError> {
34    if bytes.is_empty() {
35        return Ok(());
36    }
37
38    let input: &mut dyn std::io::Read = &mut &bytes[..];
39    let mut iter = ChangesetIter::start_strm(&input).map_err(ChangesetIdentityError::Parse)?;
40    while let Some(item) = iter.next().map_err(ChangesetIdentityError::Parse)? {
41        let op = item.op().map_err(ChangesetIdentityError::Parse)?;
42        let table_name = op.table_name();
43        let table = tables
44            .iter()
45            .find(|table| table.name() == table_name)
46            .ok_or_else(|| ChangesetIdentityError::UndeclaredTable(table_name.to_string()))?;
47        match op.code() {
48            Action::SQLITE_INSERT => {
49                let id = required_changeset_id(item, table_name, "new", ChangesetSide::New)?;
50                table.row_identity().validate(table_name, &id)?;
51            }
52            Action::SQLITE_DELETE => {
53                let id = required_changeset_id(item, table_name, "old", ChangesetSide::Old)?;
54                table.row_identity().validate(table_name, &id)?;
55            }
56            Action::SQLITE_UPDATE => {
57                let old = required_changeset_id(item, table_name, "old", ChangesetSide::Old)?;
58                let id = optional_changeset_id(item, table_name, "new", ChangesetSide::New)?
59                    .unwrap_or(old);
60                table.row_identity().validate(table_name, &id)?;
61            }
62            _ => {}
63        }
64    }
65    Ok(())
66}
67
68#[derive(Clone, Copy)]
69enum ChangesetSide {
70    Old,
71    New,
72}
73
74fn required_changeset_id(
75    item: &ChangesetItem,
76    table: &str,
77    side_name: &'static str,
78    side: ChangesetSide,
79) -> Result<String, ChangesetIdentityError> {
80    optional_changeset_id(item, table, side_name, side)?.ok_or_else(|| {
81        ChangesetIdentityError::Row(RowIdentityError::MissingPrimaryKey {
82            table: table.to_string(),
83            side: side_name,
84        })
85    })
86}
87
88fn optional_changeset_id(
89    item: &ChangesetItem,
90    table: &str,
91    side_name: &'static str,
92    side: ChangesetSide,
93) -> Result<Option<String>, ChangesetIdentityError> {
94    let value = match side {
95        ChangesetSide::Old => item.old_value(0),
96        ChangesetSide::New => item.new_value(0),
97    };
98    let value = match value {
99        Ok(value) => value,
100        Err(rusqlite::Error::InvalidColumnIndex(_)) => return Ok(None),
101        Err(error) => return Err(ChangesetIdentityError::Parse(error)),
102    };
103    let ValueRef::Text(bytes) = value else {
104        return Err(RowIdentityError::NonTextPrimaryKey {
105            table: table.to_string(),
106            side: side_name,
107        }
108        .into());
109    };
110    std::str::from_utf8(bytes)
111        .map(str::to_owned)
112        .map(Some)
113        .map_err(|error| RowIdentityError::NonUtf8PrimaryKey {
114            table: table.to_string(),
115            source: error,
116        })
117        .map_err(ChangesetIdentityError::from)
118}