mkv

Unnamed repository; edit this file 'description' to name the repository.
Log | Files | Refs | README

db.rs (7033B)


      1 use rusqlite::{Connection, OpenFlags, params};
      2 use std::sync::{Arc, Mutex};
      3 
      4 use crate::error::AppError;
      5 
      6 #[derive(Debug, Clone)]
      7 pub struct Record {
      8     pub key: String,
      9     pub volumes: Vec<String>,
     10     pub size: Option<i64>,
     11 }
     12 
     13 fn apply_pragmas(conn: &Connection) {
     14     conn.execute_batch(
     15         "PRAGMA journal_mode = WAL;
     16          PRAGMA synchronous = NORMAL;
     17          PRAGMA busy_timeout = 5000;
     18          PRAGMA temp_store = memory;
     19          PRAGMA cache_size = -64000;
     20          PRAGMA mmap_size = 268435456;",
     21     )
     22     .expect("failed to set pragmas");
     23 }
     24 
     25 fn parse_volumes(s: &str) -> Vec<String> {
     26     serde_json::from_str(s).unwrap_or_default()
     27 }
     28 
     29 fn encode_volumes(v: &[String]) -> String {
     30     serde_json::to_string(v).unwrap()
     31 }
     32 
     33 /// Examples: "abc" -> Some("abd"), "ab\xff" -> Some("ac"), "\xff\xff" -> None
     34 pub fn prefix_upper_bound(prefix: &str) -> Option<String> {
     35     let mut bytes = prefix.as_bytes().to_vec();
     36     while let Some(last) = bytes.pop() {
     37         if last < 0xFF {
     38             bytes.push(last + 1);
     39             return Some(String::from_utf8_lossy(&bytes).into_owned());
     40         }
     41     }
     42     None
     43 }
     44 
     45 #[derive(Clone)]
     46 pub struct Db {
     47     conn: Arc<Mutex<Connection>>,
     48 }
     49 
     50 impl Db {
     51     pub fn new(path: &str) -> Self {
     52         let conn = Connection::open_with_flags(
     53             path,
     54             OpenFlags::SQLITE_OPEN_READ_WRITE
     55                 | OpenFlags::SQLITE_OPEN_CREATE
     56                 | OpenFlags::SQLITE_OPEN_NO_MUTEX
     57                 | OpenFlags::SQLITE_OPEN_URI,
     58         )
     59         .expect("failed to open database");
     60         apply_pragmas(&conn);
     61         conn.execute_batch(
     62             "CREATE TABLE IF NOT EXISTS kv (
     63                 key         TEXT PRIMARY KEY,
     64                 volumes     TEXT NOT NULL,
     65                 size        INTEGER,
     66                 created_at  INTEGER DEFAULT (unixepoch())
     67             );",
     68         )
     69         .expect("failed to create tables");
     70         Self {
     71             conn: Arc::new(Mutex::new(conn)),
     72         }
     73     }
     74 
     75     pub async fn get(&self, key: &str) -> Result<Record, AppError> {
     76         let conn = self.conn.clone();
     77         let key = key.to_string();
     78         tokio::task::spawn_blocking(move || {
     79             let conn = conn.lock().unwrap();
     80             let mut stmt =
     81                 conn.prepare_cached("SELECT key, volumes, size FROM kv WHERE key = ?1")?;
     82             Ok(stmt.query_row(params![key], |row| {
     83                 let vj: String = row.get(1)?;
     84                 Ok(Record {
     85                     key: row.get(0)?,
     86                     volumes: parse_volumes(&vj),
     87                     size: row.get(2)?,
     88                 })
     89             })?)
     90         })
     91         .await
     92         .unwrap()
     93     }
     94 
     95     pub async fn list_keys(&self, prefix: &str) -> Result<Vec<String>, AppError> {
     96         let conn = self.conn.clone();
     97         let prefix = prefix.to_string();
     98         tokio::task::spawn_blocking(move || {
     99             let conn = conn.lock().unwrap();
    100             if prefix.is_empty() {
    101                 let mut stmt = conn.prepare_cached("SELECT key FROM kv ORDER BY key")?;
    102                 let keys = stmt
    103                     .query_map([], |row| row.get(0))?
    104                     .collect::<Result<Vec<String>, _>>()?;
    105                 return Ok(keys);
    106             }
    107             let upper = prefix_upper_bound(&prefix);
    108             let keys = match &upper {
    109                 Some(end) => {
    110                     let mut stmt = conn.prepare_cached(
    111                         "SELECT key FROM kv WHERE key >= ?1 AND key < ?2 ORDER BY key",
    112                     )?;
    113                     stmt.query_map(params![prefix, end], |row| row.get(0))?
    114                         .collect::<Result<Vec<String>, _>>()?
    115                 }
    116                 None => {
    117                     let mut stmt =
    118                         conn.prepare_cached("SELECT key FROM kv WHERE key >= ?1 ORDER BY key")?;
    119                     stmt.query_map(params![prefix], |row| row.get(0))?
    120                         .collect::<Result<Vec<String>, _>>()?
    121                 }
    122             };
    123             Ok(keys)
    124         })
    125         .await
    126         .unwrap()
    127     }
    128 
    129     pub async fn put(
    130         &self,
    131         key: String,
    132         volumes: Vec<String>,
    133         size: Option<i64>,
    134     ) -> Result<(), AppError> {
    135         let conn = self.conn.clone();
    136         tokio::task::spawn_blocking(move || {
    137             let conn = conn.lock().unwrap();
    138             conn.prepare_cached(
    139                 "INSERT INTO kv (key, volumes, size) VALUES (?1, ?2, ?3)
    140                  ON CONFLICT(key) DO UPDATE SET volumes = ?2, size = ?3",
    141             )?
    142             .execute(params![key, encode_volumes(&volumes), size])?;
    143             Ok(())
    144         })
    145         .await
    146         .unwrap()
    147     }
    148 
    149     pub async fn delete(&self, key: String) -> Result<(), AppError> {
    150         let conn = self.conn.clone();
    151         tokio::task::spawn_blocking(move || {
    152             let conn = conn.lock().unwrap();
    153             conn.prepare_cached("DELETE FROM kv WHERE key = ?1")?
    154                 .execute(params![key])?;
    155             Ok(())
    156         })
    157         .await
    158         .unwrap()
    159     }
    160 
    161     pub async fn bulk_put(
    162         &self,
    163         records: Vec<(String, Vec<String>, Option<i64>)>,
    164     ) -> Result<(), AppError> {
    165         let conn = self.conn.clone();
    166         tokio::task::spawn_blocking(move || {
    167             let conn = conn.lock().unwrap();
    168             conn.execute_batch("BEGIN")?;
    169             let mut stmt = conn.prepare_cached(
    170                 "INSERT INTO kv (key, volumes, size) VALUES (?1, ?2, ?3)
    171                  ON CONFLICT(key) DO UPDATE SET volumes = ?2, size = ?3",
    172             )?;
    173             for (key, volumes, size) in &records {
    174                 stmt.execute(params![key, encode_volumes(volumes), size])?;
    175             }
    176             drop(stmt);
    177             conn.execute_batch("COMMIT")?;
    178             Ok(())
    179         })
    180         .await
    181         .unwrap()
    182     }
    183 
    184     pub fn all_records_sync(&self) -> Result<Vec<Record>, AppError> {
    185         let conn = self.conn.lock().unwrap();
    186         let mut stmt = conn.prepare_cached("SELECT key, volumes, size FROM kv")?;
    187         let records = stmt
    188             .query_map([], |row| {
    189                 let vj: String = row.get(1)?;
    190                 Ok(Record {
    191                     key: row.get(0)?,
    192                     volumes: parse_volumes(&vj),
    193                     size: row.get(2)?,
    194                 })
    195             })?
    196             .collect::<Result<Vec<_>, _>>()?;
    197         Ok(records)
    198     }
    199 }
    200 
    201 #[cfg(test)]
    202 mod tests {
    203     use super::*;
    204 
    205     #[test]
    206     fn test_prefix_upper_bound_range_correctness() {
    207         let prefix = "foo";
    208         let upper = prefix_upper_bound(prefix).unwrap();
    209         let upper = upper.as_str();
    210 
    211         // in range [prefix, upper)
    212         assert!("foo" >= prefix && "foo" < upper);
    213         assert!("foo/bar" >= prefix && "foo/bar" < upper);
    214         assert!("foobar" >= prefix && "foobar" < upper);
    215         assert!("foo\x7f" >= prefix && "foo\x7f" < upper);
    216 
    217         // out of range
    218         assert!("fop" >= upper);
    219         assert!("fon" < prefix);
    220     }
    221 }