Skip to main content

synd_persistence/sqlite/feed_registry/blob/
mod.rs

1use sha2::{Digest, Sha256};
2use sqlx::{Sqlite, Transaction};
3use synd_registry::{
4    RegistryDbResult,
5    crawl::blob::{BlobRef, PutBlobCommand},
6    db::BlobDb,
7};
8
9use crate::compression::{
10    CompressionAlgo, CompressionCodec, DefaultCompressionCodec, StoredCompressedBytes,
11};
12
13use super::error::{DecodeResultExt, IntoDbResult, SqliteError, SqliteResult};
14
15const DEFAULT_COMPRESSION_ALGO: CompressionAlgo = CompressionAlgo::Zstd;
16
17async fn put<C>(
18    tx: &mut Transaction<'_, Sqlite>,
19    compression: C,
20    command: PutBlobCommand,
21) -> SqliteResult<BlobRef>
22where
23    C: CompressionCodec,
24{
25    let digest = Sha256::digest(&command.bytes);
26    let uncompressed_len = to_i64(command.bytes.len(), "uncompressed blob length")?;
27    let compressed = compression.compress(DEFAULT_COMPRESSION_ALGO, command.bytes.as_slice())?;
28
29    let row = sqlx::query_as::<_, PkRow>(
30        r#"
31            INSERT INTO blob (
32                digest,
33                compression_algo,
34                uncompressed_len,
35                bytes,
36                created_at
37            )
38            VALUES (?, ?, ?, ?, ?)
39            ON CONFLICT(digest) DO UPDATE SET
40                digest = excluded.digest
41            RETURNING pk
42            "#,
43    )
44    .bind(digest.to_vec())
45    .bind(DEFAULT_COMPRESSION_ALGO.as_str())
46    .bind(uncompressed_len)
47    .bind(compressed)
48    .bind(command.created_at)
49    .fetch_one(&mut **tx)
50    .await?;
51
52    Ok(BlobRef::new(row.pk))
53}
54
55async fn load<C>(
56    tx: &mut Transaction<'_, Sqlite>,
57    compression: C,
58    blob: BlobRef,
59) -> SqliteResult<Vec<u8>>
60where
61    C: CompressionCodec,
62{
63    let row = sqlx::query_as::<_, BlobRow>(
64        r#"
65            SELECT
66                compression_algo,
67                uncompressed_len,
68                bytes
69            FROM blob
70            WHERE pk = ?
71            "#,
72    )
73    .bind(blob.pk())
74    .fetch_optional(&mut **tx)
75    .await?;
76
77    let Some(row) = row else {
78        return Err(SqliteError::not_found("blob", blob.pk().to_string()));
79    };
80    let algo = row.compression_algo.parse::<CompressionAlgo>().decode()?;
81    let uncompressed_len = to_usize(row.uncompressed_len, "uncompressed blob length")?;
82
83    Ok(compression.decompress(
84        algo,
85        StoredCompressedBytes {
86            bytes: row.bytes.as_slice(),
87            uncompressed_len,
88        },
89    )?)
90}
91
92#[derive(sqlx::FromRow)]
93struct PkRow {
94    pk: i64,
95}
96
97#[derive(sqlx::FromRow)]
98struct BlobRow {
99    compression_algo: String,
100    uncompressed_len: i64,
101    bytes: Vec<u8>,
102}
103
104fn to_i64(value: usize, field: &'static str) -> SqliteResult<i64> {
105    i64::try_from(value)
106        .map_err(|_| SqliteError::decode_message(format!("{field} exceeds SQLite INTEGER range")))
107}
108
109fn to_usize(value: i64, field: &'static str) -> SqliteResult<usize> {
110    usize::try_from(value)
111        .map_err(|_| SqliteError::decode_message(format!("{field} must be non-negative")))
112}
113
114impl BlobDb for super::SqliteRegistryTx<'_> {
115    async fn put_blob(&mut self, command: PutBlobCommand) -> RegistryDbResult<BlobRef> {
116        put(&mut self.tx, DefaultCompressionCodec, command)
117            .await
118            .db()
119    }
120
121    async fn load_blob(&mut self, blob: BlobRef) -> RegistryDbResult<Vec<u8>> {
122        load(&mut self.tx, DefaultCompressionCodec, blob).await.db()
123    }
124}
125
126#[cfg(test)]
127mod tests;