fix(sessions): add opt-in partial database recovery
This commit is contained in:
@@ -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
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user