use crate::CliError;
use btctax_core::tax::return_inputs::ReturnInputs;
use rusqlite::{Connection, OptionalExtension};
use std::collections::BTreeMap;
pub const SCHEMA_VERSION: i64 = 2;
pub fn init_table(conn: &Connection) -> Result<(), CliError> {
conn.execute_batch(
"CREATE TABLE IF NOT EXISTS return_inputs \
(year INTEGER PRIMARY KEY, inputs_json TEXT NOT NULL, schema_version INTEGER NOT NULL DEFAULT 0);",
)?;
if let Err(e) = conn.execute_batch(
"ALTER TABLE return_inputs ADD COLUMN schema_version INTEGER NOT NULL DEFAULT 0;",
) {
let msg = e.to_string();
if !msg.contains("duplicate column name") {
return Err(e.into());
}
}
Ok(())
}
fn row_to_inputs(year: i32, json: &str, version: i64) -> Result<ReturnInputs, CliError> {
if version != SCHEMA_VERSION {
return Err(CliError::StaleReturnInputs {
year,
found: version,
expected: SCHEMA_VERSION,
});
}
serde_json::from_str(json).map_err(|e| CliError::BadConfigValue {
key: format!("return_inputs[{year}]"),
value: format!("invalid JSON: {e}"),
})
}
pub fn get(conn: &Connection, year: i32) -> Result<Option<ReturnInputs>, CliError> {
init_table(conn)?;
let json: Option<(String, i64)> = conn
.query_row(
"SELECT inputs_json, schema_version FROM return_inputs WHERE year=?1",
[year],
|r| Ok((r.get(0)?, r.get(1)?)),
)
.optional()?;
match json {
None => Ok(None),
Some((j, v)) => Ok(Some(row_to_inputs(year, &j, v)?)),
}
}
pub fn set(conn: &Connection, year: i32, ri: &ReturnInputs) -> Result<(), CliError> {
init_table(conn)?;
let j = serde_json::to_string(ri).map_err(|e| CliError::BadConfigValue {
key: format!("return_inputs[{year}]"),
value: e.to_string(),
})?;
conn.execute(
"INSERT INTO return_inputs(year,inputs_json,schema_version) VALUES(?1,?2,?3) \
ON CONFLICT(year) DO UPDATE SET inputs_json=excluded.inputs_json, \
schema_version=excluded.schema_version",
rusqlite::params![year, j, SCHEMA_VERSION],
)?;
Ok(())
}
pub fn exists(conn: &Connection, year: i32) -> Result<bool, CliError> {
init_table(conn)?;
let found: Option<i64> = conn
.query_row("SELECT 1 FROM return_inputs WHERE year=?1", [year], |r| {
r.get(0)
})
.optional()?;
Ok(found.is_some())
}
pub fn delete(conn: &Connection, year: i32) -> Result<bool, CliError> {
init_table(conn)?;
let n = conn.execute("DELETE FROM return_inputs WHERE year=?1", [year])?;
Ok(n > 0)
}
pub fn years(conn: &Connection) -> Result<Vec<i32>, CliError> {
init_table(conn)?;
let mut stmt = conn.prepare("SELECT year FROM return_inputs ORDER BY year")?;
let rows = stmt.query_map([], |r| r.get::<_, i32>(0))?;
Ok(rows.collect::<Result<Vec<_>, _>>()?)
}
pub fn all(conn: &Connection) -> Result<BTreeMap<i32, ReturnInputs>, CliError> {
init_table(conn)?;
let mut stmt =
conn.prepare("SELECT year, inputs_json, schema_version FROM return_inputs ORDER BY year")?;
let rows = stmt.query_map([], |r| {
Ok((
r.get::<_, i32>(0)?,
r.get::<_, String>(1)?,
r.get::<_, i64>(2)?,
))
})?;
let mut out = BTreeMap::new();
for row in rows {
let (y, j, v) = row?;
out.insert(y, row_to_inputs(y, &j, v)?);
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
use btctax_core::tax::return_inputs::{Owner, W2};
use btctax_core::FilingStatus;
use rust_decimal_macros::dec;
fn mem() -> Connection {
let c = Connection::open_in_memory().unwrap();
init_table(&c).unwrap();
c
}
fn inputs() -> ReturnInputs {
ReturnInputs {
filing_status: FilingStatus::Mfj,
w2s: vec![W2 {
owner: Owner::Taxpayer,
employer: "ACME".into(),
box1_wages: dec!(82000),
box2_fed_withheld: dec!(9100),
..Default::default()
}],
..Default::default()
}
}
#[test]
fn set_then_get_round_trips() {
let c = mem();
set(&c, 2024, &inputs()).unwrap();
assert_eq!(get(&c, 2024).unwrap().unwrap(), inputs());
assert_eq!(get(&c, 2025).unwrap(), None);
assert!(exists(&c, 2024).unwrap());
assert!(!exists(&c, 2025).unwrap());
}
#[test]
fn get_on_tableless_vault_is_ok_none() {
let c = Connection::open_in_memory().unwrap(); assert_eq!(get(&c, 2024).unwrap(), None);
}
#[test]
fn bad_json_is_a_typed_error_not_a_panic() {
let c = mem();
c.execute(
"INSERT INTO return_inputs(year,inputs_json,schema_version) VALUES(2024,'not json',?1)",
[SCHEMA_VERSION],
)
.unwrap();
assert!(matches!(
get(&c, 2024).unwrap_err(),
CliError::BadConfigValue { .. }
));
}
#[test]
fn all_returns_sorted_by_year() {
let c = mem();
set(&c, 2025, &inputs()).unwrap();
set(&c, 2024, &inputs()).unwrap();
assert_eq!(
all(&c).unwrap().keys().copied().collect::<Vec<_>>(),
vec![2024, 2025]
);
}
#[test]
fn delete_removes_the_row() {
let c = mem();
set(&c, 2024, &inputs()).unwrap();
assert!(exists(&c, 2024).unwrap());
assert!(delete(&c, 2024).unwrap()); assert!(!exists(&c, 2024).unwrap());
assert!(!delete(&c, 2024).unwrap()); }
}
#[cfg(test)]
mod p9_stale_row_refuses {
use super::*;
use rusqlite::Connection;
fn vault_with_row_at_version(year: i32, version: i64) -> Connection {
let conn = Connection::open_in_memory().unwrap();
init_table(&conn).unwrap();
let blob = serde_json::to_string(&ReturnInputs::default()).unwrap();
conn.execute(
"INSERT INTO return_inputs(year,inputs_json,schema_version) VALUES(?1,?2,?3)",
rusqlite::params![year, blob, version],
)
.unwrap();
conn
}
#[test]
fn a_version_0_row_refuses_stale() {
let conn = vault_with_row_at_version(2024, 0);
assert!(
matches!(get(&conn, 2024), Err(CliError::StaleReturnInputs { year: 2024, found: 0, expected }) if expected == SCHEMA_VERSION),
"a v0 row must refuse as stale, naming the version"
);
}
#[test]
fn a_version_1_row_refuses_stale() {
let conn = vault_with_row_at_version(2024, 1);
assert!(
matches!(
get(&conn, 2024),
Err(CliError::StaleReturnInputs { found: 1, .. })
),
"a v1 row must refuse as stale"
);
}
#[test]
fn all_refuses_a_stale_row_identically_to_get() {
let conn = vault_with_row_at_version(2024, 1);
assert!(
matches!(
all(&conn),
Err(CliError::StaleReturnInputs { found: 1, .. })
),
"`all()` must apply the same version gate as `get()`"
);
}
#[test]
fn a_future_version_row_refuses_too() {
let conn = vault_with_row_at_version(2024, SCHEMA_VERSION + 1);
assert!(
matches!(get(&conn, 2024), Err(CliError::StaleReturnInputs { .. })),
"a future-version row must refuse, not be half-read"
);
}
#[test]
fn a_current_version_row_reads() {
let conn = vault_with_row_at_version(2024, SCHEMA_VERSION);
assert!(
get(&conn, 2024).unwrap().is_some(),
"a current-version row must read"
);
}
#[test]
fn the_stale_message_names_the_full_three_command_remedy() {
let msg = CliError::StaleReturnInputs {
year: 2024,
found: 1,
expected: 2,
}
.to_string();
assert!(msg.contains("income clear 2024"), "names clear");
assert!(msg.contains("income import"), "names import");
assert!(
msg.contains("--write-carryover"),
"names the rebuild — disclosure is not restoration (r6 I-1)"
);
assert!(
msg.contains("2023"),
"the rebuild targets the PRIOR year (year-1)"
);
}
}