use std::fmt;
use std::fs::{File, TryLockError};
use std::path::{Path, PathBuf};
use std::sync::{Arc, Mutex, MutexGuard};
use std::time::{SystemTime, UNIX_EPOCH};
use efema_proto::{Cursor, Hash, StreamId, StreamName};
use rusqlite::{Connection, OptionalExtension, params};
use crate::Error;
const SCHEMA_VERSION: i64 = 1;
const MIGRATIONS: &[&str] = &[r#"
CREATE TABLE device (
only INTEGER PRIMARY KEY CHECK (only = 1),
id BLOB NOT NULL CHECK (length(id) = 16),
created_at INTEGER NOT NULL
) STRICT;
CREATE TABLE streams (
name TEXT PRIMARY KEY,
stream BLOB NOT NULL CHECK (length(stream) = 16),
lock BLOB NOT NULL,
seq INTEGER NOT NULL DEFAULT 0 CHECK (seq >= 0),
hash BLOB NOT NULL CHECK (length(hash) = 32),
updated_at INTEGER NOT NULL
) STRICT;
"#];
#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct DeviceId([u8; 16]);
impl DeviceId {
pub const fn from_bytes(bytes: [u8; 16]) -> Self {
Self(bytes)
}
pub const fn as_bytes(&self) -> &[u8; 16] {
&self.0
}
pub fn short(&self) -> String {
self.to_string()[..8].to_string()
}
}
impl fmt::Display for DeviceId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
for byte in self.0 {
write!(f, "{byte:02x}")?;
}
Ok(())
}
}
impl fmt::Debug for DeviceId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "DeviceId({self})")
}
}
#[derive(Clone)]
pub struct State {
inner: Arc<Inner>,
}
struct Inner {
path: PathBuf,
device: DeviceId,
conn: Mutex<Connection>,
_lock: File,
}
#[derive(Clone, Debug)]
pub(crate) struct Known {
pub(crate) stream: StreamId,
pub(crate) lock: Vec<u8>,
pub(crate) cursor: Cursor,
}
impl State {
pub fn open(path: impl AsRef<Path>) -> Result<Self, Error> {
let path = path.as_ref().to_path_buf();
let lock_path = lock_path(&path);
let lock = File::options()
.create(true)
.truncate(false)
.write(true)
.open(&lock_path)
.map_err(|source| Error::StateLock { path: lock_path.clone(), source })?;
match lock.try_lock() {
Ok(()) => {}
Err(TryLockError::WouldBlock) => return Err(Error::StateBusy(path)),
Err(TryLockError::Error(source)) => return Err(Error::StateLock { path: lock_path, source }),
}
let sqlite = |source| Error::State { path: path.clone(), source };
let mut conn = Connection::open(&path).map_err(sqlite)?;
conn.pragma_update(None, "journal_mode", "WAL").map_err(sqlite)?;
conn.pragma_update(None, "synchronous", "FULL").map_err(sqlite)?;
migrate(&mut conn, &path)?;
let existing: Option<Vec<u8>> =
conn.query_row("SELECT id FROM device", [], |row| row.get(0)).optional().map_err(sqlite)?;
let device = match existing {
Some(bytes) => DeviceId(bytes.try_into().expect("the schema keeps a device identity at 16 bytes")),
None => {
let mut bytes = [0u8; 16];
getrandom::fill(&mut bytes).map_err(|_| Error::Random)?;
conn.execute(
"INSERT INTO device (only, id, created_at) VALUES (1, ?1, ?2)",
params![bytes, unix_now()],
)
.map_err(sqlite)?;
DeviceId(bytes)
}
};
Ok(Self { inner: Arc::new(Inner { path, device, conn: Mutex::new(conn), _lock: lock }) })
}
pub fn device(&self) -> DeviceId {
self.inner.device
}
pub fn path(&self) -> &Path {
&self.inner.path
}
pub fn forget(&self, stream: &StreamName) -> Result<(), Error> {
self.conn()
.execute("DELETE FROM streams WHERE name = ?1", [stream.as_str()])
.map(drop)
.map_err(|source| self.error(source))
}
pub(crate) fn known(&self, stream: &StreamName) -> Result<Option<Known>, Error> {
self.conn()
.query_row("SELECT stream, lock, seq, hash FROM streams WHERE name = ?1", [stream.as_str()], |row| {
let id = StreamId::from_bytes(blob(row.get(0)?)?);
let seq: i64 = row.get(2)?;
Ok(Known {
stream: id,
lock: row.get(1)?,
cursor: Cursor {
stream: id,
seq: u64::try_from(seq).expect("the schema keeps positions non-negative"),
hash: Hash::from_bytes(blob(row.get(3)?)?),
},
})
})
.optional()
.map_err(|source| self.error(source))
}
pub(crate) fn join(&self, stream: &StreamName, id: StreamId, lock: &[u8]) -> Result<(), Error> {
let start = Cursor::start(id);
self.conn()
.execute(
"INSERT INTO streams (name, stream, lock, seq, hash, updated_at) VALUES (?1, ?2, ?3, 0, ?4, ?5)",
params![stream.as_str(), id.as_bytes(), lock, start.hash.as_bytes(), unix_now()],
)
.map(drop)
.map_err(|source| self.error(source))
}
pub(crate) fn advance(&self, stream: &StreamName, cursor: &Cursor) -> Result<(), Error> {
self.conn()
.execute(
"UPDATE streams SET seq = ?1, hash = ?2, updated_at = ?3
WHERE name = ?4 AND stream = ?5 AND seq < ?1",
params![
i64::try_from(cursor.seq).expect("positions stay below 2^63"),
cursor.hash.as_bytes(),
unix_now(),
stream.as_str(),
cursor.stream.as_bytes()
],
)
.map(drop)
.map_err(|source| self.error(source))
}
fn conn(&self) -> MutexGuard<'_, Connection> {
self.inner.conn.lock().unwrap_or_else(|poisoned| poisoned.into_inner())
}
fn error(&self, source: rusqlite::Error) -> Error {
Error::State { path: self.inner.path.clone(), source }
}
}
impl fmt::Debug for State {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("State").field("path", &self.inner.path).field("device", &self.inner.device).finish()
}
}
fn lock_path(path: &Path) -> PathBuf {
let mut name = path.file_name().map(|n| n.to_os_string()).unwrap_or_default();
name.push(".lock");
path.with_file_name(name)
}
fn migrate(conn: &mut Connection, path: &Path) -> Result<(), Error> {
let sqlite = |source| Error::State { path: path.to_path_buf(), source };
let version: i64 = conn.pragma_query_value(None, "user_version", |row| row.get(0)).map_err(sqlite)?;
if version > SCHEMA_VERSION {
return Err(Error::StateNewer { path: path.to_path_buf(), found: version });
}
for (index, migration) in MIGRATIONS.iter().enumerate().skip(version as usize) {
let tx = conn.transaction().map_err(sqlite)?;
tx.execute_batch(migration).map_err(sqlite)?;
tx.pragma_update(None, "user_version", index as i64 + 1).map_err(sqlite)?;
tx.commit().map_err(sqlite)?;
}
Ok(())
}
fn blob<const N: usize>(bytes: Vec<u8>) -> rusqlite::Result<[u8; N]> {
bytes.try_into().map_err(|_| rusqlite::Error::InvalidColumnType(0, "blob".into(), rusqlite::types::Type::Blob))
}
fn unix_now() -> i64 {
SystemTime::now().duration_since(UNIX_EPOCH).map_or(0, |d| d.as_secs() as i64)
}
#[cfg(test)]
mod tests {
use super::*;
fn name(text: &str) -> StreamName {
text.parse().unwrap()
}
#[test]
fn the_device_keeps_its_identity_across_runs() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("sync.sqlite3");
let first = State::open(&path).unwrap().device();
let again = State::open(&path).unwrap().device();
assert_eq!(first, again);
let other = State::open(dir.path().join("other.sqlite3")).unwrap().device();
assert_ne!(first, other);
}
#[test]
fn one_process_at_a_time() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("sync.sqlite3");
let held = State::open(&path).unwrap();
assert!(matches!(State::open(&path), Err(Error::StateBusy(_))));
drop(held);
State::open(&path).unwrap();
}
#[test]
fn a_cursor_only_moves_forward_within_its_incarnation() {
let dir = tempfile::tempdir().unwrap();
let state = State::open(dir.path().join("sync.sqlite3")).unwrap();
let id = StreamId::from_bytes([3; 16]);
state.join(&name("s"), id, b"lock").unwrap();
let at = |seq| Cursor { stream: id, seq, hash: Hash::from_bytes([seq as u8; 32]) };
state.advance(&name("s"), &at(5)).unwrap();
state.advance(&name("s"), &at(3)).unwrap();
assert_eq!(state.known(&name("s")).unwrap().unwrap().cursor, at(5));
let foreign = Cursor { stream: StreamId::from_bytes([4; 16]), ..at(9) };
state.advance(&name("s"), &foreign).unwrap();
assert_eq!(state.known(&name("s")).unwrap().unwrap().cursor, at(5));
state.forget(&name("s")).unwrap();
assert!(state.known(&name("s")).unwrap().is_none());
}
#[test]
fn a_newer_state_is_refused_rather_than_misread() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("sync.sqlite3");
drop(State::open(&path).unwrap());
Connection::open(&path).unwrap().pragma_update(None, "user_version", SCHEMA_VERSION + 1).unwrap();
assert!(matches!(State::open(&path), Err(Error::StateNewer { .. })));
}
}