package spool import ( "bytes" "context" "database/sql" "errors" "fmt" "github.com/rvbox/rvbox/internal/domain" ) var ( ErrIntegrityCheck = errors.New("client spool SQLite integrity check failed") ErrQuotaCounterMismatch = errors.New("client spool quota counter mismatch") ErrStoredPayloadChecksum = errors.New("client spool stored payload checksum mismatch") ) type RecoveryReport struct { EventsChecked uint64 ScriptsChecked uint64 } // Check verifies state that SQLite's own structural check cannot express. The // daemon runs it asynchronously during startup and exposes failure as dirty // local health; it never invents a replacement identity or counter value. func (store *Store) Check(ctx context.Context) (RecoveryReport, error) { if store == nil || store.db == nil { return RecoveryReport{}, sql.ErrConnDone } if err := quickCheck(ctx, store.db); err != nil { return RecoveryReport{}, err } if err := checkQuotaCounters(ctx, store.db); err != nil { return RecoveryReport{}, err } return checkPayloads(ctx, store.db, store.maxScriptBytes) } func quickCheck(ctx context.Context, database *sql.DB) error { rows, err := database.QueryContext(ctx, `PRAGMA quick_check(100)`) if err != nil { return fmt.Errorf("%w: query failed", ErrIntegrityCheck) } defer rows.Close() count := 0 for rows.Next() { var result string if err := rows.Scan(&result); err != nil { return fmt.Errorf("%w: unreadable result", ErrIntegrityCheck) } count++ if result != "ok" { return fmt.Errorf("%w: %s", ErrIntegrityCheck, result) } } if err := rows.Err(); err != nil || count != 1 { return fmt.Errorf("%w: incomplete result", ErrIntegrityCheck) } return nil } func checkQuotaCounters(ctx context.Context, database *sql.DB) error { var invalid int err := database.QueryRowContext(ctx, `SELECT count(*) FROM commands AS c WHERE c.total_charged_bytes != c.base_charged_bytes + COALESCE((SELECT sum(charged_bytes) FROM events WHERE issue_uuid = c.issue_uuid), 0) + COALESCE((SELECT charged_bytes FROM scripts WHERE issue_uuid = c.issue_uuid), 0) OR c.output_charged_bytes != COALESCE((SELECT sum(charged_bytes) FROM events WHERE issue_uuid = c.issue_uuid AND output = 1), 0)`).Scan(&invalid) if err != nil { return err } var stored, computed uint64 if err := database.QueryRowContext(ctx, `SELECT client_total_charged_bytes, COALESCE((SELECT sum(total_charged_bytes) FROM commands), 0) FROM spool_counters WHERE singleton = 1`).Scan(&stored, &computed); err != nil { return err } if invalid != 0 || stored != computed { return ErrQuotaCounterMismatch } return nil } func checkPayloads(ctx context.Context, database *sql.DB, maxScriptBytes uint64) (RecoveryReport, error) { var report RecoveryReport rows, err := database.QueryContext(ctx, `SELECT issue_uuid, payload, payload_sha256 FROM events ORDER BY issue_uuid, local_ordinal`) if err != nil { return report, err } for rows.Next() { var owner, payload, digest []byte if err := rows.Scan(&owner, &payload, &digest); err != nil { _ = rows.Close() return report, err } computed := immutableDigest(payload) if _, valid := copiedUUID(owner); !valid || len(digest) != 32 || !bytes.Equal(computed[:], digest) { _ = rows.Close() return report, ErrStoredPayloadChecksum } report.EventsChecked++ } if err := rows.Close(); err != nil { return report, err } scripts, err := database.QueryContext(ctx, `SELECT issue_uuid, received_raw_bytes, stored_bytes, stored_data FROM scripts ORDER BY issue_uuid`) if err != nil { return report, err } defer scripts.Close() for scripts.Next() { var owner, stored []byte var received, storedBytes uint64 if err := scripts.Scan(&owner, &received, &storedBytes, &stored); err != nil { return report, err } if _, valid := copiedUUID(owner); !valid || storedBytes != uint64(len(stored)) { return report, ErrScriptState } if _, err := decompressScript(stored, received, maxScriptBytes); err != nil { return report, err } report.ScriptsChecked++ } return report, scripts.Err() } // copiedUUID is intentionally kept here to make a corrupted BLOB type obvious // at the recovery boundary rather than relying on SQLite's dynamic conversions. func copiedUUID(value []byte) (domain.UUID, bool) { var result domain.UUID if len(value) != len(result) { return result, false } copy(result[:], value) return result, validUUID(result) }