use anyhow::{anyhow, Context, Result};
use std::collections::HashMap;
use std::path::Path;
#[derive(Debug, Clone, PartialEq)]
pub enum Value {
Null,
Int(i64),
Real(f64),
Text(String),
Blob(Vec<u8>),
}
impl Value {
pub fn literal(&self) -> String {
match self {
Value::Null => "NULL".into(),
Value::Int(i) => i.to_string(),
Value::Real(f) => {
let s = format!("{f:?}");
if s.contains(['.', 'e', 'E', 'n']) {
s
} else {
format!("{s}.0")
}
}
Value::Text(t) => format!("'{}'", t.replace('\'', "''")),
Value::Blob(b) => {
let mut out = String::with_capacity(b.len() * 2 + 3);
out.push_str("X'");
for byte in b {
out.push_str(&format!("{byte:02X}"));
}
out.push('\'');
out
}
}
}
}
#[derive(Debug, Clone)]
pub struct Row {
pub table: Option<String>,
pub page: u32,
pub rowid: Option<i64>,
pub values: Vec<Value>,
}
#[derive(Debug, Clone)]
pub struct TableDef {
pub name: String,
pub rootpage: u32,
pub columns: Vec<String>,
}
#[derive(Debug, Default)]
pub struct Recovered {
pub schema: Vec<(String, String)>,
pub tables: Vec<TableDef>,
pub rows: Vec<Row>,
pub orphan_pages: Vec<u32>,
pub notes: Vec<String>,
}
impl Recovered {
pub fn rows_for<'a>(&'a self, table: &'a str) -> impl Iterator<Item = &'a Row> {
self.rows
.iter()
.filter(move |r| r.table.as_deref() == Some(table))
}
pub fn orphans(&self) -> usize {
self.rows.iter().filter(|r| r.table.is_none()).count()
}
}
struct Db {
bytes: Vec<u8>,
page_size: usize,
usable: usize,
pages: u32,
}
impl Db {
fn open_with_wal(path: &Path) -> Result<Self> {
let mut db = Db::open(path)?;
let Some((page_size, pages, db_size)) = crate::wal::latest_pages(path) else {
return Ok(db);
};
if page_size as usize != db.page_size {
return Ok(db);
}
let needed =
(db_size.max(pages.keys().copied().max().unwrap_or(0)) as usize) * db.page_size;
if needed > db.bytes.len() {
db.bytes.resize(needed, 0);
}
for (page, image) in pages {
let start = (page as usize - 1) * db.page_size;
if start + db.page_size <= db.bytes.len() {
db.bytes[start..start + db.page_size].copy_from_slice(&image);
}
}
db.pages = (db.bytes.len() / db.page_size) as u32;
Ok(db)
}
fn open(path: &Path) -> Result<Self> {
let bytes = std::fs::read(path).with_context(|| format!("read {}", path.display()))?;
if bytes.len() < 100 {
return Err(anyhow!("too short to be a database (no header)"));
}
if &bytes[..16] != b"SQLite format 3\0" {
return Err(anyhow!("not a SQLite database (header magic)"));
}
let raw = be16(&bytes, 16) as usize;
let page_size = if raw == 1 { 65536 } else { raw };
if page_size < 512 || !page_size.is_power_of_two() {
return Err(anyhow!("page size {page_size} is not a power of two ≥ 512"));
}
let reserved = bytes[20] as usize;
let counted = be32(&bytes, 28);
let present = (bytes.len() / page_size) as u32;
Ok(Db {
page_size,
usable: page_size - reserved,
pages: present.max(1).min(counted.max(present)),
bytes,
})
}
fn page(&self, n: u32) -> Option<&[u8]> {
if n == 0 {
return None;
}
let start = (n as usize - 1) * self.page_size;
self.bytes
.get(start..(start + self.page_size).min(self.bytes.len()))
}
fn leaf_cells(&self, n: u32) -> Option<Vec<(Option<i64>, Vec<Value>)>> {
let page = self.page(n)?;
let base = if n == 1 { 100 } else { 0 };
let (with_rowid, header, skip) = match *page.get(base)? {
0x0d => (true, 8, 0),
0x0a => (false, 8, 0),
0x02 => (false, 12, 4),
_ => return None,
};
let ncells = be16(page, base + 3) as usize;
let mut out = Vec::with_capacity(ncells);
for i in 0..ncells {
let ptr = base + header + i * 2;
let off = be16(page, ptr) as usize;
if off == 0 || off >= page.len() {
continue;
}
if let Some(cell) = self.leaf_cell(page, off + skip, with_rowid) {
out.push(cell);
}
}
Some(out)
}
fn leaf_cell(
&self,
page: &[u8],
off: usize,
with_rowid: bool,
) -> Option<(Option<i64>, Vec<Value>)> {
let (payload_len, n1) = varint(page, off)?;
let (rowid, n2) = if with_rowid {
let (r, n) = varint(page, off + n1)?;
(Some(r), n)
} else {
(None, 0)
};
let head = off + n1 + n2;
let payload_len = payload_len as usize;
let max_local = if with_rowid {
self.usable - 35
} else {
((self.usable - 12) * 64 / 255) - 23
};
let min_local = ((self.usable - 12) * 32 / 255) - 23;
let local = if payload_len <= max_local {
payload_len
} else {
let candidate = min_local + (payload_len - min_local) % (self.usable - 4);
if candidate > max_local {
min_local
} else {
candidate
}
};
let mut payload = page.get(head..head + local)?.to_vec();
if local < payload_len {
let mut next = be32(page, head + local);
let mut guard = self.pages as usize + 1;
while next != 0 && payload.len() < payload_len && guard > 0 {
guard -= 1;
let ov = self.page(next)?;
let take = (payload_len - payload.len()).min(self.usable - 4);
payload.extend_from_slice(ov.get(4..4 + take)?);
next = be32(ov, 0);
}
}
Some((rowid, decode_record(&payload)))
}
fn schema_objects(&self) -> Vec<(String, String, u32)> {
let mut out = Vec::new();
for page in self.walk(1) {
for (_, values) in self.leaf_cells(page).unwrap_or_default() {
let text = |i: usize| match values.get(i) {
Some(Value::Text(t)) => Some(t.clone()),
_ => None,
};
let rootpage = match values.get(3) {
Some(Value::Int(i)) => *i as u32,
_ => 0,
};
if let (Some(kind), Some(name)) = (text(0), text(1)) {
out.push((kind, name, rootpage));
}
}
}
out
}
fn walk(&self, root: u32) -> Vec<u32> {
let mut seen = Vec::new();
let mut stack = vec![root];
let mut guard = self.pages as usize * 2 + 8;
while let Some(n) = stack.pop() {
if guard == 0 {
break;
}
guard -= 1;
if n == 0 || n > self.pages || seen.contains(&n) {
continue;
}
seen.push(n);
let page = match self.page(n) {
Some(p) => p,
None => continue,
};
let base = if n == 1 { 100 } else { 0 };
if !matches!(page.get(base).copied(), Some(0x05) | Some(0x02)) {
continue; }
let ncells = be16(page, base + 3) as usize;
for i in 0..ncells {
let off = be16(page, base + 12 + i * 2) as usize;
if off + 4 <= page.len() {
stack.push(be32(page, off));
}
}
stack.push(be32(page, base + 8));
}
seen
}
}
#[derive(Debug, Default)]
pub struct PageRows {
pub rows: Vec<(Option<i64>, Vec<Value>)>,
pub overflowed: bool,
pub kind: PageKind,
}
#[derive(Debug, Default, PartialEq, Eq, Clone, Copy)]
pub enum PageKind {
TableLeaf,
TableInterior,
IndexLeaf,
IndexInterior,
#[default]
Other,
}
impl PageKind {
pub fn label(self) -> &'static str {
match self {
PageKind::TableLeaf => "table leaf",
PageKind::TableInterior => "table interior",
PageKind::IndexLeaf => "index leaf",
PageKind::IndexInterior => "index interior",
PageKind::Other => "overflow / freelist",
}
}
}
pub fn decode_page_image(image: &[u8], page_no: u32) -> PageRows {
let base = if page_no == 1 { 100 } else { 0 };
let kind = match image.get(base) {
Some(0x0d) => PageKind::TableLeaf,
Some(0x05) => PageKind::TableInterior,
Some(0x0a) => PageKind::IndexLeaf,
Some(0x02) => PageKind::IndexInterior,
_ => PageKind::Other,
};
let mut out = PageRows {
kind,
..Default::default()
};
let (with_rowid, header, skip) = match kind {
PageKind::TableLeaf => (true, 8, 0),
PageKind::IndexLeaf => (false, 8, 0),
PageKind::IndexInterior => (false, 12, 4),
_ => return out,
};
let usable = image.len();
let max_local = if with_rowid {
usable.saturating_sub(35)
} else {
((usable.saturating_sub(12)) * 64 / 255).saturating_sub(23)
};
let min_local = ((usable.saturating_sub(12)) * 32 / 255).saturating_sub(23);
let ncells = be16(image, base + 3) as usize;
for i in 0..ncells {
let off = be16(image, base + header + i * 2) as usize + skip;
if off == 0 || off >= image.len() {
continue;
}
let Some((payload_len, n1)) = varint(image, off) else {
continue;
};
let (rowid, n2) = if with_rowid {
match varint(image, off + n1) {
Some((r, n)) => (Some(r), n),
None => continue,
}
} else {
(None, 0)
};
let head = off + n1 + n2;
let payload_len = payload_len as usize;
let local = if payload_len <= max_local {
payload_len
} else {
out.overflowed = true;
let candidate = min_local + (payload_len - min_local) % usable.saturating_sub(4);
if candidate > max_local {
min_local
} else {
candidate
}
};
let Some(payload) = image.get(head..(head + local).min(image.len())) else {
continue;
};
out.rows.push((rowid, decode_record(payload)));
}
out
}
pub fn page_owners(path: &Path) -> Result<HashMap<u32, String>> {
let db = Db::open_with_wal(path)?;
let mut out = HashMap::new();
out.insert(1, "sqlite_schema".to_string());
for (kind, name, rootpage) in db.schema_objects() {
if rootpage == 0 {
continue;
}
let label = if kind == "index" {
format!("index {name}")
} else {
name
};
for page in db.walk(rootpage) {
out.entry(page).or_insert_with(|| label.clone());
}
}
Ok(out)
}
pub fn recover(path: &Path) -> Result<Recovered> {
let db = Db::open_with_wal(path)?;
let mut out = Recovered::default();
if db.bytes.len() % db.page_size != 0 {
out.notes.push(format!(
"file is {} bytes, not a whole number of {}-byte pages — the last page is partial",
db.bytes.len(),
db.page_size
));
}
let mut schema_rows = Vec::new();
for page in db.walk(1) {
if let Some(cells) = db.leaf_cells(page) {
schema_rows.extend(cells);
}
}
if schema_rows.is_empty() {
out.notes.push(
"the schema page holds no readable rows — every row will be a lost_and_found row"
.into(),
);
}
for (_, values) in &schema_rows {
let text = |i: usize| match values.get(i) {
Some(Value::Text(t)) => Some(t.clone()),
_ => None,
};
let kind = text(0).unwrap_or_default();
let name = text(1).unwrap_or_default();
let sql = text(4).unwrap_or_default();
let rootpage = match values.get(3) {
Some(Value::Int(i)) => *i as u32,
_ => 0,
};
if sql.is_empty() || name.starts_with("sqlite_") {
continue;
}
out.schema.push((kind.clone(), sql.clone()));
if kind == "table" {
out.tables.push(TableDef {
columns: create_columns(&sql),
name,
rootpage,
});
}
}
let mut owner: HashMap<u32, String> = HashMap::new();
let mut index_pages: std::collections::HashSet<u32> = Default::default();
for (kind, name, rootpage) in db.schema_objects() {
if kind != "index" || rootpage == 0 {
continue;
}
let _ = name;
for p in db.walk(rootpage) {
index_pages.insert(p);
}
}
for t in &out.tables {
if t.rootpage == 0 {
out.notes.push(format!(
"{}: no root page recorded (a virtual table?)",
t.name
));
continue;
}
let pages = db.walk(t.rootpage);
if pages.len() == 1 && db.leaf_cells(t.rootpage).is_none() {
out.notes.push(format!(
"{}: root page {} is not a readable table page, so its rows are recovered from \
unreachable pages instead",
t.name, t.rootpage
));
}
for p in pages {
owner.insert(p, t.name.clone());
}
}
let mut attributed: HashMap<String, usize> = HashMap::new();
for n in 1..=db.pages {
let cells = match db.leaf_cells(n) {
Some(c) if !c.is_empty() => c,
_ => continue,
};
if n == 1 {
continue; }
if index_pages.contains(&n) && !owner.contains_key(&n) {
continue;
}
let table = match owner.get(&n) {
Some(t) => Some(t.clone()),
None => {
let ncols = cells[0].1.len();
let mut matches = out
.tables
.iter()
.filter(|t| t.columns.len() == ncols)
.map(|t| t.name.clone());
let first = matches.next();
match (first, matches.next()) {
(Some(name), None) => {
*attributed.entry(name.clone()).or_insert(0usize) += 1;
Some(name)
}
_ => {
out.orphan_pages.push(n);
None
}
}
}
};
for (rowid, values) in cells {
out.rows.push(Row {
table: table.clone(),
page: n,
rowid,
values,
});
}
}
let mut counts: Vec<(String, usize)> = attributed.into_iter().collect();
counts.sort();
for (name, pages) in counts {
out.notes.push(format!(
"{pages} page{} unreachable from any root matched {name} by column count alone",
if pages == 1 { "" } else { "s" }
));
}
if !out.orphan_pages.is_empty() {
out.notes.push(format!(
"{} page{} could not be attributed to a table; their rows are in lost_and_found",
out.orphan_pages.len(),
if out.orphan_pages.len() == 1 { "" } else { "s" }
));
}
Ok(out)
}
pub fn to_sql(r: &Recovered) -> String {
let mut out =
String::from("BEGIN;\nPRAGMA writable_schema = on;\nPRAGMA foreign_keys = off;\n");
for note in &r.notes {
out.push_str(&format!("-- {note}\n"));
}
for (kind, sql) in r.schema.iter().filter(|(k, _)| k == "table") {
let _ = kind;
out.push_str(sql.trim_end_matches(';'));
out.push_str(";\n");
}
for t in &r.tables {
let named = t
.columns
.iter()
.map(|c| format!("\"{}\"", c.replace('"', "\"\"")))
.collect::<Vec<_>>()
.join(", ");
for row in r.rows_for(&t.name) {
let (cols, lead) = match (row.rowid, t.columns.is_empty()) {
(_, true) => (String::new(), String::new()),
(Some(id), false) => (format!("(_rowid_, {named})"), format!("{id}, ")),
(None, false) => (format!("({named})"), String::new()),
};
out.push_str(&format!(
"INSERT OR IGNORE INTO \"{}\"{} VALUES ({}{});\n",
t.name.replace('"', "\"\""),
cols,
lead,
row.values
.iter()
.map(Value::literal)
.collect::<Vec<_>>()
.join(", ")
));
}
}
let orphans: Vec<&Row> = r.rows.iter().filter(|r| r.table.is_none()).collect();
if !orphans.is_empty() {
let widest = orphans.iter().map(|r| r.values.len()).max().unwrap_or(0);
let cols: Vec<String> = (0..widest).map(|i| format!("c{i}")).collect();
out.push_str(&format!(
"CREATE TABLE lost_and_found(pgno INTEGER, nfield INTEGER, id INTEGER{}{});\n",
if cols.is_empty() { "" } else { ", " },
cols.join(", ")
));
for row in orphans {
let mut values: Vec<String> = vec![
row.page.to_string(),
row.values.len().to_string(),
row.rowid
.map(|r| r.to_string())
.unwrap_or_else(|| "NULL".into()),
];
values.extend(row.values.iter().map(Value::literal));
for _ in row.values.len()..widest {
values.push("NULL".into());
}
out.push_str(&format!(
"INSERT INTO lost_and_found VALUES ({});\n",
values.join(", ")
));
}
}
for (kind, sql) in r.schema.iter().filter(|(k, _)| k != "table") {
let _ = kind;
out.push_str(sql.trim_end_matches(';'));
out.push_str(";\n");
}
out.push_str("PRAGMA writable_schema = off;\nCOMMIT;\n");
out
}
fn create_columns(sql: &str) -> Vec<String> {
let open = match sql.find('(') {
Some(i) => i,
None => return Vec::new(),
};
let body = &sql[open + 1..];
let mut depth = 0usize;
let mut parts: Vec<String> = Vec::new();
let mut current = String::new();
for c in body.chars() {
match c {
'(' => {
depth += 1;
current.push(c);
}
')' if depth == 0 => break,
')' => {
depth -= 1;
current.push(c);
}
',' if depth == 0 => parts.push(std::mem::take(&mut current)),
_ => current.push(c),
}
}
if !current.trim().is_empty() {
parts.push(current);
}
const CONSTRAINTS: &[&str] = &["primary", "unique", "check", "foreign", "constraint"];
parts
.iter()
.filter_map(|p| {
let name = first_identifier(p)?;
if CONSTRAINTS.contains(&name.to_lowercase().as_str()) {
return None;
}
Some(name)
})
.filter(|c| !c.is_empty())
.collect()
}
fn first_identifier(part: &str) -> Option<String> {
let text = part.trim_start();
let mut chars = text.chars();
let open = chars.next()?;
let close = match open {
'"' => '"',
'`' => '`',
'[' => ']',
'\'' => '\'',
_ => {
let end = text
.find(|c: char| c.is_whitespace() || c == '(')
.unwrap_or(text.len());
return Some(text[..end].to_string());
}
};
let rest = &text[open.len_utf8()..];
let end = rest.find(close)?;
Some(rest[..end].to_string())
}
fn decode_record(payload: &[u8]) -> Vec<Value> {
let (header_len, n) = match varint(payload, 0) {
Some(v) => v,
None => return Vec::new(),
};
let header_end = (header_len as usize).min(payload.len());
let mut types = Vec::new();
let mut at = n;
while at < header_end {
match varint(payload, at) {
Some((t, used)) => {
types.push(t);
at += used;
}
None => break,
}
}
let mut values = Vec::with_capacity(types.len());
let mut body = header_end;
for t in types {
let (value, used) = read_value(payload, body, t);
values.push(value);
body += used;
}
values
}
fn read_value(payload: &[u8], at: usize, t: i64) -> (Value, usize) {
let int = |len: usize| -> Value {
match payload.get(at..at + len) {
Some(b) => {
let mut v: i64 = if b[0] & 0x80 != 0 { -1 } else { 0 };
for byte in b {
v = (v << 8) | *byte as i64;
}
Value::Int(v)
}
None => Value::Null,
}
};
match t {
0 => (Value::Null, 0),
1 => (int(1), 1),
2 => (int(2), 2),
3 => (int(3), 3),
4 => (int(4), 4),
5 => (int(6), 6),
6 => (int(8), 8),
7 => match payload.get(at..at + 8) {
Some(b) => (
Value::Real(f64::from_be_bytes(b.try_into().unwrap_or([0; 8]))),
8,
),
None => (Value::Null, 0),
},
8 => (Value::Int(0), 0),
9 => (Value::Int(1), 0),
n if n >= 12 && n % 2 == 0 => {
let len = (n as usize - 12) / 2;
match payload.get(at..at + len) {
Some(b) => (Value::Blob(b.to_vec()), len),
None => (Value::Null, 0),
}
}
n if n >= 13 => {
let len = (n as usize - 13) / 2;
match payload.get(at..at + len) {
Some(b) => (Value::Text(String::from_utf8_lossy(b).into_owned()), len),
None => (Value::Null, 0),
}
}
_ => (Value::Null, 0),
}
}
fn varint(bytes: &[u8], at: usize) -> Option<(i64, usize)> {
let mut value: u64 = 0;
for i in 0..9 {
let byte = *bytes.get(at + i)?;
if i == 8 {
value = (value << 8) | byte as u64;
return Some((value as i64, 9));
}
value = (value << 7) | (byte & 0x7f) as u64;
if byte & 0x80 == 0 {
return Some((value as i64, i + 1));
}
}
None
}
fn be16(b: &[u8], at: usize) -> u16 {
match b.get(at..at + 2) {
Some(s) => u16::from_be_bytes([s[0], s[1]]),
None => 0,
}
}
fn be32(b: &[u8], at: usize) -> u32 {
match b.get(at..at + 4) {
Some(s) => u32::from_be_bytes([s[0], s[1], s[2], s[3]]),
None => 0,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn varints_match_the_format_spec() {
assert_eq!(varint(&[0x00], 0), Some((0, 1)));
assert_eq!(varint(&[0x7f], 0), Some((127, 1)));
assert_eq!(varint(&[0x81, 0x00], 0), Some((128, 2)));
assert_eq!(varint(&[0x82, 0x21], 0), Some((289, 2)));
assert_eq!(
varint(&[0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff], 0),
Some((-1, 9)),
"the nine-byte form is the full 64 bits, so all-ones is -1"
);
assert_eq!(
varint(&[0x81], 0),
None,
"a truncated varint is not a value"
);
}
#[test]
fn serial_types_decode_to_their_values() {
assert_eq!(read_value(&[], 0, 0).0, Value::Null);
assert_eq!(read_value(&[], 0, 8).0, Value::Int(0));
assert_eq!(read_value(&[], 0, 9).0, Value::Int(1));
assert_eq!(read_value(&[0xff], 0, 1).0, Value::Int(-1));
assert_eq!(read_value(&[0x7f, 0xff], 0, 2).0, Value::Int(32767));
let pi = std::f64::consts::PI.to_be_bytes();
assert_eq!(
read_value(&pi, 0, 7).0,
Value::Real(std::f64::consts::PI),
"type 7 is an IEEE double"
);
assert_eq!(
read_value(b"hi", 0, 13 + 2 * 2).0,
Value::Text("hi".into()),
"odd types ≥ 13 are text of (N-13)/2 bytes"
);
assert_eq!(
read_value(&[0x00, 0xff], 0, 12 + 2 * 2).0,
Value::Blob(vec![0x00, 0xff]),
"even types ≥ 12 are blobs of (N-12)/2 bytes"
);
}
#[test]
fn literals_round_trip_the_hard_cases() {
assert_eq!(Value::Null.literal(), "NULL");
assert_eq!(Value::Int(-5).literal(), "-5");
assert_eq!(Value::Real(1.0).literal(), "1.0", "a real keeps its point");
assert_eq!(Value::Text("it's".into()).literal(), "'it''s'");
assert_eq!(Value::Blob(vec![0, 255]).literal(), "X'00FF'");
}
#[test]
fn column_names_come_from_the_create_statement() {
assert_eq!(
create_columns("CREATE TABLE t (a TEXT, b INTEGER)"),
["a", "b"]
);
assert_eq!(
create_columns("CREATE TABLE t (\"odd name\" TEXT, b NUMERIC(10, 2), PRIMARY KEY (b))"),
["odd name", "b"],
"a table constraint is not a column, and a type's own commas do not split"
);
assert_eq!(
create_columns("CREATE TABLE t (id INTEGER, FOREIGN KEY (id) REFERENCES o(id))"),
["id"]
);
assert!(create_columns("CREATE VIEW v AS SELECT 1").is_empty());
}
}