use rusqlite::Connection;
use std::collections::HashSet;
pub const STAGE_EMBED: &str = "embed";
pub const STAGE_FACES: &str = "faces";
pub const STAGE_THUMBNAIL: &str = "thumbnail";
pub const FAILURE_THRESHOLD: u32 = 2;
pub fn ensure_table(conn: &Connection) -> rusqlite::Result<()> {
conn.execute_batch(
"CREATE TABLE IF NOT EXISTS decode_failures (
hash TEXT NOT NULL,
stage TEXT NOT NULL,
error TEXT NOT NULL,
fail_count INTEGER NOT NULL DEFAULT 1,
last_failed_at TEXT DEFAULT (datetime('now')),
PRIMARY KEY (hash, stage)
);",
)
}
pub fn record(conn: &Connection, hash: &str, stage: &str, error: &str) -> rusqlite::Result<()> {
conn.execute(
"INSERT INTO decode_failures (hash, stage, error, fail_count, last_failed_at)
VALUES (?1, ?2, ?3, 1, datetime('now'))
ON CONFLICT(hash, stage) DO UPDATE SET
fail_count = fail_count + 1,
error = excluded.error,
last_failed_at = excluded.last_failed_at",
rusqlite::params![hash, stage, error],
)?;
Ok(())
}
pub fn clear(conn: &Connection, hash: &str, stage: &str) -> rusqlite::Result<()> {
conn.execute(
"DELETE FROM decode_failures WHERE hash = ?1 AND stage = ?2",
rusqlite::params![hash, stage],
)?;
Ok(())
}
pub fn clear_stage(conn: &Connection, stage: &str) -> rusqlite::Result<usize> {
conn.execute(
"DELETE FROM decode_failures WHERE stage = ?1",
rusqlite::params![stage],
)
}
pub fn failed_hashes(
conn: &Connection,
stage: &str,
min_count: u32,
) -> rusqlite::Result<HashSet<String>> {
if !crate::db::table_exists(conn, "decode_failures")? {
return Ok(HashSet::new());
}
let mut stmt =
conn.prepare("SELECT hash FROM decode_failures WHERE stage = ?1 AND fail_count >= ?2")?;
let rows = stmt.query_map(rusqlite::params![stage, min_count], |r| {
r.get::<_, String>(0)
})?;
rows.collect()
}
pub fn fail_count(conn: &Connection, hash: &str, stage: &str) -> rusqlite::Result<u32> {
conn.query_row(
"SELECT fail_count FROM decode_failures WHERE hash = ?1 AND stage = ?2",
rusqlite::params![hash, stage],
|r| r.get(0),
)
.or_else(|e| match e {
rusqlite::Error::QueryReturnedNoRows => Ok(0),
other => Err(other),
})
}
#[cfg(test)]
mod tests {
use super::*;
fn db() -> Connection {
let conn = Connection::open_in_memory().unwrap();
ensure_table(&conn).unwrap();
conn
}
#[test]
fn record_creates_a_row_at_count_one() {
let conn = db();
record(&conn, "h1", STAGE_EMBED, "boom").unwrap();
assert_eq!(fail_count(&conn, "h1", STAGE_EMBED).unwrap(), 1);
}
#[test]
fn record_increments_on_repeat_and_keeps_latest_error() {
let conn = db();
record(&conn, "h1", STAGE_EMBED, "first").unwrap();
record(&conn, "h1", STAGE_EMBED, "second").unwrap();
assert_eq!(fail_count(&conn, "h1", STAGE_EMBED).unwrap(), 2);
let err: String = conn
.query_row(
"SELECT error FROM decode_failures WHERE hash='h1' AND stage='embed'",
[],
|r| r.get(0),
)
.unwrap();
assert_eq!(err, "second", "the most recent error is kept");
}
#[test]
fn failed_hashes_respects_the_threshold() {
let conn = db();
record(&conn, "once", STAGE_EMBED, "e").unwrap();
record(&conn, "twice", STAGE_EMBED, "e").unwrap();
record(&conn, "twice", STAGE_EMBED, "e").unwrap();
let set = failed_hashes(&conn, STAGE_EMBED, FAILURE_THRESHOLD).unwrap();
assert!(
!set.contains("once"),
"a single failure must not be skipped (could be transient)"
);
assert!(set.contains("twice"), "two failures crosses the threshold");
}
#[test]
fn clear_removes_a_hash_so_it_is_tried_again() {
let conn = db();
record(&conn, "h1", STAGE_EMBED, "e").unwrap();
record(&conn, "h1", STAGE_EMBED, "e").unwrap();
clear(&conn, "h1", STAGE_EMBED).unwrap();
assert_eq!(fail_count(&conn, "h1", STAGE_EMBED).unwrap(), 0);
assert!(failed_hashes(&conn, STAGE_EMBED, FAILURE_THRESHOLD)
.unwrap()
.is_empty());
}
#[test]
fn clear_stage_removes_only_that_stage() {
let conn = db();
record(&conn, "h1", STAGE_EMBED, "e").unwrap();
record(&conn, "h2", STAGE_EMBED, "e").unwrap();
record(&conn, "h1", STAGE_FACES, "e").unwrap();
let removed = clear_stage(&conn, STAGE_EMBED).unwrap();
assert_eq!(removed, 2, "both embed rows removed");
assert_eq!(
fail_count(&conn, "h1", STAGE_FACES).unwrap(),
1,
"the faces row survives"
);
}
#[test]
fn stages_are_independent() {
let conn = db();
record(&conn, "h1", STAGE_EMBED, "e").unwrap();
record(&conn, "h1", STAGE_EMBED, "e").unwrap();
assert_eq!(fail_count(&conn, "h1", STAGE_FACES).unwrap(), 0);
assert!(!failed_hashes(&conn, STAGE_FACES, FAILURE_THRESHOLD)
.unwrap()
.contains("h1"));
}
#[test]
fn failed_hashes_on_a_missing_table_is_empty_not_an_error() {
let conn = Connection::open_in_memory().unwrap();
assert!(failed_hashes(&conn, STAGE_EMBED, FAILURE_THRESHOLD)
.unwrap()
.is_empty());
}
}