1use 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#[derive(Debug, Clone)]
35pub enum DbValue {
36 Single(Value),
38 Array(Rc<Vec<Value>>),
40}
41
42impl DbValue {
43 pub fn integer(value: i64) -> Self {
45 Self::Single(Value::Integer(value))
46 }
47
48 pub fn real(value: f64) -> Self {
50 Self::Single(Value::Real(value))
51 }
52
53 pub fn text(value: String) -> Self {
55 Self::Single(Value::Text(value))
56 }
57
58 pub fn blob(value: Vec<u8>) -> Self {
60 Self::Single(Value::Blob(value))
61 }
62
63 pub fn null() -> Self {
65 Self::Single(Value::Null)
66 }
67
68 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#[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 pub fn as_i64(self) -> Result<i64, DatabaseError> {
96 self.0.as_i64().map_err(|err| DatabaseError::deserialization("i64", err))
97 }
98
99 pub fn as_blob(self) -> Result<&'a [u8], DatabaseError> {
101 self.0.as_blob().map_err(|err| DatabaseError::deserialization("blob", err))
102 }
103
104 pub fn as_str(self) -> Result<&'a str, DatabaseError> {
106 self.0.as_str().map_err(|err| DatabaseError::deserialization("str", err))
107 }
108
109 pub fn is_null(self) -> bool {
111 matches!(self.0, ValueRef::Null)
112 }
113}
114
115pub trait ToSqlValue {
120 fn to_sql_value(&self) -> DbValue;
122}
123
124pub trait FromSqlValue: Sized {
126 fn from_sql_value(value: DbValueRef<'_>) -> Result<Self, DatabaseError>;
128}
129
130impl<T: ToSqlValue + ?Sized> ToSqlValue for &T {
132 fn to_sql_value(&self) -> DbValue {
133 (**self).to_sql_value()
134 }
135}
136
137impl 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
152impl 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
252impl 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
294impl 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#[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#[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
368impl_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#[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 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 fn read_integer<T: FromSqlValue>(raw: i64) -> Result<T, DatabaseError> {
420 T::from_sql_value(DbValueRef::new(ValueRef::Integer(raw)))
421 }
422
423 #[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 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 #[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 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 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 #[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}