diff --git a/internal/storage/backfill.go b/internal/storage/backfill.go index b73fc0a..26f434d 100644 --- a/internal/storage/backfill.go +++ b/internal/storage/backfill.go @@ -3,13 +3,15 @@ package storage import ( "context" "database/sql" + "errors" "fmt" + "sort" "strings" + "time" "github.com/pablontiv/backscroll/internal/corrections" "github.com/pablontiv/backscroll/internal/models" "github.com/pablontiv/backscroll/internal/templates" - "time" ) // BackfillDerivedOpts configures BackfillDerived behavior. @@ -23,38 +25,37 @@ type BackfillDerivedOpts struct { // Incremented when template mining heuristics change (e.g., v1→v2). const CurrentNormalizationVersion = 2 -// BackfillDerived mines templates, corrections, and lossy tool_events from -// stored text for files that are EXPIRED (absent from disk) or have STALE TEMPLATES -// (with normalization_version < current). Results are inserted idempotently (INSERT OR IGNORE). -// Extraction_version=0 marks lossy (reverse-parsed) rows. On-disk files are handled by B1's -// rich re-parse path; this path avoids duplicate lossy rows. +type derivedBackfillFile struct { + SourcePath string +} + +// BackfillDerived mines derived data without cancellation. It is retained for +// compatibility with existing callers. func (d *Database) BackfillDerived(opts BackfillDerivedOpts) error { - type fileToBackfill struct { - SourcePath string - Source string + return d.BackfillDerivedContext(context.Background(), opts) +} + +// BackfillDerivedContext mines templates, corrections, and lossy tool_events +// from stored text for expired, recovered, or stale-template paths. Discovery +// and processing are deterministic by source path. Each batch is atomic: +// cancellation rolls back the active batch, preserves earlier batches, and +// prevents later batches from starting. +func (d *Database) BackfillDerivedContext(ctx context.Context, opts BackfillDerivedOpts) error { + if err := ctx.Err(); err != nil { + return err } - // First, discover and process stale templates (v1 → v2 epoch upgrade). - // These are templates with normalization_version < CurrentNormalizationVersion. - stalePaths, err := d.StaleTemplatePaths(CurrentNormalizationVersion) + stalePaths, err := d.StaleTemplatePathsContext(ctx, CurrentNormalizationVersion) if err != nil { return fmt.Errorf("query stale template paths: %w", err) } - - // Deduplicate stale paths and expired files using a map. - pathsToProcess := make(map[string]string) - for _, p := range stalePaths { - pathsToProcess[p] = "session" + pathsToProcess := make(map[string]struct{}, len(stalePaths)) + for _, path := range stalePaths { + pathsToProcess[path] = struct{}{} } - // Find files in search_items that are EXPIRED (absent from indexed_files) - // or provisionally recovered. Within those files, process only those missing - // at least one of the three derivations: - // - template_matches (templates mined), OR - // - correction_signals (corrections detected), OR - // - tool_events with extraction_version = 0 (lossy tool metadata extracted) - rows, err := d.db.Query(` - SELECT DISTINCT si.source_path, si.source + rows, err := d.db.QueryContext(ctx, ` + SELECT DISTINCT si.source_path FROM search_items si LEFT JOIN indexed_files ifx ON si.source_path = ifx.path WHERE @@ -62,129 +63,154 @@ func (d *Database) BackfillDerived(opts BackfillDerivedOpts) error { (NOT EXISTS (SELECT 1 FROM template_matches WHERE source_path = si.source_path) OR NOT EXISTS (SELECT 1 FROM correction_signals WHERE source_path = si.source_path) OR NOT EXISTS (SELECT 1 FROM tool_events WHERE source_path = si.source_path AND extraction_version = 0)) - ORDER BY si.source_path + ORDER BY si.source_path ASC `, recoveredSourceHash) if err != nil { return fmt.Errorf("query expired files: %w", err) } - defer rows.Close() - for rows.Next() { - var sourcePath, source string - if err := rows.Scan(&sourcePath, &source); err != nil { + if err := ctx.Err(); err != nil { + _ = rows.Close() + return err + } + var sourcePath string + if err := rows.Scan(&sourcePath); err != nil { + _ = rows.Close() return fmt.Errorf("scan file: %w", err) } - pathsToProcess[sourcePath] = source + pathsToProcess[sourcePath] = struct{}{} } - - // Convert map to list of files to backfill - var filesToBackfill []fileToBackfill - for path, source := range pathsToProcess { - filesToBackfill = append(filesToBackfill, fileToBackfill{path, source}) + if err := ctx.Err(); err != nil { + _ = rows.Close() + return err } if err := rows.Err(); err != nil { + _ = rows.Close() return fmt.Errorf("iterate expired files: %w", err) } + if err := rows.Close(); err != nil { + return fmt.Errorf("close expired files query: %w", err) + } - if len(filesToBackfill) == 0 { - return nil // nothing to backfill + paths := make([]string, 0, len(pathsToProcess)) + for path := range pathsToProcess { + paths = append(paths, path) + } + sort.Strings(paths) + files := make([]derivedBackfillFile, 0, len(paths)) + for _, path := range paths { + files = append(files, derivedBackfillFile{SourcePath: path}) } const batchSize = 100 - totalTemplates := 0 - totalSignals := 0 - totalEvents := 0 - - for batchStart := 0; batchStart < len(filesToBackfill); batchStart += batchSize { + var totalTemplates, totalSignals, totalEvents int + for batchStart := 0; batchStart < len(files); batchStart += batchSize { + if err := ctx.Err(); err != nil { + return err + } batchEnd := batchStart + batchSize - if batchEnd > len(filesToBackfill) { - batchEnd = len(filesToBackfill) + if batchEnd > len(files) { + batchEnd = len(files) } - batch := filesToBackfill[batchStart:batchEnd] - - tx, err := d.db.Begin() + batchTemplates, batchSignals, batchEvents, err := d.backfillDerivedBatch(ctx, files[batchStart:batchEnd]) if err != nil { - return fmt.Errorf("begin transaction: %w", err) + return err } - - batchTemplates := 0 - batchSignals := 0 - batchEvents := 0 - - for _, file := range batch { - // Load messages for this file from search_items - msgRows, err := tx.Query(` - SELECT ordinal, role, text, uuid, content_type, was_interrupted - FROM search_items - WHERE source_path = ? - ORDER BY ordinal - `, file.SourcePath) - if err != nil { - _ = tx.Rollback() - return fmt.Errorf("load messages for %s: %w", file.SourcePath, err) + totalTemplates += batchTemplates + totalSignals += batchSignals + totalEvents += batchEvents + if opts.OnProgress != nil { + opts.OnProgress(batchEnd, totalTemplates, totalSignals, totalEvents) + if err := ctx.Err(); err != nil { + return err } + } + } + if err := ctx.Err(); err != nil { + return err + } + return nil +} - var messages []IndexedMessage - for msgRows.Next() { - var m IndexedMessage - var uuid sql.NullString - var wasInterrupted sql.NullInt64 - if err := msgRows.Scan(&m.Ordinal, &m.Role, &m.Text, &uuid, - &m.ContentType, &wasInterrupted); err != nil { - msgRows.Close() - _ = tx.Rollback() - return fmt.Errorf("scan message: %w", err) - } - if uuid.Valid { - m.UUID = uuid.String - } - m.WasInterrupted = wasInterrupted.Valid && wasInterrupted.Int64 != 0 - m.Timestamp = time.Now().Format(time.RFC3339) // not needed for backfill - m.ExtractionVersion = 0 // lossy marker - messages = append(messages, m) - } - msgRows.Close() +func (d *Database) backfillDerivedBatch(ctx context.Context, batch []derivedBackfillFile) (templatesCount, signalsCount, eventsCount int, retErr error) { + tx, err := d.db.BeginTx(context.WithoutCancel(ctx), nil) + if err != nil { + return 0, 0, 0, fmt.Errorf("begin backfill batch: %w", err) + } + gate := newSyncTransactionGate(ctx, tx) + defer func() { + if rollbackErr := gate.rollback(); rollbackErr != nil && !errors.Is(retErr, rollbackErr) { + retErr = errors.Join(retErr, fmt.Errorf("rollback backfill batch: %w", rollbackErr)) + } + if cancelErr := gate.cancellationBeforeCommit(); cancelErr != nil && !errors.Is(retErr, cancelErr) { + retErr = errors.Join(cancelErr, retErr) + } + }() - // Mine templates from tool messages with is_error - miner := templates.NewMiner() - templateCount, err := d.backfillTemplatesForFile(tx, file.SourcePath, messages, miner) - if err != nil { - _ = tx.Rollback() - return fmt.Errorf("backfill templates for %s: %w", file.SourcePath, err) + for _, file := range batch { + if err := ctx.Err(); err != nil { + return 0, 0, 0, err + } + msgRows, err := tx.QueryContext(ctx, ` + SELECT ordinal, role, text, uuid, content_type, was_interrupted + FROM search_items WHERE source_path = ? ORDER BY ordinal ASC + `, file.SourcePath) + if err != nil { + return 0, 0, 0, fmt.Errorf("load messages for %s: %w", file.SourcePath, err) + } + var messages []IndexedMessage + for msgRows.Next() { + if err := ctx.Err(); err != nil { + _ = msgRows.Close() + return 0, 0, 0, err } - batchTemplates += templateCount - - // Mine corrections from user prose messages (content_type='text'|'code') - signalCount, err := d.backfillCorrectionsForFile(tx, file.SourcePath, messages) - if err != nil { - _ = tx.Rollback() - return fmt.Errorf("backfill corrections for %s: %w", file.SourcePath, err) + var message IndexedMessage + var uuid sql.NullString + var wasInterrupted sql.NullInt64 + if err := msgRows.Scan(&message.Ordinal, &message.Role, &message.Text, &uuid, &message.ContentType, &wasInterrupted); err != nil { + _ = msgRows.Close() + return 0, 0, 0, fmt.Errorf("scan message for %s: %w", file.SourcePath, err) } - batchSignals += signalCount - - // Extract lossy tool_events (uuid-NULL rows) - eventCount, err := d.backfillToolEventsForFile(tx, file.SourcePath, messages) - if err != nil { - _ = tx.Rollback() - return fmt.Errorf("backfill tool_events for %s: %w", file.SourcePath, err) + if uuid.Valid { + message.UUID = uuid.String } - batchEvents += eventCount + message.WasInterrupted = wasInterrupted.Valid && wasInterrupted.Int64 != 0 + message.Timestamp = time.Now().Format(time.RFC3339) + message.ExtractionVersion = 0 + messages = append(messages, message) } - - if err := tx.Commit(); err != nil { - return fmt.Errorf("commit backfill batch: %w", err) + if err := ctx.Err(); err != nil { + _ = msgRows.Close() + return 0, 0, 0, err + } + if err := msgRows.Err(); err != nil { + _ = msgRows.Close() + return 0, 0, 0, fmt.Errorf("iterate messages for %s: %w", file.SourcePath, err) + } + if err := msgRows.Close(); err != nil { + return 0, 0, 0, fmt.Errorf("close messages for %s: %w", file.SourcePath, err) } - totalTemplates += batchTemplates - totalSignals += batchSignals - totalEvents += batchEvents - - if opts.OnProgress != nil { - opts.OnProgress(batchEnd, totalTemplates, totalSignals, totalEvents) + count, err := d.backfillTemplatesForFileContext(ctx, tx, file.SourcePath, messages, templates.NewMiner()) + if err != nil { + return 0, 0, 0, fmt.Errorf("backfill templates for %s: %w", file.SourcePath, err) + } + templatesCount += count + count, err = d.backfillCorrectionsForFileContext(ctx, tx, file.SourcePath, messages) + if err != nil { + return 0, 0, 0, fmt.Errorf("backfill corrections for %s: %w", file.SourcePath, err) } + signalsCount += count + count, err = d.backfillToolEventsForFileContext(ctx, tx, file.SourcePath, messages) + if err != nil { + return 0, 0, 0, fmt.Errorf("backfill tool_events for %s: %w", file.SourcePath, err) + } + eventsCount += count } - - return nil + if err := gate.commit(); err != nil { + return 0, 0, 0, fmt.Errorf("commit backfill batch: %w", err) + } + return templatesCount, signalsCount, eventsCount, nil } // backfillTemplatesForFile mines templates from tool messages in the file. @@ -459,8 +485,18 @@ func (d *Database) backfillCorrectionsForFileContext(ctx context.Context, tx *sq // NOTE: outputs (tool_result text) cannot be attributed without tool_use_id linkage, // so they are skipped (ParseToolFromSerialized returns empty toolName for outputs). func (d *Database) backfillToolEventsForFile(tx *sql.Tx, sourcePath string, messages []IndexedMessage) (int, error) { + return d.backfillToolEventsForFileContext(context.Background(), tx, sourcePath, messages) +} + +func (d *Database) backfillToolEventsForFileContext(ctx context.Context, tx *sql.Tx, sourcePath string, messages []IndexedMessage) (int, error) { + if err := ctx.Err(); err != nil { + return 0, err + } count := 0 for _, m := range messages { + if err := ctx.Err(); err != nil { + return 0, err + } if m.ContentType != "tool" { continue } @@ -479,7 +515,7 @@ func (d *Database) backfillToolEventsForFile(tx *sql.Tx, sourcePath string, mess } // uuid-NULL for lossy rows (no tool_use_id linkage available) - _, err := tx.Exec(` + _, err := tx.ExecContext(ctx, ` INSERT OR IGNORE INTO tool_events (message_uuid, source_path, ordinal, tool_name, command_head, extraction_version) VALUES (?, ?, ?, ?, ?, ?) diff --git a/internal/storage/queries.go b/internal/storage/queries.go index b6ce977..9001dbb 100644 --- a/internal/storage/queries.go +++ b/internal/storage/queries.go @@ -628,24 +628,73 @@ func (d *Database) OptimizeFTS() error { return nil } -// RebuildFTS re-derives both FTS indexes from search_items using FTS5's -// external-content 'rebuild' command. It never touches search_items rows — -// the DB, not the filesystem, is the source of truth (perennity contract). +// RebuildFTS re-derives both FTS indexes without cancellation. +// It is retained for compatibility with existing callers. func (d *Database) RebuildFTS() error { - // Single transaction: either both indexes re-derive or neither does — - // a partial rebuild would leave one index stale and queries inconsistent. - tx, err := d.db.Begin() + return d.RebuildFTSContext(context.Background()) +} + +// RebuildFTSContext re-derives both FTS indexes from search_items. FTS5's +// external-content rebuild command cannot preserve the content-type routing +// contract because it indexes every content row, so each index is cleared and +// repopulated with the same exact predicates used by the schema triggers. +// +// Both indexes are replaced in one transaction: cancellation or any error +// restores both prior indexes. search_items is never mutated; the database, +// not the filesystem, remains the perennial source of truth. +func (d *Database) RebuildFTSContext(ctx context.Context) (retErr error) { + if err := ctx.Err(); err != nil { + return err + } + + // Keep transaction finalization under the gate rather than database/sql's + // automatic context rollback so cancellation and commit have one owner. + tx, err := d.db.BeginTx(context.WithoutCancel(ctx), nil) if err != nil { - return fmt.Errorf("begin transaction: %w", err) + return fmt.Errorf("begin FTS rebuild transaction: %w", err) } - defer func() { _ = tx.Rollback() }() - if _, err := tx.Exec(`INSERT INTO messages_fts(messages_fts) VALUES('rebuild')`); err != nil { - return fmt.Errorf("rebuild messages_fts: %w", err) + gate := newSyncTransactionGate(ctx, tx) + defer func() { + if rollbackErr := gate.rollback(); rollbackErr != nil && !errors.Is(retErr, rollbackErr) { + retErr = errors.Join(retErr, fmt.Errorf("rollback FTS rebuild transaction: %w", rollbackErr)) + } + if cancelErr := gate.cancellationBeforeCommit(); cancelErr != nil && !errors.Is(retErr, cancelErr) { + retErr = errors.Join(cancelErr, retErr) + } + }() + + statements := []struct { + label string + sql string + }{ + {"clear messages_fts", `INSERT INTO messages_fts(messages_fts) VALUES('delete-all')`}, + {"repopulate messages_fts", ` + INSERT INTO messages_fts(rowid, text) + SELECT id, text FROM search_items + WHERE content_type IN ('text', 'code', 'reasoning') + `}, + {"clear tool_fts", `INSERT INTO tool_fts(tool_fts) VALUES('delete-all')`}, + {"repopulate tool_fts", ` + INSERT INTO tool_fts(rowid, text) + SELECT id, text FROM search_items + WHERE content_type = 'tool' + `}, + } + for _, statement := range statements { + if err := ctx.Err(); err != nil { + return err + } + if _, err := tx.ExecContext(ctx, statement.sql); err != nil { + if ctxErr := ctx.Err(); ctxErr != nil { + return ctxErr + } + return fmt.Errorf("%s: %w", statement.label, err) + } } - if _, err := tx.Exec(`INSERT INTO tool_fts(tool_fts) VALUES('rebuild')`); err != nil { - return fmt.Errorf("rebuild tool_fts: %w", err) + if err := gate.commit(); err != nil { + return fmt.Errorf("commit FTS rebuild: %w", err) } - return tx.Commit() + return nil } // TemplateQueryOpts controls template aggregation queries. @@ -1266,69 +1315,91 @@ func (d *Database) filterEchoZeroPathsContext(ctx context.Context, paths []strin return out, nil } -// ReresolveProjects iterates all distinct source_paths where project='unknown' or project IS NULL, -// calls the resolver function for each path, and updates ALL rows for that path with the returned project ID. -// If resolver returns empty string or "unknown", the source_path is skipped and rows remain unchanged. -// Returns the count of DISTINCT source_paths that were successfully resolved (project changed). +// ReresolveProjects retains the context-taking API used by existing callers. func (d *Database) ReresolveProjects(ctx context.Context, resolver func(sourcePath string) string) (int64, error) { - // Find all distinct source_paths with unknown or NULL project + return d.ReresolveProjectsContext(ctx, resolver) +} + +// ReresolveProjectsContext resolves unknown project labels in deterministic +// source-path order. Cancellation rolls back the complete resolution set. +func (d *Database) ReresolveProjectsContext(ctx context.Context, resolver func(sourcePath string) string) (resolvedPaths int64, retErr error) { + if err := ctx.Err(); err != nil { + return 0, err + } rows, err := d.db.QueryContext(ctx, ` SELECT DISTINCT source_path FROM search_items WHERE project = 'unknown' OR project IS NULL + ORDER BY source_path ASC `) if err != nil { return 0, fmt.Errorf("query unknown source_paths: %w", err) } - defer rows.Close() - var sourcePaths []string for rows.Next() { + if err := ctx.Err(); err != nil { + _ = rows.Close() + return 0, err + } var path string if err := rows.Scan(&path); err != nil { + _ = rows.Close() return 0, fmt.Errorf("scan source_path: %w", err) } sourcePaths = append(sourcePaths, path) } + if err := ctx.Err(); err != nil { + _ = rows.Close() + return 0, err + } if err := rows.Err(); err != nil { + _ = rows.Close() return 0, fmt.Errorf("iterate source_paths: %w", err) } - + if err := rows.Close(); err != nil { + return 0, fmt.Errorf("close source_paths query: %w", err) + } if len(sourcePaths) == 0 { return 0, nil } - // Resolve each path and update in a single transaction - tx, err := d.db.BeginTx(ctx, nil) + tx, err := d.db.BeginTx(context.WithoutCancel(ctx), nil) if err != nil { - return 0, fmt.Errorf("begin transaction: %w", err) + return 0, fmt.Errorf("begin project resolution transaction: %w", err) } - defer func() { _ = tx.Rollback() }() + gate := newSyncTransactionGate(ctx, tx) + defer func() { + if rollbackErr := gate.rollback(); rollbackErr != nil && !errors.Is(retErr, rollbackErr) { + retErr = errors.Join(retErr, fmt.Errorf("rollback project resolution transaction: %w", rollbackErr)) + } + if cancelErr := gate.cancellationBeforeCommit(); cancelErr != nil && !errors.Is(retErr, cancelErr) { + retErr = errors.Join(cancelErr, retErr) + } + }() - var resolvedPaths int64 for _, sourcePath := range sourcePaths { + if err := ctx.Err(); err != nil { + return 0, err + } resolvedID := resolver(sourcePath) - - // Skip if resolver returned empty or "unknown" + if err := ctx.Err(); err != nil { + return 0, err + } if resolvedID == "" || resolvedID == "unknown" { continue } - - // Update all rows for this source_path - _, err := tx.ExecContext(ctx, ` + if _, err := tx.ExecContext(ctx, ` UPDATE search_items SET project = ? WHERE source_path = ? - `, resolvedID, sourcePath) - if err != nil { + `, resolvedID, sourcePath); err != nil { + if ctxErr := ctx.Err(); ctxErr != nil { + return 0, ctxErr + } return 0, fmt.Errorf("update source_path %s: %w", sourcePath, err) } - - // Count this source_path as successfully resolved resolvedPaths++ } - - if err := tx.Commit(); err != nil { - return 0, fmt.Errorf("commit resolution transaction: %w", err) + if err := gate.commit(); err != nil { + return 0, fmt.Errorf("commit project resolution transaction: %w", err) } - return resolvedPaths, nil } @@ -1339,80 +1410,101 @@ func (d *Database) ReresolveProjects(ctx context.Context, resolver func(sourcePa // from the stored fallback. Returns count of source_paths updated. // Only registry matches count — fallback-only paths are skipped (no churn). func (d *Database) ReresolveProjectsWithRegistry(ctx context.Context, registry projects.ProjectRegistry) (int64, error) { + return d.ReresolveProjectsWithRegistryContext(ctx, registry) +} + +// ReresolveProjectsWithRegistryContext corrects fallback project labels using +// registry matches. Cancellation rolls back all updates from this call. +func (d *Database) ReresolveProjectsWithRegistryContext(ctx context.Context, registry projects.ProjectRegistry) (updatedPaths int64, retErr error) { + if err := ctx.Err(); err != nil { + return 0, err + } if len(registry.Projects) == 0 { - return 0, nil // no registry entries, nothing to resolve + return 0, nil } - - // Find all distinct (source_path, project) tuples where we might have a fallback label rows, err := d.db.QueryContext(ctx, ` SELECT DISTINCT si.source_path, si.project FROM search_items si WHERE si.project IS NOT NULL - ORDER BY si.source_path + ORDER BY si.source_path ASC, si.project ASC `) if err != nil { return 0, fmt.Errorf("query paths for registry re-resolution: %w", err) } - defer rows.Close() - type pathLabel struct { path string project string } var pathLabels []pathLabel for rows.Next() { - var path, proj string - if err := rows.Scan(&path, &proj); err != nil { + if err := ctx.Err(); err != nil { + _ = rows.Close() + return 0, err + } + var path, project string + if err := rows.Scan(&path, &project); err != nil { + _ = rows.Close() return 0, fmt.Errorf("scan path/project: %w", err) } - pathLabels = append(pathLabels, pathLabel{path, proj}) + pathLabels = append(pathLabels, pathLabel{path: path, project: project}) + } + if err := ctx.Err(); err != nil { + _ = rows.Close() + return 0, err } if err := rows.Err(); err != nil { + _ = rows.Close() return 0, fmt.Errorf("iterate path/project: %w", err) } - + if err := rows.Close(); err != nil { + return 0, fmt.Errorf("close path/project query: %w", err) + } if len(pathLabels) == 0 { return 0, nil } - // Re-resolve each path against the registry - tx, err := d.db.BeginTx(ctx, nil) + tx, err := d.db.BeginTx(context.WithoutCancel(ctx), nil) if err != nil { - return 0, fmt.Errorf("begin transaction: %w", err) + return 0, fmt.Errorf("begin registry re-resolution transaction: %w", err) } - defer func() { _ = tx.Rollback() }() + gate := newSyncTransactionGate(ctx, tx) + defer func() { + if rollbackErr := gate.rollback(); rollbackErr != nil && !errors.Is(retErr, rollbackErr) { + retErr = errors.Join(retErr, fmt.Errorf("rollback registry re-resolution transaction: %w", rollbackErr)) + } + if cancelErr := gate.cancellationBeforeCommit(); cancelErr != nil && !errors.Is(retErr, cancelErr) { + retErr = errors.Join(cancelErr, retErr) + } + }() - var updatedPaths int64 - for _, pl := range pathLabels { - cwd := projects.DecodeCwdFromSessionPath(pl.path) + for _, path := range pathLabels { + if err := ctx.Err(); err != nil { + return 0, err + } + cwd := projects.DecodeCwdFromSessionPath(path.path) if cwd == "" { - continue // can't decode path, skip + continue } - - // Normalize cross-host equivalences cwd = projects.NormalizeRootEquivalence(cwd, registry) - - // Try registry match - id := projects.Identify(cwd, registry) - if !id.FromRegistry { - continue // registry didn't match, skip (no churn on fallback-only) + identity := projects.Identify(cwd, registry) + if err := ctx.Err(); err != nil { + return 0, err } - - // Update only if the registry ID differs from stored project - if id.ProjectID != pl.project { - _, err := tx.ExecContext(ctx, ` - UPDATE search_items SET project = ? WHERE source_path = ? - `, id.ProjectID, pl.path) - if err != nil { - return 0, fmt.Errorf("update source_path %s: %w", pl.path, err) + if !identity.FromRegistry || identity.ProjectID == path.project { + continue + } + if _, err := tx.ExecContext(ctx, ` + UPDATE search_items SET project = ? WHERE source_path = ? + `, identity.ProjectID, path.path); err != nil { + if ctxErr := ctx.Err(); ctxErr != nil { + return 0, ctxErr } - updatedPaths++ + return 0, fmt.Errorf("update source_path %s: %w", path.path, err) } + updatedPaths++ } - - if err := tx.Commit(); err != nil { + if err := gate.commit(); err != nil { return 0, fmt.Errorf("commit registry re-resolution transaction: %w", err) } - return updatedPaths, nil } diff --git a/internal/storage/rebuild_context_test.go b/internal/storage/rebuild_context_test.go new file mode 100644 index 0000000..b5d4dad --- /dev/null +++ b/internal/storage/rebuild_context_test.go @@ -0,0 +1,514 @@ +package storage + +import ( + "context" + "errors" + "fmt" + "path/filepath" + "reflect" + "sync/atomic" + "testing" + + "github.com/pablontiv/backscroll/internal/projects" +) + +func TestRebuildFTSContextPreservesExactRoutingAndPerennialRows(t *testing.T) { + db, err := Open(filepath.Join(t.TempDir(), "rebuild.db")) + if err != nil { + t.Fatal(err) + } + defer func() { _ = db.Close() }() + + messages := []IndexedMessage{ + {Ordinal: 0, UUID: "prose", Role: "user", Text: "proseonlytoken", ContentType: "text", Timestamp: "2026-01-01T00:00:00Z"}, + {Ordinal: 1, UUID: "code", Role: "assistant", Text: "codeonlytoken", ContentType: "code", Timestamp: "2026-01-01T00:00:01Z"}, + {Ordinal: 2, UUID: "reasoning", Role: "assistant", Text: "reasononlytoken", ContentType: "reasoning", Timestamp: "2026-01-01T00:00:02Z"}, + {Ordinal: 3, UUID: "tool", Role: "assistant", Text: "toolonlytoken", ContentType: "tool", Timestamp: "2026-01-01T00:00:03Z"}, + } + if err := db.SyncFiles([]IndexedFile{{Source: "session", SourcePath: "/expired/session.jsonl", Hash: "h", Messages: messages}}); err != nil { + t.Fatal(err) + } + + // Put every row in both indexes to prove rebuild repairs routing rather than + // merely making the expected rows searchable. + if _, err := db.db.Exec(`INSERT INTO messages_fts(messages_fts) VALUES('rebuild')`); err != nil { + t.Fatalf("contaminate messages index: %v", err) + } + if _, err := db.db.Exec(`INSERT INTO tool_fts(tool_fts) VALUES('rebuild')`); err != nil { + t.Fatalf("contaminate tool index: %v", err) + } + if err := db.RebuildFTSContext(context.Background()); err != nil { + t.Fatalf("RebuildFTSContext: %v", err) + } + + assertFTSHits := func(table, term string, want int) { + t.Helper() + var got int + query := fmt.Sprintf(`SELECT COUNT(*) FROM %s WHERE %s MATCH ?`, table, table) + if err := db.db.QueryRow(query, term).Scan(&got); err != nil { + t.Fatalf("query %s for %s: %v", table, term, err) + } + if got != want { + t.Fatalf("%s hits for %s = %d, want %d", table, term, got, want) + } + } + for _, term := range []string{"proseonlytoken", "codeonlytoken", "reasononlytoken"} { + assertFTSHits("messages_fts", term, 1) + assertFTSHits("tool_fts", term, 0) + } + assertFTSHits("tool_fts", "toolonlytoken", 1) + assertFTSHits("messages_fts", "toolonlytoken", 0) + + var rows int + if err := db.db.QueryRow(`SELECT COUNT(*) FROM search_items WHERE source_path = '/expired/session.jsonl'`).Scan(&rows); err != nil { + t.Fatal(err) + } + if rows != len(messages) { + t.Fatalf("perennial rows after rebuild = %d, want %d", rows, len(messages)) + } +} + +func TestRebuildStorageContextVariantsRejectPreCanceledContext(t *testing.T) { + db, cleanup := newTestDB(t) + defer cleanup() + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + tests := []struct { + name string + run func() error + }{ + {"RebuildFTSContext", func() error { return db.RebuildFTSContext(ctx) }}, + {"BackfillDerivedContext", func() error { return db.BackfillDerivedContext(ctx, BackfillDerivedOpts{}) }}, + {"ReresolveProjectsContext", func() error { + _, err := db.ReresolveProjectsContext(ctx, func(string) string { return "project" }) + return err + }}, + {"ReresolveProjectsWithRegistryContext", func() error { + _, err := db.ReresolveProjectsWithRegistryContext(ctx, projects.ProjectRegistry{}) + return err + }}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + if err := test.run(); !errors.Is(err, context.Canceled) { + t.Fatalf("error = %v, want context.Canceled", err) + } + }) + } +} + +func TestBackfillDerivedContextCancellationPreservesPriorBatchAndRollsBackActiveBatch(t *testing.T) { + db, cleanup := newTestDB(t) + defer cleanup() + + for i := 0; i < 101; i++ { + path := fmt.Sprintf("/backfill/%03d.jsonl", i) + if _, err := db.db.Exec(` + INSERT INTO search_items + (source_path, source, ordinal, role, text, timestamp, uuid, project, content_type, extraction_version) + VALUES (?, 'session', 0, 'user', 'no, eso no es un bug', '2026-01-01T00:00:00Z', ?, 'proj', 'text', 0) + `, path, fmt.Sprintf("backfill-%03d", i)); err != nil { + t.Fatalf("seed %s: %v", path, err) + } + } + + ctx, cancel := context.WithCancel(context.Background()) + setMaintenanceCancellation(t, cancel) + if _, err := db.db.Exec(fmt.Sprintf(` + CREATE TRIGGER cancel_second_backfill_batch AFTER INSERT ON correction_signals + WHEN NEW.source_path = '/backfill/100.jsonl' + BEGIN SELECT %s(); END + `, cancelMaintenanceFunction)); err != nil { + t.Fatalf("create cancellation trigger: %v", err) + } + var progress []int + err := db.BackfillDerivedContext(ctx, BackfillDerivedOpts{OnProgress: func(processed, _, _, _ int) { + progress = append(progress, processed) + }}) + if !errors.Is(err, context.Canceled) { + t.Fatalf("BackfillDerivedContext error = %v, want context.Canceled", err) + } + if len(progress) != 1 || progress[0] != 100 { + t.Fatalf("progress = %v, want only committed batch [100]", progress) + } + var committed, active int + if err := db.db.QueryRow(`SELECT COUNT(*) FROM correction_signals WHERE source_path < '/backfill/100.jsonl'`).Scan(&committed); err != nil { + t.Fatal(err) + } + if err := db.db.QueryRow(`SELECT COUNT(*) FROM correction_signals WHERE source_path = '/backfill/100.jsonl'`).Scan(&active); err != nil { + t.Fatal(err) + } + if committed != 100 || active != 0 { + t.Fatalf("signals after cancellation: prior=%d active=%d, want 100 and 0", committed, active) + } +} + +func TestReresolveProjectsContextCancellationRollsBackAllPaths(t *testing.T) { + db, cleanup := newTestDB(t) + defer cleanup() + for i, path := range []string{"/a.jsonl", "/b.jsonl", "/c.jsonl"} { + if _, err := db.db.Exec(` + INSERT INTO search_items + (source_path, source, ordinal, role, text, timestamp, uuid, project, content_type) + VALUES (?, 'session', 0, 'user', 'text', '2026-01-01T00:00:00Z', ?, 'unknown', 'text') + `, path, fmt.Sprintf("resolve-%d", i)); err != nil { + t.Fatal(err) + } + } + ctx, cancel := context.WithCancel(context.Background()) + resolved, err := db.ReresolveProjectsContext(ctx, func(path string) string { + if path == "/b.jsonl" { + cancel() + } + return "resolved" + }) + if !errors.Is(err, context.Canceled) { + t.Fatalf("error = %v, want context.Canceled", err) + } + if resolved != 0 { + t.Fatalf("resolved = %d, want 0 after rollback", resolved) + } + var unknown int + if err := db.db.QueryRow(`SELECT COUNT(*) FROM search_items WHERE project = 'unknown'`).Scan(&unknown); err != nil { + t.Fatal(err) + } + if unknown != 3 { + t.Fatalf("unknown rows after cancellation = %d, want 3", unknown) + } +} + +type cancelAfterErrChecksContext struct { + context.Context + cancelAt int64 + checks atomic.Int64 +} + +func (c *cancelAfterErrChecksContext) Done() <-chan struct{} { return nil } + +func (c *cancelAfterErrChecksContext) Err() error { + if c.checks.Add(1) >= c.cancelAt { + return context.Canceled + } + return nil +} + +type ftsRow struct { + rowID int64 + text string +} + +type ftsSnapshot struct { + messages []ftsRow + tools []ftsRow +} + +func TestRebuildFTSContextCancellationRestoresBothIndexesExactly(t *testing.T) { + for _, test := range []struct { + name string + cancelAt int64 + }{ + {name: "after_messages_repopulated", cancelAt: 4}, + {name: "at_commit_after_both_indexes_repopulated", cancelAt: 6}, + } { + t.Run(test.name, func(t *testing.T) { + db, cleanup := newTestDB(t) + defer cleanup() + seedRoutedFTSRows(t, db) + + // Deliberately index every content row in both indexes. A canceled + // selective rebuild must restore this exact pre-transaction state. + if _, err := db.db.Exec(`INSERT INTO messages_fts(messages_fts) VALUES('rebuild')`); err != nil { + t.Fatal(err) + } + if _, err := db.db.Exec(`INSERT INTO tool_fts(tool_fts) VALUES('rebuild')`); err != nil { + t.Fatal(err) + } + before := snapshotFTSRows(t, db, "snapshotall") + + ctx := &cancelAfterErrChecksContext{Context: context.Background(), cancelAt: test.cancelAt} + err := db.RebuildFTSContext(ctx) + if !errors.Is(err, context.Canceled) { + t.Fatalf("RebuildFTSContext error = %v, want context.Canceled (checks=%d)", err, ctx.checks.Load()) + } + after := snapshotFTSRows(t, db, "snapshotall") + if !reflect.DeepEqual(after, before) { + t.Fatalf("FTS indexes changed after rollback:\nbefore=%#v\nafter=%#v", before, after) + } + }) + } +} + +func TestRebuildFTSContextExactContentsRowIDsAndIdempotence(t *testing.T) { + db, cleanup := newTestDB(t) + defer cleanup() + seedRoutedFTSRows(t, db) + + // Start from known cross-contamination so the first call has real repair work. + if _, err := db.db.Exec(`INSERT INTO messages_fts(messages_fts) VALUES('rebuild')`); err != nil { + t.Fatal(err) + } + if _, err := db.db.Exec(`INSERT INTO tool_fts(tool_fts) VALUES('rebuild')`); err != nil { + t.Fatal(err) + } + if err := db.RebuildFTSContext(context.Background()); err != nil { + t.Fatalf("first rebuild: %v", err) + } + first := snapshotFTSRows(t, db, "snapshotall") + wantMessages := searchItemRows(t, db, `content_type IN ('text', 'code', 'reasoning')`) + wantTools := searchItemRows(t, db, `content_type = 'tool'`) + if !reflect.DeepEqual(first.messages, wantMessages) { + t.Fatalf("messages_fts rows = %#v, want routed search_items %#v", first.messages, wantMessages) + } + if !reflect.DeepEqual(first.tools, wantTools) { + t.Fatalf("tool_fts rows = %#v, want routed search_items %#v", first.tools, wantTools) + } + assertUniqueFTSRowIDs(t, "messages_fts", first.messages) + assertUniqueFTSRowIDs(t, "tool_fts", first.tools) + + if err := db.RebuildFTSContext(context.Background()); err != nil { + t.Fatalf("second rebuild: %v", err) + } + second := snapshotFTSRows(t, db, "snapshotall") + if !reflect.DeepEqual(second, first) { + t.Fatalf("second rebuild was not exactly idempotent:\nfirst=%#v\nsecond=%#v", first, second) + } +} + +func seedRoutedFTSRows(t *testing.T, db *Database) { + t.Helper() + messages := []IndexedMessage{ + {Ordinal: 0, UUID: "route-text", Role: "user", Text: "snapshotall proseunique", ContentType: "text", Timestamp: "2026-01-01T00:00:00Z"}, + {Ordinal: 1, UUID: "route-code", Role: "assistant", Text: "snapshotall codeunique", ContentType: "code", Timestamp: "2026-01-01T00:00:01Z"}, + {Ordinal: 2, UUID: "route-reasoning", Role: "assistant", Text: "snapshotall reasonunique", ContentType: "reasoning", Timestamp: "2026-01-01T00:00:02Z"}, + {Ordinal: 3, UUID: "route-tool-a", Role: "assistant", Text: "snapshotall tooluniquealpha", ContentType: "tool", Timestamp: "2026-01-01T00:00:03Z"}, + {Ordinal: 4, UUID: "route-tool-b", Role: "assistant", Text: "snapshotall tooluniquebeta", ContentType: "tool", Timestamp: "2026-01-01T00:00:04Z"}, + } + if err := db.SyncFiles([]IndexedFile{{Source: "session", SourcePath: "/expired/routed.jsonl", Hash: "route-hash", Messages: messages}}); err != nil { + t.Fatal(err) + } +} + +func snapshotFTSRows(t *testing.T, db *Database, token string) ftsSnapshot { + t.Helper() + return ftsSnapshot{ + messages: matchedFTSRows(t, db, "messages_fts", token), + tools: matchedFTSRows(t, db, "tool_fts", token), + } +} + +func matchedFTSRows(t *testing.T, db *Database, table, token string) []ftsRow { + t.Helper() + query := fmt.Sprintf(`SELECT rowid, text FROM %s WHERE %s MATCH ? ORDER BY rowid, text`, table, table) + rows, err := db.db.Query(query, token) + if err != nil { + t.Fatalf("query %s: %v", table, err) + } + defer func() { _ = rows.Close() }() + var out []ftsRow + for rows.Next() { + var row ftsRow + if err := rows.Scan(&row.rowID, &row.text); err != nil { + t.Fatalf("scan %s: %v", table, err) + } + out = append(out, row) + } + if err := rows.Err(); err != nil { + t.Fatalf("iterate %s: %v", table, err) + } + return out +} + +func searchItemRows(t *testing.T, db *Database, predicate string) []ftsRow { + t.Helper() + rows, err := db.db.Query(`SELECT id, text FROM search_items WHERE ` + predicate + ` ORDER BY id, text`) + if err != nil { + t.Fatal(err) + } + defer func() { _ = rows.Close() }() + var out []ftsRow + for rows.Next() { + var row ftsRow + if err := rows.Scan(&row.rowID, &row.text); err != nil { + t.Fatal(err) + } + out = append(out, row) + } + if err := rows.Err(); err != nil { + t.Fatal(err) + } + return out +} + +func assertUniqueFTSRowIDs(t *testing.T, table string, rows []ftsRow) { + t.Helper() + seen := make(map[int64]struct{}, len(rows)) + for _, row := range rows { + if _, duplicate := seen[row.rowID]; duplicate { + t.Fatalf("%s contains duplicate rowid %d: %#v", table, row.rowID, rows) + } + seen[row.rowID] = struct{}{} + } +} + +func TestBackfillDerivedContextCancellationFromFinalProgressCallback(t *testing.T) { + db, cleanup := newTestDB(t) + defer cleanup() + const path = "/backfill/final.jsonl" + if _, err := db.db.Exec(` + INSERT INTO search_items + (source_path, source, ordinal, role, text, timestamp, uuid, project, content_type, extraction_version) + VALUES (?, 'session', 0, 'user', 'no, eso no es un bug', '2026-01-01T00:00:00Z', 'final-progress', 'proj', 'text', 0) + `, path); err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + progressCalls := 0 + err := db.BackfillDerivedContext(ctx, BackfillDerivedOpts{OnProgress: func(processed, _, _, _ int) { + progressCalls++ + if processed != 1 { + t.Errorf("processed = %d, want 1", processed) + } + cancel() + }}) + if !errors.Is(err, context.Canceled) { + t.Fatalf("BackfillDerivedContext error = %v, want context.Canceled", err) + } + if progressCalls != 1 { + t.Fatalf("progress calls = %d, want 1", progressCalls) + } + var committed int + if err := db.db.QueryRow(`SELECT COUNT(*) FROM correction_signals WHERE source_path = ?`, path).Scan(&committed); err != nil { + t.Fatal(err) + } + if committed == 0 { + t.Fatal("final batch did not remain committed after progress callback cancellation") + } +} + +func TestReresolveProjectsWithRegistryContextCancellationRollsBack(t *testing.T) { + for _, test := range []struct { + name string + cancelAt int64 + }{ + {name: "during_scan", cancelAt: 2}, + {name: "during_resolve_after_first_update", cancelAt: 7}, + {name: "at_commit", cancelAt: 9}, + } { + t.Run(test.name, func(t *testing.T) { + db, cleanup := newTestDB(t) + defer cleanup() + registry := seedRegistryResolutionRows(t, db) + before := projectRows(t, db) + ctx := &cancelAfterErrChecksContext{Context: context.Background(), cancelAt: test.cancelAt} + updated, err := db.ReresolveProjectsWithRegistryContext(ctx, registry) + if !errors.Is(err, context.Canceled) { + t.Fatalf("error = %v, want context.Canceled (updated=%d checks=%d)", err, updated, ctx.checks.Load()) + } + if updated != 0 { + t.Fatalf("updated = %d, want 0 after rollback", updated) + } + after := projectRows(t, db) + if !reflect.DeepEqual(after, before) { + t.Fatalf("projects changed after cancellation: before=%v after=%v", before, after) + } + }) + } +} + +func TestReresolveLegacyWrappersMatchContextVariants(t *testing.T) { + t.Run("resolver", func(t *testing.T) { + contextDB, contextCleanup := newTestDB(t) + defer contextCleanup() + legacyDB, legacyCleanup := newTestDB(t) + defer legacyCleanup() + for _, db := range []*Database{contextDB, legacyDB} { + seedUnknownResolutionRows(t, db) + } + resolver := func(path string) string { return "resolved-" + filepath.Base(path) } + got, err := contextDB.ReresolveProjectsContext(context.Background(), resolver) + if err != nil { + t.Fatal(err) + } + want, err := legacyDB.ReresolveProjects(context.Background(), resolver) + if err != nil { + t.Fatal(err) + } + if got != want || !reflect.DeepEqual(projectRows(t, contextDB), projectRows(t, legacyDB)) { + t.Fatalf("context count/rows = %d/%v, legacy = %d/%v", got, projectRows(t, contextDB), want, projectRows(t, legacyDB)) + } + }) + + t.Run("registry", func(t *testing.T) { + contextDB, contextCleanup := newTestDB(t) + defer contextCleanup() + legacyDB, legacyCleanup := newTestDB(t) + defer legacyCleanup() + contextRegistry := seedRegistryResolutionRows(t, contextDB) + legacyRegistry := seedRegistryResolutionRows(t, legacyDB) + got, err := contextDB.ReresolveProjectsWithRegistryContext(context.Background(), contextRegistry) + if err != nil { + t.Fatal(err) + } + want, err := legacyDB.ReresolveProjectsWithRegistry(context.Background(), legacyRegistry) + if err != nil { + t.Fatal(err) + } + if got != want || !reflect.DeepEqual(projectRows(t, contextDB), projectRows(t, legacyDB)) { + t.Fatalf("context count/rows = %d/%v, legacy = %d/%v", got, projectRows(t, contextDB), want, projectRows(t, legacyDB)) + } + }) +} + +func seedUnknownResolutionRows(t *testing.T, db *Database) { + t.Helper() + for i, path := range []string{"/wrapper/a.jsonl", "/wrapper/b.jsonl"} { + if _, err := db.db.Exec(` + INSERT INTO search_items + (source_path, source, ordinal, role, text, timestamp, uuid, project, content_type) + VALUES (?, 'session', 0, 'user', 'text', '2026-01-01T00:00:00Z', ?, 'unknown', 'text') + `, path, fmt.Sprintf("wrapper-%d", i)); err != nil { + t.Fatal(err) + } + } +} + +func seedRegistryResolutionRows(t *testing.T, db *Database) projects.ProjectRegistry { + t.Helper() + for i, path := range []string{ + "/home/test/.claude/projects/-registryalpha/a.jsonl", + "/home/test/.claude/projects/-registryalpha/b.jsonl", + } { + if _, err := db.db.Exec(` + INSERT INTO search_items + (source_path, source, ordinal, role, text, timestamp, uuid, project, content_type) + VALUES (?, 'session', 0, 'user', 'text', '2026-01-01T00:00:00Z', ?, 'fallback', 'text') + `, path, fmt.Sprintf("registry-%d", i)); err != nil { + t.Fatal(err) + } + } + return projects.ProjectRegistry{Projects: []projects.ProjectConfig{{ + ID: "canonical-registry", Roots: []string{"registryalpha"}, + }}} +} + +func projectRows(t *testing.T, db *Database) []string { + t.Helper() + rows, err := db.db.Query(`SELECT source_path || '=' || COALESCE(project, '') FROM search_items ORDER BY source_path, ordinal`) + if err != nil { + t.Fatal(err) + } + defer func() { _ = rows.Close() }() + var out []string + for rows.Next() { + var value string + if err := rows.Scan(&value); err != nil { + t.Fatal(err) + } + out = append(out, value) + } + if err := rows.Err(); err != nil { + t.Fatal(err) + } + return out +}