Skip to main content

miden_node_db/sqlite/
codec.rs

1//! Column codec for the rusqlite-based SQLite framework.
2//!
3//! [`ToSqlValue`] and [`FromSqlValue`] are the per-column write/read codec for our domain types.
4//! They operate on [`DbValue`]/[`DbValueRef`], thin wrappers over rusqlite's value types, so that
5//! crates implementing a codec for their own types never have to name `rusqlite` directly.
6//!
7//! Structured BLOBs use protobuf. Fixed-width keys retain their native byte encoding. Scalar
8//! types map onto an SQLite `INTEGER`/`TEXT` and implement the traits directly (see the impls ported
9//! from the legacy `SqlTypeConvert` below).
10//!
11//! Integer primitives read back range-checked rather than cast, so a column holding a value outside
12//! the target type's range errors instead of silently truncating. The one place a lossy conversion
13//! is deliberate is `Felt`.
14
15use std::rc::Rc;
16
17use miden_protocol::Felt;
18use miden_protocol::account::StorageSlotName;
19use miden_protocol::block::BlockNumber;
20use miden_protocol::note::NoteTag;
21use rusqlite::ToSql;
22use rusqlite::types::{ToSqlOutput, Value, ValueRef};
23
24use crate::DatabaseError;
25
26// DB VALUE WRAPPERS
27// =================================================================================================
28
29/// An owned SQL value produced when binding a Rust value as a query parameter.
30///
31/// Wraps `rusqlite`'s value types so codec implementors never name `rusqlite`. A value is either a
32/// single column value or a list bound for a `rarray(?)` table-valued parameter (used by the
33/// cacheable IN-list helpers in [`in_list`](crate::sqlite::in_list)).
34#[derive(Debug, Clone)]
35pub enum DbValue {
36    /// A single SQL column value.
37    Single(Value),
38    /// A list of values bound via rusqlite's `array` extension for use with `rarray(?)`.
39    Array(Rc<Vec<Value>>),
40}
41
42impl DbValue {
43    /// Builds an `INTEGER` value.
44    pub fn integer(value: i64) -> Self {
45        Self::Single(Value::Integer(value))
46    }
47
48    /// Builds a `REAL` value.
49    pub fn real(value: f64) -> Self {
50        Self::Single(Value::Real(value))
51    }
52
53    /// Builds a `TEXT` value.
54    pub fn text(value: String) -> Self {
55        Self::Single(Value::Text(value))
56    }
57
58    /// Builds a `BLOB` value.
59    pub fn blob(value: Vec<u8>) -> Self {
60        Self::Single(Value::Blob(value))
61    }
62
63    /// Builds a `NULL` value.
64    pub fn null() -> Self {
65        Self::Single(Value::Null)
66    }
67
68    /// Builds a list value bound for a `rarray(?)` table-valued parameter.
69    pub(crate) fn array(values: Vec<Value>) -> Self {
70        Self::Array(Rc::new(values))
71    }
72}
73
74impl ToSql for DbValue {
75    fn to_sql(&self) -> rusqlite::Result<ToSqlOutput<'_>> {
76        match self {
77            Self::Single(value) => value.to_sql(),
78            Self::Array(values) => values.to_sql(),
79        }
80    }
81}
82
83/// A borrowed SQL value handed to [`FromSqlValue`] when reading a column.
84///
85/// Wraps `rusqlite::types::ValueRef` so codec implementors never name `rusqlite`.
86#[derive(Debug, Clone, Copy)]
87pub struct DbValueRef<'a>(ValueRef<'a>);
88
89impl<'a> DbValueRef<'a> {
90    pub(crate) fn new(value: ValueRef<'a>) -> Self {
91        Self(value)
92    }
93
94    /// Reads the value as an `i64`.
95    pub fn as_i64(self) -> Result<i64, DatabaseError> {
96        self.0.as_i64().map_err(|err| DatabaseError::deserialization("i64", err))
97    }
98
99    /// Reads the value as a borrowed BLOB.
100    pub fn as_blob(self) -> Result<&'a [u8], DatabaseError> {
101        self.0.as_blob().map_err(|err| DatabaseError::deserialization("blob", err))
102    }
103
104    /// Reads the value as a borrowed string.
105    pub fn as_str(self) -> Result<&'a str, DatabaseError> {
106        self.0.as_str().map_err(|err| DatabaseError::deserialization("str", err))
107    }
108
109    /// Returns `true` if the value is SQL `NULL`.
110    pub fn is_null(self) -> bool {
111        matches!(self.0, ValueRef::Null)
112    }
113}
114
115// CODEC TRAITS
116// =================================================================================================
117
118/// Converts a Rust value into its SQL parameter representation (the write side of the codec).
119pub trait ToSqlValue {
120    /// Returns the SQL value bound for this Rust value.
121    fn to_sql_value(&self) -> DbValue;
122}
123
124/// Builds a Rust value from a SQL column value (the read side of the codec).
125pub trait FromSqlValue: Sized {
126    /// Reads `Self` from a SQL column value.
127    fn from_sql_value(value: DbValueRef<'_>) -> Result<Self, DatabaseError>;
128}
129
130// Forward `ToSqlValue` through references so callers can pass `&value` in a parameter slice.
131impl<T: ToSqlValue + ?Sized> ToSqlValue for &T {
132    fn to_sql_value(&self) -> DbValue {
133        (**self).to_sql_value()
134    }
135}
136
137// PRIMITIVE IMPLS
138// =================================================================================================
139
140impl ToSqlValue for i64 {
141    fn to_sql_value(&self) -> DbValue {
142        DbValue::integer(*self)
143    }
144}
145
146impl FromSqlValue for i64 {
147    fn from_sql_value(value: DbValueRef<'_>) -> Result<Self, DatabaseError> {
148        value.as_i64()
149    }
150}
151
152// The unsigned integers widen losslessly on the write side and are range-checked on the read side,
153// so a column holding a value outside the type's range errors instead of silently truncating.
154
155impl ToSqlValue for u8 {
156    fn to_sql_value(&self) -> DbValue {
157        DbValue::integer(i64::from(*self))
158    }
159}
160
161impl FromSqlValue for u8 {
162    fn from_sql_value(value: DbValueRef<'_>) -> Result<Self, DatabaseError> {
163        Self::try_from(value.as_i64()?).map_err(|err| DatabaseError::deserialization("u8", err))
164    }
165}
166
167impl ToSqlValue for u16 {
168    fn to_sql_value(&self) -> DbValue {
169        DbValue::integer(i64::from(*self))
170    }
171}
172
173impl FromSqlValue for u16 {
174    fn from_sql_value(value: DbValueRef<'_>) -> Result<Self, DatabaseError> {
175        Self::try_from(value.as_i64()?).map_err(|err| DatabaseError::deserialization("u16", err))
176    }
177}
178
179impl ToSqlValue for u32 {
180    fn to_sql_value(&self) -> DbValue {
181        DbValue::integer(i64::from(*self))
182    }
183}
184
185impl FromSqlValue for u32 {
186    fn from_sql_value(value: DbValueRef<'_>) -> Result<Self, DatabaseError> {
187        Self::try_from(value.as_i64()?).map_err(|err| DatabaseError::deserialization("u32", err))
188    }
189}
190
191impl ToSqlValue for bool {
192    fn to_sql_value(&self) -> DbValue {
193        DbValue::integer(i64::from(*self))
194    }
195}
196
197impl FromSqlValue for bool {
198    fn from_sql_value(value: DbValueRef<'_>) -> Result<Self, DatabaseError> {
199        Ok(value.as_i64()? != 0)
200    }
201}
202
203impl ToSqlValue for Vec<u8> {
204    fn to_sql_value(&self) -> DbValue {
205        DbValue::blob(self.clone())
206    }
207}
208
209impl FromSqlValue for Vec<u8> {
210    fn from_sql_value(value: DbValueRef<'_>) -> Result<Self, DatabaseError> {
211        Ok(value.as_blob()?.to_vec())
212    }
213}
214
215impl ToSqlValue for str {
216    fn to_sql_value(&self) -> DbValue {
217        DbValue::text(self.to_owned())
218    }
219}
220
221impl ToSqlValue for String {
222    fn to_sql_value(&self) -> DbValue {
223        DbValue::text(self.clone())
224    }
225}
226
227impl FromSqlValue for String {
228    fn from_sql_value(value: DbValueRef<'_>) -> Result<Self, DatabaseError> {
229        Ok(value.as_str()?.to_owned())
230    }
231}
232
233impl<T: ToSqlValue> ToSqlValue for Option<T> {
234    fn to_sql_value(&self) -> DbValue {
235        match self {
236            Some(value) => value.to_sql_value(),
237            None => DbValue::null(),
238        }
239    }
240}
241
242impl<T: FromSqlValue> FromSqlValue for Option<T> {
243    fn from_sql_value(value: DbValueRef<'_>) -> Result<Self, DatabaseError> {
244        if value.is_null() {
245            Ok(None)
246        } else {
247            Ok(Some(T::from_sql_value(value)?))
248        }
249    }
250}
251
252// DOMAIN SCALAR IMPLS
253// =================================================================================================
254//
255// Domain types stored in an `INTEGER`/`TEXT` column rather than as a BLOB.
256
257impl ToSqlValue for BlockNumber {
258    fn to_sql_value(&self) -> DbValue {
259        DbValue::integer(i64::from(self.as_u32()))
260    }
261}
262
263impl FromSqlValue for BlockNumber {
264    fn from_sql_value(value: DbValueRef<'_>) -> Result<Self, DatabaseError> {
265        u32::from_sql_value(value).map(BlockNumber::from)
266    }
267}
268
269impl ToSqlValue for NoteTag {
270    fn to_sql_value(&self) -> DbValue {
271        DbValue::integer(i64::from(self.as_u32()))
272    }
273}
274
275impl FromSqlValue for NoteTag {
276    fn from_sql_value(value: DbValueRef<'_>) -> Result<Self, DatabaseError> {
277        u32::from_sql_value(value).map(NoteTag::new)
278    }
279}
280
281impl ToSqlValue for StorageSlotName {
282    fn to_sql_value(&self) -> DbValue {
283        DbValue::text(self.as_str().to_owned())
284    }
285}
286
287impl FromSqlValue for StorageSlotName {
288    fn from_sql_value(value: DbValueRef<'_>) -> Result<Self, DatabaseError> {
289        StorageSlotName::new(value.as_str()?)
290            .map_err(|err| DatabaseError::deserialization("StorageSlotName", err))
291    }
292}
293
294/// A field element is stored as the bit reinterpretation of its canonical `u64`.
295impl ToSqlValue for Felt {
296    #[expect(
297        clippy::cast_possible_wrap,
298        reason = "canonical field elements are stored as the wrapped i64 bit pattern"
299    )]
300    fn to_sql_value(&self) -> DbValue {
301        DbValue::integer(self.as_canonical_u64() as i64)
302    }
303}
304
305impl FromSqlValue for Felt {
306    #[expect(clippy::cast_sign_loss, reason = "reverses the u64 -> i64 wrap applied on write")]
307    fn from_sql_value(value: DbValueRef<'_>) -> Result<Self, DatabaseError> {
308        Felt::new(value.as_i64()? as u64).map_err(|err| DatabaseError::deserialization("Felt", err))
309    }
310}
311
312// BLOB CODEC MACRO
313// =================================================================================================
314
315/// Generates [`ToSqlValue`](crate::sqlite::ToSqlValue) and
316/// [`FromSqlValue`](crate::sqlite::FromSqlValue) for types stored as a BLOB via their
317/// `Serializable`/`Deserializable` impls.
318///
319/// Use this codec for fixed-width keys and indexed values.
320#[macro_export]
321macro_rules! impl_raw_blob_codec {
322    ($($t:ty),+ $(,)?) => {
323        $(
324            impl $crate::sqlite::ToSqlValue for $t {
325                fn to_sql_value(&self) -> $crate::sqlite::DbValue {
326                    $crate::sqlite::DbValue::blob(
327                        ::miden_protocol::utils::serde::Serializable::to_bytes(self),
328                    )
329                }
330            }
331
332            impl $crate::sqlite::FromSqlValue for $t {
333                fn from_sql_value(
334                    value: $crate::sqlite::DbValueRef<'_>,
335                ) -> ::core::result::Result<Self, $crate::DatabaseError> {
336                    let bytes = value.as_blob()?;
337                    <$t as ::miden_protocol::utils::serde::Deserializable>::read_from_bytes(bytes)
338                        .map_err(|err| {
339                            $crate::DatabaseError::deserialization(::core::stringify!($t), err)
340                        })
341                }
342            }
343        )+
344    };
345}
346
347/// Implements SQLite conversion for a structured protobuf value.
348#[macro_export]
349macro_rules! impl_protobuf_codec {
350    ($($t:ty),+ $(,)?) => {
351        $(
352            impl $crate::sqlite::ToSqlValue for $t {
353                fn to_sql_value(&self) -> $crate::sqlite::DbValue {
354                    $crate::sqlite::DbValue::blob($crate::persistence::encode(self))
355                }
356            }
357            impl $crate::sqlite::FromSqlValue for $t {
358                fn from_sql_value(value: $crate::sqlite::DbValueRef<'_>) -> Result<Self, $crate::DatabaseError> {
359                    $crate::persistence::decode(value.as_blob()?).map_err(|err| {
360                        $crate::DatabaseError::deserialization(stringify!($t), err)
361                    })
362                }
363            }
364        )+
365    };
366}
367
368// Codec for the common protocol types stored as BLOBs. Shared by all node crates so that the orphan
369// rule does not force each consumer to redeclare them.
370impl_raw_blob_codec!(
371    miden_protocol::account::AccountId,
372    miden_protocol::account::StorageMapKey,
373    miden_protocol::transaction::TransactionId,
374    miden_protocol::note::NoteId,
375    miden_protocol::note::Nullifier,
376    miden_protocol::Word,
377);
378
379impl_protobuf_codec!(
380    miden_protocol::block::BlockHeader,
381    miden_protocol::block::BlockSignatures,
382    miden_protocol::account::Account,
383    miden_protocol::account::AccountCode,
384    miden_protocol::account::AccountStorageHeader,
385    miden_protocol::asset::Asset,
386    miden_protocol::note::Note,
387    miden_protocol::note::NoteAssets,
388    miden_protocol::note::NoteAttachments,
389    miden_protocol::note::NoteHeader,
390    miden_protocol::note::NoteDetails,
391    miden_protocol::note::NoteScript,
392    miden_protocol::note::NoteStorage,
393    miden_protocol::protocol_config::ProtocolConfig,
394    miden_protocol::crypto::merkle::SparseMerklePath,
395    miden_protocol::crypto::merkle::mmr::PartialMmr,
396);
397
398// TESTS
399// =================================================================================================
400
401#[cfg(test)]
402mod tests {
403    use miden_protocol::Word;
404    use miden_protocol::block::BlockNumber;
405    use rusqlite::types::{Value, ValueRef};
406
407    use super::*;
408    use crate::SqlTypeConvert;
409
410    /// Returns the `i64` a value binds to, failing the test for non-integer values.
411    fn bound_integer(value: &impl ToSqlValue) -> i64 {
412        match value.to_sql_value() {
413            DbValue::Single(Value::Integer(raw)) => raw,
414            other => panic!("expected an INTEGER binding, got {other:?}"),
415        }
416    }
417
418    /// Reads a value back from the `i64` a column holds.
419    fn read_integer<T: FromSqlValue>(raw: i64) -> Result<T, DatabaseError> {
420        T::from_sql_value(DbValueRef::new(ValueRef::Integer(raw)))
421    }
422
423    // ENCODING PARITY WITH `SqlTypeConvert`
424    // ---------------------------------------------------------------------------------------------
425
426    #[test]
427    fn block_number_roundtrip() {
428        for block_num in [
429            BlockNumber::GENESIS,
430            BlockNumber::from(1),
431            BlockNumber::from(u32::MAX - 1),
432            BlockNumber::from(u32::MAX),
433        ] {
434            let raw = bound_integer(&block_num);
435            assert_eq!(raw, block_num.to_raw_sql(), "write side diverged for {block_num}");
436            assert_eq!(
437                read_integer::<BlockNumber>(raw).unwrap(),
438                BlockNumber::from_raw_sql(raw).unwrap(),
439                "read side diverged for {block_num}",
440            );
441        }
442    }
443
444    #[test]
445    fn note_tag_roundtrip() {
446        for tag in [
447            NoteTag::new(0),
448            NoteTag::new(1),
449            NoteTag::new((1 << 31) - 1),
450            NoteTag::new(1 << 31),
451            NoteTag::new(u32::MAX),
452        ] {
453            let raw = bound_integer(&tag);
454            assert_eq!(raw, i64::from(tag.as_u32()), "tags are stored unsigned: {tag:?}");
455            assert_eq!(
456                read_integer::<NoteTag>(raw).unwrap(),
457                tag,
458                "read side diverged for {tag:?}"
459            );
460        }
461    }
462
463    #[test]
464    fn felt_roundtrip() {
465        // `Felt::MAX` is the largest canonical element; it exceeds `i64::MAX` and is therefore
466        // stored as a negative integer.
467        for felt in [Felt::ZERO, Felt::ONE, Felt::from_u32(u32::MAX), Felt::MAX] {
468            let raw = bound_integer(&felt);
469            #[expect(clippy::cast_possible_wrap, reason = "mirrors the legacy nonce encoding")]
470            let legacy = felt.as_canonical_u64() as i64;
471            assert_eq!(raw, legacy, "write side diverged for {felt}");
472            assert_eq!(read_integer::<Felt>(raw).unwrap(), felt, "read side diverged for {felt}");
473        }
474    }
475
476    #[test]
477    fn storage_slot_name_roundtrip() {
478        let name = StorageSlotName::new("some_component::some_slot").unwrap();
479        let DbValue::Single(Value::Text(text)) = name.to_sql_value() else {
480            panic!("storage slot names are stored as TEXT");
481        };
482        assert_eq!(text, String::from(name.clone()));
483        assert_eq!(
484            StorageSlotName::from_sql_value(DbValueRef::new(ValueRef::Text(text.as_bytes())))
485                .unwrap(),
486            name,
487        );
488    }
489
490    // RANGE CHECKING
491    // ---------------------------------------------------------------------------------------------
492
493    #[test]
494    fn unsigned_int_roundtrip() {
495        assert_eq!(read_integer::<u8>(bound_integer(&u8::MAX)).unwrap(), u8::MAX);
496        assert_eq!(read_integer::<u16>(bound_integer(&u16::MAX)).unwrap(), u16::MAX);
497        assert_eq!(read_integer::<u32>(bound_integer(&u32::MAX)).unwrap(), u32::MAX);
498        assert_eq!(read_integer::<u8>(0).unwrap(), 0);
499    }
500
501    #[test]
502    fn out_of_range_ints_error_instead_of_truncating() {
503        // A cast would have yielded 0, 0, and `u32::MAX` respectively.
504        assert_matches::assert_matches!(
505            read_integer::<u8>(256),
506            Err(DatabaseError::ConversionSqlToRust { to: "u8", .. })
507        );
508        assert_matches::assert_matches!(
509            read_integer::<u16>(65_536),
510            Err(DatabaseError::ConversionSqlToRust { to: "u16", .. })
511        );
512        assert_matches::assert_matches!(
513            read_integer::<u32>(-1),
514            Err(DatabaseError::ConversionSqlToRust { to: "u32", .. })
515        );
516    }
517
518    #[test]
519    fn out_of_range_block_number_errors() {
520        // `BlockNumber` is a u32 on the wire; a wider column value is corruption, not a wrap.
521        assert_matches::assert_matches!(
522            read_integer::<BlockNumber>(i64::from(u32::MAX) + 1),
523            Err(DatabaseError::ConversionSqlToRust { to: "u32", .. })
524        );
525        assert_matches::assert_matches!(
526            read_integer::<BlockNumber>(-1),
527            Err(DatabaseError::ConversionSqlToRust { to: "u32", .. })
528        );
529    }
530
531    // BLOB CODEC
532    // ---------------------------------------------------------------------------------------------
533
534    #[test]
535    fn blob_roundtrip() {
536        let word = Word::from([1u32, 2, 3, 4]);
537        let DbValue::Single(Value::Blob(bytes)) = word.to_sql_value() else {
538            panic!("words are stored as BLOBs");
539        };
540        assert_eq!(Word::from_sql_value(DbValueRef::new(ValueRef::Blob(&bytes))).unwrap(), word);
541    }
542
543    #[test]
544    fn blob_deserialization_error_names_the_type() {
545        assert_matches::assert_matches!(
546            Word::from_sql_value(DbValueRef::new(ValueRef::Blob(&[0xff]))),
547            Err(DatabaseError::ConversionSqlToRust { to: "miden_protocol::Word", .. })
548        );
549    }
550    #[test]
551    fn structured_blobs_are_protobuf() {
552        use miden_protocol::crypto::merkle::mmr::PartialMmr;
553        let value = PartialMmr::default();
554        let DbValue::Single(Value::Blob(bytes)) = value.to_sql_value() else {
555            panic!("expected blob")
556        };
557        let decoded: PartialMmr = miden_node_persistence::decode(&bytes).unwrap();
558        assert_eq!(decoded, value);
559        assert_eq!(
560            PartialMmr::from_sql_value(DbValueRef::new(ValueRef::Blob(&bytes))).unwrap(),
561            value
562        );
563    }
564
565    #[test]
566    fn raw_keys_keep_their_byte_encoding_and_order() {
567        use miden_protocol::utils::serde::Serializable;
568        fn check<T: ToSqlValue + FromSqlValue + Serializable + std::fmt::Debug + PartialEq>(
569            values: Vec<T>,
570        ) {
571            let conn = rusqlite::Connection::open_in_memory().unwrap();
572            conn.execute_batch("CREATE TABLE keys (value BLOB PRIMARY KEY);").unwrap();
573            let mut expected = Vec::new();
574            for value in values {
575                let DbValue::Single(Value::Blob(bytes)) = value.to_sql_value() else {
576                    panic!("expected blob")
577                };
578                assert_eq!(bytes, value.to_bytes());
579                assert_eq!(
580                    T::from_sql_value(DbValueRef::new(ValueRef::Blob(&bytes))).unwrap(),
581                    value
582                );
583                conn.execute("INSERT INTO keys VALUES (?1)", [&bytes]).unwrap();
584                expected.push(bytes);
585            }
586            expected.sort();
587            let mut stmt = conn.prepare("SELECT value FROM keys ORDER BY value").unwrap();
588            let actual = stmt
589                .query_map([], |row| row.get::<_, Vec<u8>>(0))
590                .unwrap()
591                .collect::<Result<Vec<_>, _>>()
592                .unwrap();
593            assert_eq!(actual, expected);
594        }
595        let words: Vec<_> =
596            [256, 1, 42_u32].into_iter().map(|n| Word::from([n, 0, 0, 0])).collect();
597        check(words.clone());
598        check::<miden_protocol::account::AccountId>(
599            [
600                miden_protocol::testing::account_id::ACCOUNT_ID_MAX_ZEROES,
601                miden_protocol::testing::account_id::ACCOUNT_ID_MAX_ONES,
602            ]
603            .into_iter()
604            .map(|id| id.try_into().unwrap())
605            .collect(),
606        );
607        check(
608            words
609                .iter()
610                .copied()
611                .map(miden_protocol::transaction::TransactionId::from_raw)
612                .collect(),
613        );
614        check(words.iter().copied().map(miden_protocol::note::NoteId::from_raw).collect());
615        check(words.iter().copied().map(miden_protocol::note::Nullifier::from_raw).collect());
616        check(
617            words
618                .into_iter()
619                .map(miden_protocol::account::StorageMapKey::from_raw)
620                .collect(),
621        );
622    }
623}