fix(sessions): add opt-in partial database recovery

This commit is contained in:
Gille
2026-07-25 22:37:53 -06:00
committed by Teknium
parent fe431651c5
commit 508764d384
3 changed files with 787 additions and 10 deletions
+27
View File
@@ -16050,6 +16050,14 @@ def main():
default=1000,
help="Rows committed per recovery batch (default: 1000)",
)
sessions_recover.add_argument(
"--allow-partial",
action="store_true",
help=(
"Best-effort salvage across damaged row ranges; the output remains "
"separate and every skipped range is recorded"
),
)
sessions_recover.add_argument(
"--report",
type=Path,
@@ -16159,10 +16167,14 @@ def main():
source = args.source
output = getattr(args, "output", None)
inspect_only = bool(getattr(args, "inspect_only", False))
allow_partial = bool(getattr(args, "allow_partial", False))
report_path = getattr(args, "report", None)
if inspect_only and output is not None:
print("Error: --output cannot be used with --inspect-only.")
return 2
if inspect_only and allow_partial:
print("Error: --allow-partial cannot be used with --inspect-only.")
return 2
if not inspect_only and output is None:
print("Error: --output is required unless --inspect-only is used.")
return 2
@@ -16203,6 +16215,7 @@ def main():
work_dir=getattr(args, "work_dir", None),
chunk_size=getattr(args, "chunk_size", 1000),
progress_cb=_recovery_progress,
allow_partial=allow_partial,
)
if last_progress["table"] is not None:
print()
@@ -16228,6 +16241,20 @@ def main():
print(" The active session database was not changed.")
print(" Review the JSON report before installing this database.")
return 0
if allow_partial and report.get("verified"):
counts = report.get("verification", {}).get("table_counts", {})
print(f"✓ Partial recovery output verified at: {output}")
print(
" Recovered "
f"{int(counts.get('sessions') or 0):,} sessions and "
f"{int(counts.get('messages') or 0):,} messages."
)
print(" The active session database was not changed.")
print(
" This output is incomplete. Review every skipped range "
"and orphan count in the JSON report before installing it."
)
return 0
print("✗ Recovery output did not pass every verification check.")
print(" Do not install it. Review the JSON report for partial data or errors.")
return 1
+398 -10
View File
@@ -59,6 +59,9 @@ _GENERATED_META_KEYS = frozenset({
_SIDECAR_SUFFIXES = ("", "-wal", "-shm", "-journal")
_MINIMUM_SPACE_HEADROOM = 256 * 1024 * 1024
_MAX_SALVAGE_RANGE_QUERIES = 10_000
_MIN_SQLITE_ROWID = -(2**63)
_MAX_SQLITE_ROWID = 2**63 - 1
class SessionRecoveryError(RuntimeError):
@@ -421,6 +424,224 @@ def _copy_table(
return result
def _append_skipped_range(
ranges: list[dict[str, Any]],
low: int,
high: int,
error: str,
) -> None:
"""Record skipped rowid ranges without producing one entry per row."""
if (
ranges
and ranges[-1]["high"] + 1 == low
and ranges[-1]["error"] == error
):
ranges[-1]["high"] = high
return
ranges.append({"low": low, "high": high, "error": error})
def _salvage_rowid_bounds(
source: sqlite3.Connection,
table: str,
) -> dict[str, Any]:
"""Find the readable rowid edges without scanning the complete table."""
result: dict[str, Any] = {"errors": [], "fallback_edges": []}
rows: dict[str, Optional[int]] = {"low": None, "high": None}
directions = (("low", "ASC"), ("high", "DESC"))
for edge, direction in directions:
try:
row = source.execute(
f'SELECT rowid FROM "{table}" ORDER BY rowid {direction} LIMIT 1'
).fetchone()
if row is not None:
rows[edge] = int(row[0])
except sqlite3.DatabaseError as exc:
result["errors"].append(f"{edge} rowid: {exc}")
if rows["low"] is None and rows["high"] is None and not result["errors"]:
result["empty"] = True
return result
if rows["low"] is None and rows["high"] is None:
result["unavailable"] = True
return result
# A damaged edge can prevent one of the ordered probes from completing.
# Keep the other readable edge and bound the missing side by SQLite's
# rowid domain. Range bisection can then approach the surviving data
# without assuming that user-created databases contain only positive IDs.
if rows["low"] is None:
rows["low"] = _MIN_SQLITE_ROWID
result["fallback_edges"].append("low")
if rows["high"] is None:
rows["high"] = _MAX_SQLITE_ROWID
result["fallback_edges"].append("high")
result["low"] = rows["low"]
result["high"] = rows["high"]
return result
def _copy_table_salvage(
source: sqlite3.Connection,
destination: sqlite3.Connection,
table: str,
*,
chunk_size: int,
progress_cb: Optional[ProgressCallback],
source_rows: Optional[int],
insert_prefix: str = "INSERT",
row_filter: Optional[
Callable[[tuple[Any, ...], tuple[str, ...]], bool]
] = None,
) -> dict[str, Any]:
"""Best-effort rowid-range copy that continues past damaged source pages."""
source_columns = _table_columns(source, table)
destination_columns = _table_columns(destination, table)
columns = [column for column in destination_columns if column in source_columns]
result: dict[str, Any] = {
"mode": "rowid_range_salvage",
"source_rows": source_rows,
"copied_rows": 0,
"excluded_rows": 0,
"columns": columns,
"range_queries": 0,
"skipped_rowid_ranges": [],
}
if not source_columns:
result["status"] = "missing"
return result
if not columns:
result["status"] = "failed"
result["error"] = "source and destination have no compatible columns"
return result
bounds = _salvage_rowid_bounds(source, table)
result["rowid_bounds"] = bounds
if bounds.get("empty"):
result["status"] = "complete"
return result
if bounds.get("low") is None or bounds.get("high") is None:
result["status"] = "failed"
details = "; ".join(bounds.get("errors") or [])
result["error"] = "could not determine a rowid range for salvage"
if details:
result["error"] += f": {details}"
return result
quoted = ", ".join(f'"{column}"' for column in columns)
placeholders = ", ".join("?" for _ in columns)
select_sql = (
f'SELECT rowid, {quoted} FROM "{table}" '
"WHERE rowid BETWEEN ? AND ? ORDER BY rowid"
)
insert_sql = (
f'{insert_prefix} INTO "{table}" ({quoted}) VALUES ({placeholders})'
)
column_names = tuple(columns)
stopped_at_query_limit = False
def copy_range(low: int, high: int) -> None:
nonlocal stopped_at_query_limit
if low > high:
return
if result["range_queries"] >= _MAX_SALVAGE_RANGE_QUERIES:
stopped_at_query_limit = True
_append_skipped_range(
result["skipped_rowid_ranges"],
low,
high,
"salvage range query limit reached",
)
return
result["range_queries"] += 1
last_committed_rowid: Optional[int] = None
try:
cursor = source.execute(select_sql, (low, high))
while True:
fetched = cursor.fetchmany(chunk_size)
if not fetched:
return
values = [tuple(row[1:]) for row in fetched]
if row_filter is not None:
included = [
row for row in values if row_filter(row, column_names)
]
excluded_count = len(values) - len(included)
else:
included = values
excluded_count = 0
if included:
destination.execute("BEGIN IMMEDIATE")
try:
destination.executemany(insert_sql, included)
destination.execute("COMMIT")
except BaseException:
destination.execute("ROLLBACK")
raise
result["copied_rows"] += len(included)
result["excluded_rows"] += excluded_count
last_committed_rowid = int(fetched[-1][0])
if progress_cb is not None:
progress_cb({
"table": table,
"copied_rows": result["copied_rows"],
"source_rows": source_rows,
"skipped_ranges": len(result["skipped_rowid_ranges"]),
})
except sqlite3.DatabaseError as exc:
retry_low = (
last_committed_rowid + 1
if last_committed_rowid is not None
else low
)
if retry_low > high:
return
if retry_low == high:
_append_skipped_range(
result["skipped_rowid_ranges"],
retry_low,
high,
str(exc),
)
return
midpoint = retry_low + (high - retry_low) // 2
copy_range(retry_low, midpoint)
copy_range(midpoint + 1, high)
copy_range(int(bounds["low"]), int(bounds["high"]))
skipped_ranges = result["skipped_rowid_ranges"]
result["skipped_rowid_span"] = sum(
item["high"] - item["low"] + 1 for item in skipped_ranges
)
result["query_limit_reached"] = stopped_at_query_limit
if skipped_ranges:
result["status"] = "partial" if result["copied_rows"] else "failed"
result["error"] = (
f"{len(skipped_ranges)} rowid range(s) skipped"
)
elif (
source_rows is not None
and result["copied_rows"] + result["excluded_rows"] != source_rows
):
result["status"] = "partial"
result["error"] = (
f"copied {result['copied_rows']} and excluded "
f"{result['excluded_rows']} of {source_rows} source rows"
)
else:
result["status"] = "complete"
return result
def _copy_state_meta(
source: sqlite3.Connection,
destination: sqlite3.Connection,
@@ -501,13 +722,121 @@ def _copy_state_meta(
return result
def _copy_state_meta_salvage(
source: sqlite3.Connection,
destination: sqlite3.Connection,
*,
chunk_size: int,
progress_cb: Optional[ProgressCallback],
source_rows: Optional[int],
) -> dict[str, Any]:
"""Salvage readable user metadata while regenerating derived FTS state."""
def keep_user_meta(
row: tuple[Any, ...],
columns: tuple[str, ...],
) -> bool:
return str(row[columns.index("key")]) not in _GENERATED_META_KEYS
result = _copy_table_salvage(
source,
destination,
"state_meta",
chunk_size=chunk_size,
progress_cb=progress_cb,
source_rows=source_rows,
insert_prefix="INSERT OR REPLACE",
row_filter=keep_user_meta,
)
result["source_meta_rows"] = result.pop("source_rows")
result["excluded_keys"] = sorted(_GENERATED_META_KEYS)
return result
def _cleanup_partial_orphans(
destination: sqlite3.Connection,
) -> dict[str, Any]:
"""Remove references to sessions that could not be salvaged."""
result: dict[str, Any] = {
"sessions_parent_cleared": 0,
"messages_removed": 0,
"session_model_usage_removed": 0,
"compression_locks_removed": 0,
"telegram_dm_topic_bindings_removed": 0,
}
destination.execute("BEGIN IMMEDIATE")
try:
parent_count = int(
destination.execute(
"SELECT COUNT(*) FROM sessions AS child "
"WHERE child.parent_session_id IS NOT NULL "
"AND NOT EXISTS ("
"SELECT 1 FROM sessions AS parent "
"WHERE parent.id = child.parent_session_id)"
).fetchone()[0]
)
if parent_count:
destination.execute(
"UPDATE sessions SET parent_session_id = NULL "
"WHERE parent_session_id IS NOT NULL "
"AND NOT EXISTS ("
"SELECT 1 FROM sessions AS parent "
"WHERE parent.id = sessions.parent_session_id)"
)
result["sessions_parent_cleared"] = parent_count
dependent_tables = (
("messages", "messages_removed"),
("session_model_usage", "session_model_usage_removed"),
("compression_locks", "compression_locks_removed"),
(
"telegram_dm_topic_bindings",
"telegram_dm_topic_bindings_removed",
),
)
for table, report_key in dependent_tables:
if not _table_columns(destination, table):
continue
orphan_count = int(
destination.execute(
f'SELECT COUNT(*) FROM "{table}" AS dependent '
"WHERE NOT EXISTS ("
"SELECT 1 FROM sessions "
"WHERE sessions.id = dependent.session_id)"
).fetchone()[0]
)
if orphan_count:
destination.execute(
f'DELETE FROM "{table}" '
"WHERE NOT EXISTS ("
"SELECT 1 FROM sessions "
f'WHERE sessions.id = "{table}".session_id)'
)
result[report_key] = orphan_count
destination.execute("COMMIT")
except BaseException:
destination.execute("ROLLBACK")
raise
result["total_removed_or_relinked"] = sum(
int(value) for value in result.values()
)
return result
def _verify_recovered_database(
output: Path,
*,
expected_counts: dict[str, Optional[int]],
copy_report: dict[str, dict[str, Any]],
allow_partial: bool = False,
orphan_cleanup: Optional[dict[str, Any]] = None,
) -> dict[str, Any]:
verification: dict[str, Any] = {"errors": []}
verification: dict[str, Any] = {
"errors": [],
"warnings": [],
"loss_detected": False,
}
open_error = _db_opens_cleanly(output)
verification["opens_cleanly"] = open_error is None
@@ -588,15 +917,37 @@ def _verify_recovered_database(
for table in ("sessions", "messages"):
expected = expected_counts.get(table)
if expected is not None and counts.get(table) != expected:
verification["errors"].append(
message = (
f"{table} count is {counts.get(table)}, expected {expected}"
)
if allow_partial:
verification["warnings"].append(message)
verification["loss_detected"] = True
else:
verification["errors"].append(message)
for table, table_report in copy_report.items():
if table_report.get("status") in {"failed", "partial"}:
verification["errors"].append(
f"{table} copy status is {table_report.get('status')}"
status = table_report.get("status")
if status not in {"failed", "partial"}:
continue
message = f"{table} copy status is {status}"
if allow_partial and (
status == "partial" or table not in {"sessions", "messages"}
):
verification["warnings"].append(message)
verification["loss_detected"] = True
else:
verification["errors"].append(message)
if orphan_cleanup:
orphan_count = int(
orphan_cleanup.get("total_removed_or_relinked") or 0
)
if orphan_count:
verification["warnings"].append(
f"{orphan_count} orphaned reference(s) were removed or relinked"
)
verification["loss_detected"] = True
fts_checks: dict[str, str] = {}
for table in ("messages_fts", "messages_fts_trigram", "messages_fts_cjk"):
@@ -619,7 +970,10 @@ def _verify_recovered_database(
finally:
conn.close()
verification["complete"] = not verification["errors"]
verification["healthy"] = not verification["errors"]
verification["complete"] = bool(
verification["healthy"] and not verification["loss_detected"]
)
return verification
@@ -666,6 +1020,7 @@ def recover_session_database(
work_dir: Optional[Path] = None,
chunk_size: int = 1_000,
progress_cb: Optional[ProgressCallback] = None,
allow_partial: bool = False,
) -> dict[str, Any]:
"""Recover canonical rows into a separate current-schema database.
@@ -686,11 +1041,22 @@ def recover_session_database(
temp_dir, snapshot_source, inspection = _snapshot_and_inspect(source, work_root)
try:
if not inspection.get("recoverable"):
if not inspection.get("recoverable") and not allow_partial:
reasons = "; ".join(inspection.get("errors") or ["unknown source error"])
raise SessionRecoverySourceError(
f"Required canonical tables are not readable: {reasons}"
)
if allow_partial:
missing_required = [
table
for table in ("sessions", "messages")
if not inspection["tables"][table].get("available")
]
if missing_required:
raise SessionRecoverySourceError(
"Partial recovery still requires readable table schemas for: "
+ ", ".join(missing_required)
)
source_conn = sqlite3.connect(
str(snapshot_source),
@@ -721,7 +1087,10 @@ def recover_session_database(
copy_report: dict[str, dict[str, Any]] = {}
for table in _CANONICAL_TABLES:
table_inspection = inspection["tables"][table]
copy_report[table] = _copy_table(
copy_function = (
_copy_table_salvage if allow_partial else _copy_table
)
copy_report[table] = copy_function(
source_conn,
destination_conn,
table,
@@ -732,7 +1101,12 @@ def recover_session_database(
state_meta_inspection = inspection["tables"]["state_meta"]
if state_meta_inspection.get("available"):
copy_report["state_meta"] = _copy_state_meta(
state_meta_copy_function = (
_copy_state_meta_salvage
if allow_partial
else _copy_state_meta
)
copy_report["state_meta"] = state_meta_copy_function(
source_conn,
destination_conn,
chunk_size=chunk_size,
@@ -750,7 +1124,10 @@ def recover_session_database(
"copied_rows": 0,
}
continue
copy_report[table] = _copy_table(
copy_function = (
_copy_table_salvage if allow_partial else _copy_table
)
copy_report[table] = copy_function(
source_conn,
destination_conn,
table,
@@ -758,6 +1135,11 @@ def recover_session_database(
progress_cb=progress_cb,
source_rows=table_inspection.get("rows"),
)
orphan_cleanup = (
_cleanup_partial_orphans(destination_conn)
if allow_partial
else None
)
derived_metadata = _finalize_derived_metadata(destination_conn)
finally:
source_conn.close()
@@ -773,6 +1155,8 @@ def recover_session_database(
for table in _CANONICAL_TABLES
},
copy_report=copy_report,
allow_partial=allow_partial,
orphan_cleanup=orphan_cleanup,
)
source_unchanged = (
_source_fingerprint(source) == inspection["source_fingerprint"]
@@ -785,6 +1169,7 @@ def recover_session_database(
return {
"operation": "recover",
"allow_partial": allow_partial,
"source": str(source),
"output": str(output),
"source_bundle": inspection["source_bundle"],
@@ -798,9 +1183,12 @@ def recover_session_database(
"warnings": inspection["warnings"],
},
"copy": copy_report,
"orphan_cleanup": orphan_cleanup,
"derived_metadata": derived_metadata,
"verification": verification,
"complete": bool(verification.get("complete") and source_unchanged),
"partial": bool(verification.get("loss_detected")),
"verified": bool(verification.get("healthy") and source_unchanged),
"installed": False,
}
finally:
+362
View File
@@ -114,6 +114,180 @@ def _orphan_fts_schema(path: Path) -> None:
conn.close()
def _make_page_spanning_source(
path: Path,
message_count: int = 320,
) -> tuple[int, int | None]:
db = SessionDB(db_path=path)
try:
db.create_session(
"partial-recovery-session",
"cli",
cwd="/tmp/partial-recovery",
)
for message_number in range(message_count):
db.append_message(
"partial-recovery-session",
"user" if message_number % 2 == 0 else "assistant",
(
f"partial recovery payload {message_number:04d} "
+ chr(65 + message_number % 26) * 1_500
),
)
finally:
db.close()
conn = sqlite3.connect(str(path), isolation_level=None)
try:
conn.execute("PRAGMA wal_checkpoint(TRUNCATE)")
conn.execute("PRAGMA journal_mode=DELETE")
conn.execute("VACUUM")
plan = " ".join(
str(row[3])
for row in conn.execute(
"EXPLAIN QUERY PLAN SELECT COUNT(*) FROM messages"
).fetchall()
)
count_index = next(
(
str(row[0])
for row in conn.execute(
"SELECT name FROM sqlite_master "
"WHERE type = 'index' AND tbl_name = 'messages'"
).fetchall()
if plan.endswith(str(row[0]))
),
None,
)
names = ["messages"]
if count_index is not None:
names.append(count_index)
placeholders = ", ".join("?" for _ in names)
roots = {
str(row[0]): int(row[1])
for row in conn.execute(
"SELECT name, rootpage FROM sqlite_master "
f"WHERE name IN ({placeholders})",
tuple(names),
).fetchall()
}
return roots["messages"], (
roots[count_index] if count_index is not None else None
)
finally:
conn.close()
def _make_many_sessions_source(
path: Path,
session_count: int = 180,
) -> int:
db = SessionDB(db_path=path)
try:
for session_number in range(session_count):
session_id = f"partial-session-{session_number:04d}"
db.create_session(
session_id,
"cli",
cwd=f"/tmp/partial-session-{session_number:04d}",
system_prompt=(
f"session payload {session_number:04d} "
+ chr(65 + session_number % 26) * 1_500
),
)
db.append_message(session_id, "user", f"message {session_number}")
finally:
db.close()
conn = sqlite3.connect(str(path), isolation_level=None)
try:
conn.execute("PRAGMA wal_checkpoint(TRUNCATE)")
conn.execute("PRAGMA journal_mode=DELETE")
conn.execute("VACUUM")
row = conn.execute(
"SELECT rootpage FROM sqlite_master "
"WHERE type = 'table' AND name = 'sessions'"
).fetchone()
assert row is not None
return int(row[0])
finally:
conn.close()
def _btree_leaf_pages(path: Path, root_page: int) -> tuple[int, list[int]]:
data = path.read_bytes()
page_size = int.from_bytes(data[16:18], "big")
if page_size == 1:
page_size = 65_536
leaf_pages: list[int] = []
visited: set[int] = set()
def visit(page_number: int) -> None:
if page_number in visited:
return
visited.add(page_number)
page_start = (page_number - 1) * page_size
header_offset = page_start + (100 if page_number == 1 else 0)
page_type = data[header_offset]
cell_count = int.from_bytes(
data[header_offset + 3 : header_offset + 5],
"big",
)
if page_type in {0x0A, 0x0D}:
leaf_pages.append(page_number)
return
assert page_type in {0x02, 0x05}, (
f"unexpected table b-tree page type {page_type:#x} "
f"on page {page_number}"
)
pointer_array = header_offset + 12
for cell_number in range(cell_count):
pointer_offset = pointer_array + cell_number * 2
cell_offset = int.from_bytes(
data[pointer_offset : pointer_offset + 2],
"big",
)
child_offset = page_start + cell_offset
child_page = int.from_bytes(
data[child_offset : child_offset + 4],
"big",
)
visit(child_page)
rightmost_page = int.from_bytes(
data[header_offset + 8 : header_offset + 12],
"big",
)
visit(rightmost_page)
visit(root_page)
return page_size, leaf_pages
def _corrupt_middle_table_leaf(
path: Path,
root_page: int,
*,
require_interior: bool = True,
) -> int:
page_size, leaf_pages = _btree_leaf_pages(path, root_page)
assert leaf_pages
if require_interior:
assert len(leaf_pages) >= 3
leaf_page = leaf_pages[len(leaf_pages) // 2]
page_start = (leaf_page - 1) * page_size
header_offset = page_start + (100 if leaf_page == 1 else 0)
data = bytearray(path.read_bytes())
assert data[header_offset] in {0x0A, 0x0D}
# An impossible cell count damages this one middle leaf while preserving
# the table root and leaves on both sides. This is a physical SQLite page
# failure, not a mocked cursor exception.
data[header_offset + 3 : header_offset + 5] = b"\xff\xff"
path.write_bytes(data)
return leaf_page
def test_recovery_rebuilds_canonical_data_without_opening_source(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
@@ -281,6 +455,194 @@ def test_recovery_requires_readable_sessions_and_messages(tmp_path: Path) -> Non
assert not output.exists()
def test_allow_partial_still_reports_a_complete_healthy_copy(
tmp_path: Path,
) -> None:
source = tmp_path / "healthy-state.db"
output = tmp_path / "healthy-recovered.db"
expected = _make_source(source)
report = recover_session_database(
source,
output,
work_dir=tmp_path,
chunk_size=4,
allow_partial=True,
)
assert report["verified"] is True
assert report["complete"] is True
assert report["partial"] is False
assert report["verification"]["warnings"] == []
assert report["verification"]["table_counts"]["sessions"] == expected["sessions"]
assert report["verification"]["table_counts"]["messages"] == expected["messages"]
assert report["orphan_cleanup"]["total_removed_or_relinked"] == 0
def test_cli_allow_partial_salvages_rows_across_a_corrupt_leaf(
tmp_path: Path,
) -> None:
source = tmp_path / "corrupt-state.db"
rejected_output = tmp_path / "rejected.db"
output = tmp_path / "partial-recovered.db"
message_count = 320
messages_root, count_index_root = _make_page_spanning_source(
source,
message_count,
)
corrupt_page = _corrupt_middle_table_leaf(source, messages_root)
if count_index_root is not None:
_corrupt_middle_table_leaf(
source,
count_index_root,
require_interior=False,
)
source_hash = _sha256(source)
inspection = inspect_session_database(source, work_dir=tmp_path)
assert inspection["recoverable"] is False
assert inspection["tables"]["messages"]["rows"] is None
with pytest.raises(SessionRecoverySourceError, match="messages"):
recover_session_database(
source,
rejected_output,
work_dir=tmp_path,
)
assert not rejected_output.exists()
env = os.environ.copy()
env["HERMES_HOME"] = str(tmp_path / "isolated-hermes-home")
result = subprocess.run(
[
sys.executable,
"-m",
"hermes_cli.main",
"sessions",
"recover",
"--source",
str(source),
"--output",
str(output),
"--work-dir",
str(tmp_path),
"--chunk-size",
"8",
"--allow-partial",
],
cwd=Path(__file__).resolve().parents[2],
env=env,
capture_output=True,
text=True,
timeout=60,
check=False,
)
assert result.returncode == 0, result.stdout + result.stderr
assert "Partial recovery output verified" in result.stdout
assert "active session database was not changed" in result.stdout
assert _sha256(source) == source_hash
report_path = output.with_name(output.name + ".recovery.json")
report = json.loads(report_path.read_text(encoding="utf-8"))
assert report["allow_partial"] is True
assert report["verified"] is True
assert report["complete"] is False
assert report["partial"] is True
assert report["installed"] is False
assert report["source_unchanged"] is True
assert report["verification"]["healthy"] is True
assert report["verification"]["integrity_check"] == ["ok"]
assert report["verification"]["foreign_key_check"] == []
assert report["verification"]["table_counts"]["sessions"] == 1
copied_messages = report["copy"]["messages"]
assert copied_messages["status"] == "partial"
assert copied_messages["copied_rows"] < message_count
assert copied_messages["copied_rows"] > 0
assert copied_messages["skipped_rowid_ranges"]
assert any(
item["low"] <= message_count and item["high"] >= 1
for item in copied_messages["skipped_rowid_ranges"]
)
assert copied_messages["query_limit_reached"] is False
conn = sqlite3.connect(str(output))
try:
recovered_ids = {
int(row[0]) for row in conn.execute("SELECT id FROM messages")
}
assert 1 in recovered_ids
assert message_count in recovered_ids
assert len(recovered_ids) == copied_messages["copied_rows"]
assert conn.execute("PRAGMA integrity_check").fetchall() == [("ok",)]
finally:
conn.close()
# Prove the helper damaged an interior data leaf, so successful recovery of
# the first and last message IDs really crossed the corrupted region.
assert corrupt_page not in {
min(_btree_leaf_pages(source, messages_root)[1]),
max(_btree_leaf_pages(source, messages_root)[1]),
}
def test_partial_recovery_removes_messages_for_unreadable_sessions(
tmp_path: Path,
) -> None:
source = tmp_path / "corrupt-sessions.db"
output = tmp_path / "partial-sessions.db"
session_count = 180
sessions_root = _make_many_sessions_source(
source,
session_count,
)
_corrupt_middle_table_leaf(source, sessions_root)
source_hash = _sha256(source)
report = recover_session_database(
source,
output,
work_dir=tmp_path,
chunk_size=8,
allow_partial=True,
)
assert report["verified"] is True
assert report["complete"] is False
assert report["partial"] is True
assert report["source_unchanged"] is True
assert _sha256(source) == source_hash
assert report["copy"]["sessions"]["status"] == "partial"
assert report["copy"]["messages"]["status"] == "complete"
removed_messages = report["orphan_cleanup"]["messages_removed"]
assert removed_messages > 0
assert report["orphan_cleanup"]["total_removed_or_relinked"] >= removed_messages
assert report["verification"]["foreign_key_check"] == []
conn = sqlite3.connect(str(output))
try:
recovered_sessions = {
str(row[0]) for row in conn.execute("SELECT id FROM sessions")
}
assert "partial-session-0000" in recovered_sessions
assert f"partial-session-{session_count - 1:04d}" in recovered_sessions
assert 0 < len(recovered_sessions) < session_count
assert (
conn.execute(
"SELECT COUNT(*) FROM messages AS message "
"WHERE NOT EXISTS ("
"SELECT 1 FROM sessions "
"WHERE sessions.id = message.session_id)"
).fetchone()[0]
== 0
)
assert conn.execute("SELECT COUNT(*) FROM messages").fetchone()[0] == len(
recovered_sessions
)
finally:
conn.close()
def test_cli_recover_writes_verified_report_without_touching_source(
tmp_path: Path,
) -> None: