use std::fs;
use std::path::{Path, PathBuf};
use std::time::Duration;
use rusqlite::Connection;
use crate::error::Result;
const TMP_SUFFIX: &str = ".zkvtmp";
#[derive(Debug)]
pub struct Database {
conn: Connection,
backing_tmp: Option<PathBuf>,
}
impl Database {
pub fn conn(&self) -> &Connection {
&self.conn
}
pub fn open_in_memory() -> Result<Database> {
let conn = Connection::open_in_memory()?;
let db = Database {
conn,
backing_tmp: None,
};
db.migrate()?;
db.migrate_data()?;
Ok(db)
}
pub fn migrate(&self) -> Result<()> {
let c = &self.conn;
c.execute_batch(
r#"
-- 分类:支持层级(parent_id 自引用)
CREATE TABLE IF NOT EXISTS categories (
id INTEGER PRIMARY KEY,
name TEXT NOT NULL,
parent_id INTEGER REFERENCES categories(id) ON DELETE SET NULL,
sort_order INTEGER NOT NULL DEFAULT 0,
created_at INTEGER NOT NULL
);
-- 标签
CREATE TABLE IF NOT EXISTS tags (
id INTEGER PRIMARY KEY,
name TEXT NOT NULL UNIQUE
);
-- 条目:核心表,统一承载所有类型
CREATE TABLE IF NOT EXISTS items (
id INTEGER PRIMARY KEY,
type TEXT NOT NULL,
title TEXT NOT NULL,
category_id INTEGER REFERENCES categories(id) ON DELETE SET NULL,
data TEXT NOT NULL,
favorite INTEGER NOT NULL DEFAULT 0,
search_text TEXT NOT NULL DEFAULT '',
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_items_category ON items(category_id);
CREATE INDEX IF NOT EXISTS idx_items_type ON items(type);
-- 条目 ↔ 标签(多对多)
CREATE TABLE IF NOT EXISTS item_tags (
item_id INTEGER REFERENCES items(id) ON DELETE CASCADE,
tag_id INTEGER REFERENCES tags(id) ON DELETE CASCADE,
PRIMARY KEY (item_id, tag_id)
);
-- 附件:图片/电子档内嵌为 BLOB
CREATE TABLE IF NOT EXISTS attachments (
id INTEGER PRIMARY KEY,
item_id INTEGER REFERENCES items(id) ON DELETE CASCADE,
filename TEXT NOT NULL,
mime_type TEXT,
size INTEGER NOT NULL,
blob BLOB NOT NULL,
created_at INTEGER NOT NULL
);
-- 字段模板(A1):承载内置/自定义模板。fields 为 FieldSpec 数组的 JSON。
CREATE TABLE IF NOT EXISTS templates (
id TEXT PRIMARY KEY,
name TEXT NOT NULL,
fields TEXT NOT NULL,
built_in INTEGER NOT NULL DEFAULT 0,
created_at INTEGER NOT NULL
);
"#,
)?;
c.execute_batch(
r#"
CREATE VIRTUAL TABLE IF NOT EXISTS items_fts USING fts5(
title, search_text,
content='items', content_rowid='id'
);
-- items 增/改时重建对应行的 FTS 索引
CREATE TRIGGER IF NOT EXISTS items_ai AFTER INSERT ON items BEGIN
INSERT INTO items_fts(rowid, title, search_text)
VALUES (new.id, new.title, new.search_text);
END;
CREATE TRIGGER IF NOT EXISTS items_ad AFTER DELETE ON items BEGIN
INSERT INTO items_fts(items_fts, rowid, title, search_text)
VALUES ('delete', old.id, old.title, old.search_text);
END;
CREATE TRIGGER IF NOT EXISTS items_au AFTER UPDATE ON items BEGIN
INSERT INTO items_fts(items_fts, rowid, title, search_text)
VALUES ('delete', old.id, old.title, old.search_text);
INSERT INTO items_fts(rowid, title, search_text)
VALUES (new.id, new.title, new.search_text);
END;
"#,
)?;
Ok(())
}
pub fn migrate_data(&self) -> Result<()> {
let c = &self.conn;
let v: i64 = c.query_row("PRAGMA user_version", [], |r| r.get(0))?;
if v >= 1 {
return Ok(());
}
let mut stmt = c.prepare("SELECT id, type, data FROM items")?;
let rows: Vec<(i64, String, String)> = stmt
.query_map([], |r| {
Ok((
r.get::<_, i64>(0)?,
r.get::<_, String>(1)?,
r.get::<_, String>(2)?,
))
})?
.filter_map(|r| r.ok())
.collect();
drop(stmt);
c.execute_batch("BEGIN IMMEDIATE")?;
let tx_result: Result<()> = (|| {
for (id, ty, data) in &rows {
if let Ok(fields) = serde_json::from_str::<Vec<crate::model::Field>>(data) {
let new_json = serde_json::to_string(&fields)?;
let st = crate::model::fields_search_text(&fields);
c.execute(
"UPDATE items SET type=?1, data=?2, search_text=?3 WHERE id=?4",
rusqlite::params![ty, new_json, st, id],
)?;
continue;
}
if let Ok(legacy) = serde_json::from_str::<crate::model::LegacyItemData>(data) {
let (tpl, fields) = crate::model::legacy_to_fields(legacy);
let new_json = serde_json::to_string(&fields)?;
let st = crate::model::fields_search_text(&fields);
c.execute(
"UPDATE items SET type=?1, data=?2, search_text=?3 WHERE id=?4",
rusqlite::params![tpl, new_json, st, id],
)?;
continue;
}
}
Ok(())
})();
match tx_result {
Ok(()) => {
c.execute_batch("COMMIT")?;
c.execute_batch("PRAGMA user_version = 1")?;
Ok(())
}
Err(e) => {
let _ = c.execute_batch("ROLLBACK");
Err(e)
}
}
}
pub fn dump_bytes(&self) -> Result<Vec<u8>> {
let tmp = secure_tmp_path("zkv_dump");
let sql = format!("VACUUM INTO '{}'", tmp.display());
let result = (|| -> Result<Vec<u8>> {
self.conn.execute_batch(&sql)?;
let bytes = fs::read(&tmp)?;
Ok(bytes)
})();
let _ = fs::remove_file(&tmp);
result
}
pub fn from_bytes(bytes: &[u8]) -> Result<Database> {
let tmp = secure_tmp_path("zkv_load");
{
use std::io::Write;
let mut f = open_secure(&tmp)?;
f.write_all(bytes)?;
f.sync_all()?;
}
let res = (|| -> Result<Database> {
let src = Connection::open(&tmp)?;
let mut dst = Connection::open_in_memory()?;
{
let b = rusqlite::backup::Backup::new(&src, &mut dst)?;
b.run_to_completion(100, Duration::from_millis(250), None)?;
}
drop(src);
let db = Database {
conn: dst,
backing_tmp: None,
};
db.migrate()?;
db.migrate_data()?;
Ok(db)
})();
let _ = fs::remove_file(&tmp);
res
}
}
impl Drop for Database {
fn drop(&mut self) {
if let Some(p) = self.backing_tmp.take() {
let _ = fs::remove_file(&p);
}
}
}
fn secure_tmp_base(prefix: &str) -> PathBuf {
let mut buf = [0u8; 16];
getrandom::fill(&mut buf).expect("getrandom::fill failed for temp name");
let hex: String = buf.iter().map(|b| format!("{b:02x}")).collect();
let mut name = String::from(prefix);
name.push('-');
name.push_str(&hex);
std::env::temp_dir().join(name)
}
fn secure_tmp_path(prefix: &str) -> PathBuf {
let mut p = secure_tmp_base(prefix);
p.set_extension(TMP_SUFFIX.trim_start_matches('.'));
p
}
#[cfg(unix)]
fn open_secure(path: &Path) -> Result<std::fs::File> {
use std::os::unix::fs::OpenOptionsExt;
Ok(std::fs::OpenOptions::new()
.write(true)
.create_new(true)
.truncate(true)
.mode(0o600)
.open(path)?)
}
#[cfg(not(unix))]
fn open_secure(path: &Path) -> Result<std::fs::File> {
Ok(std::fs::OpenOptions::new()
.write(true)
.create_new(true)
.truncate(true)
.open(path)?)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn open_in_memory_and_migrate_ok() {
let db = Database::open_in_memory().expect("open_in_memory");
let names: Vec<String> = db
.conn()
.prepare("SELECT name FROM sqlite_master WHERE type='table' ORDER BY name")
.unwrap()
.query_map([], |r| r.get::<_, String>(0))
.unwrap()
.filter_map(|r| r.ok())
.collect();
for expected in [
"categories",
"tags",
"items",
"item_tags",
"attachments",
"items_fts",
] {
assert!(
names.iter().any(|n| n == expected),
"missing table: {expected} (have {:?})",
names
);
}
let trigs: Vec<String> = db
.conn()
.prepare("SELECT name FROM sqlite_master WHERE type='trigger' ORDER BY name")
.unwrap()
.query_map([], |r| r.get::<_, String>(0))
.unwrap()
.filter_map(|r| r.ok())
.collect();
assert!(trigs.iter().any(|n| n == "items_ai"));
assert!(trigs.iter().any(|n| n == "items_ad"));
assert!(trigs.iter().any(|n| n == "items_au"));
}
#[test]
fn insert_item_and_fts_hit() {
let db = Database::open_in_memory().unwrap();
let now = 1_700_000_000i64;
db.conn()
.execute(
"INSERT INTO items(type,title,data,search_text,created_at,updated_at)
VALUES (?1,?2,?3,?4,?5,?6)",
rusqlite::params![
"password",
"GitHub Login",
"{}",
"github.com myuser",
now,
now,
],
)
.unwrap();
let hit: bool = db
.conn()
.prepare("SELECT 1 FROM items_fts WHERE items_fts MATCH ?1 LIMIT 1")
.unwrap()
.query_map(["github"], |r| r.get::<_, i64>(0))
.unwrap()
.filter_map(|r| r.ok())
.next()
.is_some();
assert!(hit, "FTS5 应命中 'github'");
let miss: bool = db
.conn()
.prepare("SELECT 1 FROM items_fts WHERE items_fts MATCH ?1 LIMIT 1")
.unwrap()
.query_map(["nomatchxyz"], |r| r.get::<_, i64>(0))
.unwrap()
.filter_map(|r| r.ok())
.next()
.is_some();
assert!(!miss);
}
#[test]
fn dump_from_bytes_roundtrip_preserves_data() {
let db = Database::open_in_memory().unwrap();
db.conn().execute_batch("PRAGMA user_version = 0").unwrap();
let now = 1_700_000_000i64;
db.conn()
.execute(
"INSERT INTO items(type,title,data,search_text,created_at,updated_at)
VALUES (?1,?2,?3,?4,?5,?6)",
rusqlite::params![
"note",
"My Note",
"{\"type\":\"note\",\"format\":\"text\",\"content\":\"hi\"}",
"hi body",
now,
now
],
)
.unwrap();
db.conn()
.execute(
"INSERT INTO tags(name) VALUES ('personal')",
[],
)
.unwrap();
let bytes = db.dump_bytes().expect("dump");
assert!(!bytes.is_empty(), "dump 应产出非空字节");
let db2 = Database::from_bytes(&bytes).expect("from_bytes");
let cnt: i64 = db2
.conn()
.query_row("SELECT COUNT(*) FROM items", [], |r| r.get::<_, i64>(0))
.unwrap();
assert_eq!(cnt, 1);
let tag_cnt: i64 = db2
.conn()
.query_row("SELECT COUNT(*) FROM tags", [], |r| r.get::<_, i64>(0))
.unwrap();
assert_eq!(tag_cnt, 1);
let data: String = db2
.conn()
.query_row(
"SELECT data FROM items WHERE title='My Note'",
[],
|r| r.get::<_, String>(0),
)
.unwrap();
assert!(data.contains("\"name\":\"content\""));
assert!(data.contains("\"value\":\"hi\""));
let fts_hit: bool = db2
.conn()
.prepare("SELECT 1 FROM items_fts WHERE items_fts MATCH ?1 LIMIT 1")
.unwrap()
.query_map(["note"], |r| r.get::<_, i64>(0))
.unwrap()
.filter_map(|r| r.ok())
.next()
.is_some();
assert!(fts_hit, "round-trip 后 FTS5 仍应命中");
}
fn read_item_row(conn: &Connection, id: i64) -> (String, String, String) {
conn.query_row(
"SELECT type, data, search_text FROM items WHERE id = ?1",
rusqlite::params![id],
|r| {
Ok((
r.get::<_, String>(0)?,
r.get::<_, String>(1)?,
r.get::<_, String>(2)?,
))
},
)
.unwrap()
}
fn user_version(conn: &Connection) -> i64 {
conn.query_row("PRAGMA user_version", [], |r| r.get::<_, i64>(0))
.unwrap()
}
#[test]
fn migrate_data_converts_three_variants_and_corrupt() {
let db = Database::open_in_memory().unwrap();
let conn = db.conn();
conn.execute_batch("PRAGMA user_version = 0").unwrap();
assert_eq!(user_version(conn), 0);
conn.execute(
"INSERT INTO items(type,title,data,search_text,created_at,updated_at)
VALUES ('password','t','{\"type\":\"password\",\"username\":\"u\",\"password\":\"s3cret\",\"url\":\"x\",\"totp_secret\":\"T\",\"notes\":\"n\"}','old',1,1)",
[],
).unwrap();
conn.execute(
"INSERT INTO items(type,title,data,search_text,created_at,updated_at)
VALUES ('note','n','{\"type\":\"note\",\"format\":\"text\",\"content\":\"body\"}','old',1,1)",
[],
).unwrap();
conn.execute(
"INSERT INTO items(type,title,data,search_text,created_at,updated_at)
VALUES ('card','c','{\"type\":\"card\",\"holder\":\"h\",\"number\":\"4111\",\"expiry\":\"01/30\",\"cvv\":\"9\",\"bank\":\"b\",\"notes\":\"cn\"}','old',1,1)",
[],
).unwrap();
conn.execute(
"INSERT INTO items(type,title,data,search_text,created_at,updated_at)
VALUES ('password','broken','not json','old',1,1)",
[],
).unwrap();
db.migrate_data().unwrap();
assert_eq!(user_version(conn), 1);
let pw_id: i64 = conn
.query_row("SELECT id FROM items WHERE title='t'", [], |r| r.get(0))
.unwrap();
let (ty, data, st) = read_item_row(conn, pw_id);
assert_eq!(ty, "password");
let fields: Vec<crate::model::Field> = serde_json::from_str(&data).unwrap();
assert_eq!(fields.len(), 5);
assert!(st.contains("u"));
assert!(st.contains("x"));
assert!(st.contains("n"));
assert!(!st.contains("s3cret"), "Secret 不应进入 search_text");
let note_id: i64 = conn
.query_row("SELECT id FROM items WHERE title='n'", [], |r| r.get(0))
.unwrap();
let (_, nd, nst) = read_item_row(conn, note_id);
let nf: Vec<crate::model::Field> = serde_json::from_str(&nd).unwrap();
assert_eq!(nf.len(), 2);
assert!(nst.contains("body"));
let card_id: i64 = conn
.query_row("SELECT id FROM items WHERE title='c'", [], |r| r.get(0))
.unwrap();
let (_, cd, cst) = read_item_row(conn, card_id);
let cf: Vec<crate::model::Field> = serde_json::from_str(&cd).unwrap();
assert_eq!(cf.len(), 6);
assert!(cst.contains("h"));
assert!(!cst.contains("4111"));
let broken_id: i64 = conn
.query_row("SELECT id FROM items WHERE title='broken'", [], |r| r.get(0))
.unwrap();
let (bt, bd, bst) = read_item_row(conn, broken_id);
assert_eq!(bt, "password");
assert_eq!(bd, "not json", "损坏行 data 应保持不变");
assert_eq!(bst, "old", "损坏行 search_text 应保持不变");
}
#[test]
fn migrate_data_is_idempotent() {
let db = Database::open_in_memory().unwrap();
let conn = db.conn();
conn.execute_batch("PRAGMA user_version = 0").unwrap();
conn.execute(
"INSERT INTO items(type,title,data,search_text,created_at,updated_at)
VALUES ('password','t','{\"type\":\"password\",\"username\":\"u\",\"password\":\"s\",\"url\":\"\",\"totp_secret\":\"\",\"notes\":\"n\"}','old',1,1)",
[],
).unwrap();
db.migrate_data().unwrap();
assert_eq!(user_version(conn), 1);
let pw_id: i64 = conn
.query_row("SELECT id FROM items WHERE title='t'", [], |r| r.get(0))
.unwrap();
let (_, data_before, st_before) = read_item_row(conn, pw_id);
db.migrate_data().unwrap();
assert_eq!(user_version(conn), 1);
let (_, data_after, st_after) = read_item_row(conn, pw_id);
assert_eq!(data_before, data_after);
assert_eq!(st_before, st_after);
}
#[test]
fn migrate_data_handles_mixed_new_and_legacy_shapes() {
let db = Database::open_in_memory().unwrap();
let conn = db.conn();
conn.execute_batch("PRAGMA user_version = 0").unwrap();
let new_shape = serde_json::to_string(&vec![
crate::model::Field {
name: "ssid".into(),
value: "HomeNet".into(),
kind: crate::model::FieldKind::Text,
protected: false,
},
crate::model::Field {
name: "password".into(),
value: "wifipass".into(),
kind: crate::model::FieldKind::Secret,
protected: true,
},
])
.unwrap();
conn.execute(
"INSERT INTO items(type,title,data,search_text,created_at,updated_at)
VALUES ('wifi','w',?1,'old',1,1)",
rusqlite::params![new_shape],
)
.unwrap();
db.migrate_data().unwrap();
assert_eq!(user_version(conn), 1);
let w_id: i64 = conn
.query_row("SELECT id FROM items WHERE title='w'", [], |r| r.get(0))
.unwrap();
let (ty, data, st) = read_item_row(conn, w_id);
assert_eq!(ty, "wifi");
let fields: Vec<crate::model::Field> = serde_json::from_str(&data).unwrap();
assert_eq!(fields.len(), 2);
assert!(st.contains("HomeNet"));
assert!(!st.contains("wifipass"));
}
#[test]
fn migrate_data_noop_on_fresh_empty_db() {
let db = Database::open_in_memory().unwrap();
db.migrate_data().unwrap();
assert_eq!(user_version(db.conn()), 1);
}
#[test]
fn templates_table_exists_after_migrate() {
let db = Database::open_in_memory().unwrap();
let names: Vec<String> = db
.conn()
.prepare("SELECT name FROM sqlite_master WHERE type='table' AND name='templates'")
.unwrap()
.query_map([], |r| r.get::<_, String>(0))
.unwrap()
.filter_map(|r| r.ok())
.collect();
assert_eq!(names, vec!["templates".to_string()]);
}
}